refactor(event cache): using proper locking for the pinned event cache

This commit is contained in:
Benjamin Bouvier
2026-02-03 17:57:41 +01:00
parent ca027f7eb6
commit da5ac9e3e6
3 changed files with 222 additions and 55 deletions
@@ -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]);
@@ -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<Event>, Receiver<RoomEventCacheUpdate>) {
) -> Result<(Vec<Event>, Receiver<RoomEventCacheUpdate>), EventCacheError> {
let pinned_event_cache = self.state.pinned_event_cache.get_or_init(|| {
PinnedEventCache::new(
room,
@@ -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<OwnedEventId> {
self.chunk.events().filter_map(|(_position, event)| event.event_id()).collect()
struct PinnedEventCacheStateLock {
/// The per-thread lock around the real state.
locked_state: RwLock<PinnedEventCacheState>,
/// 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<PinnedEventCacheStateLockReadGuard<'_>> {
// 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<PinnedEventCacheStateLockWriteGuard<'_>> {
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<Event>) -> 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::<BTreeSet<_>>()
== 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<OwnedEventId> {
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<RwLock<PinnedEventCacheState>>,
state: Arc<PinnedEventCacheStateLock>,
/// 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<Event>, Receiver<RoomEventCacheUpdate>) {
let state = self.state.read().await;
let events = state.chunk.events().map(|(_position, item)| item.clone()).collect();
pub async fn subscribe(&self) -> Result<(Vec<Event>, Receiver<RoomEventCacheUpdate>)> {
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<OwnedEventId> =
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<RwLock<PinnedEventCacheState>>,
state: Arc<PinnedEventCacheStateLock>,
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::<BTreeSet<_>>();
let current_set = guard.state.current_event_ids().into_iter().collect::<BTreeSet<_>>();
if !new_list.is_empty()
&& new_list.iter().all(|event_id| current_set.contains(event_id))