Skip to main content

tuwunel_service/rooms/state_compressor/
mod.rs

1use std::{
2	collections::{BTreeSet, HashMap},
3	fmt::{Debug, Write},
4	sync::{Arc, Mutex},
5};
6
7use async_trait::async_trait;
8use futures::{Stream, StreamExt};
9use lru_cache::LruCache;
10use ruma::{EventId, RoomId};
11use tuwunel_core::{
12	Result,
13	arrayvec::ArrayVec,
14	at, checked, err, expected, implement, utils,
15	utils::{bytes, math::usize_from_f64, stream::IterStream},
16};
17use tuwunel_database::{Map, Txn};
18
19use crate::rooms::short::{ShortEventId, ShortId, ShortStateHash, ShortStateKey};
20
21pub struct Service {
22	pub stateinfo_cache: Mutex<StateInfoLruCache>,
23	db: Data,
24	services: Arc<crate::services::OnceServices>,
25}
26
27struct Data {
28	shortstatehash_statediff: Arc<Map>,
29}
30
31/// One state as a delta against a parent state.
32///
33/// `added` and `removed` are the compressed entries this state adds to and
34/// removes from its parent chain's accumulation; a `None` parent makes
35/// `added` the full state.
36#[derive(Clone)]
37pub(crate) struct StateDiff {
38	pub(crate) parent: Option<ShortStateHash>,
39	pub(crate) added: Arc<CompressedState>,
40	pub(crate) removed: Arc<CompressedState>,
41}
42
43#[derive(Clone, Default)]
44pub struct ShortStateInfo {
45	pub shortstatehash: ShortStateHash,
46	pub full_state: Arc<CompressedState>,
47	pub added: Arc<CompressedState>,
48	pub removed: Arc<CompressedState>,
49}
50
51#[derive(Clone, Default)]
52pub struct HashSetCompressStateEvent {
53	pub shortstatehash: ShortStateHash,
54	pub added: Arc<CompressedState>,
55	pub removed: Arc<CompressedState>,
56}
57
58type StateInfoLruCache = LruCache<ShortStateHash, ShortStateInfoVec>;
59type ShortStateInfoVec = Vec<ShortStateInfo>;
60type ParentStatesVec = Vec<ShortStateInfo>;
61
62pub type CompressedState = BTreeSet<CompressedStateEvent>;
63pub type CompressedStateEvent = [u8; 2 * size_of::<ShortId>()];
64
65#[async_trait]
66impl crate::Service for Service {
67	fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
68		let config = &args.server.config;
69		let cache_capacity =
70			f64::from(config.stateinfo_cache_capacity) * config.cache_capacity_modifier;
71		Ok(Arc::new(Self {
72			stateinfo_cache: LruCache::new(usize_from_f64(cache_capacity)?).into(),
73			db: Data {
74				shortstatehash_statediff: args.db["shortstatehash_statediff"].clone(),
75			},
76			services: args.services.clone(),
77		}))
78	}
79
80	async fn memory_usage(&self, out: &mut (dyn Write + Send)) -> Result {
81		let (cache_len, ents) = {
82			let cache = self.stateinfo_cache.lock().expect("locked");
83			let ents = cache
84				.iter()
85				.map(at!(1))
86				.flat_map(|vec| vec.iter())
87				.fold(HashMap::new(), |mut ents, ssi| {
88					for cs in &[&ssi.added, &ssi.removed, &ssi.full_state] {
89						ents.insert(Arc::as_ptr(cs), compressed_state_size(cs));
90					}
91
92					ents
93				});
94
95			(cache.len(), ents)
96		};
97
98		let ents_len = ents.len();
99		let bytes = ents
100			.values()
101			.copied()
102			.fold(0_usize, usize::saturating_add);
103
104		let bytes = bytes::pretty(bytes);
105		writeln!(out, "- stateinfo_cache: {cache_len} entries, {ents_len} states ({bytes})")?;
106
107		Ok(())
108	}
109
110	async fn clear_cache(&self) {
111		self.stateinfo_cache
112			.lock()
113			.expect("locked")
114			.clear();
115	}
116
117	fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
118}
119
120/// Returns a stack with info on shortstatehash, full state, added diff and
121/// removed diff for the selected shortstatehash and each parent layer.
122#[implement(Service)]
123#[tracing::instrument(name = "load", level = "debug", skip(self))]
124pub async fn load_shortstatehash_info(
125	&self,
126	shortstatehash: ShortStateHash,
127) -> Result<ShortStateInfoVec> {
128	if let Some(r) = self
129		.stateinfo_cache
130		.lock()?
131		.get_mut(&shortstatehash)
132	{
133		return Ok(r.clone());
134	}
135
136	let stack = self
137		.new_shortstatehash_info(shortstatehash)
138		.await?;
139
140	self.cache_shortstatehash_info(shortstatehash, stack.clone())
141		.await?;
142
143	Ok(stack)
144}
145
146/// Returns a stack with info on shortstatehash, full state, added diff and
147/// removed diff for the selected shortstatehash and each parent layer.
148#[implement(Service)]
149#[tracing::instrument(
150		name = "cache",
151		level = "debug",
152		skip_all,
153		fields(
154			?shortstatehash,
155			stack = stack.len(),
156		),
157	)]
158async fn cache_shortstatehash_info(
159	&self,
160	shortstatehash: ShortStateHash,
161	stack: ShortStateInfoVec,
162) -> Result {
163	self.stateinfo_cache
164		.lock()?
165		.insert(shortstatehash, stack);
166
167	Ok(())
168}
169
170#[implement(Service)]
171async fn new_shortstatehash_info(
172	&self,
173	shortstatehash: ShortStateHash,
174) -> Result<ShortStateInfoVec> {
175	let StateDiff { parent, added, removed } = self.get_statediff(shortstatehash).await?;
176
177	let Some(parent) = parent else {
178		return Ok(vec![ShortStateInfo {
179			shortstatehash,
180			full_state: added.clone(),
181			added,
182			removed,
183		}]);
184	};
185
186	let mut stack = Box::pin(self.load_shortstatehash_info(parent)).await?;
187	let top = stack.last().expect("at least one frame");
188
189	let mut full_state = (*top.full_state).clone();
190	full_state.extend(added.iter().copied());
191
192	let removed = (*removed).clone();
193	for r in &removed {
194		full_state.remove(r);
195	}
196
197	stack.push(ShortStateInfo {
198		shortstatehash,
199		added,
200		removed: Arc::new(removed),
201		full_state: Arc::new(full_state),
202	});
203
204	Ok(stack)
205}
206
207#[implement(Service)]
208pub fn compress_state_events<'a, I>(
209	&'a self,
210	state: I,
211) -> impl Stream<Item = CompressedStateEvent> + Send + 'a
212where
213	I: Iterator<Item = (&'a ShortStateKey, &'a EventId)> + Clone + Debug + Send + 'a,
214{
215	let event_ids = state.clone().map(at!(1));
216
217	let short_event_ids = self
218		.services
219		.short
220		.multi_get_or_create_shorteventid(event_ids);
221
222	state
223		.stream()
224		.map(at!(0))
225		.zip(short_event_ids)
226		.map(|(shortstatekey, shorteventid)| compress_state_event(*shortstatekey, shorteventid))
227}
228
229#[implement(Service)]
230pub async fn compress_state_event(
231	&self,
232	shortstatekey: ShortStateKey,
233	event_id: &EventId,
234) -> CompressedStateEvent {
235	let shorteventid = self
236		.services
237		.short
238		.get_or_create_shorteventid(event_id)
239		.await;
240
241	compress_state_event(shortstatekey, shorteventid)
242}
243
244/// Creates a new shortstatehash that often is just a diff to an already
245/// existing shortstatehash and therefore very efficient.
246///
247/// There are multiple layers of diffs. The bottom layer 0 always contains
248/// the full state. Layer 1 contains diffs to states of layer 0, layer 2
249/// diffs to layer 1 and so on. If layer n > 0 grows too big, it will be
250/// combined with layer n-1 to create a new diff on layer n-1 that's
251/// based on layer n-2. If that layer is also too big, it will recursively
252/// fix above layers too.
253///
254/// * `txn` - Caller-owned transaction that receives the StateDiff without
255///   executing it
256/// * `shortstatehash` - Shortstatehash of this state
257/// * `statediffnew` - Added to base. Each vec is shortstatekey+shorteventid
258/// * `statediffremoved` - Removed from base. Each vec is
259///   shortstatekey+shorteventid
260/// * `diff_to_sibling` - Approximately how much the diff grows each time for
261///   this layer
262/// * `parent_states` - A stack with info on shortstatehash, full state, added
263///   diff and removed diff for each parent layer
264#[implement(Service)]
265pub fn save_state_from_diff(
266	&self,
267	txn: &mut Txn,
268	shortstatehash: ShortStateHash,
269	statediffnew: Arc<CompressedState>,
270	statediffremoved: Arc<CompressedState>,
271	diff_to_sibling: usize,
272	mut parent_states: ParentStatesVec,
273) -> Result {
274	let statediffnew_len = statediffnew.len();
275	let statediffremoved_len = statediffremoved.len();
276	let diffsum = checked!(statediffnew_len + statediffremoved_len)?;
277
278	if parent_states.len() > 3 {
279		// Number of layers
280		// To many layers, we have to go deeper
281		let parent = parent_states
282			.pop()
283			.expect("parent must have a state");
284
285		let mut parent_new = (*parent.added).clone();
286		let mut parent_removed = (*parent.removed).clone();
287
288		for removed in statediffremoved.iter() {
289			if !parent_new.remove(removed) {
290				// It was not added in the parent and we removed it
291				parent_removed.insert(*removed);
292			}
293			// Else it was added in the parent and we removed it again. We
294			// can forget this change
295		}
296
297		for new in statediffnew.iter() {
298			if !parent_removed.remove(new) {
299				// It was not touched in the parent and we added it
300				parent_new.insert(*new);
301			}
302			// Else it was removed in the parent and we added it again. We
303			// can forget this change
304		}
305
306		self.save_state_from_diff(
307			txn,
308			shortstatehash,
309			Arc::new(parent_new),
310			Arc::new(parent_removed),
311			diffsum,
312			parent_states,
313		)?;
314
315		return Ok(());
316	}
317
318	if parent_states.is_empty() {
319		// There is no parent layer, create a new state
320		self.save_statediff(txn, shortstatehash, &StateDiff {
321			parent: None,
322			added: statediffnew,
323			removed: statediffremoved,
324		});
325
326		return Ok(());
327	}
328
329	// Else we have two options.
330	// 1. We add the current diff on top of the parent layer.
331	// 2. We replace a layer above
332
333	let parent = parent_states
334		.pop()
335		.expect("parent must have a state");
336
337	let parent_added_len = parent.added.len();
338	let parent_removed_len = parent.removed.len();
339	let parent_diff = checked!(parent_added_len + parent_removed_len)?;
340
341	if checked!(diffsum * diffsum)? >= checked!(2 * diff_to_sibling * parent_diff)? {
342		// Diff too big, we replace above layer(s)
343		let mut parent_new = (*parent.added).clone();
344		let mut parent_removed = (*parent.removed).clone();
345
346		for removed in statediffremoved.iter() {
347			if !parent_new.remove(removed) {
348				// It was not added in the parent and we removed it
349				parent_removed.insert(*removed);
350			}
351			// Else it was added in the parent and we removed it again. We
352			// can forget this change
353		}
354
355		for new in statediffnew.iter() {
356			if !parent_removed.remove(new) {
357				// It was not touched in the parent and we added it
358				parent_new.insert(*new);
359			}
360			// Else it was removed in the parent and we added it again. We
361			// can forget this change
362		}
363
364		self.save_state_from_diff(
365			txn,
366			shortstatehash,
367			Arc::new(parent_new),
368			Arc::new(parent_removed),
369			diffsum,
370			parent_states,
371		)?;
372	} else {
373		// Diff small enough, we add diff as layer on top of parent
374		self.save_statediff(txn, shortstatehash, &StateDiff {
375			parent: Some(parent.shortstatehash),
376			added: statediffnew,
377			removed: statediffremoved,
378		});
379	}
380
381	Ok(())
382}
383
384/// Returns the new shortstatehash, and the state diff from the previous
385/// room state
386#[implement(Service)]
387#[tracing::instrument(skip(self, new_state_ids_compressed), level = "debug")]
388pub async fn save_state(
389	&self,
390	room_id: &RoomId,
391	new_state_ids_compressed: Arc<CompressedState>,
392) -> Result<HashSetCompressStateEvent> {
393	let previous_shortstatehash = self
394		.services
395		.state
396		.get_room_shortstatehash(room_id)
397		.await
398		.ok();
399
400	let state_hash = utils::calculate_hash(
401		new_state_ids_compressed
402			.iter()
403			.map(|bytes| &bytes[..]),
404	);
405
406	let existing_shortstatehash = self
407		.services
408		.short
409		.get_shortstatehash(&state_hash)
410		.await
411		.ok();
412
413	if let Some(new_shortstatehash) = existing_shortstatehash
414		.filter(|&new_shortstatehash| previous_shortstatehash.eq(&Some(new_shortstatehash)))
415	{
416		return Ok(HashSetCompressStateEvent {
417			shortstatehash: new_shortstatehash,
418			..Default::default()
419		});
420	}
421
422	let states_parents = if let Some(p) = previous_shortstatehash {
423		self.load_shortstatehash_info(p)
424			.await
425			.unwrap_or_default()
426	} else {
427		ShortStateInfoVec::new()
428	};
429
430	let (statediffnew, statediffremoved) = if let Some(parent_stateinfo) = states_parents.last() {
431		let statediffnew: CompressedState = new_state_ids_compressed
432			.difference(&parent_stateinfo.full_state)
433			.copied()
434			.collect();
435
436		let statediffremoved: CompressedState = parent_stateinfo
437			.full_state
438			.difference(&new_state_ids_compressed)
439			.copied()
440			.collect();
441
442		(Arc::new(statediffnew), Arc::new(statediffremoved))
443	} else {
444		(new_state_ids_compressed, Arc::new(CompressedState::new()))
445	};
446
447	let new_shortstatehash = if let Some(new_shortstatehash) = existing_shortstatehash {
448		new_shortstatehash
449	} else {
450		self.services
451			.short
452			.get_or_create_shortstatehash(&state_hash, |txn, shortstatehash| {
453				self.save_state_from_diff(
454					txn,
455					shortstatehash,
456					statediffnew.clone(),
457					statediffremoved.clone(),
458					2, // every state change is 2 event changes on average
459					states_parents,
460				)
461			})
462			.await?
463			.0
464	};
465
466	Ok(HashSetCompressStateEvent {
467		shortstatehash: new_shortstatehash,
468		added: statediffnew,
469		removed: statediffremoved,
470	})
471}
472
473/// Reads one state's delta row into its typed form.
474///
475/// Rows round-trip through [`save_statediff`], the pair being the only
476/// codec for the statediff encoding.
477#[implement(Service)]
478#[tracing::instrument(skip(self), level = "debug", name = "get")]
479pub(crate) async fn get_statediff(&self, shortstatehash: ShortStateHash) -> Result<StateDiff> {
480	const BUFSIZE: usize = size_of::<ShortStateHash>();
481	const STRIDE: usize = size_of::<ShortStateHash>();
482
483	let value = self
484		.db
485		.shortstatehash_statediff
486		.aqry::<BUFSIZE, _>(&shortstatehash)
487		.await
488		.map_err(|e| {
489			err!(Database("Failed to find StateDiff from short {shortstatehash:?}: {e}"))
490		})?;
491
492	let parent = utils::u64_from_bytes(&value[0..size_of::<u64>()])
493		.ok()
494		.take_if(|parent| *parent != 0);
495
496	debug_assert!(value.len().is_multiple_of(STRIDE), "value not aligned to stride");
497	let _num_values = value.len() / STRIDE;
498
499	let mut add_mode = true;
500	let mut added = CompressedState::new();
501	let mut removed = CompressedState::new();
502
503	let mut i = STRIDE;
504	while let Some(v) = value.get(i..expected!(i + 2 * STRIDE)) {
505		if add_mode && v.starts_with(&0_u64.to_be_bytes()) {
506			add_mode = false;
507			i = expected!(i + STRIDE);
508			continue;
509		}
510		if add_mode {
511			added.insert(v.try_into()?);
512		} else {
513			removed.insert(v.try_into()?);
514		}
515		i = expected!(i + 2 * STRIDE);
516	}
517
518	Ok(StateDiff {
519		parent,
520		added: Arc::new(added),
521		removed: Arc::new(removed),
522	})
523}
524
525/// Serializes one state's delta into the caller's transaction.
526///
527/// The one writer of the statediff encoding: entries emit sorted, the
528/// removed run behind its sentinel only when nonempty.
529#[implement(Service)]
530pub(crate) fn save_statediff(
531	&self,
532	txn: &mut Txn,
533	shortstatehash: ShortStateHash,
534	diff: &StateDiff,
535) {
536	let event_count = diff
537		.added
538		.len()
539		.saturating_add(diff.removed.len());
540
541	let event_bytes = event_count.saturating_mul(size_of::<CompressedStateEvent>());
542	let separator_bytes =
543		usize::from(!diff.removed.is_empty()).saturating_mul(size_of::<ShortStateHash>());
544
545	let capacity = size_of::<ShortStateHash>()
546		.saturating_add(event_bytes)
547		.saturating_add(separator_bytes);
548
549	let parent = diff.parent.unwrap_or(0_u64);
550	let mut value = Vec::<u8>::with_capacity(capacity);
551	value.extend_from_slice(&parent.to_be_bytes());
552
553	for new in diff.added.iter() {
554		value.extend_from_slice(&new[..]);
555	}
556
557	if !diff.removed.is_empty() {
558		value.extend_from_slice(&0_u64.to_be_bytes());
559		for removed in diff.removed.iter() {
560			value.extend_from_slice(&removed[..]);
561		}
562	}
563
564	txn.insert_raw(&self.db.shortstatehash_statediff, shortstatehash.to_be_bytes(), value);
565}
566
567#[inline]
568#[must_use]
569pub(crate) fn compress_state_event(
570	shortstatekey: ShortStateKey,
571	shorteventid: ShortEventId,
572) -> CompressedStateEvent {
573	const SIZE: usize = size_of::<CompressedStateEvent>();
574
575	let mut v = ArrayVec::<u8, SIZE>::new();
576	v.extend(shortstatekey.to_be_bytes());
577	v.extend(shorteventid.to_be_bytes());
578	v.as_ref()
579		.try_into()
580		.expect("failed to create CompressedStateEvent")
581}
582
583#[inline]
584#[must_use]
585pub(crate) fn parse_compressed_state_event(
586	compressed_event: CompressedStateEvent,
587) -> (ShortStateKey, ShortEventId) {
588	use utils::u64_from_u8;
589
590	let shortstatekey = u64_from_u8(&compressed_event[0..size_of::<ShortStateKey>()]);
591	let shorteventid = u64_from_u8(&compressed_event[size_of::<ShortStateKey>()..]);
592
593	(shortstatekey, shorteventid)
594}
595
596#[inline]
597fn compressed_state_size(compressed_state: &CompressedState) -> usize {
598	compressed_state
599		.len()
600		.checked_mul(size_of::<CompressedStateEvent>())
601		.expect("CompressedState size overflow")
602}