tuwunel_service/rooms/state_compressor/
mod.rs1use 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#[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#[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#[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#[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 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 parent_removed.insert(*removed);
292 }
293 }
296
297 for new in statediffnew.iter() {
298 if !parent_removed.remove(new) {
299 parent_new.insert(*new);
301 }
302 }
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 self.save_statediff(txn, shortstatehash, &StateDiff {
321 parent: None,
322 added: statediffnew,
323 removed: statediffremoved,
324 });
325
326 return Ok(());
327 }
328
329 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 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 parent_removed.insert(*removed);
350 }
351 }
354
355 for new in statediffnew.iter() {
356 if !parent_removed.remove(new) {
357 parent_new.insert(*new);
359 }
360 }
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 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#[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, 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#[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#[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}