Skip to main content

matrix_sdk_base/store/
memory_store.rs

1// Copyright 2021 The Matrix.org Foundation C.I.C.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::{
16    cmp::Reverse,
17    collections::{BTreeMap, BTreeSet, HashMap},
18    sync::RwLock,
19};
20
21use async_trait::async_trait;
22use growable_bloom_filter::GrowableBloom;
23use matrix_sdk_common::{ROOM_VERSION_FALLBACK, ROOM_VERSION_RULES_FALLBACK, ttl::TtlValue};
24use ruma::{
25    CanonicalJsonObject, EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedMxcUri,
26    OwnedRoomId, OwnedTransactionId, OwnedUserId, RoomId, TransactionId, UserId,
27    api::client::discovery::get_capabilities::v3::Capabilities,
28    canonical_json::{RedactedBecause, redact},
29    events::{
30        AnyGlobalAccountDataEvent, AnyRoomAccountDataEvent, AnyStrippedStateEvent,
31        AnySyncStateEvent, GlobalAccountDataEventType, RoomAccountDataEventType, StateEventType,
32        presence::PresenceEvent,
33        receipt::{Receipt, ReceiptThread, ReceiptType},
34        room::member::{MembershipState, StrippedRoomMemberEvent, SyncRoomMemberEvent},
35    },
36    profile::{UserProfile, UserProfileUpdate},
37    serde::Raw,
38    time::Instant,
39};
40use tracing::{debug, instrument, warn};
41
42use super::{
43    DependentQueuedRequest, DependentQueuedRequestKind, QueuedRequestKind, Result, RoomInfo,
44    RoomLoadSettings, StateChanges, StateStore, StoreError, SupportedVersionsResponse,
45    WellKnownResponse,
46    send_queue::{ChildTransactionId, QueuedRequest, SentRequestKey},
47    traits::ComposerDraft,
48};
49use crate::{
50    MinimalRoomMemberEvent, RoomMemberships, StateStoreDataKey, StateStoreDataValue,
51    deserialized_responses::{DisplayName, RawAnySyncOrStrippedState},
52    store::{
53        QueueWedgeError, StoredThreadSubscription,
54        traits::{ThreadSubscriptionCatchupToken, compare_thread_subscription_bump_stamps},
55    },
56};
57
58#[derive(Debug, Default)]
59#[allow(clippy::type_complexity)]
60struct MemoryStoreInner {
61    recently_visited_rooms: HashMap<OwnedUserId, Vec<OwnedRoomId>>,
62    composer_drafts: HashMap<(OwnedRoomId, Option<OwnedEventId>), ComposerDraft>,
63    user_avatar_url: HashMap<OwnedUserId, OwnedMxcUri>,
64    sync_token: Option<String>,
65    supported_versions: Option<TtlValue<SupportedVersionsResponse>>,
66    well_known: Option<TtlValue<Option<WellKnownResponse>>>,
67    filters: HashMap<String, String>,
68    utd_hook_manager_data: Option<GrowableBloom>,
69    one_time_key_uploaded_error: bool,
70    account_data: HashMap<GlobalAccountDataEventType, Raw<AnyGlobalAccountDataEvent>>,
71    profiles: HashMap<OwnedRoomId, HashMap<OwnedUserId, MinimalRoomMemberEvent>>,
72    display_names: HashMap<OwnedRoomId, HashMap<DisplayName, BTreeSet<OwnedUserId>>>,
73    members: HashMap<OwnedRoomId, HashMap<OwnedUserId, MembershipState>>,
74    room_info: HashMap<OwnedRoomId, RoomInfo>,
75    room_state:
76        HashMap<OwnedRoomId, HashMap<StateEventType, HashMap<String, Raw<AnySyncStateEvent>>>>,
77    room_account_data:
78        HashMap<OwnedRoomId, HashMap<RoomAccountDataEventType, Raw<AnyRoomAccountDataEvent>>>,
79    stripped_room_state:
80        HashMap<OwnedRoomId, HashMap<StateEventType, HashMap<String, Raw<AnyStrippedStateEvent>>>>,
81    stripped_members: HashMap<OwnedRoomId, HashMap<OwnedUserId, MembershipState>>,
82    presence: HashMap<OwnedUserId, Raw<PresenceEvent>>,
83    room_user_receipts: HashMap<
84        OwnedRoomId,
85        HashMap<(String, Option<String>), HashMap<OwnedUserId, (OwnedEventId, Receipt)>>,
86    >,
87    room_event_receipts: HashMap<
88        OwnedRoomId,
89        HashMap<(String, Option<String>), HashMap<OwnedEventId, HashMap<OwnedUserId, Receipt>>>,
90    >,
91    custom: HashMap<Vec<u8>, Vec<u8>>,
92    send_queue_events: BTreeMap<OwnedRoomId, Vec<QueuedRequest>>,
93    dependent_send_queue_events: BTreeMap<OwnedRoomId, Vec<DependentQueuedRequest>>,
94    seen_knock_requests: BTreeMap<OwnedRoomId, BTreeMap<OwnedEventId, OwnedUserId>>,
95    thread_subscriptions: BTreeMap<OwnedRoomId, BTreeMap<OwnedEventId, StoredThreadSubscription>>,
96    thread_subscriptions_catchup_tokens: Option<Vec<ThreadSubscriptionCatchupToken>>,
97    global_profiles: HashMap<OwnedUserId, UserProfile>,
98    homeserver_capabilities: Option<TtlValue<Capabilities>>,
99}
100
101/// In-memory, non-persistent implementation of the `StateStore`.
102///
103/// Default if no other is configured at startup.
104#[derive(Debug, Default)]
105pub struct MemoryStore {
106    inner: RwLock<MemoryStoreInner>,
107}
108
109impl MemoryStore {
110    /// Create a new empty MemoryStore
111    pub fn new() -> Self {
112        Self::default()
113    }
114
115    fn get_user_room_receipt_event_impl(
116        &self,
117        room_id: &RoomId,
118        receipt_type: ReceiptType,
119        thread: ReceiptThread,
120        user_id: &UserId,
121    ) -> Option<(OwnedEventId, Receipt)> {
122        self.inner
123            .read()
124            .unwrap()
125            .room_user_receipts
126            .get(room_id)?
127            .get(&(receipt_type.to_string(), thread.as_str().map(ToOwned::to_owned)))?
128            .get(user_id)
129            .cloned()
130    }
131
132    fn get_event_room_receipt_events_impl(
133        &self,
134        room_id: &RoomId,
135        receipt_type: ReceiptType,
136        thread: ReceiptThread,
137        event_id: &EventId,
138    ) -> Option<Vec<(OwnedUserId, Receipt)>> {
139        Some(
140            self.inner
141                .read()
142                .unwrap()
143                .room_event_receipts
144                .get(room_id)?
145                .get(&(receipt_type.to_string(), thread.as_str().map(ToOwned::to_owned)))?
146                .get(event_id)?
147                .iter()
148                .map(|(key, value)| (key.clone(), value.clone()))
149                .collect(),
150        )
151    }
152}
153
154#[cfg_attr(target_family = "wasm", async_trait(?Send))]
155#[cfg_attr(not(target_family = "wasm"), async_trait)]
156impl StateStore for MemoryStore {
157    type Error = StoreError;
158
159    async fn close(&self) -> Result<(), Self::Error> {
160        Ok(())
161    }
162
163    async fn reopen(&self) -> Result<(), Self::Error> {
164        Ok(())
165    }
166
167    async fn get_kv_data(&self, key: StateStoreDataKey<'_>) -> Result<Option<StateStoreDataValue>> {
168        let inner = self.inner.read().unwrap();
169
170        Ok(match key {
171            StateStoreDataKey::SyncToken => {
172                inner.sync_token.clone().map(StateStoreDataValue::SyncToken)
173            }
174            StateStoreDataKey::SupportedVersions => {
175                inner.supported_versions.clone().map(StateStoreDataValue::SupportedVersions)
176            }
177            StateStoreDataKey::WellKnown => {
178                inner.well_known.clone().map(StateStoreDataValue::WellKnown)
179            }
180            StateStoreDataKey::Filter(filter_name) => {
181                inner.filters.get(filter_name).cloned().map(StateStoreDataValue::Filter)
182            }
183            StateStoreDataKey::UserAvatarUrl(user_id) => {
184                inner.user_avatar_url.get(user_id).cloned().map(StateStoreDataValue::UserAvatarUrl)
185            }
186            StateStoreDataKey::RecentlyVisitedRooms(user_id) => inner
187                .recently_visited_rooms
188                .get(user_id)
189                .cloned()
190                .map(StateStoreDataValue::RecentlyVisitedRooms),
191            StateStoreDataKey::UtdHookManagerData => {
192                inner.utd_hook_manager_data.clone().map(StateStoreDataValue::UtdHookManagerData)
193            }
194            StateStoreDataKey::OneTimeKeyAlreadyUploaded => inner
195                .one_time_key_uploaded_error
196                .then_some(StateStoreDataValue::OneTimeKeyAlreadyUploaded),
197            StateStoreDataKey::ComposerDraft(room_id, thread_root) => {
198                let key = (room_id.to_owned(), thread_root.map(ToOwned::to_owned));
199                inner.composer_drafts.get(&key).cloned().map(StateStoreDataValue::ComposerDraft)
200            }
201            StateStoreDataKey::SeenKnockRequests(room_id) => inner
202                .seen_knock_requests
203                .get(room_id)
204                .cloned()
205                .map(StateStoreDataValue::SeenKnockRequests),
206            StateStoreDataKey::ThreadSubscriptionsCatchupTokens => inner
207                .thread_subscriptions_catchup_tokens
208                .clone()
209                .map(StateStoreDataValue::ThreadSubscriptionsCatchupTokens),
210            StateStoreDataKey::HomeserverCapabilities => inner
211                .homeserver_capabilities
212                .clone()
213                .map(StateStoreDataValue::HomeserverCapabilities),
214        })
215    }
216
217    async fn set_kv_data(
218        &self,
219        key: StateStoreDataKey<'_>,
220        value: StateStoreDataValue,
221    ) -> Result<()> {
222        let mut inner = self.inner.write().unwrap();
223        match key {
224            StateStoreDataKey::SyncToken => {
225                inner.sync_token =
226                    Some(value.into_sync_token().expect("Session data not a sync token"))
227            }
228            StateStoreDataKey::Filter(filter_name) => {
229                inner.filters.insert(
230                    filter_name.to_owned(),
231                    value.into_filter().expect("Session data not a filter"),
232                );
233            }
234            StateStoreDataKey::UserAvatarUrl(user_id) => {
235                inner.user_avatar_url.insert(
236                    user_id.to_owned(),
237                    value.into_user_avatar_url().expect("Session data not a user avatar url"),
238                );
239            }
240            StateStoreDataKey::RecentlyVisitedRooms(user_id) => {
241                inner.recently_visited_rooms.insert(
242                    user_id.to_owned(),
243                    value
244                        .into_recently_visited_rooms()
245                        .expect("Session data not a list of recently visited rooms"),
246                );
247            }
248            StateStoreDataKey::UtdHookManagerData => {
249                inner.utd_hook_manager_data = Some(
250                    value
251                        .into_utd_hook_manager_data()
252                        .expect("Session data not the hook manager data"),
253                );
254            }
255            StateStoreDataKey::OneTimeKeyAlreadyUploaded => {
256                inner.one_time_key_uploaded_error = true;
257            }
258            StateStoreDataKey::ComposerDraft(room_id, thread_root) => {
259                inner.composer_drafts.insert(
260                    (room_id.to_owned(), thread_root.map(ToOwned::to_owned)),
261                    value.into_composer_draft().expect("Session data not a composer draft"),
262                );
263            }
264            StateStoreDataKey::SupportedVersions => {
265                inner.supported_versions = Some(
266                    value
267                        .into_supported_versions()
268                        .expect("Session data not containing supported versions"),
269                );
270            }
271            StateStoreDataKey::WellKnown => {
272                inner.well_known =
273                    Some(value.into_well_known().expect("Session data not containing well-known"));
274            }
275            StateStoreDataKey::SeenKnockRequests(room_id) => {
276                inner.seen_knock_requests.insert(
277                    room_id.to_owned(),
278                    value
279                        .into_seen_knock_requests()
280                        .expect("Session data is not a set of seen join request ids"),
281                );
282            }
283            StateStoreDataKey::ThreadSubscriptionsCatchupTokens => {
284                inner.thread_subscriptions_catchup_tokens =
285                    Some(value.into_thread_subscriptions_catchup_tokens().expect(
286                        "Session data is not a list of thread subscription catchup tokens",
287                    ));
288            }
289            StateStoreDataKey::HomeserverCapabilities => {
290                inner.homeserver_capabilities = Some(
291                    value
292                        .into_homeserver_capabilities()
293                        .expect("Session data is not a homeserver capabilities"),
294                );
295            }
296        }
297
298        Ok(())
299    }
300
301    async fn remove_kv_data(&self, key: StateStoreDataKey<'_>) -> Result<()> {
302        let mut inner = self.inner.write().unwrap();
303        match key {
304            StateStoreDataKey::SyncToken => inner.sync_token = None,
305            StateStoreDataKey::SupportedVersions => inner.supported_versions = None,
306            StateStoreDataKey::WellKnown => inner.well_known = None,
307            StateStoreDataKey::Filter(filter_name) => {
308                inner.filters.remove(filter_name);
309            }
310            StateStoreDataKey::UserAvatarUrl(user_id) => {
311                inner.user_avatar_url.remove(user_id);
312            }
313            StateStoreDataKey::RecentlyVisitedRooms(user_id) => {
314                inner.recently_visited_rooms.remove(user_id);
315            }
316            StateStoreDataKey::UtdHookManagerData => inner.utd_hook_manager_data = None,
317            StateStoreDataKey::OneTimeKeyAlreadyUploaded => {
318                inner.one_time_key_uploaded_error = false
319            }
320            StateStoreDataKey::ComposerDraft(room_id, thread_root) => {
321                let key = (room_id.to_owned(), thread_root.map(ToOwned::to_owned));
322                inner.composer_drafts.remove(&key);
323            }
324            StateStoreDataKey::SeenKnockRequests(room_id) => {
325                inner.seen_knock_requests.remove(room_id);
326            }
327            StateStoreDataKey::ThreadSubscriptionsCatchupTokens => {
328                inner.thread_subscriptions_catchup_tokens = None;
329            }
330            StateStoreDataKey::HomeserverCapabilities => inner.homeserver_capabilities = None,
331        }
332        Ok(())
333    }
334
335    #[instrument(skip(self, changes))]
336    async fn save_changes(&self, changes: &StateChanges) -> Result<()> {
337        let now = Instant::now();
338
339        let mut inner = self.inner.write().unwrap();
340
341        if let Some(s) = &changes.sync_token {
342            inner.sync_token = Some(s.to_owned());
343        }
344
345        for (room, users) in &changes.profiles_to_delete {
346            let Some(room_profiles) = inner.profiles.get_mut(room) else {
347                continue;
348            };
349            for user in users {
350                room_profiles.remove(user);
351            }
352        }
353
354        for (room, users) in &changes.profiles {
355            for (user_id, profile) in users {
356                inner
357                    .profiles
358                    .entry(room.clone())
359                    .or_default()
360                    .insert(user_id.clone(), profile.clone());
361            }
362        }
363
364        for (room, map) in &changes.ambiguity_maps {
365            for (display_name, display_names) in map {
366                inner
367                    .display_names
368                    .entry(room.clone())
369                    .or_default()
370                    .insert(display_name.clone(), display_names.clone());
371            }
372        }
373
374        for (event_type, event) in &changes.account_data {
375            inner.account_data.insert(event_type.clone(), event.clone());
376        }
377
378        for (room, events) in &changes.room_account_data {
379            for (event_type, event) in events {
380                inner
381                    .room_account_data
382                    .entry(room.clone())
383                    .or_default()
384                    .insert(event_type.clone(), event.clone());
385            }
386        }
387
388        for (room, event_types) in &changes.state {
389            for (event_type, events) in event_types {
390                for (state_key, raw_event) in events {
391                    inner
392                        .room_state
393                        .entry(room.clone())
394                        .or_default()
395                        .entry(event_type.clone())
396                        .or_default()
397                        .insert(state_key.to_owned(), raw_event.clone());
398                    inner.stripped_room_state.remove(room);
399
400                    if *event_type == StateEventType::RoomMember {
401                        let event =
402                            match raw_event.deserialize_as_unchecked::<SyncRoomMemberEvent>() {
403                                Ok(ev) => ev,
404                                Err(e) => {
405                                    let event_id: Option<String> =
406                                        raw_event.get_field("event_id").ok().flatten();
407                                    debug!(event_id, "Failed to deserialize member event: {e}");
408                                    continue;
409                                }
410                            };
411
412                        inner.stripped_members.remove(room);
413
414                        inner
415                            .members
416                            .entry(room.clone())
417                            .or_default()
418                            .insert(event.state_key().to_owned(), event.membership().clone());
419                    }
420                }
421            }
422        }
423
424        for (room_id, info) in &changes.room_infos {
425            inner.room_info.insert(room_id.clone(), info.clone());
426        }
427
428        for (sender, event) in &changes.presence {
429            inner.presence.insert(sender.clone(), event.clone());
430        }
431
432        for (room, event_types) in &changes.stripped_state {
433            for (event_type, events) in event_types {
434                for (state_key, raw_event) in events {
435                    inner
436                        .stripped_room_state
437                        .entry(room.clone())
438                        .or_default()
439                        .entry(event_type.clone())
440                        .or_default()
441                        .insert(state_key.to_owned(), raw_event.clone());
442
443                    if *event_type == StateEventType::RoomMember {
444                        let event =
445                            match raw_event.deserialize_as_unchecked::<StrippedRoomMemberEvent>() {
446                                Ok(ev) => ev,
447                                Err(e) => {
448                                    let event_id: Option<String> =
449                                        raw_event.get_field("event_id").ok().flatten();
450                                    debug!(
451                                        event_id,
452                                        "Failed to deserialize stripped member event: {e}"
453                                    );
454                                    continue;
455                                }
456                            };
457
458                        inner
459                            .stripped_members
460                            .entry(room.clone())
461                            .or_default()
462                            .insert(event.state_key, event.content.membership.clone());
463                    }
464                }
465            }
466        }
467
468        for (room, content) in &changes.receipts {
469            for (event_id, receipts) in &content.0 {
470                for (receipt_type, receipts) in receipts {
471                    for (user_id, receipt) in receipts {
472                        let thread = receipt.thread.as_str().map(ToOwned::to_owned);
473                        // Add the receipt to the room user receipts
474                        if let Some((old_event, _)) = inner
475                            .room_user_receipts
476                            .entry(room.clone())
477                            .or_default()
478                            .entry((receipt_type.to_string(), thread.clone()))
479                            .or_default()
480                            .insert(user_id.clone(), (event_id.clone(), receipt.clone()))
481                        {
482                            // Remove the old receipt from the room event receipts
483                            if let Some(receipt_map) = inner.room_event_receipts.get_mut(room)
484                                && let Some(event_map) =
485                                    receipt_map.get_mut(&(receipt_type.to_string(), thread.clone()))
486                                && let Some(user_map) = event_map.get_mut(&old_event)
487                            {
488                                user_map.remove(user_id);
489                            }
490                        }
491
492                        // Add the receipt to the room event receipts
493                        inner
494                            .room_event_receipts
495                            .entry(room.clone())
496                            .or_default()
497                            .entry((receipt_type.to_string(), thread))
498                            .or_default()
499                            .entry(event_id.clone())
500                            .or_default()
501                            .insert(user_id.clone(), receipt.clone());
502                    }
503                }
504            }
505        }
506
507        let make_redaction_rules = |room_info: &HashMap<OwnedRoomId, RoomInfo>, room_id| {
508            room_info.get(room_id).map(|info| info.room_version_rules_or_default()).unwrap_or_else(|| {
509                warn!(
510                    ?room_id,
511                    "Unable to get the room version rules, defaulting to rules for room version {ROOM_VERSION_FALLBACK}"
512                );
513                ROOM_VERSION_RULES_FALLBACK
514            }).redaction
515        };
516
517        let inner = &mut *inner;
518        for (room_id, redactions) in &changes.redactions {
519            let mut redaction_rules = None;
520
521            if let Some(room) = inner.room_state.get_mut(room_id) {
522                for ref_room_mu in room.values_mut() {
523                    for raw_evt in ref_room_mu.values_mut() {
524                        if let Ok(Some(event_id)) = raw_evt.get_field::<OwnedEventId>("event_id")
525                            && let Some(redaction) = redactions.get(&event_id)
526                        {
527                            let redacted = redact(
528                                raw_evt.deserialize_as::<CanonicalJsonObject>()?,
529                                redaction_rules.get_or_insert_with(|| {
530                                    make_redaction_rules(&inner.room_info, room_id)
531                                }),
532                                Some(RedactedBecause::from_raw_event(redaction)?),
533                            )
534                            .map_err(StoreError::Redaction)?;
535                            *raw_evt = Raw::new(&redacted)?.cast_unchecked();
536                        }
537                    }
538                }
539            }
540        }
541
542        for (user_id, profile_update) in &changes.global_profiles {
543            match profile_update {
544                UserProfileUpdate::Updated(profile_changes) => {
545                    inner
546                        .global_profiles
547                        .entry(user_id.clone())
548                        .or_default()
549                        .apply(profile_changes.clone());
550                }
551                UserProfileUpdate::Dropped => {
552                    inner.global_profiles.remove(user_id);
553                }
554                _ => {
555                    warn!(%user_id, "Unhandled UserProfileUpdate variant; ignoring");
556                }
557            }
558        }
559
560        debug!("Saved changes in {:?}", now.elapsed());
561
562        Ok(())
563    }
564
565    async fn get_presence_event(&self, user_id: &UserId) -> Result<Option<Raw<PresenceEvent>>> {
566        Ok(self.inner.read().unwrap().presence.get(user_id).cloned())
567    }
568
569    async fn get_presence_events(
570        &self,
571        user_ids: &[OwnedUserId],
572    ) -> Result<Vec<Raw<PresenceEvent>>> {
573        let presence = &self.inner.read().unwrap().presence;
574        Ok(user_ids.iter().filter_map(|user_id| presence.get(user_id).cloned()).collect())
575    }
576
577    async fn get_state_event(
578        &self,
579        room_id: &RoomId,
580        event_type: StateEventType,
581        state_key: &str,
582    ) -> Result<Option<RawAnySyncOrStrippedState>> {
583        Ok(self
584            .get_state_events_for_keys(room_id, event_type, &[state_key])
585            .await?
586            .into_iter()
587            .next())
588    }
589
590    async fn get_state_events(
591        &self,
592        room_id: &RoomId,
593        event_type: StateEventType,
594    ) -> Result<Vec<RawAnySyncOrStrippedState>> {
595        fn get_events<T>(
596            state_map: &HashMap<OwnedRoomId, HashMap<StateEventType, HashMap<String, Raw<T>>>>,
597            room_id: &RoomId,
598            event_type: &StateEventType,
599            to_enum: fn(Raw<T>) -> RawAnySyncOrStrippedState,
600        ) -> Option<Vec<RawAnySyncOrStrippedState>> {
601            let state_events = state_map.get(room_id)?.get(event_type)?;
602            Some(state_events.values().cloned().map(to_enum).collect())
603        }
604
605        let inner = self.inner.read().unwrap();
606        Ok(get_events(
607            &inner.stripped_room_state,
608            room_id,
609            &event_type,
610            RawAnySyncOrStrippedState::Stripped,
611        )
612        .or_else(|| {
613            get_events(&inner.room_state, room_id, &event_type, RawAnySyncOrStrippedState::Sync)
614        })
615        .unwrap_or_default())
616    }
617
618    async fn get_state_events_for_keys(
619        &self,
620        room_id: &RoomId,
621        event_type: StateEventType,
622        state_keys: &[&str],
623    ) -> Result<Vec<RawAnySyncOrStrippedState>, Self::Error> {
624        let inner = self.inner.read().unwrap();
625
626        if let Some(stripped_state_events) =
627            inner.stripped_room_state.get(room_id).and_then(|events| events.get(&event_type))
628        {
629            Ok(state_keys
630                .iter()
631                .filter_map(|k| {
632                    stripped_state_events
633                        .get(*k)
634                        .map(|e| RawAnySyncOrStrippedState::Stripped(e.clone()))
635                })
636                .collect())
637        } else if let Some(sync_state_events) =
638            inner.room_state.get(room_id).and_then(|events| events.get(&event_type))
639        {
640            Ok(state_keys
641                .iter()
642                .filter_map(|k| {
643                    sync_state_events.get(*k).map(|e| RawAnySyncOrStrippedState::Sync(e.clone()))
644                })
645                .collect())
646        } else {
647            Ok(Vec::new())
648        }
649    }
650
651    async fn get_profile(
652        &self,
653        room_id: &RoomId,
654        user_id: &UserId,
655    ) -> Result<Option<MinimalRoomMemberEvent>> {
656        Ok(self
657            .inner
658            .read()
659            .unwrap()
660            .profiles
661            .get(room_id)
662            .and_then(|room_profiles| room_profiles.get(user_id))
663            .cloned())
664    }
665
666    async fn get_profiles<'a>(
667        &self,
668        room_id: &RoomId,
669        user_ids: &'a [OwnedUserId],
670    ) -> Result<BTreeMap<&'a UserId, MinimalRoomMemberEvent>> {
671        if user_ids.is_empty() {
672            return Ok(BTreeMap::new());
673        }
674
675        let profiles = &self.inner.read().unwrap().profiles;
676        let Some(room_profiles) = profiles.get(room_id) else {
677            return Ok(BTreeMap::new());
678        };
679
680        Ok(user_ids
681            .iter()
682            .filter_map(|user_id| room_profiles.get(user_id).map(|p| (&**user_id, p.clone())))
683            .collect())
684    }
685
686    #[instrument(skip(self, memberships))]
687    async fn get_user_ids(
688        &self,
689        room_id: &RoomId,
690        memberships: RoomMemberships,
691    ) -> Result<Vec<OwnedUserId>> {
692        /// Get the user IDs for the given room with the given memberships and
693        /// stripped state.
694        ///
695        /// If `memberships` is empty, returns all user IDs in the room with the
696        /// given stripped state.
697        fn get_user_ids_inner(
698            members: &HashMap<OwnedRoomId, HashMap<OwnedUserId, MembershipState>>,
699            room_id: &RoomId,
700            memberships: RoomMemberships,
701        ) -> Vec<OwnedUserId> {
702            members
703                .get(room_id)
704                .map(|members| {
705                    members
706                        .iter()
707                        .filter_map(|(user_id, membership)| {
708                            memberships.matches(membership).then_some(user_id)
709                        })
710                        .cloned()
711                        .collect()
712                })
713                .unwrap_or_default()
714        }
715        let inner = self.inner.read().unwrap();
716        let v = get_user_ids_inner(&inner.stripped_members, room_id, memberships);
717        if !v.is_empty() {
718            return Ok(v);
719        }
720        Ok(get_user_ids_inner(&inner.members, room_id, memberships))
721    }
722
723    async fn get_room_infos(&self, room_load_settings: &RoomLoadSettings) -> Result<Vec<RoomInfo>> {
724        let memory_store_inner = self.inner.read().unwrap();
725        let room_infos = &memory_store_inner.room_info;
726
727        Ok(match room_load_settings {
728            RoomLoadSettings::All => room_infos.values().cloned().collect(),
729
730            RoomLoadSettings::One(room_id) => match room_infos.get(room_id) {
731                Some(room_info) => vec![room_info.clone()],
732                None => vec![],
733            },
734        })
735    }
736
737    async fn get_users_with_display_name(
738        &self,
739        room_id: &RoomId,
740        display_name: &DisplayName,
741    ) -> Result<BTreeSet<OwnedUserId>> {
742        Ok(self
743            .inner
744            .read()
745            .unwrap()
746            .display_names
747            .get(room_id)
748            .and_then(|room_names| room_names.get(display_name).cloned())
749            .unwrap_or_default())
750    }
751
752    async fn get_users_with_display_names<'a>(
753        &self,
754        room_id: &RoomId,
755        display_names: &'a [DisplayName],
756    ) -> Result<HashMap<&'a DisplayName, BTreeSet<OwnedUserId>>> {
757        if display_names.is_empty() {
758            return Ok(HashMap::new());
759        }
760
761        let inner = self.inner.read().unwrap();
762        let Some(room_names) = inner.display_names.get(room_id) else {
763            return Ok(HashMap::new());
764        };
765
766        Ok(display_names.iter().filter_map(|n| room_names.get(n).map(|d| (n, d.clone()))).collect())
767    }
768
769    async fn get_account_data_event(
770        &self,
771        event_type: GlobalAccountDataEventType,
772    ) -> Result<Option<Raw<AnyGlobalAccountDataEvent>>> {
773        Ok(self.inner.read().unwrap().account_data.get(&event_type).cloned())
774    }
775
776    async fn get_room_account_data_event(
777        &self,
778        room_id: &RoomId,
779        event_type: RoomAccountDataEventType,
780    ) -> Result<Option<Raw<AnyRoomAccountDataEvent>>> {
781        Ok(self
782            .inner
783            .read()
784            .unwrap()
785            .room_account_data
786            .get(room_id)
787            .and_then(|m| m.get(&event_type))
788            .cloned())
789    }
790
791    async fn get_user_room_receipt_event(
792        &self,
793        room_id: &RoomId,
794        receipt_type: ReceiptType,
795        thread: ReceiptThread,
796        user_id: &UserId,
797    ) -> Result<Option<(OwnedEventId, Receipt)>> {
798        Ok(self.get_user_room_receipt_event_impl(room_id, receipt_type, thread, user_id))
799    }
800
801    async fn get_event_room_receipt_events(
802        &self,
803        room_id: &RoomId,
804        receipt_type: ReceiptType,
805        thread: ReceiptThread,
806        event_id: &EventId,
807    ) -> Result<Vec<(OwnedUserId, Receipt)>> {
808        Ok(self
809            .get_event_room_receipt_events_impl(room_id, receipt_type, thread, event_id)
810            .unwrap_or_default())
811    }
812
813    async fn get_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
814        Ok(self.inner.read().unwrap().custom.get(key).cloned())
815    }
816
817    async fn set_custom_value(&self, key: &[u8], value: Vec<u8>) -> Result<Option<Vec<u8>>> {
818        Ok(self.inner.write().unwrap().custom.insert(key.to_vec(), value))
819    }
820
821    async fn remove_custom_value(&self, key: &[u8]) -> Result<Option<Vec<u8>>> {
822        Ok(self.inner.write().unwrap().custom.remove(key))
823    }
824
825    async fn remove_room(&self, room_id: &RoomId) -> Result<()> {
826        let mut inner = self.inner.write().unwrap();
827
828        inner.profiles.remove(room_id);
829        inner.display_names.remove(room_id);
830        inner.members.remove(room_id);
831        inner.room_info.remove(room_id);
832        inner.room_state.remove(room_id);
833        inner.room_account_data.remove(room_id);
834        inner.stripped_room_state.remove(room_id);
835        inner.stripped_members.remove(room_id);
836        inner.room_user_receipts.remove(room_id);
837        inner.room_event_receipts.remove(room_id);
838        inner.send_queue_events.remove(room_id);
839        inner.dependent_send_queue_events.remove(room_id);
840        inner.thread_subscriptions.remove(room_id);
841
842        Ok(())
843    }
844
845    async fn save_send_queue_request(
846        &self,
847        room_id: &RoomId,
848        transaction_id: OwnedTransactionId,
849        created_at: MilliSecondsSinceUnixEpoch,
850        kind: QueuedRequestKind,
851        priority: usize,
852    ) -> Result<(), Self::Error> {
853        self.inner
854            .write()
855            .unwrap()
856            .send_queue_events
857            .entry(room_id.to_owned())
858            .or_default()
859            .push(QueuedRequest { kind, transaction_id, error: None, priority, created_at });
860        Ok(())
861    }
862
863    async fn update_send_queue_request(
864        &self,
865        room_id: &RoomId,
866        transaction_id: &TransactionId,
867        kind: QueuedRequestKind,
868    ) -> Result<bool, Self::Error> {
869        if let Some(entry) = self
870            .inner
871            .write()
872            .unwrap()
873            .send_queue_events
874            .entry(room_id.to_owned())
875            .or_default()
876            .iter_mut()
877            .find(|item| item.transaction_id == transaction_id)
878        {
879            entry.kind = kind;
880            entry.error = None;
881            Ok(true)
882        } else {
883            Ok(false)
884        }
885    }
886
887    async fn remove_send_queue_request(
888        &self,
889        room_id: &RoomId,
890        transaction_id: &TransactionId,
891    ) -> Result<bool, Self::Error> {
892        let mut inner = self.inner.write().unwrap();
893        let q = &mut inner.send_queue_events;
894
895        let entry = q.get_mut(room_id);
896        if let Some(entry) = entry {
897            // Find the event by id in its room queue, and remove it if present.
898            if let Some(pos) = entry.iter().position(|item| item.transaction_id == transaction_id) {
899                entry.remove(pos);
900                // And if this was the last event before removal, remove the entire room entry.
901                if entry.is_empty() {
902                    q.remove(room_id);
903                }
904                return Ok(true);
905            }
906        }
907
908        Ok(false)
909    }
910
911    async fn load_send_queue_requests(
912        &self,
913        room_id: &RoomId,
914    ) -> Result<Vec<QueuedRequest>, Self::Error> {
915        let mut ret = self
916            .inner
917            .write()
918            .unwrap()
919            .send_queue_events
920            .entry(room_id.to_owned())
921            .or_default()
922            .clone();
923        // Inverted order of priority, use stable sort to keep insertion order.
924        ret.sort_by_key(|item| Reverse(item.priority));
925        Ok(ret)
926    }
927
928    async fn update_send_queue_request_status(
929        &self,
930        room_id: &RoomId,
931        transaction_id: &TransactionId,
932        error: Option<QueueWedgeError>,
933    ) -> Result<(), Self::Error> {
934        if let Some(entry) = self
935            .inner
936            .write()
937            .unwrap()
938            .send_queue_events
939            .entry(room_id.to_owned())
940            .or_default()
941            .iter_mut()
942            .find(|item| item.transaction_id == transaction_id)
943        {
944            entry.error = error;
945        }
946        Ok(())
947    }
948
949    async fn load_rooms_with_unsent_requests(&self) -> Result<Vec<OwnedRoomId>, Self::Error> {
950        Ok(self.inner.read().unwrap().send_queue_events.keys().cloned().collect())
951    }
952
953    async fn save_dependent_queued_request(
954        &self,
955        room: &RoomId,
956        parent_transaction_id: &TransactionId,
957        own_transaction_id: ChildTransactionId,
958        created_at: MilliSecondsSinceUnixEpoch,
959        content: DependentQueuedRequestKind,
960    ) -> Result<(), Self::Error> {
961        self.inner
962            .write()
963            .unwrap()
964            .dependent_send_queue_events
965            .entry(room.to_owned())
966            .or_default()
967            .push(DependentQueuedRequest {
968                kind: content,
969                parent_transaction_id: parent_transaction_id.to_owned(),
970                own_transaction_id,
971                parent_key: None,
972                created_at,
973            });
974        Ok(())
975    }
976
977    async fn mark_dependent_queued_requests_as_ready(
978        &self,
979        room: &RoomId,
980        parent_txn_id: &TransactionId,
981        sent_parent_key: SentRequestKey,
982    ) -> Result<usize, Self::Error> {
983        let mut inner = self.inner.write().unwrap();
984        let dependents = inner.dependent_send_queue_events.entry(room.to_owned()).or_default();
985        let mut num_updated = 0;
986        for d in dependents.iter_mut().filter(|item| item.parent_transaction_id == parent_txn_id) {
987            d.parent_key = Some(sent_parent_key.clone());
988            num_updated += 1;
989        }
990        Ok(num_updated)
991    }
992
993    async fn update_dependent_queued_request(
994        &self,
995        room: &RoomId,
996        own_transaction_id: &ChildTransactionId,
997        new_content: DependentQueuedRequestKind,
998    ) -> Result<bool, Self::Error> {
999        let mut inner = self.inner.write().unwrap();
1000        let dependents = inner.dependent_send_queue_events.entry(room.to_owned()).or_default();
1001        for d in dependents.iter_mut() {
1002            if d.own_transaction_id == *own_transaction_id {
1003                d.kind = new_content;
1004                return Ok(true);
1005            }
1006        }
1007        Ok(false)
1008    }
1009
1010    async fn remove_dependent_queued_request(
1011        &self,
1012        room: &RoomId,
1013        txn_id: &ChildTransactionId,
1014    ) -> Result<bool, Self::Error> {
1015        let mut inner = self.inner.write().unwrap();
1016        let dependents = inner.dependent_send_queue_events.entry(room.to_owned()).or_default();
1017        if let Some(pos) = dependents.iter().position(|item| item.own_transaction_id == *txn_id) {
1018            dependents.remove(pos);
1019            Ok(true)
1020        } else {
1021            Ok(false)
1022        }
1023    }
1024
1025    async fn load_dependent_queued_requests(
1026        &self,
1027        room: &RoomId,
1028    ) -> Result<Vec<DependentQueuedRequest>, Self::Error> {
1029        Ok(self
1030            .inner
1031            .read()
1032            .unwrap()
1033            .dependent_send_queue_events
1034            .get(room)
1035            .cloned()
1036            .unwrap_or_default())
1037    }
1038
1039    async fn upsert_thread_subscriptions(
1040        &self,
1041        updates: Vec<(&RoomId, &EventId, StoredThreadSubscription)>,
1042    ) -> Result<(), Self::Error> {
1043        let mut inner = self.inner.write().unwrap();
1044
1045        for (room_id, thread_id, mut new) in updates {
1046            let room_subs = inner.thread_subscriptions.entry(room_id.to_owned()).or_default();
1047
1048            if let Some(previous) = room_subs.get(thread_id) {
1049                if *previous == new {
1050                    continue;
1051                }
1052                if !compare_thread_subscription_bump_stamps(
1053                    previous.bump_stamp,
1054                    &mut new.bump_stamp,
1055                ) {
1056                    continue;
1057                }
1058            }
1059
1060            room_subs.insert(thread_id.to_owned(), new);
1061        }
1062
1063        Ok(())
1064    }
1065
1066    async fn load_thread_subscription(
1067        &self,
1068        room: &RoomId,
1069        thread_id: &EventId,
1070    ) -> Result<Option<StoredThreadSubscription>, Self::Error> {
1071        let inner = self.inner.read().unwrap();
1072        Ok(inner
1073            .thread_subscriptions
1074            .get(room)
1075            .and_then(|subscriptions| subscriptions.get(thread_id))
1076            .copied())
1077    }
1078
1079    async fn remove_thread_subscription(
1080        &self,
1081        room: &RoomId,
1082        thread_id: &EventId,
1083    ) -> Result<(), Self::Error> {
1084        let mut inner = self.inner.write().unwrap();
1085
1086        let Some(room_subs) = inner.thread_subscriptions.get_mut(room) else {
1087            return Ok(());
1088        };
1089
1090        room_subs.remove(thread_id);
1091
1092        if room_subs.is_empty() {
1093            // If there are no more subscriptions for this room, remove the room entry.
1094            inner.thread_subscriptions.remove(room);
1095        }
1096
1097        Ok(())
1098    }
1099
1100    async fn get_global_profile(
1101        &self,
1102        user_id: &UserId,
1103    ) -> Result<Option<UserProfile>, Self::Error> {
1104        let inner = self.inner.read().unwrap();
1105        Ok(inner.global_profiles.get(user_id).cloned())
1106    }
1107
1108    async fn get_global_profiles<'a>(
1109        &self,
1110        user_ids: &'a [OwnedUserId],
1111    ) -> Result<BTreeMap<&'a UserId, UserProfile>, Self::Error> {
1112        let inner = self.inner.read().unwrap();
1113        Ok(user_ids
1114            .iter()
1115            .filter_map(|user_id| {
1116                inner.global_profiles.get(user_id).map(|profile| (&**user_id, profile.clone()))
1117            })
1118            .collect())
1119    }
1120
1121    async fn optimize(&self) -> Result<(), Self::Error> {
1122        Ok(())
1123    }
1124
1125    async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
1126        Ok(None)
1127    }
1128}
1129
1130#[cfg(test)]
1131mod tests {
1132    use super::{MemoryStore, Result, StateStore};
1133
1134    async fn get_store() -> Result<impl StateStore> {
1135        Ok(MemoryStore::new())
1136    }
1137
1138    statestore_integration_tests!();
1139}