mod updates;
use std::{cmp::Ordering, collections::BTreeSet, fmt, sync::Arc};
use eyeball_im::VectorDiff;
use futures_util::{StreamExt as _, stream};
use matrix_sdk_base::{
apply_redaction,
event_cache::{Event, Gap},
linked_chunk::{LinkedChunkId, OwnedLinkedChunkId, Position, Update},
serde_helpers::extract_redaction_target,
sync::Timeline,
task_monitor::BackgroundTaskHandle,
};
use matrix_sdk_common::executor::spawn;
use ruma::{
EventId, MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedRoomId, OwnedUserId,
events::{relation::RelationType, room::redaction::SyncRoomRedactionEvent},
room_version_rules::RoomVersionRules,
};
use tokio::sync::broadcast::{Receiver, Sender};
use tracing::{debug, instrument, trace, warn};
pub(super) use self::updates::PinnedEventsCacheUpdateSender;
#[cfg(feature = "e2e-encryption")]
use super::super::redecryptor::MaybeResolvedEvent;
use super::{
super::{
EventCacheError, EventsOrigin, Result,
deduplicator::{DeduplicationOutcome, filter_duplicate_events},
persistence::{find_event, send_updates_to_store},
states::{
CacheStateLock, ReloadPreprocessing, StateLock, StateLockWriteGuard,
selectors::PinnedEventsStateSelector,
},
},
EventLocation, TimelineVectorDiffs,
event_linked_chunk::{EventLinkedChunk, sort_positions_descending},
room::RoomEventCacheLinkedChunkUpdate,
};
use crate::{Room, client::WeakClient, config::RequestConfig, room::WeakRoom};
pub struct PinnedEventsCacheState {
room_id: OwnedRoomId,
own_user_id: OwnedUserId,
room_version_rules: RoomVersionRules,
chunk: EventLinkedChunk,
pub update_sender: PinnedEventsCacheUpdateSender,
linked_chunk_update_sender: Sender<RoomEventCacheLinkedChunkUpdate>,
}
#[cfg(not(tarpaulin_include))]
impl fmt::Debug for PinnedEventsCacheState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PinnedEventsCacheState")
.field("room_id", &self.room_id)
.field("chunk", &self.chunk)
.finish_non_exhaustive()
}
}
impl<'a> StateLockWriteGuard<'a, PinnedEventsCacheState> {
#[must_use = "Propagate `VectorDiff` updates via `TimelineVectorDiffs`"]
pub async fn reload(
&mut self,
preprocessing: ReloadPreprocessing,
) -> Result<Vec<VectorDiff<Event>>> {
match preprocessing {
ReloadPreprocessing::ForgetAll => {
self.state.chunk.reset();
self.propagate_changes().await?;
}
ReloadPreprocessing::None => {}
}
self.reload_from_storage().await?;
Ok(self.state.chunk.updates_as_vector_diffs())
}
async fn handle_sync(&mut self, timeline: Timeline) -> Result<()> {
let DeduplicationOutcome {
all_events: events,
in_memory_duplicated_event_ids,
in_store_duplicated_event_ids,
non_empty_all_duplicates: all_duplicates,
} = filter_duplicate_events(
&self.state.own_user_id,
&self.store,
LinkedChunkId::PinnedEvents(&self.state.room_id),
&self.state.chunk,
timeline.events,
)
.await?;
if all_duplicates {
return Ok(());
}
self.remove_events(in_memory_duplicated_event_ids, in_store_duplicated_event_ids).await?;
self.state.chunk.push_live_events(None, &events);
self.propagate_changes().await?;
self.notify_subscribers(EventsOrigin::Sync);
for event in &events {
self.maybe_apply_new_redaction(event).await?;
}
Ok(())
}
#[instrument(skip_all)]
pub async fn remove_events(
&mut self,
in_memory_events: Vec<(OwnedEventId, Position)>,
in_store_events: Vec<(OwnedEventId, Position)>,
) -> Result<()> {
if !in_store_events.is_empty() {
let mut positions = in_store_events
.into_iter()
.map(|(_event_id, position)| position)
.collect::<Vec<_>>();
sort_positions_descending(&mut positions);
let updates =
positions.into_iter().map(|pos| Update::RemoveItem { at: pos }).collect::<Vec<_>>();
self.apply_store_only_updates(updates).await?;
}
if in_memory_events.is_empty() {
return Ok(());
}
self.state
.chunk
.remove_events_by_position(
in_memory_events.into_iter().map(|(_event_id, position)| position).collect(),
)
.expect("failed to remove an event");
self.propagate_changes().await
}
async fn apply_store_only_updates(&mut self, updates: Vec<Update<Event, Gap>>) -> Result<()> {
self.send_updates_to_store(updates).await
}
#[instrument(skip_all)]
async fn maybe_apply_new_redaction(&mut self, event: &Event) -> Result<()> {
let Some(event_id) =
extract_redaction_target(event.raw(), &self.room_version_rules.redaction)
else {
return Ok(());
};
let Some((location, mut target_event)) = self.find_event(&event_id).await? else {
trace!("redacted event is missing from the linked chunk");
return Ok(());
};
let target_event_raw = target_event.raw();
if let Ok(deserialized) = target_event_raw.deserialize()
&& deserialized.is_redacted()
{
return Ok(());
}
if let Some(redacted_event) = apply_redaction(
target_event_raw,
event.raw().cast_ref_unchecked::<SyncRoomRedactionEvent>(),
&self.room_version_rules.redaction,
) {
target_event.replace_raw(redacted_event.cast_unchecked());
self.replace_event_at(location, target_event.clone()).await?;
}
Ok(())
}
pub(super) async fn find_event(
&self,
event_id: &EventId,
) -> Result<Option<(EventLocation, Event)>> {
find_event(event_id, &self.room_id, &self.chunk, &self.store).await
}
pub async fn replace_event_at(
&mut self,
location: EventLocation,
new_event: Event,
) -> Result<()> {
match location {
EventLocation::Memory(position) => {
self.state
.chunk
.replace_event_at(position, new_event)
.expect("should have been a valid position of an item");
self.propagate_changes().await?;
}
EventLocation::Store => {
self.save_events([new_event]).await?;
}
}
Ok(())
}
pub async fn save_events(&mut self, events: impl IntoIterator<Item = Event>) -> Result<()> {
let store = self.store.clone();
let room_id = self.state.room_id.clone();
let events = events.into_iter().collect::<Vec<_>>();
spawn(async move {
for event in events {
store.save_event(&room_id, event).await?;
}
Result::Ok(())
})
.await
.expect("joining failed")?;
Ok(())
}
async fn reload_from_storage(&mut self) -> Result<()> {
let room_id = self.state.room_id.clone();
let linked_chunk_id = LinkedChunkId::PinnedEvents(&room_id);
let (last_chunk, chunk_id_gen) = self.store.load_last_chunk(linked_chunk_id).await?;
let Some(last_chunk) = last_chunk else {
if self.state.chunk.events().next().is_some() {
self.state.chunk.reset();
self.notify_subscribers(EventsOrigin::Sync);
}
return Ok(());
};
{
let mut current_chunk_identifier = last_chunk.identifier;
self.state.chunk.shrink_to_last_reloaded_chunk(
Some(last_chunk),
chunk_id_gen,
None,
)?;
while let Some(previous_chunk) =
self.store.load_previous_chunk(linked_chunk_id, current_chunk_identifier).await?
{
current_chunk_identifier = previous_chunk.identifier;
self.state.chunk.insert_new_chunk_as_first(previous_chunk)?;
}
}
self.state.chunk.store_updates().take();
self.notify_subscribers(EventsOrigin::Cache);
Ok(())
}
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.state.current_event_ids();
if new_events
.iter()
.filter_map(|e| e.event_id())
.map(ToOwned::to_owned)
.collect::<BTreeSet<_>>()
== previous_pinned_event_ids.into_iter().collect()
{
return Ok(());
}
if self.state.chunk.events().next().is_some() {
self.state.chunk.reset();
}
self.state.chunk.push_live_events(None, &new_events);
self.propagate_changes().await?;
self.notify_subscribers(EventsOrigin::Sync);
Ok(())
}
pub async fn propagate_changes(&mut self) -> Result<()> {
let updates = self.state.chunk.store_updates().take();
self.send_updates_to_store(updates).await
}
async fn send_updates_to_store(&mut self, updates: Vec<Update<Event, Gap>>) -> Result<()> {
let linked_chunk_id = OwnedLinkedChunkId::PinnedEvents(self.room_id.clone());
send_updates_to_store(
&self.store,
linked_chunk_id,
&self.state.linked_chunk_update_sender,
updates,
)
.await
}
fn notify_subscribers(&mut self, origin: EventsOrigin) {
let diffs = self.state.chunk.updates_as_vector_diffs();
if !diffs.is_empty() {
self.update_sender.send(TimelineVectorDiffs { diffs, origin });
}
}
}
impl PinnedEventsCacheState {
pub(super) fn current_event_ids(&self) -> Vec<OwnedEventId> {
self.chunk
.events()
.filter_map(|(_position, event)| event.event_id().map(ToOwned::to_owned))
.collect()
}
}
#[derive(Clone)]
pub struct PinnedEventsCache {
inner: Arc<PinnedEventsCacheInner>,
_task: Arc<BackgroundTaskHandle>,
}
struct PinnedEventsCacheInner {
room_id: OwnedRoomId,
state: CacheStateLock<PinnedEventsStateSelector>,
}
impl PinnedEventsCache {
pub(in super::super) async fn new(
weak_room: &WeakRoom,
own_user_id: OwnedUserId,
room_version_rules: RoomVersionRules,
linked_chunk_update_sender: Sender<RoomEventCacheLinkedChunkUpdate>,
state: &StateLock,
) -> Result<Self> {
let room = weak_room.get().ok_or(EventCacheError::ClientDropped)?;
let room_id = room.room_id().to_owned();
let cache_state = state
.try_insert_once_with(
PinnedEventsStateSelector::new(room_id.clone()),
|_store_guard| async {
Ok(PinnedEventsCacheState {
room_id: room_id.clone(),
own_user_id,
room_version_rules,
chunk: EventLinkedChunk::new(),
update_sender: PinnedEventsCacheUpdateSender::new(),
linked_chunk_update_sender,
})
},
)
.await?;
let inner = Arc::new(PinnedEventsCacheInner { room_id, state: cache_state });
let task = room
.client()
.task_monitor()
.spawn_infinite_task(
"pinned_event_listener_task",
Self::pinned_event_listener_task(room, inner.clone()),
)
.abort_on_drop();
Ok(Self { inner, _task: Arc::new(task) })
}
pub(super) fn state(&self) -> &CacheStateLock<PinnedEventsStateSelector> {
&self.inner.state
}
pub async fn subscribe(&self) -> Result<(Vec<Event>, Receiver<TimelineVectorDiffs>)> {
let guard = self.inner.state.read().await?;
let events = guard.state.chunk.events().map(|(_position, item)| item.clone()).collect();
let recv = guard.state.update_sender.new_pinned_events_receiver();
Ok((events, recv))
}
#[cfg(feature = "e2e-encryption")]
pub(in super::super) async fn replace_in_memory_utds(
&self,
resolved_events: &[MaybeResolvedEvent],
) -> Result<()> {
let mut state = self.inner.state.write().await?;
let _ = state.state.chunk.store_updates().take();
if state.state.chunk.replace_utds(resolved_events) {
state.propagate_changes().await?;
state.notify_subscribers(EventsOrigin::Cache);
}
Ok(())
}
#[instrument(skip_all, fields(room_id = %self.inner.room_id))]
pub(super) async fn handle_joined_room_update(&self, timeline: Timeline) -> Result<()> {
self.handle_timeline(timeline).await
}
#[instrument(skip_all, fields(room_id = %self.inner.room_id))]
pub(super) async fn handle_left_room_update(&self, timeline: Timeline) -> Result<()> {
self.handle_timeline(timeline).await
}
async fn handle_timeline(&self, timeline: Timeline) -> Result<()> {
if timeline.events.is_empty() {
return Ok(());
}
trace!("adding new {} events", timeline.events.len());
self.inner.state.write().await?.handle_sync(timeline).await
}
#[instrument(fields(%room_id = room.room_id()), skip(room, inner))]
async fn pinned_event_listener_task(room: Room, inner: Arc<PinnedEventsCacheInner>) {
debug!("pinned events listener task started");
let reload_from_network = async |room: Room| {
let events = match Self::reload_pinned_events(room).await {
Ok(Some(events)) => events,
Ok(None) => Vec::new(),
Err(err) => {
warn!("error when loading pinned events: {err}");
return;
}
};
match inner.state.write().await {
Ok(mut guard) => {
guard.replace_all_events(events).await.unwrap_or_else(|err| {
warn!("error when replacing pinned events: {err}");
});
}
Err(err) => {
warn!("error when acquiring write lock to replace pinned events: {err}");
}
}
};
match inner.state.write().await {
Ok(mut guard) => {
guard.reload_from_storage().await.unwrap_or_else(|err| {
warn!("error when reloading pinned events from storage, at start: {err}");
});
let actual_pinned_events = room.pinned_event_ids().unwrap_or_default();
let reloaded_set =
guard.state.current_event_ids().into_iter().collect::<BTreeSet<_>>();
if actual_pinned_events.len() != reloaded_set.len()
|| actual_pinned_events.iter().any(|event_id| !reloaded_set.contains(event_id))
{
drop(guard);
reload_from_network(room.clone()).await;
}
}
Err(err) => {
warn!("error when acquiring write lock to initialize pinned events: {err}");
}
}
let weak_room =
WeakRoom::new(WeakClient::from_client(&room.client()), room.room_id().to_owned());
let mut stream = room.pinned_event_ids_stream();
drop(room);
while let Some(new_list) = stream.next().await {
trace!("handling update");
let guard = match inner.state.read().await {
Ok(guard) => guard,
Err(err) => {
warn!("error when acquiring read lock to handle pinned events update: {err}");
break;
}
};
let current_set = guard.state.current_event_ids().into_iter().collect::<BTreeSet<_>>();
if !new_list.is_empty()
&& new_list.len() == current_set.len()
&& new_list.iter().all(|event_id| current_set.contains(event_id))
{
continue;
}
let Some(room) = weak_room.get() else {
debug!("room has been dropped, ending pinned events listener task");
break;
};
drop(guard);
reload_from_network(room).await;
}
debug!("pinned events listener task ended");
}
async fn reload_pinned_events(room: Room) -> Result<Option<Vec<Event>>> {
let (max_events_to_load, max_concurrent_requests) = {
let client = room.client();
let config = client.event_cache().config();
(config.max_pinned_events_to_load, config.max_pinned_events_concurrent_requests)
};
let pinned_event_ids: Vec<OwnedEventId> = room
.pinned_event_ids()
.unwrap_or_default()
.into_iter()
.rev()
.take(max_events_to_load)
.rev()
.collect();
if pinned_event_ids.is_empty() {
return Ok(Some(Vec::new()));
}
let mut num_successful_loads = 0;
let mut loaded_events: Vec<Event> =
stream::iter(pinned_event_ids.clone().into_iter().map(|event_id| {
let room = room.clone();
let filter = vec![RelationType::Annotation, RelationType::Replacement];
let request_config = RequestConfig::default().retry_limit(3);
async move {
let (target, mut relations) = room
.load_or_fetch_event_with_relations(
&event_id,
Some(filter),
Some(request_config),
)
.await?;
relations.insert(0, target);
Ok::<_, crate::Error>(relations)
}
}))
.buffer_unordered(max_concurrent_requests)
.inspect(|result| {
if result.is_ok() {
num_successful_loads += 1;
}
})
.flat_map(stream::iter)
.flat_map(stream::iter)
.collect()
.await;
if num_successful_loads != pinned_event_ids.len() {
warn!(
"only successfully loaded {} out of {} pinned events",
num_successful_loads,
pinned_event_ids.len()
);
}
if loaded_events.is_empty() {
return Err(EventCacheError::UnableToLoadPinnedEvents);
}
loaded_events.sort_by(compare_pinned_items);
Ok(Some(loaded_events))
}
}
impl fmt::Debug for PinnedEventsCache {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PinnedEventsCache").finish_non_exhaustive()
}
}
fn compare_pinned_items(a: &Event, b: &Event) -> Ordering {
let a_time: Option<MilliSecondsSinceUnixEpoch> = a.timestamp_raw();
let b_time: Option<MilliSecondsSinceUnixEpoch> = b.timestamp_raw();
compare_by_optional_timestamp(a_time, b_time)
}
fn compare_by_optional_timestamp(
a: Option<MilliSecondsSinceUnixEpoch>,
b: Option<MilliSecondsSinceUnixEpoch>,
) -> Ordering {
match (a, b) {
(None, None) => Ordering::Equal,
(None, Some(_)) => Ordering::Greater,
(Some(_), None) => Ordering::Less,
(Some(a), Some(b)) => a.cmp(&b),
}
}
#[cfg(not(target_family = "wasm"))]
#[cfg(test)]
mod tests {
use proptest::prelude::*;
use ruma::UInt;
use super::*;
fn any_timestamp() -> impl Strategy<Value = Option<MilliSecondsSinceUnixEpoch>> {
prop::option::of(
any::<u32>().prop_map(|value| MilliSecondsSinceUnixEpoch(UInt::from(value))),
)
}
#[test]
fn sort_pinned_events_never_panics_only_nones() {
let mut vec = vec![None; 100_000];
vec.sort_by(|a, b| compare_by_optional_timestamp(*a, *b))
}
proptest! {
#[test]
fn sort_pinned_events_never_panics(mut v in prop::collection::vec(any_timestamp(), 0..1000)) {
v.sort_by(
|a, b| compare_by_optional_timestamp(*a, *b))
}
#[test]
fn compare_pinned_events_reflexive(a in any_timestamp()) {
prop_assert_eq!(compare_by_optional_timestamp(a, a), Ordering::Equal);
}
#[test]
fn compare_pinned_events_antisymmetric(a in any_timestamp(), b in any_timestamp()) {
let ab = compare_by_optional_timestamp(a, b);
let ba = compare_by_optional_timestamp(b, a);
prop_assert_eq!(ab, ba.reverse());
}
#[test]
fn compare_pinned_events_transitive(
a in any_timestamp(),
b in any_timestamp(),
c in any_timestamp()
) {
let ab = compare_by_optional_timestamp(a, b);
let bc = compare_by_optional_timestamp(b, c);
let ac = compare_by_optional_timestamp(a, c);
if ab == Ordering::Less && bc == Ordering::Less {
prop_assert_eq!(ac, Ordering::Less);
}
if ab == Ordering::Equal && bc == Ordering::Equal {
prop_assert_eq!(ac, Ordering::Equal);
}
if ab == Ordering::Greater && bc == Ordering::Greater {
prop_assert_eq!(ac, Ordering::Greater);
}
}
}
}