use std::{
collections::HashMap,
fmt,
ops::{Deref, DerefMut},
sync::Arc,
};
use matrix_sdk_base::{
event_cache::store::{EventCacheStoreLock, EventCacheStoreLockGuard, EventCacheStoreLockState},
timer,
tracing_timer::TracingTimer,
};
use ruma::{OwnedEventId, OwnedRoomId, RoomId};
use tokio::sync::{Mutex, RwLock, RwLockMappedWriteGuard, RwLockReadGuard, RwLockWriteGuard};
use tracing::{instrument, trace};
use super::{
CachesByRoom, EventCacheError, EventsOrigin, Result,
caches::{
TimelineVectorDiffs,
event_focused::{EventFocusedCacheKey, EventFocusedCacheState},
pinned_events::PinnedEventsCacheState,
room::{self, RoomEventCacheState},
thread::{self, ThreadEventCacheState},
},
};
pub(in super::super) mod selectors;
pub struct State {
store: EventCacheStoreLock,
by_room: HashMap<OwnedRoomId, StateForRoom>,
}
#[derive(Default)]
pub(super) struct StateForRoom {
room: Option<RoomEventCacheState>,
threads: HashMap<OwnedEventId, ThreadEventCacheState>,
pinned_events: Option<PinnedEventsCacheState>,
event_focused: HashMap<EventFocusedCacheKey, EventFocusedCacheState>,
}
#[derive(Clone)]
pub struct StateLock {
inner: Arc<StateLockInner>,
}
struct StateLockInner {
locked_state: RwLock<State>,
state_lock_upgrade_mutex: Mutex<()>,
}
impl StateLock {
pub fn new(store: EventCacheStoreLock) -> Self {
Self {
inner: Arc::new(StateLockInner {
locked_state: RwLock::new(State { store, by_room: HashMap::new() }),
state_lock_upgrade_mutex: Mutex::new(()),
}),
}
}
#[instrument(skip_all)]
pub(super) async fn read<'state>(&'state self) -> Result<StateLockReadGuard<'state, State>> {
trace!("Acquiring the lock");
let tracing_timer = timer!("`read` lock");
let _state_lock_upgrade_guard = self.inner.state_lock_upgrade_mutex.lock().await;
let state_guard = self.inner.locked_state.read().await;
Ok(match state_guard.store.lock().await? {
EventCacheStoreLockState::Clean(store_guard) => {
trace!("Lock acquired (from clean)");
StateLockReadGuard {
state: StateLockReadGuardKind::Owned(state_guard),
store: store_guard,
tracing_timer: Some(tracing_timer),
}
}
EventCacheStoreLockState::Dirty(store_guard) => {
drop(state_guard);
let mut guard = ReloadableStateLockWriteGuard {
state: self.inner.locked_state.write().await,
store: store_guard,
tracing_timer,
};
guard.reload(ReloadPreprocessing::None).await?;
EventCacheStoreLockGuard::clear_dirty(&guard.store);
trace!("Lock acquired (from dirty)");
guard.downgrade()
}
})
}
#[instrument(skip_all)]
async fn write<'state>(&'state self) -> Result<ReloadableStateLockWriteGuard<'state>> {
trace!("Acquiring lock");
let tracing_timer = timer!("`write` lock");
let state_guard = self.inner.locked_state.write().await;
Ok(match state_guard.store.lock().await? {
EventCacheStoreLockState::Clean(store_guard) => {
trace!("Lock acquired (from clean)");
ReloadableStateLockWriteGuard {
state: state_guard,
store: store_guard,
tracing_timer,
}
}
EventCacheStoreLockState::Dirty(store_guard) => {
let mut guard = ReloadableStateLockWriteGuard {
state: state_guard,
store: store_guard,
tracing_timer,
};
guard.reload(ReloadPreprocessing::None).await?;
EventCacheStoreLockGuard::clear_dirty(&guard.store);
trace!("Lock acquired (from dirty)");
guard
}
})
}
#[instrument(skip_all)]
pub(super) async fn clear_and_reload(
&self,
_caches_for_all_rooms_exclusive_lock_guard: &RwLockWriteGuard<'_, CachesByRoom>,
room_id: Option<&RoomId>,
) -> Result<()> {
let tracing_timer = timer!("`clear_and_reload` lock");
let state_guard = self.inner.locked_state.write().await;
let mut guard = match state_guard.store.lock().await? {
EventCacheStoreLockState::Clean(store_guard)
| EventCacheStoreLockState::Dirty(store_guard) => ReloadableStateLockWriteGuard {
state: state_guard,
store: store_guard,
tracing_timer,
},
};
guard.store.clear_all_events(room_id).await?;
guard.reload(ReloadPreprocessing::ForgetAll).await?;
if EventCacheStoreLockGuard::is_dirty(&guard.store) {
EventCacheStoreLockGuard::clear_dirty(&guard.store);
}
Ok(())
}
#[instrument(skip_all)]
pub(super) async fn try_insert_once_with<Selector, Constructor>(
&self,
cache_state_selector: Selector,
cache_constructor: Constructor,
) -> Result<CacheStateLock<Selector>>
where
Selector: selectors::CacheState,
Constructor: AsyncFnOnce(EventCacheStoreLockGuard) -> Result<Selector::Item>,
{
let mut state = self.write().await?;
let cache_state = cache_constructor(state.store).await?;
cache_state_selector
.insert_once(&mut state.state, cache_state)
.then(|| CacheStateLock::new(cache_state_selector, self.clone()))
.ok_or_else(|| EventCacheError::CacheStateAlreadyExists)
}
}
impl fmt::Debug for StateLock {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("StateLock").finish_non_exhaustive()
}
}
pub struct StateLockReadGuard<'state, S> {
pub state: StateLockReadGuardKind<'state, S>,
pub store: EventCacheStoreLockGuard,
tracing_timer: Option<TracingTimer>,
}
impl<'state> StateLockReadGuard<'state, State> {
fn try_map_into_cache_state<'selector, Selector>(
self,
cache_state_selector: &'selector Selector,
) -> Result<StateLockReadGuard<'state, Selector::Item>>
where
Selector: selectors::CacheState,
EventCacheError: From<&'selector Selector>,
{
Ok(StateLockReadGuard {
state: match self.state {
StateLockReadGuardKind::Reference(state) => StateLockReadGuardKind::Reference(
cache_state_selector
.select(state)
.ok_or_else(|| EventCacheError::from(cache_state_selector))?,
),
StateLockReadGuardKind::Owned(state) => StateLockReadGuardKind::Owned(
RwLockReadGuard::try_map(state, |state| cache_state_selector.select(state))
.map_err(|_| EventCacheError::from(cache_state_selector))?,
),
},
store: self.store,
tracing_timer: self.tracing_timer,
})
}
}
impl<'state> StateLockReadGuard<'state, StateForRoom> {
pub(super) fn room(&'state self) -> Option<StateLockReadGuard<'state, RoomEventCacheState>> {
self.state.room.as_ref().map(|room| StateLockReadGuard {
state: StateLockReadGuardKind::Reference(room),
store: self.store.clone(),
tracing_timer: None,
})
}
pub(super) fn threads(
&'state self,
) -> StateLockReadGuard<'state, HashMap<OwnedEventId, ThreadEventCacheState>> {
StateLockReadGuard {
state: StateLockReadGuardKind::Reference(&self.state.threads),
store: self.store.clone(),
tracing_timer: None,
}
}
}
impl<'state> StateLockReadGuard<'state, HashMap<OwnedEventId, ThreadEventCacheState>> {
pub(super) fn values(
&'state self,
) -> impl Iterator<Item = StateLockReadGuard<'state, ThreadEventCacheState>> {
self.state.values().map(|item| StateLockReadGuard {
state: StateLockReadGuardKind::Reference(item),
store: self.store.clone(),
tracing_timer: None,
})
}
}
impl<'state, S> Deref for StateLockReadGuard<'state, S> {
type Target = S;
fn deref(&self) -> &Self::Target {
&self.state
}
}
pub enum StateLockReadGuardKind<'state, S> {
Reference(&'state S),
Owned(RwLockReadGuard<'state, S>),
}
impl<'state, S> Deref for StateLockReadGuardKind<'state, S> {
type Target = S;
fn deref(&self) -> &Self::Target {
match self {
Self::Reference(state) => state,
Self::Owned(state) => state.deref(),
}
}
}
struct ReloadableStateLockWriteGuard<'state> {
state: RwLockWriteGuard<'state, State>,
store: EventCacheStoreLockGuard,
tracing_timer: TracingTimer,
}
impl<'state> ReloadableStateLockWriteGuard<'state> {
fn try_map_into_cache_state<'selector, Selector>(
self,
cache_state_selector: &'selector Selector,
) -> Result<StateLockWriteGuard<'state, Selector::Item>>
where
Selector: selectors::CacheState,
EventCacheError: From<&'selector Selector>,
{
Ok(StateLockWriteGuard {
state: StateLockWriteGuardKind::Owned(
RwLockWriteGuard::try_map(self.state, |state| {
cache_state_selector.select_mut(state)
})
.map_err(|_| EventCacheError::from(cache_state_selector))?,
),
store: self.store,
_tracing_timer: Some(self.tracing_timer),
})
}
fn downgrade(self) -> StateLockReadGuard<'state, State> {
StateLockReadGuard {
state: StateLockReadGuardKind::Owned(self.state.downgrade()),
store: self.store,
tracing_timer: Some(self.tracing_timer),
}
}
async fn reload(&mut self, preprocessing: ReloadPreprocessing) -> Result<()> {
trace!("Reloading the state");
for (room_id, StateForRoom { room, threads, pinned_events, event_focused }) in
self.state.by_room.iter_mut()
{
if let Some(room_state) = room {
let mut room_state = StateLockWriteGuard {
state: StateLockWriteGuardKind::Reference(room_state),
store: self.store.clone(),
_tracing_timer: None,
};
let updates_as_vector_diffs = room_state.reload(preprocessing).await?;
room_state.update_sender.send(
room::RoomEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs {
diffs: updates_as_vector_diffs,
origin: EventsOrigin::Cache,
}),
Some(room::RoomEventCacheGenericUpdate { room_id: room_id.clone() }),
);
}
for thread_state in threads.values_mut() {
let mut thread_state = StateLockWriteGuard {
state: StateLockWriteGuardKind::Reference(thread_state),
store: self.store.clone(),
_tracing_timer: None,
};
let updates_as_vector_diffs = thread_state.reload(preprocessing).await?;
thread_state.update_sender.send(
thread::ThreadEventCacheUpdate::UpdateTimelineEvents(TimelineVectorDiffs {
diffs: updates_as_vector_diffs,
origin: EventsOrigin::Cache,
}),
Some(room::RoomEventCacheGenericUpdate { room_id: room_id.clone() }),
);
}
if let Some(pinned_events_state) = pinned_events {
let mut pinned_events_state = StateLockWriteGuard {
state: StateLockWriteGuardKind::Reference(pinned_events_state),
store: self.store.clone(),
_tracing_timer: None,
};
let updates_as_vector_diffs = pinned_events_state.reload(preprocessing).await?;
pinned_events_state.update_sender.send(TimelineVectorDiffs {
diffs: updates_as_vector_diffs,
origin: EventsOrigin::Cache,
});
}
for event_focused_state in event_focused.values_mut() {
let mut event_focused_state = StateLockWriteGuard {
state: StateLockWriteGuardKind::Reference(event_focused_state),
store: self.store.clone(),
_tracing_timer: None,
};
let updates_as_vector_diffs = event_focused_state.reload(preprocessing).await?;
let _ = event_focused_state.update_sender.send(TimelineVectorDiffs {
diffs: updates_as_vector_diffs,
origin: EventsOrigin::Cache,
});
}
}
Ok(())
}
}
pub struct StateLockWriteGuard<'state, S> {
pub state: StateLockWriteGuardKind<'state, S>,
pub store: EventCacheStoreLockGuard,
_tracing_timer: Option<TracingTimer>,
}
impl<'state, S> Deref for StateLockWriteGuard<'state, S> {
type Target = S;
fn deref(&self) -> &Self::Target {
&self.state
}
}
impl<'state, S> DerefMut for StateLockWriteGuard<'state, S> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.state
}
}
pub enum StateLockWriteGuardKind<'state, S> {
Reference(&'state mut S),
Owned(RwLockMappedWriteGuard<'state, S>),
}
impl<'state, S> Deref for StateLockWriteGuardKind<'state, S> {
type Target = S;
fn deref(&self) -> &Self::Target {
match self {
Self::Reference(state) => state,
Self::Owned(state) => state.deref(),
}
}
}
impl<'state, S> DerefMut for StateLockWriteGuardKind<'state, S> {
fn deref_mut(&mut self) -> &mut Self::Target {
match self {
Self::Reference(state) => state,
Self::Owned(state) => state.deref_mut(),
}
}
}
pub struct CacheStateLock<Selector> {
cache_state_selector: Selector,
state_lock: StateLock,
}
impl<Selector> CacheStateLock<Selector>
where
Selector: selectors::CacheState,
{
pub(super) fn new(cache_state_selector: Selector, state_lock: StateLock) -> Self {
Self { cache_state_selector, state_lock }
}
}
impl<Selector> CacheStateLock<Selector>
where
Selector: selectors::CacheState,
EventCacheError: for<'a> From<&'a Selector>,
{
pub async fn read(&self) -> Result<StateLockReadGuard<'_, Selector::Item>> {
self.state_lock.read().await?.try_map_into_cache_state(&self.cache_state_selector)
}
pub async fn write(&self) -> Result<StateLockWriteGuard<'_, Selector::Item>> {
self.state_lock.write().await?.try_map_into_cache_state(&self.cache_state_selector)
}
#[cfg(test)]
pub async fn reload_no_preprocessing(&self) -> Result<()> {
self.state_lock.write().await?.reload(ReloadPreprocessing::None).await
}
}
#[derive(Clone, Copy)]
pub enum ReloadPreprocessing {
ForgetAll,
None,
}