tuwunel_service/rooms/state_compressor/
mod.rs1use 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
27pub struct Service {
33 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#[derive(Clone)]
49pub(crate) struct StateDiff {
50 pub(crate) parent: Option<ShortStateHash>,
52
53 pub(crate) added: Arc<CompressedState>,
55
56 pub(crate) removed: Arc<CompressedState>,
58}
59
60#[derive(Clone, Default)]
65pub struct ShortStateInfo {
66 pub shortstatehash: ShortStateHash,
68
69 pub full_state: Arc<CompressedState>,
71
72 pub added: Arc<CompressedState>,
74
75 pub removed: Arc<CompressedState>,
77}
78
79#[derive(Clone, Default)]
84pub struct HashSetCompressStateEvent {
85 pub shortstatehash: ShortStateHash,
87
88 pub added: Arc<CompressedState>,
90
91 pub removed: Arc<CompressedState>,
93}
94
95type StateInfoLruCache = LruCache<ShortStateHash, ShortStateInfoVec>;
96type ShortStateInfoVec = Vec<ShortStateInfo>;
97type ParentStatesVec = Vec<ShortStateInfo>;
98
99pub type CompressedState = BTreeSet<CompressedStateEvent>;
104
105pub 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#[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#[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?; 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#[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#[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#[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 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 parent_removed.insert(*removed);
335 }
336 }
339
340 for new in statediffnew.iter() {
341 if !parent_removed.remove(new) {
342 parent_new.insert(*new);
344 }
345 }
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 self.save_statediff(txn, shortstatehash, &StateDiff {
364 parent: None,
365 added: statediffnew,
366 removed: statediffremoved,
367 });
368
369 return Ok(());
370 }
371
372 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 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 parent_removed.insert(*removed);
393 }
394 }
397
398 for new in statediffnew.iter() {
399 if !parent_removed.remove(new) {
400 parent_new.insert(*new);
402 }
403 }
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 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#[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, 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#[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#[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#[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#[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}