1use std::{
117 collections::{BTreeMap, BTreeSet},
118 pin::Pin,
119 sync::Weak,
120};
121
122use as_variant::as_variant;
123use futures_core::Stream;
124use futures_util::{StreamExt, future::try_join_all, pin_mut};
125#[cfg(doc)]
126use matrix_sdk_base::{BaseClient, crypto::OlmMachine};
127use matrix_sdk_base::{
128 crypto::{
129 store::types::{RoomKeyInfo, RoomKeyWithheldInfo},
130 types::events::room::encrypted::EncryptedEvent,
131 },
132 deserialized_responses::{DecryptedRoomEvent, TimelineEvent, TimelineEventKind},
133 locks::Mutex,
134 task_monitor::BackgroundTaskHandle,
135 timer,
136};
137#[cfg(doc)]
138use matrix_sdk_common::deserialized_responses::EncryptionInfo;
139use ruma::{
140 OwnedEventId, OwnedRoomId, RoomId,
141 events::{AnySyncTimelineEvent, room::encrypted::OriginalSyncRoomEncryptedEvent},
142 push::Action,
143 serde::Raw,
144};
145use tokio::sync::{
146 broadcast::{self, Sender},
147 mpsc::{UnboundedReceiver, UnboundedSender, unbounded_channel},
148};
149use tokio_stream::wrappers::{
150 BroadcastStream, UnboundedReceiverStream, errors::BroadcastStreamRecvError,
151};
152use tracing::{info, instrument, trace, warn};
153
154#[cfg(doc)]
155use super::RoomEventCache;
156use super::{
157 EventCache, EventCacheError, EventCacheInner, EventsOrigin, RoomEventCacheGenericUpdate,
158 RoomEventCacheUpdate, TimelineVectorDiffs, caches::room::RoomEventCacheLinkedChunkUpdate,
159};
160use crate::{Client, Result, Room, encryption::backups::BackupState, room::PushContext};
161
162type SessionId<'a> = &'a str;
163type OwnedSessionId = String;
164
165type EventIdAndUtd = (OwnedEventId, Raw<AnySyncTimelineEvent>);
166type EventIdAndEvent = (OwnedEventId, DecryptedRoomEvent);
167pub(in crate::event_cache) type ResolvedUtd =
168 (OwnedEventId, DecryptedRoomEvent, Option<Vec<Action>>);
169
170#[derive(Debug, Clone)]
173pub struct DecryptionRetryRequest {
174 pub room_id: OwnedRoomId,
176 pub utd_session_ids: BTreeSet<OwnedSessionId>,
178 pub refresh_info_session_ids: BTreeSet<OwnedSessionId>,
181}
182
183#[derive(Debug, Clone)]
185pub enum RedecryptorReport {
186 ResolvedUtds {
188 room_id: OwnedRoomId,
190 events: BTreeSet<OwnedEventId>,
192 },
193 Lagging,
196 BackupAvailable,
201}
202
203pub(super) struct RedecryptorChannels {
204 utd_reporter: Sender<RedecryptorReport>,
205 pub(super) decryption_request_sender: UnboundedSender<DecryptionRetryRequest>,
206 pub(super) decryption_request_receiver:
207 Mutex<Option<UnboundedReceiver<DecryptionRetryRequest>>>,
208}
209
210impl RedecryptorChannels {
211 pub(super) fn new() -> Self {
212 let (utd_reporter, _) = broadcast::channel(100);
213 let (decryption_request_sender, decryption_request_receiver) = unbounded_channel();
214
215 Self {
216 utd_reporter,
217 decryption_request_sender,
218 decryption_request_receiver: Mutex::new(Some(decryption_request_receiver)),
219 }
220 }
221}
222
223fn filter_timeline_event_to_utd(
228 event: TimelineEvent,
229) -> Option<(OwnedEventId, Raw<AnySyncTimelineEvent>)> {
230 let event_id = event.event_id().map(ToOwned::to_owned);
231
232 let event = as_variant!(event.kind, TimelineEventKind::UnableToDecrypt { event, .. } => event);
235 event_id.zip(event)
238}
239
240fn filter_timeline_event_to_decrypted(
246 event: TimelineEvent,
247) -> Option<(OwnedEventId, DecryptedRoomEvent)> {
248 let event_id = event.event_id().map(ToOwned::to_owned);
249
250 let event = as_variant!(event.kind, TimelineEventKind::Decrypted(event) => event);
251 event_id.zip(event)
254}
255
256impl EventCache {
257 async fn all_encrypted_events(
265 &self,
266 room_id: &RoomId,
267 session_id: SessionId<'_>,
268 ) -> Result<Vec<EventIdAndUtd>, EventCacheError> {
269 let caches = self.inner.all_caches_for_room(room_id).await?;
270
271 Ok(caches
272 .all_events_of_type(Some("m.room.encrypted"), Some(session_id))
273 .await?
274 .filter_map(filter_timeline_event_to_utd)
275 .collect())
276 }
277
278 async fn all_in_memory_encrypted_events(&self) -> BTreeMap<OwnedRoomId, Vec<EventIdAndUtd>> {
281 let mut utds = BTreeMap::new();
282
283 for (room_id, caches) in self.inner.by_room.read().await.iter() {
284 let room_utds: Vec<_> = caches
285 .all_in_memory_events()
286 .await
287 .into_iter()
288 .flatten()
289 .filter_map(filter_timeline_event_to_utd)
290 .collect();
291
292 utds.insert(room_id.to_owned(), room_utds);
293 }
294
295 utds
296 }
297
298 async fn all_decrypted_events(
299 &self,
300 room_id: &RoomId,
301 session_id: SessionId<'_>,
302 ) -> Result<Vec<EventIdAndEvent>, EventCacheError> {
303 let caches = self.inner.all_caches_for_room(room_id).await?;
304
305 Ok(caches
306 .all_events_of_type(None, Some(session_id))
307 .await?
308 .filter_map(filter_timeline_event_to_decrypted)
309 .collect())
310 }
311
312 async fn all_in_memory_decrypted_events(&self) -> BTreeMap<OwnedRoomId, Vec<EventIdAndEvent>> {
313 let mut decrypted_events = BTreeMap::new();
314
315 for (room_id, caches) in self.inner.by_room.read().await.iter() {
316 let room_utds: Vec<_> = caches
317 .all_in_memory_events()
318 .await
319 .into_iter()
320 .flatten()
321 .filter_map(filter_timeline_event_to_decrypted)
322 .collect();
323
324 decrypted_events.insert(room_id.to_owned(), room_utds);
325 }
326
327 decrypted_events
328 }
329
330 #[instrument(skip_all, fields(room_id))]
342 async fn on_resolved_utds(
343 &self,
344 room_id: &RoomId,
345 events: Vec<ResolvedUtd>,
346 ) -> Result<(), EventCacheError> {
347 if events.is_empty() {
348 trace!("No events were redecrypted or updated, nothing to replace");
349 return Ok(());
350 }
351
352 timer!("Resolving UTDs");
353
354 let event_ids: BTreeSet<_> =
355 events.iter().cloned().map(|(event_id, _, _)| event_id).collect();
356
357 let all_caches = self.inner.all_caches_for_room(room_id).await?;
358
359 {
361 let room_cache = &all_caches.room;
362 let mut state = room_cache.state().write().await?;
363
364 let mut new_events = Vec::with_capacity(events.len());
365
366 for (event_id, decrypted, actions) in &events {
367 if let Some((location, mut target_event)) = state.find_event(event_id).await?
368 && (
369 matches!(target_event.kind, TimelineEventKind::UnableToDecrypt { .. })
380 || target_event.encryption_info() != Some(&decrypted.encryption_info)
381 )
382 {
383 target_event.kind = TimelineEventKind::Decrypted(decrypted.clone());
384
385 if let Some(actions) = actions {
386 target_event.set_push_actions(actions.clone());
387 }
388
389 state.replace_event_at(location, target_event.clone()).await?;
392 new_events.push(target_event);
393 }
394 }
395
396 let receipt_event = None;
404
405 state.post_process_new_events(new_events, receipt_event).await?;
406
407 let updates_as_vector_diffs = state.room_linked_chunk_mut().updates_as_vector_diffs();
408
409 if !updates_as_vector_diffs.is_empty() {
410 room_cache.update_sender().send(
411 RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs {
412 diffs: updates_as_vector_diffs,
413 origin: EventsOrigin::Cache,
414 }),
415 Some(RoomEventCacheGenericUpdate { room_id: room_id.to_owned() }),
416 );
417 }
418 }
419
420 {
422 for (thread_id, thread_cache) in try_join_all(
430 all_caches.threads.read().await.iter().map(|(thread_id, thread_cache)| async {
431 Result::<_, EventCacheError>::Ok(
432 thread_cache
435 .replace_utds(&events)
436 .await?
437 .then(|| (thread_id.clone(), thread_cache.clone())),
438 )
439 }),
440 )
441 .await?
442 .into_iter()
443 .flatten()
445 {
446 let new_thread_summary =
447 thread_cache.state().read().await?.compute_thread_summary().await?;
448
449 all_caches.room.update_thread_summary(&thread_id, new_thread_summary).await?;
450 }
451 }
452
453 if let Some(pinned_events_cache) = all_caches.pinned_events.get() {
455 pinned_events_cache.replace_utds(&events).await?;
456 }
457
458 {
460 try_join_all(
465 all_caches
466 .event_focused
467 .read()
468 .await
469 .values()
470 .map(|event_focused_cache| event_focused_cache.replace_utds(&events)),
471 )
472 .await?;
473 }
474
475 let report =
476 RedecryptorReport::ResolvedUtds { room_id: room_id.to_owned(), events: event_ids };
477 let _ = self.inner.redecryption_channels.utd_reporter.send(report);
478
479 Ok(())
480 }
481
482 async fn decrypt_event(
484 &self,
485 room_id: &RoomId,
486 room: Option<&Room>,
487 push_context: Option<&PushContext>,
488 event: &Raw<EncryptedEvent>,
489 ) -> Option<(DecryptedRoomEvent, Option<Vec<Action>>)> {
490 if let Some(room) = room {
491 match room
492 .decrypt_event(
493 event.cast_ref_unchecked::<OriginalSyncRoomEncryptedEvent>(),
494 push_context,
495 )
496 .await
497 {
498 Ok(maybe_decrypted) => {
499 let actions = maybe_decrypted.push_actions().map(|a| a.to_vec());
500
501 if let TimelineEventKind::Decrypted(decrypted) = maybe_decrypted.kind {
502 Some((decrypted, actions))
503 } else {
504 warn!(
505 "Failed to redecrypt an event despite receiving a room key or request to redecrypt"
506 );
507 None
508 }
509 }
510 Err(e) => {
511 warn!(
512 "Failed to redecrypt an event despite receiving a room key or request to redecrypt {e:?}"
513 );
514 None
515 }
516 }
517 } else {
518 let client = self.inner.client().ok()?;
519 let machine = client.olm_machine().await;
520 let machine = machine.as_ref()?;
521
522 match machine.decrypt_room_event(event, room_id, client.decryption_settings()).await {
523 Ok(decrypted) => Some((decrypted, None)),
524 Err(e) => {
525 warn!(
526 "Failed to redecrypt an event despite receiving a room key or a request to redecrypt {e:?}"
527 );
528 None
529 }
530 }
531 }
532 }
533
534 #[instrument(skip_all, fields(room_id, session_id))]
537 async fn retry_decryption(
538 &self,
539 room_id: &RoomId,
540 session_id: SessionId<'_>,
541 ) -> Result<(), EventCacheError> {
542 let events = self.all_encrypted_events(room_id, session_id).await?;
544 self.retry_decryption_for_events(room_id, events).await
545 }
546
547 #[instrument(skip_all, fields(updates.linked_chunk_id))]
549 async fn retry_decryption_for_event_cache_updates(
550 &self,
551 updates: RoomEventCacheLinkedChunkUpdate,
552 ) -> Result<(), EventCacheError> {
553 let room_id = updates.linked_chunk_id.room_id();
554 let events: Vec<_> = updates
555 .updates
556 .into_iter()
557 .flat_map(|updates| updates.into_items())
558 .filter_map(filter_timeline_event_to_utd)
559 .collect();
560
561 self.retry_decryption_for_events(room_id, events).await
562 }
563
564 async fn retry_decryption_for_in_memory_events(&self) {
565 let utds = self.all_in_memory_encrypted_events().await;
566
567 for (room_id, utds) in utds.into_iter() {
568 if let Err(e) = self.retry_decryption_for_events(&room_id, utds).await {
569 warn!(%room_id, "Failed to redecrypt in-memory events {e:?}");
570 }
571 }
572 }
573
574 #[instrument(skip_all, fields(room_id, session_id))]
576 async fn retry_decryption_for_events(
577 &self,
578 room_id: &RoomId,
579 events: Vec<EventIdAndUtd>,
580 ) -> Result<(), EventCacheError> {
581 trace!("Retrying to decrypt");
582
583 if events.is_empty() {
584 trace!("No relevant events found.");
585 return Ok(());
586 }
587
588 let room = self.inner.client().ok().and_then(|client| client.get_room(room_id));
589 let push_context =
590 if let Some(room) = &room { room.push_context().await.ok().flatten() } else { None };
591
592 let mut decrypted_events = Vec::with_capacity(events.len());
594
595 for (event_id, event) in events {
596 if let Some((decrypted, actions)) = self
599 .decrypt_event(
600 room_id,
601 room.as_ref(),
602 push_context.as_ref(),
603 event.cast_ref_unchecked(),
604 )
605 .await
606 {
607 decrypted_events.push((event_id, decrypted, actions));
608 }
609 }
610
611 let event_ids: BTreeSet<_> =
612 decrypted_events.iter().map(|(event_id, _, _)| event_id).collect();
613
614 if !event_ids.is_empty() {
615 trace!(?event_ids, "Successfully redecrypted events");
616 }
617
618 self.on_resolved_utds(room_id, decrypted_events).await?;
621
622 Ok(())
623 }
624
625 async fn update_encryption_info_for_events(
627 &self,
628 room: &Room,
629 events: Vec<EventIdAndEvent>,
630 ) -> Result<(), EventCacheError> {
631 let mut updated_events = Vec::with_capacity(events.len());
633
634 for (event_id, mut event) in events {
635 if let Some(session_id) = event.encryption_info.session_id() {
636 let new_encryption_info =
637 room.get_encryption_info(session_id, &event.encryption_info.sender).await;
638
639 if let Some(new_encryption_info) = new_encryption_info
641 && event.encryption_info != new_encryption_info
642 {
643 event.encryption_info = new_encryption_info;
644 updated_events.push((event_id, event, None));
645 }
646 }
647 }
648
649 let event_ids: BTreeSet<_> =
650 updated_events.iter().map(|(event_id, _, _)| event_id).collect();
651
652 if !event_ids.is_empty() {
653 trace!(?event_ids, "Replacing the encryption info of some events");
654 }
655
656 self.on_resolved_utds(room.room_id(), updated_events).await
657 }
658
659 #[instrument(skip_all, fields(room_id, session_id))]
660 async fn update_encryption_info(
661 &self,
662 room_id: &RoomId,
663 session_id: SessionId<'_>,
664 ) -> Result<(), EventCacheError> {
665 trace!("Updating encryption info");
666
667 let Ok(client) = self.inner.client() else {
668 return Ok(());
669 };
670
671 let Some(room) = client.get_room(room_id) else {
672 return Ok(());
673 };
674
675 let events = self.all_decrypted_events(room_id, session_id).await?;
677
678 if events.is_empty() {
679 trace!("No relevant events found.");
680 return Ok(());
681 }
682
683 self.update_encryption_info_for_events(&room, events).await
685 }
686
687 async fn retry_update_encryption_info_for_in_memory_events(&self) {
688 let decrypted_events = self.all_in_memory_decrypted_events().await;
689
690 for (room_id, events) in decrypted_events.into_iter() {
691 let Some(room) = self.inner.client().ok().and_then(|c| c.get_room(&room_id)) else {
692 continue;
693 };
694
695 if let Err(e) = self.update_encryption_info_for_events(&room, events).await {
696 warn!(
697 %room_id,
698 "Failed to replace the encryption info for in-memory events {e:?}"
699 );
700 }
701 }
702 }
703
704 async fn retry_in_memory_events(&self) {
715 self.retry_decryption_for_in_memory_events().await;
716 self.retry_update_encryption_info_for_in_memory_events().await;
717 }
718
719 pub fn request_decryption(&self, request: DecryptionRetryRequest) {
760 let _ =
761 self.inner.redecryption_channels.decryption_request_sender.send(request).inspect_err(
762 |_| warn!("Requesting a decryption while the redecryption task has been shut down"),
763 );
764 }
765
766 pub fn subscribe_to_decryption_reports(
817 &self,
818 ) -> impl Stream<Item = Result<RedecryptorReport, BroadcastStreamRecvError>> {
819 BroadcastStream::new(self.inner.redecryption_channels.utd_reporter.subscribe())
820 }
821}
822
823#[inline(always)]
824fn upgrade_event_cache(cache: &Weak<EventCacheInner>) -> Option<EventCache> {
825 cache.upgrade().map(|inner| EventCache { inner })
826}
827
828async fn send_report_and_retry_memory_events(
829 cache: &Weak<EventCacheInner>,
830 report: RedecryptorReport,
831) -> Result<(), ()> {
832 let Some(cache) = upgrade_event_cache(cache) else {
833 return Err(());
834 };
835
836 cache.retry_in_memory_events().await;
837 let _ = cache.inner.redecryption_channels.utd_reporter.send(report);
838
839 Ok(())
840}
841
842pub(crate) struct Redecryptor {
849 _task: BackgroundTaskHandle,
850}
851
852impl Redecryptor {
853 pub(super) fn new(
858 client: &Client,
859 cache: Weak<EventCacheInner>,
860 receiver: UnboundedReceiver<DecryptionRetryRequest>,
861 linked_chunk_update_sender: &Sender<RoomEventCacheLinkedChunkUpdate>,
862 ) -> Self {
863 let linked_chunk_stream = BroadcastStream::new(linked_chunk_update_sender.subscribe());
864 let backup_state_stream = client.encryption().backups().state_stream();
865
866 let task = client
867 .task_monitor()
868 .spawn_infinite_task("event_cache::redecryptor", async {
869 let request_redecryption_stream = UnboundedReceiverStream::new(receiver);
870
871 Self::listen_for_room_keys_task(
872 cache,
873 request_redecryption_stream,
874 linked_chunk_stream,
875 backup_state_stream,
876 )
877 .await;
878 })
879 .abort_on_drop();
880
881 Self { _task: task }
882 }
883
884 async fn subscribe_to_room_key_stream(
889 cache: &Weak<EventCacheInner>,
890 ) -> Option<(
891 impl Stream<Item = Result<Vec<RoomKeyInfo>, BroadcastStreamRecvError>>,
892 impl Stream<Item = Vec<RoomKeyWithheldInfo>>,
893 )> {
894 let event_cache = cache.upgrade()?;
895 let client = event_cache.client().ok()?;
896 let machine = client.olm_machine().await;
897
898 machine.as_ref().map(|m| {
899 (m.store().room_keys_received_stream(), m.store().room_keys_withheld_received_stream())
900 })
901 }
902
903 async fn redecryption_loop(
904 cache: &Weak<EventCacheInner>,
905 decryption_request_stream: &mut Pin<&mut impl Stream<Item = DecryptionRetryRequest>>,
906 events_stream: &mut Pin<
907 &mut impl Stream<Item = Result<RoomEventCacheLinkedChunkUpdate, BroadcastStreamRecvError>>,
908 >,
909 backup_state_stream: &mut Pin<
910 &mut impl Stream<Item = Result<BackupState, BroadcastStreamRecvError>>,
911 >,
912 ) -> bool {
913 let Some((room_key_stream, withheld_stream)) =
914 Self::subscribe_to_room_key_stream(cache).await
915 else {
916 return false;
917 };
918
919 pin_mut!(room_key_stream);
920 pin_mut!(withheld_stream);
921
922 loop {
923 tokio::select! {
924 Some(request) = decryption_request_stream.next() => {
927 let Some(cache) = upgrade_event_cache(cache) else {
928 break false;
929 };
930
931 trace!(?request, "Received a redecryption request");
932
933 for session_id in request.utd_session_ids {
934 let _ = cache
935 .retry_decryption(&request.room_id, &session_id)
936 .await
937 .inspect_err(|e| warn!("Error redecrypting after an explicit request was received {e:?}"));
938 }
939
940 for session_id in request.refresh_info_session_ids {
941 let _ = cache.update_encryption_info(&request.room_id, &session_id).await.inspect_err(|e|
942 warn!(
943 room_id = %request.room_id,
944 session_id = session_id,
945 "Unable to update the encryption info {e:?}",
946 ));
947 }
948 }
949 room_keys = room_key_stream.next() => {
952 match room_keys {
953 Some(Ok(room_keys)) => {
954 let Some(cache) = upgrade_event_cache(cache) else {
958 break false;
959 };
960
961 trace!(?room_keys, "Received new room keys");
962
963 for key in &room_keys {
964 let _ = cache
965 .retry_decryption(&key.room_id, &key.session_id)
966 .await
967 .inspect_err(|e| warn!("Error redecrypting {e:?}"));
968 }
969
970 for key in room_keys {
971 let _ = cache.update_encryption_info(&key.room_id, &key.session_id).await.inspect_err(|e|
972 warn!(
973 room_id = %key.room_id,
974 session_id = key.session_id,
975 "Unable to update the encryption info {e:?}",
976 ));
977 }
978 },
979 Some(Err(_)) => {
980 warn!("The room key stream lagged, reporting the lag to our listeners");
987
988 if send_report_and_retry_memory_events(cache, RedecryptorReport::Lagging).await.is_err() {
989 break false;
990 }
991 },
992 None => {
995 break true;
996 }
997 }
998 }
999 withheld_info = withheld_stream.next() => {
1000 match withheld_info {
1001 Some(infos) => {
1002 let Some(cache) = upgrade_event_cache(cache) else {
1003 break false;
1004 };
1005
1006 trace!(?infos, "Received new withheld infos");
1007
1008 for RoomKeyWithheldInfo { room_id, session_id, .. } in &infos {
1009 let _ = cache.update_encryption_info(room_id, session_id).await.inspect_err(|e|
1010 warn!(
1011 room_id = %room_id,
1012 session_id = session_id,
1013 "Unable to update the encryption info {e:?}",
1014 ));
1015 }
1016 }
1017 None => break true,
1020 }
1021 }
1022 Some(event_updates) = events_stream.next() => {
1026 match event_updates {
1027 Ok(updates) => {
1028 let Some(cache) = upgrade_event_cache(cache) else {
1029 break false;
1030 };
1031
1032 let linked_chunk_id = updates.linked_chunk_id.to_owned();
1033
1034 let _ = cache.retry_decryption_for_event_cache_updates(updates).await.inspect_err(|e|
1035 warn!(
1036 %linked_chunk_id,
1037 "Unable to handle UTDs from event cache updates {e:?}",
1038 )
1039 );
1040 }
1041 Err(_) => {
1042 if send_report_and_retry_memory_events(cache, RedecryptorReport::Lagging).await.is_err() {
1043 break false;
1044 }
1045 }
1046 }
1047 }
1048 Some(backup_state_update) = backup_state_stream.next() => {
1049 match backup_state_update {
1050 Ok(state) => {
1051 match state {
1052 BackupState::Unknown |
1053 BackupState::Creating |
1054 BackupState::Enabling |
1055 BackupState::Resuming |
1056 BackupState::Downloading |
1057 BackupState::Disabling =>{
1058 }
1061 BackupState::Enabled => {
1062 if send_report_and_retry_memory_events(cache, RedecryptorReport::BackupAvailable).await.is_err() {
1067 break false;
1068 }
1069 }
1070 }
1071 }
1072 Err(_) => {
1073 if send_report_and_retry_memory_events(cache, RedecryptorReport::Lagging).await.is_err() {
1074 break false;
1075 }
1076 }
1077 }
1078 }
1079 else => break false,
1080 }
1081 }
1082 }
1083
1084 async fn listen_for_room_keys_task(
1085 cache: Weak<EventCacheInner>,
1086 decryption_request_stream: UnboundedReceiverStream<DecryptionRetryRequest>,
1087 events_stream: BroadcastStream<RoomEventCacheLinkedChunkUpdate>,
1088 backup_state_stream: impl Stream<Item = Result<BackupState, BroadcastStreamRecvError>>,
1089 ) {
1090 pin_mut!(decryption_request_stream);
1094 pin_mut!(events_stream);
1095 pin_mut!(backup_state_stream);
1096
1097 while Self::redecryption_loop(
1098 &cache,
1099 &mut decryption_request_stream,
1100 &mut events_stream,
1101 &mut backup_state_stream,
1102 )
1103 .await
1104 {
1105 info!("Regenerating the re-decryption streams");
1106
1107 if send_report_and_retry_memory_events(&cache, RedecryptorReport::Lagging)
1110 .await
1111 .is_err()
1112 {
1113 break;
1114 }
1115 }
1116
1117 info!("Shutting down the event cache redecryptor");
1118 }
1119}
1120
1121#[cfg(not(target_family = "wasm"))]
1122#[cfg(test)]
1123mod tests {
1124 use std::{
1125 collections::BTreeSet,
1126 sync::{
1127 Arc,
1128 atomic::{AtomicBool, Ordering},
1129 },
1130 time::Duration,
1131 };
1132
1133 use assert_matches2::assert_matches;
1134 use async_trait::async_trait;
1135 use eyeball_im::VectorDiff;
1136 use matrix_sdk_base::{
1137 cross_process_lock::CrossProcessLockGeneration,
1138 crypto::types::events::{ToDeviceEvent, room::encrypted::ToDeviceEncryptedEventContent},
1139 deserialized_responses::{TimelineEventKind, VerificationState},
1140 event_cache::{
1141 Event, Gap,
1142 store::{EventCacheStore, EventCacheStoreError, MemoryStore},
1143 },
1144 linked_chunk::{
1145 ChunkIdentifier, ChunkIdentifierGenerator, ChunkMetadata, LinkedChunkId, Position,
1146 RawChunk, Update,
1147 },
1148 locks::Mutex,
1149 sleep::sleep,
1150 store::StoreConfig,
1151 };
1152 use matrix_sdk_common::cross_process_lock::CrossProcessLockConfig;
1153 use matrix_sdk_test::{JoinedRoomBuilder, async_test, event_factory::EventFactory};
1154 use ruma::{
1155 EventId, OwnedEventId, RoomId, RoomVersionId, device_id, event_id,
1156 events::{AnySyncTimelineEvent, relation::RelationType},
1157 room_id,
1158 serde::Raw,
1159 user_id,
1160 };
1161 use serde_json::json;
1162 use tokio::sync::oneshot::{self, Sender};
1163 use tracing::{Instrument, info};
1164
1165 use crate::{
1166 Client, assert_let_timeout,
1167 encryption::EncryptionSettings,
1168 event_cache::{
1169 DecryptionRetryRequest, RoomEventCacheGenericUpdate, RoomEventCacheUpdate,
1170 TimelineVectorDiffs,
1171 },
1172 test_utils::mocks::MatrixMockServer,
1173 };
1174
1175 #[derive(Debug, Clone)]
1180 struct DelayingStore {
1181 memory_store: MemoryStore,
1182 delaying: Arc<AtomicBool>,
1183 foo: Arc<Mutex<Option<Sender<()>>>>,
1184 }
1185
1186 impl DelayingStore {
1187 fn new() -> Self {
1188 Self {
1189 memory_store: MemoryStore::new(),
1190 delaying: AtomicBool::new(true).into(),
1191 foo: Arc::new(Mutex::new(None)),
1192 }
1193 }
1194
1195 async fn stop_delaying(&self) {
1196 let (sender, receiver) = oneshot::channel();
1197
1198 {
1199 *self.foo.lock() = Some(sender);
1200 }
1201
1202 self.delaying.store(false, Ordering::SeqCst);
1203
1204 receiver.await.expect("We should be able to receive a response")
1205 }
1206 }
1207
1208 #[cfg_attr(target_family = "wasm", async_trait(?Send))]
1209 #[cfg_attr(not(target_family = "wasm"), async_trait)]
1210 impl EventCacheStore for DelayingStore {
1211 type Error = EventCacheStoreError;
1212
1213 async fn close(&self) -> Result<(), EventCacheStoreError> {
1214 self.memory_store.close().await
1215 }
1216
1217 async fn reopen(&self) -> Result<(), EventCacheStoreError> {
1218 self.memory_store.reopen().await
1219 }
1220
1221 async fn try_take_leased_lock(
1222 &self,
1223 lease_duration_ms: u32,
1224 key: &str,
1225 holder: &str,
1226 ) -> Result<Option<CrossProcessLockGeneration>, Self::Error> {
1227 self.memory_store.try_take_leased_lock(lease_duration_ms, key, holder).await
1228 }
1229
1230 async fn handle_linked_chunk_updates(
1231 &self,
1232 linked_chunk_id: LinkedChunkId<'_>,
1233 updates: Vec<Update<Event, Gap>>,
1234 ) -> Result<(), Self::Error> {
1235 while self.delaying.load(Ordering::SeqCst) {
1241 sleep(Duration::from_millis(10)).await;
1242 }
1243
1244 let sender = self.foo.lock().take();
1245 let ret = self.memory_store.handle_linked_chunk_updates(linked_chunk_id, updates).await;
1246
1247 if let Some(sender) = sender {
1248 sender.send(()).expect("We should be able to notify the other side that we're done with the storage operation");
1249 }
1250
1251 ret
1252 }
1253
1254 async fn load_all_chunks(
1255 &self,
1256 linked_chunk_id: LinkedChunkId<'_>,
1257 ) -> Result<Vec<RawChunk<Event, Gap>>, Self::Error> {
1258 self.memory_store.load_all_chunks(linked_chunk_id).await
1259 }
1260
1261 async fn load_all_chunks_metadata(
1262 &self,
1263 linked_chunk_id: LinkedChunkId<'_>,
1264 ) -> Result<Vec<ChunkMetadata>, Self::Error> {
1265 self.memory_store.load_all_chunks_metadata(linked_chunk_id).await
1266 }
1267
1268 async fn load_last_chunk(
1269 &self,
1270 linked_chunk_id: LinkedChunkId<'_>,
1271 ) -> Result<(Option<RawChunk<Event, Gap>>, ChunkIdentifierGenerator), Self::Error> {
1272 self.memory_store.load_last_chunk(linked_chunk_id).await
1273 }
1274
1275 async fn load_previous_chunk(
1276 &self,
1277 linked_chunk_id: LinkedChunkId<'_>,
1278 before_chunk_identifier: ChunkIdentifier,
1279 ) -> Result<Option<RawChunk<Event, Gap>>, Self::Error> {
1280 self.memory_store.load_previous_chunk(linked_chunk_id, before_chunk_identifier).await
1281 }
1282
1283 async fn remember_thread(
1284 &self,
1285 room_id: &RoomId,
1286 thread_id: &EventId,
1287 ) -> Result<(), Self::Error> {
1288 self.memory_store.remember_thread(room_id, thread_id).await
1289 }
1290
1291 async fn clear_all_events(&self, room_id: Option<&RoomId>) -> Result<(), Self::Error> {
1292 self.memory_store.clear_all_events(room_id).await
1293 }
1294
1295 async fn filter_duplicated_events(
1296 &self,
1297 linked_chunk_id: LinkedChunkId<'_>,
1298 events: Vec<OwnedEventId>,
1299 ) -> Result<Vec<(OwnedEventId, Position)>, Self::Error> {
1300 self.memory_store.filter_duplicated_events(linked_chunk_id, events).await
1301 }
1302
1303 async fn find_event(
1304 &self,
1305 room_id: &RoomId,
1306 event_id: &EventId,
1307 ) -> Result<Option<Event>, Self::Error> {
1308 self.memory_store.find_event(room_id, event_id).await
1309 }
1310
1311 async fn find_event_relations(
1312 &self,
1313 room_id: &RoomId,
1314 event_id: &EventId,
1315 filters: Option<&[RelationType]>,
1316 ) -> Result<Vec<(Event, Option<Position>)>, Self::Error> {
1317 self.memory_store.find_event_relations(room_id, event_id, filters).await
1318 }
1319
1320 async fn get_room_events(
1321 &self,
1322 room_id: &RoomId,
1323 event_type: Option<&str>,
1324 session_id: Option<&str>,
1325 ) -> Result<Vec<Event>, Self::Error> {
1326 self.memory_store.get_room_events(room_id, event_type, session_id).await
1327 }
1328
1329 async fn save_event(&self, room_id: &RoomId, event: Event) -> Result<(), Self::Error> {
1330 self.memory_store.save_event(room_id, event).await
1331 }
1332
1333 async fn optimize(&self) -> Result<(), Self::Error> {
1334 self.memory_store.optimize().await
1335 }
1336
1337 async fn get_size(&self) -> Result<Option<usize>, Self::Error> {
1338 self.memory_store.get_size().await
1339 }
1340 }
1341
1342 async fn set_up_clients(
1343 room_id: &RoomId,
1344 alice_enables_cross_signing: bool,
1345 use_delayed_store: bool,
1346 ) -> (Client, Client, MatrixMockServer, Option<DelayingStore>) {
1347 let alice_span = tracing::info_span!("alice");
1348 let bob_span = tracing::info_span!("bob");
1349
1350 let alice_user_id = user_id!("@alice:localhost");
1351 let alice_device_id = device_id!("ALICEDEVICE");
1352 let bob_user_id = user_id!("@bob:localhost");
1353 let bob_device_id = device_id!("BOBDEVICE");
1354
1355 let matrix_mock_server = MatrixMockServer::new().await;
1356 matrix_mock_server.mock_crypto_endpoints_preset().await;
1357
1358 let encryption_settings = EncryptionSettings {
1359 auto_enable_cross_signing: alice_enables_cross_signing,
1360 ..Default::default()
1361 };
1362
1363 let alice = matrix_mock_server
1366 .client_builder_for_crypto_end_to_end(alice_user_id, alice_device_id)
1367 .on_builder(|builder| {
1368 builder
1369 .with_enable_share_history_on_invite(true)
1370 .with_encryption_settings(encryption_settings)
1371 })
1372 .build()
1373 .instrument(alice_span.clone())
1374 .await;
1375
1376 let encryption_settings =
1377 EncryptionSettings { auto_enable_cross_signing: true, ..Default::default() };
1378
1379 let (store_config, store) = if use_delayed_store {
1380 let store = DelayingStore::new();
1381
1382 (
1383 StoreConfig::new(CrossProcessLockConfig::multi_process(
1384 "delayed_store_event_cache_test",
1385 ))
1386 .event_cache_store(store.clone()),
1387 Some(store),
1388 )
1389 } else {
1390 (
1391 StoreConfig::new(CrossProcessLockConfig::multi_process(
1392 "normal_store_event_cache_test",
1393 )),
1394 None,
1395 )
1396 };
1397
1398 let bob = matrix_mock_server
1399 .client_builder_for_crypto_end_to_end(bob_user_id, bob_device_id)
1400 .on_builder(|builder| {
1401 builder
1402 .with_enable_share_history_on_invite(true)
1403 .with_encryption_settings(encryption_settings)
1404 .store_config(store_config)
1405 })
1406 .build()
1407 .instrument(bob_span.clone())
1408 .await;
1409
1410 bob.event_cache().subscribe().expect("Bob should be able to enable the event cache");
1411
1412 matrix_mock_server.exchange_e2ee_identities(&alice, &bob).await;
1414
1415 let event_factory = EventFactory::new().room(room_id).sender(alice_user_id);
1416
1417 let room_builder = JoinedRoomBuilder::new(room_id)
1419 .add_state_event(event_factory.create(alice_user_id, RoomVersionId::V1))
1420 .add_state_event(event_factory.room_encryption());
1421
1422 matrix_mock_server
1423 .mock_sync()
1424 .ok_and_run(&alice, |builder| {
1425 builder.add_joined_room(room_builder.clone());
1426 })
1427 .instrument(alice_span)
1428 .await;
1429
1430 matrix_mock_server
1431 .mock_sync()
1432 .ok_and_run(&bob, |builder| {
1433 builder.add_joined_room(room_builder);
1434 })
1435 .instrument(bob_span)
1436 .await;
1437
1438 (alice, bob, matrix_mock_server, store)
1439 }
1440
1441 async fn prepare_room(
1442 matrix_mock_server: &MatrixMockServer,
1443 event_factory: &EventFactory,
1444 alice: &Client,
1445 bob: &Client,
1446 room_id: &RoomId,
1447 ) -> (Raw<AnySyncTimelineEvent>, Raw<ToDeviceEvent<ToDeviceEncryptedEventContent>>) {
1448 let alice_user_id = alice.user_id().unwrap();
1449 let bob_user_id = bob.user_id().unwrap();
1450
1451 let alice_member_event = event_factory.member(alice_user_id).into_raw();
1452 let bob_member_event = event_factory.member(bob_user_id).into_raw();
1453
1454 let room = alice
1455 .get_room(room_id)
1456 .expect("Alice should have access to the room now that we synced");
1457
1458 let event_type = "m.room.message";
1463 let content = json!({"body": "It's a secret to everybody", "msgtype": "m.text"});
1464
1465 let event_id = event_id!("$some_id");
1466 let (event_receiver, mock) =
1467 matrix_mock_server.mock_room_send().ok_with_capture(event_id, alice_user_id);
1468 let (_guard, room_key) = matrix_mock_server.mock_capture_put_to_device(alice_user_id).await;
1469
1470 {
1471 let _guard = mock.mock_once().mount_as_scoped().await;
1472
1473 matrix_mock_server
1474 .mock_get_members()
1475 .ok(vec![alice_member_event.clone(), bob_member_event.clone()])
1476 .mock_once()
1477 .mount()
1478 .await;
1479
1480 room.send_raw(event_type, content)
1481 .await
1482 .expect("We should be able to send an initial message");
1483 };
1484
1485 let event = event_receiver.await.expect("Alice should have sent the event by now");
1487 let room_key = room_key.await;
1488
1489 (event, room_key)
1490 }
1491
1492 #[async_test]
1493 async fn test_redecryptor() {
1494 let room_id = room_id!("!test:localhost");
1495
1496 let event_factory = EventFactory::new().room(room_id);
1497 let (alice, bob, matrix_mock_server, _) = set_up_clients(room_id, true, false).await;
1498
1499 let (event, room_key) =
1500 prepare_room(&matrix_mock_server, &event_factory, &alice, &bob, room_id).await;
1501
1502 let event_cache = bob.event_cache();
1505 let (room_cache, _) = event_cache
1506 .room(room_id)
1507 .await
1508 .expect("We should be able to get to the event cache for a specific room");
1509
1510 let (_, mut subscriber) = room_cache.subscribe().await.unwrap();
1511 let mut generic_stream = event_cache.subscribe_to_room_generic_updates();
1512
1513 bob.inner
1516 .base_client
1517 .regenerate_olm(None)
1518 .await
1519 .expect("We should be able to regenerate the Olm machine");
1520
1521 matrix_mock_server
1523 .mock_sync()
1524 .ok_and_run(&bob, |builder| {
1525 builder.add_joined_room(JoinedRoomBuilder::new(room_id).add_timeline_event(event));
1526 })
1527 .await;
1528
1529 assert_let_timeout!(
1532 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1533 subscriber.recv()
1534 );
1535
1536 assert_eq!(diffs.len(), 1);
1539 assert_matches!(&diffs[0], VectorDiff::Append { values });
1540 assert_eq!(values.len(), 1);
1541 assert_matches!(&values[0].kind, TimelineEventKind::UnableToDecrypt { .. });
1542
1543 assert_let_timeout!(
1544 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1545 );
1546 assert_eq!(expected_room_id, room_id);
1547 assert!(generic_stream.is_empty());
1548
1549 matrix_mock_server
1551 .mock_sync()
1552 .ok_and_run(&bob, |builder| {
1553 builder.add_to_device_event(
1554 room_key
1555 .deserialize_as()
1556 .expect("We should be able to deserialize the room key"),
1557 );
1558 })
1559 .await;
1560
1561 assert_let_timeout!(
1563 Duration::from_secs(1),
1564 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1565 subscriber.recv()
1566 );
1567
1568 assert_eq!(diffs.len(), 1);
1570 assert_matches!(&diffs[0], VectorDiff::Set { index, value });
1571 assert_eq!(*index, 0);
1572 assert_matches!(&value.kind, TimelineEventKind::Decrypted { .. });
1573
1574 assert_let_timeout!(
1575 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1576 );
1577 assert_eq!(expected_room_id, room_id);
1578 assert!(generic_stream.is_empty());
1579 }
1580
1581 #[async_test]
1582 async fn test_redecryptor_updating_encryption_info() {
1583 let bob_span = tracing::info_span!("bob");
1584
1585 let room_id = room_id!("!test:localhost");
1586
1587 let event_factory = EventFactory::new().room(room_id);
1588 let (alice, bob, matrix_mock_server, _) = set_up_clients(room_id, false, false).await;
1589
1590 let (event, room_key) =
1591 prepare_room(&matrix_mock_server, &event_factory, &alice, &bob, room_id).await;
1592
1593 let event_cache = bob.event_cache();
1596 let (room_cache, _) = event_cache
1597 .room(room_id)
1598 .instrument(bob_span.clone())
1599 .await
1600 .expect("We should be able to get to the event cache for a specific room");
1601
1602 let (_, mut subscriber) = room_cache.subscribe().await.unwrap();
1603 let mut generic_stream = event_cache.subscribe_to_room_generic_updates();
1604
1605 matrix_mock_server
1607 .mock_sync()
1608 .ok_and_run(&bob, |builder| {
1609 builder.add_joined_room(JoinedRoomBuilder::new(room_id).add_timeline_event(event));
1610 })
1611 .instrument(bob_span.clone())
1612 .await;
1613
1614 assert_let_timeout!(
1617 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1618 subscriber.recv()
1619 );
1620
1621 assert_eq!(diffs.len(), 1);
1624 assert_matches!(&diffs[0], VectorDiff::Append { values });
1625 assert_eq!(values.len(), 1);
1626 assert_matches!(&values[0].kind, TimelineEventKind::UnableToDecrypt { .. });
1627
1628 assert_let_timeout!(
1629 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1630 );
1631 assert_eq!(expected_room_id, room_id);
1632 assert!(generic_stream.is_empty());
1633
1634 matrix_mock_server
1636 .mock_sync()
1637 .ok_and_run(&bob, |builder| {
1638 builder.add_to_device_event(
1639 room_key
1640 .deserialize_as()
1641 .expect("We should be able to deserialize the room key"),
1642 );
1643 })
1644 .instrument(bob_span.clone())
1645 .await;
1646
1647 assert_let_timeout!(
1649 Duration::from_secs(1),
1650 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1651 subscriber.recv()
1652 );
1653
1654 assert_eq!(diffs.len(), 1);
1656 assert_matches!(&diffs[0], VectorDiff::Set { index: 0, value });
1657 assert_matches!(&value.kind, TimelineEventKind::Decrypted { .. });
1658
1659 let encryption_info = value.encryption_info().unwrap();
1660 assert_matches!(&encryption_info.verification_state, VerificationState::Unverified(_));
1661
1662 assert_let_timeout!(
1663 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1664 );
1665 assert_eq!(expected_room_id, room_id);
1666 assert!(generic_stream.is_empty());
1667
1668 let session_id = encryption_info.session_id().unwrap().to_owned();
1669 let alice_user_id = alice.user_id().unwrap();
1670
1671 alice
1673 .encryption()
1674 .bootstrap_cross_signing(None)
1675 .await
1676 .expect("Alice should be able to create the cross-signing keys");
1677
1678 bob.update_tracked_users_for_testing([alice_user_id]).instrument(bob_span.clone()).await;
1679 matrix_mock_server
1680 .mock_sync()
1681 .ok_and_run(&bob, |builder| {
1682 builder.add_change_device(alice_user_id);
1683 })
1684 .instrument(bob_span.clone())
1685 .await;
1686
1687 bob.event_cache().request_decryption(DecryptionRetryRequest {
1688 room_id: room_id.into(),
1689 utd_session_ids: BTreeSet::new(),
1690 refresh_info_session_ids: BTreeSet::from([session_id]),
1691 });
1692
1693 assert_let_timeout!(
1696 Duration::from_secs(1),
1697 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1698 subscriber.recv()
1699 );
1700
1701 assert_eq!(diffs.len(), 1);
1702 assert_matches!(&diffs[0], VectorDiff::Set { index: 0, value });
1703 assert_matches!(&value.kind, TimelineEventKind::Decrypted { .. });
1704 let encryption_info = value.encryption_info().unwrap();
1705
1706 assert_matches!(
1707 &encryption_info.verification_state,
1708 VerificationState::Unverified(_),
1709 "The event should now know about the identity but still be unverified"
1710 );
1711
1712 assert_let_timeout!(
1713 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1714 );
1715 assert_eq!(expected_room_id, room_id);
1716 assert!(generic_stream.is_empty());
1717 }
1718
1719 #[async_test]
1720 async fn test_event_is_redecrypted_even_if_key_arrives_while_event_processing() {
1721 let room_id = room_id!("!test:localhost");
1722
1723 let event_factory = EventFactory::new().room(room_id);
1724 let (alice, bob, matrix_mock_server, delayed_store) =
1725 set_up_clients(room_id, true, true).await;
1726
1727 let delayed_store = delayed_store.unwrap();
1728
1729 let (event, room_key) =
1730 prepare_room(&matrix_mock_server, &event_factory, &alice, &bob, room_id).await;
1731
1732 let event_cache = bob.event_cache();
1733
1734 let (room_cache, _) = event_cache
1736 .room(room_id)
1737 .await
1738 .expect("We should be able to get to the event cache for a specific room");
1739
1740 let (_, mut subscriber) = room_cache.subscribe().await.unwrap();
1741 let mut generic_stream = event_cache.subscribe_to_room_generic_updates();
1742
1743 matrix_mock_server
1745 .mock_sync()
1746 .ok_and_run(&bob, |builder| {
1747 builder.add_joined_room(JoinedRoomBuilder::new(room_id).add_timeline_event(event));
1748 })
1749 .await;
1750
1751 matrix_mock_server
1753 .mock_sync()
1754 .ok_and_run(&bob, |builder| {
1755 builder.add_to_device_event(
1756 room_key
1757 .deserialize_as()
1758 .expect("We should be able to deserialize the room key"),
1759 );
1760 })
1761 .await;
1762
1763 info!("Stopping the delay");
1764 delayed_store.stop_delaying().await;
1765
1766 assert_let_timeout!(
1772 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1773 subscriber.recv()
1774 );
1775
1776 assert_eq!(diffs.len(), 1);
1779 assert_matches!(&diffs[0], VectorDiff::Append { values });
1780 assert_eq!(values.len(), 1);
1781 assert_matches!(&values[0].kind, TimelineEventKind::UnableToDecrypt { .. });
1782
1783 assert_let_timeout!(
1785 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1786 );
1787 assert_eq!(expected_room_id, room_id);
1788
1789 assert_let_timeout!(
1791 Duration::from_secs(1),
1792 Ok(RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs { diffs, .. })) =
1793 subscriber.recv()
1794 );
1795
1796 assert_eq!(diffs.len(), 1);
1798 assert_matches!(&diffs[0], VectorDiff::Set { index, value });
1799 assert_eq!(*index, 0);
1800 assert_matches!(&value.kind, TimelineEventKind::Decrypted { .. });
1801
1802 assert_let_timeout!(
1804 Ok(RoomEventCacheGenericUpdate { room_id: expected_room_id }) = generic_stream.recv()
1805 );
1806 assert_eq!(expected_room_id, room_id);
1807 assert!(generic_stream.is_empty());
1808 }
1809}