1mod updates;
16
17use std::{cmp::Ordering, collections::BTreeSet, fmt, sync::Arc};
18
19use eyeball_im::VectorDiff;
20use futures_util::{StreamExt as _, stream};
21use matrix_sdk_base::{
22 apply_redaction,
23 event_cache::{Event, Gap},
24 linked_chunk::{LinkedChunkId, OwnedLinkedChunkId, Position, Update},
25 serde_helpers::extract_redaction_target,
26 sync::{JoinedRoomUpdate, LeftRoomUpdate, Timeline},
27 task_monitor::BackgroundTaskHandle,
28};
29use matrix_sdk_common::executor::spawn;
30use ruma::{
31 EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedRoomId, OwnedUserId,
32 events::{relation::RelationType, room::redaction::SyncRoomRedactionEvent},
33 room_version_rules::RoomVersionRules,
34};
35use tokio::sync::broadcast::{Receiver, Sender};
36use tracing::{debug, instrument, trace, warn};
37
38pub(super) use self::updates::PinnedEventsCacheUpdateSender;
39#[cfg(feature = "e2e-encryption")]
40use super::super::redecryptor::ResolvedUtd;
41use super::{
42 super::{
43 EventCacheError, EventsOrigin, Result,
44 deduplicator::{DeduplicationOutcome, filter_duplicate_events},
45 persistence::{find_event, send_updates_to_store},
46 states::{
47 CacheStateLock, ReloadPreprocessing, StateLock, StateLockWriteGuard,
48 selectors::PinnedEventsStateSelector,
49 },
50 },
51 EventLocation, TimelineVectorDiffs,
52 event_linked_chunk::{EventLinkedChunk, sort_positions_descending},
53 room::RoomEventCacheLinkedChunkUpdate,
54};
55use crate::{Room, client::WeakClient, config::RequestConfig, room::WeakRoom};
56
57pub struct PinnedEventsCacheState {
58 room_id: OwnedRoomId,
60
61 own_user_id: OwnedUserId,
63
64 room_version_rules: RoomVersionRules,
66
67 chunk: EventLinkedChunk,
74
75 pub update_sender: PinnedEventsCacheUpdateSender,
77
78 linked_chunk_update_sender: Sender<RoomEventCacheLinkedChunkUpdate>,
83}
84
85#[cfg(not(tarpaulin_include))]
86impl fmt::Debug for PinnedEventsCacheState {
87 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
88 f.debug_struct("PinnedEventsCacheState")
89 .field("room_id", &self.room_id)
90 .field("chunk", &self.chunk)
91 .finish_non_exhaustive()
92 }
93}
94
95impl<'a> StateLockWriteGuard<'a, PinnedEventsCacheState> {
96 #[must_use = "Propagate `VectorDiff` updates via `TimelineVectorDiffs`"]
102 pub async fn reload(
103 &mut self,
104 preprocessing: ReloadPreprocessing,
105 ) -> Result<Vec<VectorDiff<Event>>> {
106 match preprocessing {
107 ReloadPreprocessing::ForgetAll => {
108 self.state.chunk.reset();
110 self.propagate_changes().await?;
111 }
112
113 ReloadPreprocessing::None => {}
114 }
115
116 self.reload_from_storage().await?;
119
120 Ok(self.state.chunk.updates_as_vector_diffs())
121 }
122
123 async fn handle_sync(&mut self, timeline: Timeline) -> Result<()> {
124 let DeduplicationOutcome {
125 all_events: events,
126 in_memory_duplicated_event_ids,
127 in_store_duplicated_event_ids,
128 non_empty_all_duplicates: all_duplicates,
129 } = filter_duplicate_events(
130 &self.state.own_user_id,
131 &self.store,
132 LinkedChunkId::PinnedEvents(&self.state.room_id),
133 &self.state.chunk,
134 timeline.events,
135 )
136 .await?;
137
138 if all_duplicates {
139 return Ok(());
142 }
143
144 self.remove_events(in_memory_duplicated_event_ids, in_store_duplicated_event_ids).await?;
149
150 self.state.chunk.push_live_events(None, &events);
152
153 self.propagate_changes().await?;
154 self.notify_subscribers(EventsOrigin::Sync);
155
156 for event in &events {
158 self.maybe_apply_new_redaction(event).await?;
160 }
161
162 Ok(())
163 }
164
165 #[instrument(skip_all)]
170 pub async fn remove_events(
171 &mut self,
172 in_memory_events: Vec<(OwnedEventId, Position)>,
173 in_store_events: Vec<(OwnedEventId, Position)>,
174 ) -> Result<()> {
175 if !in_store_events.is_empty() {
177 let mut positions = in_store_events
178 .into_iter()
179 .map(|(_event_id, position)| position)
180 .collect::<Vec<_>>();
181
182 sort_positions_descending(&mut positions);
183
184 let updates =
185 positions.into_iter().map(|pos| Update::RemoveItem { at: pos }).collect::<Vec<_>>();
186
187 self.apply_store_only_updates(updates).await?;
188 }
189
190 if in_memory_events.is_empty() {
192 return Ok(());
194 }
195
196 self.state
198 .chunk
199 .remove_events_by_position(
200 in_memory_events.into_iter().map(|(_event_id, position)| position).collect(),
201 )
202 .expect("failed to remove an event");
203
204 self.propagate_changes().await
205 }
206
207 async fn apply_store_only_updates(&mut self, updates: Vec<Update<Event, Gap>>) -> Result<()> {
213 self.send_updates_to_store(updates).await
214 }
215
216 #[instrument(skip_all)]
220 async fn maybe_apply_new_redaction(&mut self, event: &Event) -> Result<()> {
221 let Some(event_id) =
222 extract_redaction_target(event.raw(), &self.room_version_rules.redaction)
223 else {
224 return Ok(());
225 };
226
227 let Some((location, mut target_event)) = self.find_event(&event_id).await? else {
229 trace!("redacted event is missing from the linked chunk");
230 return Ok(());
231 };
232
233 let target_event_raw = target_event.raw();
234
235 if let Ok(deserialized) = target_event_raw.deserialize()
237 && deserialized.is_redacted()
238 {
239 return Ok(());
240 }
241
242 if let Some(redacted_event) = apply_redaction(
243 target_event_raw,
244 event.raw().cast_ref_unchecked::<SyncRoomRedactionEvent>(),
245 &self.room_version_rules.redaction,
246 ) {
247 target_event.replace_raw(redacted_event.cast_unchecked());
252
253 self.replace_event_at(location, target_event.clone()).await?;
254 }
255
256 Ok(())
257 }
258
259 pub(super) async fn find_event(
261 &self,
262 event_id: &EventId,
263 ) -> Result<Option<(EventLocation, Event)>> {
264 find_event(event_id, &self.room_id, &self.chunk, &self.store).await
265 }
266
267 pub async fn replace_event_at(
274 &mut self,
275 location: EventLocation,
276 new_event: Event,
277 ) -> Result<()> {
278 match location {
279 EventLocation::Memory(position) => {
280 self.state
281 .chunk
282 .replace_event_at(position, new_event)
283 .expect("should have been a valid position of an item");
284 self.propagate_changes().await?;
287 }
288 EventLocation::Store => {
289 self.save_events([new_event]).await?;
290 }
291 }
292
293 Ok(())
294 }
295
296 pub async fn save_events(&mut self, events: impl IntoIterator<Item = Event>) -> Result<()> {
298 let store = self.store.clone();
299 let room_id = self.state.room_id.clone();
300 let events = events.into_iter().collect::<Vec<_>>();
301
302 spawn(async move {
304 for event in events {
305 store.save_event(&room_id, event).await?;
306 }
307
308 Result::Ok(())
309 })
310 .await
311 .expect("joining failed")?;
312
313 Ok(())
314 }
315
316 async fn reload_from_storage(&mut self) -> Result<()> {
319 let room_id = self.state.room_id.clone();
320 let linked_chunk_id = LinkedChunkId::PinnedEvents(&room_id);
321
322 let (last_chunk, chunk_id_gen) = self.store.load_last_chunk(linked_chunk_id).await?;
323
324 let Some(last_chunk) = last_chunk else {
325 if self.state.chunk.events().next().is_some() {
328 self.state.chunk.reset();
329 self.notify_subscribers(EventsOrigin::Sync);
330 }
331
332 return Ok(());
333 };
334
335 {
336 let mut current_chunk_identifier = last_chunk.identifier;
337 self.state.chunk.shrink_to_last_reloaded_chunk(
338 Some(last_chunk),
339 chunk_id_gen,
340 None,
342 )?;
343
344 while let Some(previous_chunk) =
346 self.store.load_previous_chunk(linked_chunk_id, current_chunk_identifier).await?
347 {
348 current_chunk_identifier = previous_chunk.identifier;
349 self.state.chunk.insert_new_chunk_as_first(previous_chunk)?;
350 }
351 }
352
353 self.state.chunk.store_updates().take();
355
356 self.notify_subscribers(EventsOrigin::Cache);
358
359 Ok(())
360 }
361
362 async fn replace_all_events(&mut self, new_events: Vec<Event>) -> Result<()> {
363 trace!("resetting all pinned events in linked chunk");
364
365 let previous_pinned_event_ids = self.state.current_event_ids();
366
367 if new_events
368 .iter()
369 .filter_map(|e| e.event_id())
370 .map(ToOwned::to_owned)
371 .collect::<BTreeSet<_>>()
372 == previous_pinned_event_ids.into_iter().collect()
373 {
374 return Ok(());
376 }
377
378 if self.state.chunk.events().next().is_some() {
379 self.state.chunk.reset();
380 }
381
382 self.state.chunk.push_live_events(None, &new_events);
383 self.propagate_changes().await?;
384 self.notify_subscribers(EventsOrigin::Sync);
385
386 Ok(())
387 }
388
389 pub async fn propagate_changes(&mut self) -> Result<()> {
392 let updates = self.state.chunk.store_updates().take();
393
394 self.send_updates_to_store(updates).await
395 }
396
397 async fn send_updates_to_store(&mut self, updates: Vec<Update<Event, Gap>>) -> Result<()> {
398 let linked_chunk_id = OwnedLinkedChunkId::PinnedEvents(self.room_id.clone());
399
400 send_updates_to_store(
401 &self.store,
402 linked_chunk_id,
403 &self.state.linked_chunk_update_sender,
404 updates,
405 )
406 .await
407 }
408
409 fn notify_subscribers(&mut self, origin: EventsOrigin) {
411 let diffs = self.state.chunk.updates_as_vector_diffs();
412
413 if !diffs.is_empty() {
414 self.update_sender.send(TimelineVectorDiffs { diffs, origin });
415 }
416 }
417}
418
419impl PinnedEventsCacheState {
420 pub(super) fn current_event_ids(&self) -> Vec<OwnedEventId> {
422 self.chunk
423 .events()
424 .filter_map(|(_position, event)| event.event_id().map(ToOwned::to_owned))
425 .collect()
426 }
427}
428
429#[derive(Clone)]
433pub struct PinnedEventsCache {
434 inner: Arc<PinnedEventsCacheInner>,
435
436 _task: Arc<BackgroundTaskHandle>,
439}
440
441struct PinnedEventsCacheInner {
443 room_id: OwnedRoomId,
445
446 state: CacheStateLock<PinnedEventsStateSelector>,
450}
451
452impl PinnedEventsCache {
453 pub(in super::super) async fn new(
455 weak_room: &WeakRoom,
456 own_user_id: OwnedUserId,
457 room_version_rules: RoomVersionRules,
458 linked_chunk_update_sender: Sender<RoomEventCacheLinkedChunkUpdate>,
459 state: &StateLock,
460 ) -> Result<Self> {
461 let room = weak_room.get().ok_or(EventCacheError::ClientDropped)?;
462 let room_id = room.room_id().to_owned();
463
464 let cache_state = state
465 .try_insert_once_with(
466 PinnedEventsStateSelector::new(room_id.clone()),
467 |_store_guard| async {
468 Ok(PinnedEventsCacheState {
469 room_id: room_id.clone(),
470 own_user_id,
471 room_version_rules,
472 chunk: EventLinkedChunk::new(),
473 update_sender: PinnedEventsCacheUpdateSender::new(),
474 linked_chunk_update_sender,
475 })
476 },
477 )
478 .await?;
479
480 let inner = Arc::new(PinnedEventsCacheInner { room_id, state: cache_state });
481
482 let task = room
483 .client()
484 .task_monitor()
485 .spawn_infinite_task(
486 "pinned_event_listener_task",
487 Self::pinned_event_listener_task(room, inner.clone()),
488 )
489 .abort_on_drop();
490
491 Ok(Self { inner, _task: Arc::new(task) })
492 }
493
494 pub(super) fn state(&self) -> &CacheStateLock<PinnedEventsStateSelector> {
496 &self.inner.state
497 }
498
499 pub async fn subscribe(&self) -> Result<(Vec<Event>, Receiver<TimelineVectorDiffs>)> {
501 let guard = self.inner.state.read().await?;
502 let events = guard.state.chunk.events().map(|(_position, item)| item.clone()).collect();
503
504 let recv = guard.state.update_sender.new_pinned_events_receiver();
505
506 Ok((events, recv))
507 }
508
509 #[cfg(feature = "e2e-encryption")]
513 pub(in super::super) async fn replace_utds(&self, events: &[ResolvedUtd]) -> Result<()> {
514 let mut guard = self.inner.state.write().await?;
515
516 if guard.state.chunk.replace_utds(events) {
517 guard.propagate_changes().await?;
518 guard.notify_subscribers(EventsOrigin::Cache);
519 }
520
521 Ok(())
522 }
523
524 #[instrument(skip_all, fields(room_id = %self.inner.room_id))]
526 pub(super) async fn handle_joined_room_update(&self, updates: JoinedRoomUpdate) -> Result<()> {
527 self.handle_timeline(updates.timeline).await?;
528
529 Ok(())
530 }
531
532 #[instrument(skip_all, fields(room_id = %self.inner.room_id))]
534 pub(super) async fn handle_left_room_update(&self, updates: LeftRoomUpdate) -> Result<()> {
535 self.handle_timeline(updates.timeline).await?;
536
537 Ok(())
538 }
539
540 async fn handle_timeline(&self, timeline: Timeline) -> Result<()> {
543 if timeline.events.is_empty() {
544 return Ok(());
545 }
546
547 trace!("adding new {} events", timeline.events.len());
548
549 self.inner.state.write().await?.handle_sync(timeline).await?;
550
551 Ok(())
552 }
553
554 #[instrument(fields(%room_id = room.room_id()), skip(room, inner))]
555 async fn pinned_event_listener_task(room: Room, inner: Arc<PinnedEventsCacheInner>) {
556 debug!("pinned events listener task started");
557
558 let reload_from_network = async |room: Room| {
559 let events = match Self::reload_pinned_events(room).await {
560 Ok(Some(events)) => events,
561 Ok(None) => Vec::new(),
562 Err(err) => {
563 warn!("error when loading pinned events: {err}");
564 return;
565 }
566 };
567
568 match inner.state.write().await {
571 Ok(mut guard) => {
572 guard.replace_all_events(events).await.unwrap_or_else(|err| {
573 warn!("error when replacing pinned events: {err}");
574 });
575 }
576
577 Err(err) => {
578 warn!("error when acquiring write lock to replace pinned events: {err}");
579 }
580 }
581 };
582
583 match inner.state.write().await {
585 Ok(mut guard) => {
586 guard.reload_from_storage().await.unwrap_or_else(|err| {
588 warn!("error when reloading pinned events from storage, at start: {err}");
589 });
590
591 let actual_pinned_events = room.pinned_event_ids().unwrap_or_default();
593 let reloaded_set =
594 guard.state.current_event_ids().into_iter().collect::<BTreeSet<_>>();
595
596 if actual_pinned_events.len() != reloaded_set.len()
597 || actual_pinned_events.iter().any(|event_id| !reloaded_set.contains(event_id))
598 {
599 drop(guard);
601 reload_from_network(room.clone()).await;
602 }
603 }
604
605 Err(err) => {
606 warn!("error when acquiring write lock to initialize pinned events: {err}");
607 }
608 }
609
610 let weak_room =
611 WeakRoom::new(WeakClient::from_client(&room.client()), room.room_id().to_owned());
612
613 let mut stream = room.pinned_event_ids_stream();
614
615 drop(room);
616
617 while let Some(new_list) = stream.next().await {
619 trace!("handling update");
620
621 let guard = match inner.state.read().await {
622 Ok(guard) => guard,
623 Err(err) => {
624 warn!("error when acquiring read lock to handle pinned events update: {err}");
625 break;
626 }
627 };
628
629 let current_set = guard.state.current_event_ids().into_iter().collect::<BTreeSet<_>>();
631
632 if !new_list.is_empty()
633 && new_list.iter().all(|event_id| current_set.contains(event_id))
634 {
635 continue;
637 }
638
639 let Some(room) = weak_room.get() else {
640 debug!("room has been dropped, ending pinned events listener task");
641 break;
642 };
643
644 drop(guard);
645
646 reload_from_network(room).await;
648 }
649
650 debug!("pinned events listener task ended");
651 }
652
653 async fn reload_pinned_events(room: Room) -> Result<Option<Vec<Event>>> {
662 let (max_events_to_load, max_concurrent_requests) = {
663 let client = room.client();
664 let config = client.event_cache().config();
665 (config.max_pinned_events_to_load, config.max_pinned_events_concurrent_requests)
666 };
667
668 let pinned_event_ids: Vec<OwnedEventId> = room
669 .pinned_event_ids()
670 .unwrap_or_default()
671 .into_iter()
672 .rev()
673 .take(max_events_to_load)
674 .rev()
675 .collect();
676
677 if pinned_event_ids.is_empty() {
678 return Ok(Some(Vec::new()));
679 }
680
681 let mut num_successful_loads = 0;
682
683 let mut loaded_events: Vec<Event> =
684 stream::iter(pinned_event_ids.clone().into_iter().map(|event_id| {
685 let room = room.clone();
686 let filter = vec![RelationType::Annotation, RelationType::Replacement];
687 let request_config = RequestConfig::default().retry_limit(3);
688
689 async move {
690 let (target, mut relations) = room
691 .load_or_fetch_event_with_relations(
692 &event_id,
693 Some(filter),
694 Some(request_config),
695 )
696 .await?;
697
698 relations.insert(0, target);
699 Ok::<_, crate::Error>(relations)
700 }
701 }))
702 .buffer_unordered(max_concurrent_requests)
703 .inspect(|result| {
705 if result.is_ok() {
706 num_successful_loads += 1;
707 }
708 })
709 .flat_map(stream::iter)
711 .flat_map(stream::iter)
713 .collect()
714 .await;
715
716 if num_successful_loads != pinned_event_ids.len() {
717 warn!(
718 "only successfully loaded {} out of {} pinned events",
719 num_successful_loads,
720 pinned_event_ids.len()
721 );
722 }
723
724 if loaded_events.is_empty() {
725 return Err(EventCacheError::UnableToLoadPinnedEvents);
728 }
729
730 loaded_events.sort_by(compare_pinned_items);
735
736 Ok(Some(loaded_events))
737 }
738}
739
740impl fmt::Debug for PinnedEventsCache {
741 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
742 f.debug_struct("PinnedEventsCache").finish_non_exhaustive()
743 }
744}
745
746fn compare_pinned_items(a: &Event, b: &Event) -> Ordering {
747 let a_time: Option<MilliSecondsSinceUnixEpoch> = a.timestamp_raw();
748 let b_time: Option<MilliSecondsSinceUnixEpoch> = b.timestamp_raw();
749
750 compare_by_optional_timestamp(a_time, b_time)
751}
752
753fn compare_by_optional_timestamp(
754 a: Option<MilliSecondsSinceUnixEpoch>,
755 b: Option<MilliSecondsSinceUnixEpoch>,
756) -> Ordering {
757 match (a, b) {
758 (None, None) => Ordering::Equal,
759 (None, Some(_)) => Ordering::Greater,
760 (Some(_), None) => Ordering::Less,
761 (Some(a), Some(b)) => a.cmp(&b),
762 }
763}
764
765#[cfg(not(target_family = "wasm"))]
766#[cfg(test)]
767mod tests {
768 use proptest::prelude::*;
769 use ruma::UInt;
770
771 use super::*;
772
773 fn any_timestamp() -> impl Strategy<Value = Option<MilliSecondsSinceUnixEpoch>> {
774 prop::option::of(
775 any::<u32>().prop_map(|value| MilliSecondsSinceUnixEpoch(UInt::from(value))),
776 )
777 }
778
779 #[test]
780 fn sort_pinned_events_never_panics_only_nones() {
781 let mut vec = vec![None; 100_000];
782 vec.sort_by(|a, b| compare_by_optional_timestamp(*a, *b))
783 }
784
785 proptest! {
786 #[test]
787 fn sort_pinned_events_never_panics(mut v in prop::collection::vec(any_timestamp(), 0..1000)) {
788 v.sort_by(
789 |a, b| compare_by_optional_timestamp(*a, *b))
790 }
791
792 #[test]
793 fn compare_pinned_events_reflexive(a in any_timestamp()) {
794 prop_assert_eq!(compare_by_optional_timestamp(a, a), Ordering::Equal);
795 }
796
797 #[test]
798 fn compare_pinned_events_antisymmetric(a in any_timestamp(), b in any_timestamp()) {
799 let ab = compare_by_optional_timestamp(a, b);
800 let ba = compare_by_optional_timestamp(b, a);
801
802 prop_assert_eq!(ab, ba.reverse());
803 }
804
805 #[test]
806 fn compare_pinned_events_transitive(
807 a in any_timestamp(),
808 b in any_timestamp(),
809 c in any_timestamp()
810 ) {
811 let ab = compare_by_optional_timestamp(a, b);
812 let bc = compare_by_optional_timestamp(b, c);
813 let ac = compare_by_optional_timestamp(a, c);
814
815 if ab == Ordering::Less && bc == Ordering::Less {
816 prop_assert_eq!(ac, Ordering::Less);
817 }
818
819 if ab == Ordering::Equal && bc == Ordering::Equal {
820 prop_assert_eq!(ac, Ordering::Equal);
821 }
822
823 if ab == Ordering::Greater && bc == Ordering::Greater {
824 prop_assert_eq!(ac, Ordering::Greater);
825 }
826 }
827 }
828}