Skip to main content

tuwunel_service/rooms/state_compressor/
mod.rs

1//! Encodes room state snapshots as compact parent-linked deltas.
2//!
3//! Each state entry combines a short state key with a short event ID in a
4//! fixed-width record. Reconstructed parent chains are cached to make repeated
5//! state resolution inexpensive while bounding persistent diff depth.
6
7use std::{
8	collections::{BTreeSet, HashMap},
9	fmt::{Debug, Write},
10	sync::{Arc, Mutex},
11};
12
13use async_trait::async_trait;
14use futures::{Stream, StreamExt};
15use lru_cache::LruCache;
16use ruma::{EventId, RoomId};
17use tuwunel_core::{
18	Result,
19	arrayvec::ArrayVec,
20	at, checked, err, expected, implement, utils,
21	utils::{bytes, math::usize_from_f64, stream::IterStream},
22};
23use tuwunel_database::{Map, Txn};
24
25use crate::rooms::short::{ShortEventId, ShortId, ShortStateHash, ShortStateKey};
26
27/// Persists, reconstructs, and caches compressed room state snapshots.
28///
29/// New snapshots are stored as bounded delta chains and flattened when their
30/// depth or relative size becomes inefficient. Cached chain entries include
31/// both each frame's delta and its fully materialized state.
32pub struct Service {
33	/// Reconstructed state chains keyed by their requested short state hash.
34	pub stateinfo_cache: Mutex<StateInfoLruCache>,
35	db: Data,
36	services: Arc<crate::services::OnceServices>,
37}
38
39struct Data {
40	shortstatehash_statediff: Arc<Map>,
41}
42
43/// One state as a delta against a parent state.
44///
45/// `added` and `removed` are the compressed entries this state adds to and
46/// removes from its parent chain's accumulation; a `None` parent makes
47/// `added` the full state.
48#[derive(Clone)]
49pub(crate) struct StateDiff {
50	/// Parent snapshot against which this delta is applied, if any.
51	pub(crate) parent: Option<ShortStateHash>,
52
53	/// Compressed entries added to the parent snapshot.
54	pub(crate) added: Arc<CompressedState>,
55
56	/// Compressed entries removed from the parent snapshot.
57	pub(crate) removed: Arc<CompressedState>,
58}
59
60/// Describes one materialized frame in a compressed state chain.
61///
62/// Frames are ordered from the root snapshot toward the requested snapshot.
63/// Each frame retains its local delta alongside the resulting full state.
64#[derive(Clone, Default)]
65pub struct ShortStateInfo {
66	/// Short hash identifying this state frame.
67	pub shortstatehash: ShortStateHash,
68
69	/// Fully materialized state after applying this frame.
70	pub full_state: Arc<CompressedState>,
71
72	/// Entries added by this frame relative to its parent.
73	pub added: Arc<CompressedState>,
74
75	/// Entries removed by this frame relative to its parent.
76	pub removed: Arc<CompressedState>,
77}
78
79/// Reports a saved snapshot and its change from the room's previous state.
80///
81/// An unchanged snapshot reuses its short hash and returns empty added and
82/// removed sets.
83#[derive(Clone, Default)]
84pub struct HashSetCompressStateEvent {
85	/// Short hash identifying the saved snapshot.
86	pub shortstatehash: ShortStateHash,
87
88	/// Entries present only in the saved snapshot.
89	pub added: Arc<CompressedState>,
90
91	/// Entries present only in the previous snapshot.
92	pub removed: Arc<CompressedState>,
93}
94
95type StateInfoLruCache = LruCache<ShortStateHash, ShortStateInfoVec>;
96type ShortStateInfoVec = Vec<ShortStateInfo>;
97type ParentStatesVec = Vec<ShortStateInfo>;
98
99/// Ordered set of compressed state-key and event-ID pairs.
100///
101/// Ordering makes hashing, differences, and persistent serialization
102/// deterministic for the same logical state.
103pub type CompressedState = BTreeSet<CompressedStateEvent>;
104
105/// Fixed-width encoding of one short state key and short event ID.
106///
107/// The first eight big-endian bytes hold the state key and the remaining eight
108/// hold the event ID.
109pub type CompressedStateEvent = [u8; 2 * size_of::<ShortId>()];
110
111#[async_trait]
112impl crate::Service for Service {
113	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
114		let config = &args.server.config;
115		let cache_capacity =
116			f64::from(config.stateinfo_cache_capacity) * config.cache_capacity_modifier;
117		Ok(Arc::new(Self {
118			stateinfo_cache: LruCache::new(usize_from_f64(cache_capacity)?).into(),
119			db: Data {
120				shortstatehash_statediff: args.db["shortstatehash_statediff"].clone(),
121			},
122			services: args.services.clone(),
123		}))
124	}
125
126	async fn memory_usage(&self, out: &mut (dyn Write + Send)) -> Result {
127		let (cache_len, ents) = {
128			let cache = self.stateinfo_cache.lock().expect("locked");
129			let ents = cache
130				.iter()
131				.map(at!(1))
132				.flat_map(|vec| vec.iter())
133				.fold(HashMap::new(), |mut ents, ssi| {
134					for cs in &[&ssi.added, &ssi.removed, &ssi.full_state] {
135						ents.insert(Arc::as_ptr(cs), compressed_state_size(cs));
136					}
137
138					ents
139				});
140
141			(cache.len(), ents)
142		};
143
144		let ents_len = ents.len();
145		let bytes = ents
146			.values()
147			.copied()
148			.fold(0_usize, usize::saturating_add);
149
150		let bytes = bytes::pretty(bytes);
151		writeln!(out, "- stateinfo_cache: {cache_len} entries, {ents_len} states ({bytes})")?;
152
153		Ok(())
154	}
155
156	async fn clear_cache(&self) {
157		self.stateinfo_cache
158			.lock()
159			.expect("locked")
160			.clear();
161	}
162
163	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
164}
165
166/// Loads and materializes the parent chain for a short state hash.
167///
168/// The returned frames are ordered root-first and include each frame's full
169/// state plus its added and removed entries. A previously reconstructed chain
170/// is returned from the LRU cache.
171#[implement(Service)]
172#[tracing::instrument(name = "load", level = "debug", skip(self))]
173pub async fn load_shortstatehash_info(
174	&self,
175	shortstatehash: ShortStateHash,
176) -> Result<ShortStateInfoVec> {
177	if let Some(r) = self
178		.stateinfo_cache
179		.lock()?
180		.get_mut(&shortstatehash)
181	{
182		return Ok(r.clone());
183	}
184
185	let stack = self
186		.new_shortstatehash_info(shortstatehash)
187		.await?;
188
189	self.cache_shortstatehash_info(shortstatehash, stack.clone())
190		.await?;
191
192	Ok(stack)
193}
194
195/// Caches a reconstructed state chain under its requested short hash.
196///
197/// Lock poisoning is reported without modifying the cache.
198#[implement(Service)]
199#[tracing::instrument(
200		name = "cache",
201		level = "debug",
202		skip_all,
203		fields(
204			?shortstatehash,
205			stack = stack.len(),
206		),
207	)]
208async fn cache_shortstatehash_info(
209	&self,
210	shortstatehash: ShortStateHash,
211	stack: ShortStateInfoVec,
212) -> Result {
213	self.stateinfo_cache
214		.lock()?
215		.insert(shortstatehash, stack);
216
217	Ok(())
218}
219
220#[implement(Service)]
221async fn new_shortstatehash_info(
222	&self,
223	shortstatehash: ShortStateHash,
224) -> Result<ShortStateInfoVec> {
225	let StateDiff { parent, added, removed } = self.get_statediff(shortstatehash).await?;
226
227	let Some(parent) = parent else {
228		return Ok(vec![ShortStateInfo {
229			shortstatehash,
230			full_state: added.clone(),
231			added,
232			removed,
233		}]);
234	};
235
236	let mut stack = Box::pin(self.load_shortstatehash_info(parent)).await?; // recursion cycle
237
238	let top = stack.last().expect("at least one frame");
239
240	let mut full_state = (*top.full_state).clone();
241	full_state.extend(added.iter().copied());
242
243	let removed = (*removed).clone();
244	for r in &removed {
245		full_state.remove(r);
246	}
247
248	stack.push(ShortStateInfo {
249		shortstatehash,
250		added,
251		removed: Arc::new(removed),
252		full_state: Arc::new(full_state),
253	});
254
255	Ok(stack)
256}
257
258/// Compresses a stream of state-key and event-ID pairs.
259///
260/// Missing short event IDs are allocated as the returned stream is polled, and
261/// each result packs both short IDs into the fixed-width representation.
262#[implement(Service)]
263pub fn compress_state_events<'a, I>(
264	&'a self,
265	state: I,
266) -> impl Stream<Item = CompressedStateEvent> + Send + 'a
267where
268	I: Iterator<Item = (&'a ShortStateKey, &'a EventId)> + Clone + Debug + Send + 'a,
269{
270	let event_ids = state.clone().map(at!(1));
271
272	let short_event_ids = self
273		.services
274		.short
275		.multi_get_or_create_shorteventid(event_ids);
276
277	state
278		.stream()
279		.map(at!(0))
280		.zip(short_event_ids)
281		.map(|(shortstatekey, shorteventid)| compress_state_event(*shortstatekey, shorteventid))
282}
283
284/// Compresses one state key and event ID into its fixed-width representation.
285///
286/// A short event ID is allocated first when the event has not been seen.
287#[implement(Service)]
288pub async fn compress_state_event(
289	&self,
290	shortstatekey: ShortStateKey,
291	event_id: &EventId,
292) -> CompressedStateEvent {
293	let shorteventid = self
294		.services
295		.short
296		.get_or_create_shorteventid(event_id)
297		.await;
298
299	compress_state_event(shortstatekey, shorteventid)
300}
301
302/// Stages a compressed state delta under a new short state hash.
303///
304/// The caller-owned transaction receives the row but is not executed here.
305/// Chains deeper than three parent frames, or deltas too large relative to
306/// their parent, are recursively flattened into an earlier layer.
307#[implement(Service)]
308pub fn save_state_from_diff(
309	&self,
310	txn: &mut Txn,
311	shortstatehash: ShortStateHash,
312	statediffnew: Arc<CompressedState>,
313	statediffremoved: Arc<CompressedState>,
314	diff_to_sibling: usize,
315	mut parent_states: ParentStatesVec,
316) -> Result {
317	let statediffnew_len = statediffnew.len();
318	let statediffremoved_len = statediffremoved.len();
319	let diffsum = checked!(statediffnew_len + statediffremoved_len)?;
320
321	if parent_states.len() > 3 {
322		// Number of layers
323		// To many layers, we have to go deeper
324		let parent = parent_states
325			.pop()
326			.expect("parent must have a state");
327
328		let mut parent_new = (*parent.added).clone();
329		let mut parent_removed = (*parent.removed).clone();
330
331		for removed in statediffremoved.iter() {
332			if !parent_new.remove(removed) {
333				// It was not added in the parent and we removed it
334				parent_removed.insert(*removed);
335			}
336			// Else it was added in the parent and we removed it again. We
337			// can forget this change
338		}
339
340		for new in statediffnew.iter() {
341			if !parent_removed.remove(new) {
342				// It was not touched in the parent and we added it
343				parent_new.insert(*new);
344			}
345			// Else it was removed in the parent and we added it again. We
346			// can forget this change
347		}
348
349		self.save_state_from_diff(
350			txn,
351			shortstatehash,
352			Arc::new(parent_new),
353			Arc::new(parent_removed),
354			diffsum,
355			parent_states,
356		)?;
357
358		return Ok(());
359	}
360
361	if parent_states.is_empty() {
362		// There is no parent layer, create a new state
363		self.save_statediff(txn, shortstatehash, &StateDiff {
364			parent: None,
365			added: statediffnew,
366			removed: statediffremoved,
367		});
368
369		return Ok(());
370	}
371
372	// Else we have two options.
373	// 1. We add the current diff on top of the parent layer.
374	// 2. We replace a layer above
375
376	let parent = parent_states
377		.pop()
378		.expect("parent must have a state");
379
380	let parent_added_len = parent.added.len();
381	let parent_removed_len = parent.removed.len();
382	let parent_diff = checked!(parent_added_len + parent_removed_len)?;
383
384	if checked!(diffsum * diffsum)? >= checked!(2 * diff_to_sibling * parent_diff)? {
385		// Diff too big, we replace above layer(s)
386		let mut parent_new = (*parent.added).clone();
387		let mut parent_removed = (*parent.removed).clone();
388
389		for removed in statediffremoved.iter() {
390			if !parent_new.remove(removed) {
391				// It was not added in the parent and we removed it
392				parent_removed.insert(*removed);
393			}
394			// Else it was added in the parent and we removed it again. We
395			// can forget this change
396		}
397
398		for new in statediffnew.iter() {
399			if !parent_removed.remove(new) {
400				// It was not touched in the parent and we added it
401				parent_new.insert(*new);
402			}
403			// Else it was removed in the parent and we added it again. We
404			// can forget this change
405		}
406
407		self.save_state_from_diff(
408			txn,
409			shortstatehash,
410			Arc::new(parent_new),
411			Arc::new(parent_removed),
412			diffsum,
413			parent_states,
414		)?;
415	} else {
416		// Diff small enough, we add diff as layer on top of parent
417		self.save_statediff(txn, shortstatehash, &StateDiff {
418			parent: Some(parent.shortstatehash),
419			added: statediffnew,
420			removed: statediffremoved,
421		});
422	}
423
424	Ok(())
425}
426
427/// Saves a complete compressed snapshot and reports its previous-state delta.
428///
429/// An existing content hash is reused; otherwise the short hash and delta row
430/// are created together. Failure to reconstruct the previous chain is treated
431/// as an absent parent, while an exactly unchanged snapshot returns empty sets.
432#[implement(Service)]
433#[tracing::instrument(skip(self, new_state_ids_compressed), level = "debug")]
434pub async fn save_state(
435	&self,
436	room_id: &RoomId,
437	new_state_ids_compressed: Arc<CompressedState>,
438) -> Result<HashSetCompressStateEvent> {
439	let previous_shortstatehash = self
440		.services
441		.state
442		.get_room_shortstatehash(room_id)
443		.await
444		.ok();
445
446	let state_hash = utils::calculate_hash(
447		new_state_ids_compressed
448			.iter()
449			.map(|bytes| &bytes[..]),
450	);
451
452	let existing_shortstatehash = self
453		.services
454		.short
455		.get_shortstatehash(&state_hash)
456		.await
457		.ok();
458
459	if let Some(new_shortstatehash) = existing_shortstatehash
460		.filter(|&new_shortstatehash| previous_shortstatehash.eq(&Some(new_shortstatehash)))
461	{
462		return Ok(HashSetCompressStateEvent {
463			shortstatehash: new_shortstatehash,
464			..Default::default()
465		});
466	}
467
468	let states_parents = if let Some(p) = previous_shortstatehash {
469		self.load_shortstatehash_info(p)
470			.await
471			.unwrap_or_default()
472	} else {
473		ShortStateInfoVec::new()
474	};
475
476	let (statediffnew, statediffremoved) = if let Some(parent_stateinfo) = states_parents.last() {
477		let statediffnew: CompressedState = new_state_ids_compressed
478			.difference(&parent_stateinfo.full_state)
479			.copied()
480			.collect();
481
482		let statediffremoved: CompressedState = parent_stateinfo
483			.full_state
484			.difference(&new_state_ids_compressed)
485			.copied()
486			.collect();
487
488		(Arc::new(statediffnew), Arc::new(statediffremoved))
489	} else {
490		(new_state_ids_compressed, Arc::new(CompressedState::new()))
491	};
492
493	let new_shortstatehash = if let Some(new_shortstatehash) = existing_shortstatehash {
494		new_shortstatehash
495	} else {
496		self.services
497			.short
498			.get_or_create_shortstatehash(&state_hash, |txn, shortstatehash| {
499				self.save_state_from_diff(
500					txn,
501					shortstatehash,
502					statediffnew.clone(),
503					statediffremoved.clone(),
504					2, // every state change is 2 event changes on average
505					states_parents,
506				)
507			})
508			.await?
509			.0
510	};
511
512	Ok(HashSetCompressStateEvent {
513		shortstatehash: new_shortstatehash,
514		added: statediffnew,
515		removed: statediffremoved,
516	})
517}
518
519/// Reads one state's delta row into its typed form.
520///
521/// Rows round-trip through [`save_statediff`], the pair being the only
522/// codec for the statediff encoding.
523///
524/// # Panics
525///
526/// Panics if a stored delta row is shorter than its eight-byte parent prefix.
527#[implement(Service)]
528#[tracing::instrument(skip(self), level = "debug", name = "get")]
529pub(crate) async fn get_statediff(&self, shortstatehash: ShortStateHash) -> Result<StateDiff> {
530	const BUFSIZE: usize = size_of::<ShortStateHash>();
531	const STRIDE: usize = size_of::<ShortStateHash>();
532
533	let value = self
534		.db
535		.shortstatehash_statediff
536		.aqry::<BUFSIZE, _>(&shortstatehash)
537		.await
538		.map_err(|e| {
539			err!(Database("Failed to find StateDiff from short {shortstatehash:?}: {e}"))
540		})?;
541
542	let parent = utils::u64_from_bytes(&value[0..size_of::<u64>()])
543		.ok()
544		.take_if(|parent| *parent != 0);
545
546	debug_assert!(value.len().is_multiple_of(STRIDE), "value not aligned to stride");
547	let _num_values = value.len() / STRIDE;
548
549	let mut add_mode = true;
550	let mut added = CompressedState::new();
551	let mut removed = CompressedState::new();
552
553	let mut i = STRIDE;
554	while let Some(v) = value.get(i..expected!(i + 2 * STRIDE)) {
555		if add_mode && v.starts_with(&0_u64.to_be_bytes()) {
556			add_mode = false;
557			i = expected!(i + STRIDE);
558			continue;
559		}
560		if add_mode {
561			added.insert(v.try_into()?);
562		} else {
563			removed.insert(v.try_into()?);
564		}
565		i = expected!(i + 2 * STRIDE);
566	}
567
568	Ok(StateDiff {
569		parent,
570		added: Arc::new(added),
571		removed: Arc::new(removed),
572	})
573}
574
575/// Serializes one state's delta into the caller's transaction.
576///
577/// Added and removed entries are emitted in sorted set order. The removed run
578/// follows an all-zero sentinel only when it is nonempty, and this method does
579/// not execute the transaction.
580#[implement(Service)]
581pub(crate) fn save_statediff(
582	&self,
583	txn: &mut Txn,
584	shortstatehash: ShortStateHash,
585	diff: &StateDiff,
586) {
587	let event_count = diff
588		.added
589		.len()
590		.saturating_add(diff.removed.len());
591
592	let event_bytes = event_count.saturating_mul(size_of::<CompressedStateEvent>());
593	let separator_bytes =
594		usize::from(!diff.removed.is_empty()).saturating_mul(size_of::<ShortStateHash>());
595
596	let capacity = size_of::<ShortStateHash>()
597		.saturating_add(event_bytes)
598		.saturating_add(separator_bytes);
599
600	let parent = diff.parent.unwrap_or(0_u64);
601	let mut value = Vec::<u8>::with_capacity(capacity);
602	value.extend_from_slice(&parent.to_be_bytes());
603
604	for new in diff.added.iter() {
605		value.extend_from_slice(&new[..]);
606	}
607
608	if !diff.removed.is_empty() {
609		value.extend_from_slice(&0_u64.to_be_bytes());
610		for removed in diff.removed.iter() {
611			value.extend_from_slice(&removed[..]);
612		}
613	}
614
615	txn.insert_raw(&self.db.shortstatehash_statediff, shortstatehash.to_be_bytes(), value);
616}
617
618/// Packs a short state key and short event ID into one compressed record.
619///
620/// Both IDs use big-endian encoding so byte ordering follows numeric ordering.
621#[inline]
622#[must_use]
623pub(crate) fn compress_state_event(
624	shortstatekey: ShortStateKey,
625	shorteventid: ShortEventId,
626) -> CompressedStateEvent {
627	const SIZE: usize = size_of::<CompressedStateEvent>();
628
629	let mut v = ArrayVec::<u8, SIZE>::new();
630	v.extend(shortstatekey.to_be_bytes());
631	v.extend(shorteventid.to_be_bytes());
632	v.as_ref()
633		.try_into()
634		.expect("failed to create CompressedStateEvent")
635}
636
637/// Unpacks a compressed state record into its two short IDs.
638///
639/// This is the inverse of [`compress_state_event`].
640#[inline]
641#[must_use]
642pub(crate) fn parse_compressed_state_event(
643	compressed_event: CompressedStateEvent,
644) -> (ShortStateKey, ShortEventId) {
645	use utils::u64_from_u8;
646
647	let shortstatekey = u64_from_u8(&compressed_event[0..size_of::<ShortStateKey>()]);
648	let shorteventid = u64_from_u8(&compressed_event[size_of::<ShortStateKey>()..]);
649
650	(shortstatekey, shorteventid)
651}
652
653#[inline]
654fn compressed_state_size(compressed_state: &CompressedState) -> usize {
655	compressed_state
656		.len()
657		.checked_mul(size_of::<CompressedStateEvent>())
658		.expect("CompressedState size overflow")
659}