diff --git a/crates/matrix-sdk-ui/tests/integration/timeline/pinned_event.rs b/crates/matrix-sdk-ui/tests/integration/timeline/pinned_event.rs index 0b923e0ef..a2ab00502 100644 --- a/crates/matrix-sdk-ui/tests/integration/timeline/pinned_event.rs +++ b/crates/matrix-sdk-ui/tests/integration/timeline/pinned_event.rs @@ -352,7 +352,10 @@ async fn test_new_pinned_event_ids_reload_the_timeline() { .await .expect("Sync failed"); - assert_let_timeout!(Some(timeline_updates) = timeline_stream.next()); + assert_let_timeout!( + Duration::from_millis(300), + Some(timeline_updates) = timeline_stream.next() + ); assert_eq!(timeline_updates.len(), 4); assert_let!(VectorDiff::Clear = &timeline_updates[0]); diff --git a/crates/matrix-sdk/src/event_cache/room/mod.rs b/crates/matrix-sdk/src/event_cache/room/mod.rs index 56798b9a0..9c395cebb 100644 --- a/crates/matrix-sdk/src/event_cache/room/mod.rs +++ b/crates/matrix-sdk/src/event_cache/room/mod.rs @@ -252,7 +252,7 @@ impl RoomEventCache { let mut state = self.inner.state.write().await?; let event_cache = self.clone(); - Ok(state.subscribe_to_pinned_events(room, event_cache).await) + state.subscribe_to_pinned_events(room, event_cache).await } /// Paginate backwards in a thread, given its root event ID. @@ -1867,7 +1867,7 @@ mod private { &mut self, room: Room, event_cache: RoomEventCache, - ) -> (Vec, Receiver) { + ) -> Result<(Vec, Receiver), EventCacheError> { let pinned_event_cache = self.state.pinned_event_cache.get_or_init(|| { PinnedEventCache::new( room, diff --git a/crates/matrix-sdk/src/event_cache/room/pinned_events.rs b/crates/matrix-sdk/src/event_cache/room/pinned_events.rs index eeb5c92d6..516b8a2c0 100644 --- a/crates/matrix-sdk/src/event_cache/room/pinned_events.rs +++ b/crates/matrix-sdk/src/event_cache/room/pinned_events.rs @@ -17,7 +17,10 @@ use std::{collections::BTreeSet, sync::Arc}; use futures_util::{StreamExt as _, stream}; use matrix_sdk_base::{ deserialized_responses::TimelineEventKind, - event_cache::{Event, Gap, store::EventCacheStoreLock}, + event_cache::{ + Event, Gap, + store::{EventCacheStoreLock, EventCacheStoreLockGuard, EventCacheStoreLockState}, + }, linked_chunk::{LinkedChunkId, OwnedLinkedChunkId, Update}, }; use ruma::{ @@ -30,13 +33,9 @@ use ruma::{ uint, }; use serde::Deserialize; -use tokio::{ - spawn, - sync::{ - RwLock, - broadcast::{Receiver, Sender}, - }, - task::JoinHandle, +use tokio::sync::{ + Mutex, RwLock, RwLockReadGuard, RwLockWriteGuard, + broadcast::{Receiver, Sender}, }; use tracing::{debug, instrument, trace, warn}; @@ -64,6 +63,7 @@ struct PinnedEventCacheState { chunk: EventLinkedChunk, /// Reference to the underlying backing store. + // TODO: can be removed? store: EventCacheStoreLock, /// A sender for the globally observable linked chunk updates that happened @@ -83,16 +83,163 @@ impl std::fmt::Debug for PinnedEventCacheState { } } -impl PinnedEventCacheState { - /// Return a list of the current event IDs in this linked chunk. - fn current_event_ids(&self) -> Vec { - self.chunk.events().filter_map(|(_position, event)| event.event_id()).collect() +struct PinnedEventCacheStateLock { + /// The per-thread lock around the real state. + locked_state: RwLock, + + /// A lock to guard against races when upgrading a read lock into a write + /// lock, after noticing the cross-process lock has been dirtied. + state_lock_upgrade_mutex: Mutex<()>, +} + +impl PinnedEventCacheStateLock { + /// Lock this [`PinnedEventCacheStateLock`] with per-thread shared access. + /// + /// This method locks the per-thread lock over the state, and then locks + /// the cross-process lock over the store. It returns an RAII guard + /// which will drop the read access to the state and to the store when + /// dropped. + /// + /// If the cross-process lock over the store is dirty (see + /// [`EventCacheStoreLockState`]), the state is reset to the last chunk. + async fn read(&self) -> Result> { + // Se comment in [`RoomEventCacheStateLock::read`] for explanation. + let _state_lock_upgrade_guard = self.state_lock_upgrade_mutex.lock().await; + + // Obtain a read lock. + let state_guard = self.locked_state.read().await; + + match state_guard.store.lock().await? { + EventCacheStoreLockState::Clean(store_guard) => { + Ok(PinnedEventCacheStateLockReadGuard { state: state_guard, _store: store_guard }) + } + + EventCacheStoreLockState::Dirty(store_guard) => { + // Drop the read lock, and take a write lock to modify the state. + // This is safe because only one reader at a time (see + // `Self::state_lock_upgrade_mutex`) is allowed. + drop(state_guard); + let state_guard = self.locked_state.write().await; + + let guard = + PinnedEventCacheStateLockWriteGuard { state: state_guard, store: store_guard }; + + // Force to reload by shrinking to the last chunk. + // TODO: reload the full pinned events list from the store. + //let updates_as_vector_diffs = guard.force_shrink_to_last_chunk().await?; + let updates_as_vector_diffs = Vec::new(); + + // All good now, mark the cross-process lock as non-dirty. + EventCacheStoreLockGuard::clear_dirty(&guard.store); + + // Downgrade the guard as soon as possible. + let guard = guard.downgrade(); + + // Now let the world know about the reload. + if !updates_as_vector_diffs.is_empty() { + // Notify observers about the update. + let _ = guard.state.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { + diffs: updates_as_vector_diffs, + origin: EventsOrigin::Cache, + }); + } + + Ok(guard) + } + } } + /// Lock this [`PinnedEventCacheStateLock`] with exclusive per-thread + /// write access. + /// + /// This method locks the per-thread lock over the state, and then locks + /// the cross-process lock over the store. It returns an RAII guard + /// which will drop the write access to the state and to the store when + /// dropped. + /// + /// If the cross-process lock over the store is dirty (see + /// [`EventCacheStoreLockState`]), the state is reset to the last chunk. + async fn write(&self) -> Result> { + let state_guard = self.locked_state.write().await; + + match state_guard.store.lock().await? { + EventCacheStoreLockState::Clean(store_guard) => { + Ok(PinnedEventCacheStateLockWriteGuard { state: state_guard, store: store_guard }) + } + + EventCacheStoreLockState::Dirty(store_guard) => { + let guard = + PinnedEventCacheStateLockWriteGuard { state: state_guard, store: store_guard }; + + // TODO: reload the full pinned events list from the store. + //let updates_as_vector_diffs = guard.force_shrink_to_last_chunk().await?; + let updates_as_vector_diffs = Vec::new(); + + // All good now, mark the cross-process lock as non-dirty. + EventCacheStoreLockGuard::clear_dirty(&guard.store); + + // Now let the world know about the reload. + if !updates_as_vector_diffs.is_empty() { + // Notify observers about the update. + let _ = guard.state.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { + diffs: updates_as_vector_diffs, + origin: EventsOrigin::Cache, + }); + } + + Ok(guard) + } + } + } +} + +#[cfg(not(tarpaulin_include))] +impl std::fmt::Debug for PinnedEventCacheStateLock { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PinnedEventCacheStateLock") + .field("locked_state", &self.locked_state) + .finish_non_exhaustive() + } +} + +/// The read lock guard returned by [`PinnedEventCacheStateLock::read`]. +pub struct PinnedEventCacheStateLockReadGuard<'a> { + /// The per-thread read lock guard over the + /// [`PinnedEventCacheState`]. + state: RwLockReadGuard<'a, PinnedEventCacheState>, + + /// The cross-process lock guard over the store. + _store: EventCacheStoreLockGuard, +} + +/// The write lock guard return by [`PinnedEventCacheStateLock::write`]. +struct PinnedEventCacheStateLockWriteGuard<'a> { + /// The per-thread write lock guard over the + /// [`PinnedEventCacheState`]. + state: RwLockWriteGuard<'a, PinnedEventCacheState>, + + /// The cross-process lock guard over the store. + store: EventCacheStoreLockGuard, +} + +impl<'a> PinnedEventCacheStateLockWriteGuard<'a> { + /// Synchronously downgrades a write lock into a read lock. + /// + /// The per-thread/state lock is downgraded atomically, without allowing + /// any writers to take exclusive access of the lock in the meantime. + /// + /// It returns an RAII guard which will drop the write access to the + /// state and to the store when dropped. + fn downgrade(self) -> PinnedEventCacheStateLockReadGuard<'a> { + PinnedEventCacheStateLockReadGuard { state: self.state.downgrade(), _store: self.store } + } +} + +impl<'a> PinnedEventCacheStateLockWriteGuard<'a> { async fn replace_all_events(&mut self, new_events: Vec) -> Result<()> { trace!("resetting all pinned events in linked chunk"); - let previous_pinned_event_ids = self.current_event_ids(); + let previous_pinned_event_ids = self.state.current_event_ids(); if new_events.iter().filter_map(|e| e.event_id()).collect::>() == previous_pinned_event_ids.iter().cloned().collect() @@ -101,14 +248,14 @@ impl PinnedEventCacheState { return Ok(()); } - self.chunk.reset(); - self.chunk.push_live_events(None, &new_events); + self.state.chunk.reset(); + self.state.chunk.push_live_events(None, &new_events); self.propagate_changes().await?; - let diffs = self.chunk.updates_as_vector_diffs(); + let diffs = self.state.chunk.updates_as_vector_diffs(); if !diffs.is_empty() { - let _ = self.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { + let _ = self.state.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { diffs, origin: EventsOrigin::Sync, }); @@ -120,7 +267,7 @@ impl PinnedEventCacheState { /// Propagate the changes in this linked chunk to observers, and save the /// changes on disk. async fn propagate_changes(&mut self) -> Result<()> { - let updates = self.chunk.store_updates().take(); + let updates = self.state.chunk.store_updates().take(); self.send_updates_to_store(updates).await } @@ -154,18 +301,13 @@ impl PinnedEventCacheState { // storing updates happens in the expected order. let store = self.store.clone(); - let room_id = self.room_id.clone(); + let room_id = self.state.room_id.clone(); let cloned_updates = updates.clone(); spawn(async move { trace!(updates = ?cloned_updates, "sending linked chunk updates to the store"); let linked_chunk_id = LinkedChunkId::PinnedEvents(&room_id); - // TODO: do not get the lock here, but above in the program! - let locked_store = store.lock().await?; - // TODO: better be clean! - let store = locked_store.as_clean().unwrap(); - store.handle_linked_chunk_updates(linked_chunk_id, cloned_updates).await?; trace!("linked chunk updates applied"); @@ -175,8 +317,8 @@ impl PinnedEventCacheState { .expect("joining failed")?; // Forward that the store got updated to observers. - let _ = self.linked_chunk_update_sender.send(RoomEventCacheLinkedChunkUpdate { - linked_chunk_id: OwnedLinkedChunkId::PinnedEvents(self.room_id.clone()), + let _ = self.state.linked_chunk_update_sender.send(RoomEventCacheLinkedChunkUpdate { + linked_chunk_id: OwnedLinkedChunkId::PinnedEvents(self.state.room_id.clone()), updates, }); @@ -231,9 +373,16 @@ impl PinnedEventCacheState { } } +impl PinnedEventCacheState { + /// Return a list of the current event IDs in this linked chunk. + fn current_event_ids(&self) -> Vec { + self.chunk.events().filter_map(|(_position, event)| event.event_id()).collect() + } +} + /// All the information related to a room's pinned events cache. pub struct PinnedEventCache { - state: Arc>, + state: Arc, /// The task handling the refreshing of pinned events for this specific /// room. @@ -263,14 +412,12 @@ impl PinnedEventCache { let chunk = EventLinkedChunk::new(); - let state = PinnedEventCacheState { - room_id: room_id.clone(), - chunk, - sender, - linked_chunk_update_sender, - store, - }; - let state = Arc::new(RwLock::new(state)); + let state = + PinnedEventCacheState { room_id, chunk, sender, linked_chunk_update_sender, store }; + let state = Arc::new(PinnedEventCacheStateLock { + locked_state: RwLock::new(state), + state_lock_upgrade_mutex: Mutex::new(()), + }); let task = Arc::new(spawn(Self::pinned_event_listener_task(room, state.clone(), event_cache))); @@ -279,13 +426,13 @@ impl PinnedEventCache { } /// Subscribe to live events from this room's pinned events cache. - pub async fn subscribe(&self) -> (Vec, Receiver) { - let state = self.state.read().await; - let events = state.chunk.events().map(|(_position, item)| item.clone()).collect(); + pub async fn subscribe(&self) -> Result<(Vec, Receiver)> { + let guard = self.state.read().await?; + let events = guard.state.chunk.events().map(|(_position, item)| item.clone()).collect(); - let recv = state.sender.subscribe(); + let recv = guard.state.sender.subscribe(); - (events, recv) + Ok((events, recv)) } /// Given a raw event, try to extract the target event ID of a relation as @@ -345,10 +492,10 @@ impl PinnedEventCache { room_redaction_rules: &RedactionRules, ) -> Result<()> { trace!("checking live events for relations to pinned events"); - let mut state = self.state.write().await; + let mut guard = self.state.write().await?; let pinned_event_ids: BTreeSet = - state.current_event_ids().await.into_iter().collect(); + guard.state.current_event_ids().into_iter().collect(); if pinned_event_ids.is_empty() { return Ok(()); @@ -381,13 +528,13 @@ impl PinnedEventCache { trace!("found {} new related events to pinned events", new_relations.len()); // We've found new relations; append them to the linked chunk. - state.chunk.push_live_events(None, &new_relations); + guard.state.chunk.push_live_events(None, &new_relations); - state.propagate_changes().await?; + guard.propagate_changes().await?; - let diffs = state.chunk.updates_as_vector_diffs(); + let diffs = guard.state.chunk.updates_as_vector_diffs(); if !diffs.is_empty() { - let _ = state.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { + let _ = guard.state.sender.send(RoomEventCacheUpdate::UpdateTimelineEvents { diffs, origin: EventsOrigin::Sync, }); @@ -400,7 +547,7 @@ impl PinnedEventCache { #[instrument(fields(%room_id = room.room_id()), skip(room, state, event_cache))] async fn pinned_event_listener_task( room: Room, - state: Arc>, + state: Arc, event_cache: RoomEventCache, ) { debug!("pinned events listener task started"); @@ -417,9 +564,19 @@ impl PinnedEventCache { // Replace the whole linked chunk with those new events, and propagate updates // to the observers. - state.write().await.replace_all_events(events).await.unwrap_or_else(|err| { - warn!("error when replacing initial pinned events: {err}"); - }); + match state.write().await { + Ok(mut guard) => { + guard.replace_all_events(events).await.unwrap_or_else(|err| { + warn!("error when replacing initial pinned events: {err}"); + }); + } + + Err(err) => { + warn!( + "error when acquiring write lock to replace initial pinned events: {err}" + ); + } + } }; // TODO: reload from persisted cache! @@ -438,9 +595,16 @@ impl PinnedEventCache { while let Some(new_list) = stream.next().await { trace!("handling update"); + let guard = match state.read().await { + Ok(guard) => guard, + Err(err) => { + warn!("error when acquiring read lock to handle pinned events update: {err}"); + break; + } + }; + // Compare to the current linked chunk. - let current_set = - state.read().await.current_event_ids().await.into_iter().collect::>(); + let current_set = guard.state.current_event_ids().into_iter().collect::>(); if !new_list.is_empty() && new_list.iter().all(|event_id| current_set.contains(event_id))