1use 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#[derive(Debug, Default)]
105pub struct MemoryStore {
106 inner: RwLock<MemoryStoreInner>,
107}
108
109impl MemoryStore {
110 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 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 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 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 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 if let Some(pos) = entry.iter().position(|item| item.transaction_id == transaction_id) {
899 entry.remove(pos);
900 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 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 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}