1mod fetch_state;
8mod prune;
9
10use std::{collections::HashMap, fmt::Write, iter::once, sync::Arc};
11
12use async_trait::async_trait;
13pub(crate) use fetch_state::IdMapState;
17use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt, future::join_all};
18pub(crate) use prune::prune_goal;
22pub use prune::{PruneSummary, Trigger};
27use ruma::{
28 CanonicalJsonObject, EventId, OwnedEventId, OwnedRoomId, RoomId, RoomVersionId, UserId,
29 events::{
30 AnyStrippedStateEvent, StateEventType, TimelineEventType,
31 room::member::{MembershipState, RoomMemberEventContent},
32 },
33 room_version_rules::AuthorizationRules,
34 serde::Raw,
35};
36use serde_json::value::RawValue as RawJsonValue;
37use tuwunel_core::{
38 Event, PduEvent, Result, err,
39 error::inspect_debug_log,
40 implement,
41 matrix::{PduCount, RoomVersionRules, StateKey, TypeStateKey, room_version},
42 result::{AndThenRef, FlatOk, NotFound},
43 smallvec::SmallVec,
44 trace,
45 utils::{
46 BoolExt, IterStream, MutexMap, MutexMapGuard, ReadyExt, TryReadyExt, calculate_hash,
47 mutex_map::Guard,
48 stream::{TryBroadbandExt, TryIgnore, WidebandExt},
49 },
50 warn,
51};
52use tuwunel_database::{Deserialized, Ignore, Interfix, Map, Txn};
53
54use crate::{
55 rooms::{
56 short::{ShortEventId, ShortStateHash, ShortStateKey},
57 state_cache::{MembershipUpdate, StrippedRoomState},
58 state_compressor::{CompressedState, parse_compressed_state_event},
59 state_res::{StateMap, auth_types_for_event},
60 },
61 services::OnceServices,
62};
63
64pub struct Service {
70 pub mutex: RoomMutexMap,
76 services: Arc<OnceServices>,
77 db: Data,
78}
79
80struct Data {
81 shorteventid_shortstatehash: Arc<Map>,
82 roomid_shortstatehash: Arc<Map>,
83 roomid_pduleaves: Arc<Map>,
84}
85
86type RoomMutexMap = MutexMap<OwnedRoomId, ()>;
87pub type RoomMutexGuard = MutexMapGuard<OwnedRoomId, ()>;
92type ForwardExtremities = SmallVec<[OwnedEventId; 1]>;
93
94#[async_trait]
95impl crate::Service for Service {
96 fn build(args: &crate::Args<'_>) -> Result<Arc<Self>> {
97 Ok(Arc::new(Self {
98 mutex: RoomMutexMap::new(),
99 services: args.services.clone(),
100 db: Data {
101 shorteventid_shortstatehash: args.db["shorteventid_shortstatehash"].clone(),
102 roomid_shortstatehash: args.db["roomid_shortstatehash"].clone(),
103 roomid_pduleaves: args.db["roomid_pduleaves"].clone(),
104 },
105 }))
106 }
107
108 async fn memory_usage(&self, out: &mut (dyn Write + Send)) -> Result {
109 let mutex = self.mutex.len();
110 writeln!(out, "- state_mutex: {mutex}")?;
111
112 Ok(())
113 }
114
115 fn name(&self) -> &str { crate::service::make_name(std::module_path!()) }
116}
117
118#[implement(Service)]
126#[tracing::instrument(
127 name = "force",
128 level = "debug",
129 skip_all,
130 fields(
131 count = ?self.services.globals.pending_count(),
132 %shortstatehash,
133 )
134)]
135pub async fn force_state(
136 &self,
137 room_id: &RoomId,
138 shortstatehash: u64,
139 statediffnew: Arc<CompressedState>,
140 _statediffremoved: Arc<CompressedState>,
141 state_lock: &RoomMutexGuard,
142) -> Result {
143 statediffnew
144 .iter()
145 .stream()
146 .map(|&new| parse_compressed_state_event(new).1)
147 .wide_filter_map(async |shorteventid| {
148 let event_id: OwnedEventId = self
149 .services
150 .short
151 .get_eventid_from_short(shorteventid)
152 .inspect_err(inspect_debug_log)
153 .await
154 .ok()?;
155
156 self.services
157 .timeline
158 .get_pdu(&event_id)
159 .await
160 .ok()
161 })
162 .map(Ok)
163 .try_for_each(async |pdu| match pdu.kind {
164 | TimelineEventType::RoomMember => self.force_member_effects(room_id, &pdu).await,
165 | _ => Ok(()),
166 })
167 .boxed() .await?;
169
170 self.services
171 .state_cache
172 .update_joined_count(room_id)
173 .await;
174
175 self.set_room_state(room_id, shortstatehash, state_lock);
176
177 self.services.spaces.cache_evict(room_id);
179
180 Ok(())
181}
182
183#[implement(Service)]
189async fn force_member_effects(&self, room_id: &RoomId, pdu: &PduEvent) -> Result {
190 let Some(user_id) = pdu
191 .state_key
192 .as_ref()
193 .map(UserId::parse)
194 .flat_ok()
195 else {
196 return Ok(());
197 };
198
199 let Ok(membership_event): Result<RoomMemberEventContent> = pdu.get_content() else {
200 return Ok(());
201 };
202
203 let last_state = membership_event
204 .membership
205 .eq(&MembershipState::Invite)
206 .and_is(self.services.globals.user_is_local(&user_id))
207 .then_async(|| self.replayed_invite_state(room_id, &user_id, pdu))
208 .map(Option::transpose)
209 .map_ok(Option::flatten)
210 .await?;
211
212 let count = self.services.globals.next_count();
213
214 self.services
215 .state_cache
216 .update_membership(MembershipUpdate {
217 room_id,
218 user_id: &user_id,
219 membership_event,
220 sender: &pdu.sender,
221 last_state,
222 invite_via: None,
223 update_joined_count: false,
224 count: PduCount::Normal(*count),
225 })
226 .await
227}
228
229#[implement(Service)]
236async fn replayed_invite_state(
237 &self,
238 room_id: &RoomId,
239 user_id: &UserId,
240 pdu: &PduEvent,
241) -> Result<StrippedRoomState> {
242 self.services
243 .state_cache
244 .has_invite_state(user_id, room_id)
245 .await?
246 .is_false()
247 .then_async(|| self.summary_stripped(pdu))
248 .map(Ok)
249 .await
250}
251
252#[implement(Service)]
258#[tracing::instrument(
259 name = "set",
260 level = "debug",
261 skip(self, state_ids_compressed),
262 fields(
263 count = ?self.services.globals.pending_count(),
264 )
265)]
266pub async fn set_event_state(
267 &self,
268 event_id: &EventId,
269 room_id: &RoomId,
270 state_ids_compressed: Arc<CompressedState>,
271) -> Result<ShortStateHash> {
272 const KEY_LEN: usize = size_of::<ShortEventId>();
273 const VAL_LEN: usize = size_of::<ShortStateHash>();
274
275 let shorteventid = self
276 .services
277 .short
278 .get_or_create_shorteventid(event_id)
279 .await;
280
281 let state_hash = calculate_hash(state_ids_compressed.iter().map(|s| &s[..]));
282
283 if let Ok(shortstatehash) = self
284 .services
285 .short
286 .get_shortstatehash(&state_hash)
287 .await
288 {
289 self.db
290 .shorteventid_shortstatehash
291 .aput::<KEY_LEN, VAL_LEN, _, _>(shorteventid, shortstatehash);
292
293 return Ok(shortstatehash);
294 }
295
296 let previous_shortstatehash = self.get_room_shortstatehash(room_id).await;
297 let states_parents = match previous_shortstatehash {
298 | Ok(p) =>
299 self.services
300 .state_compressor
301 .load_shortstatehash_info(p)
302 .await?,
303 | _ => Vec::new(),
304 };
305
306 let (statediffnew, statediffremoved) = if let Some(parent_stateinfo) = states_parents.last() {
307 let statediffnew: CompressedState = state_ids_compressed
308 .difference(&parent_stateinfo.full_state)
309 .copied()
310 .collect();
311
312 let statediffremoved: CompressedState = parent_stateinfo
313 .full_state
314 .difference(&state_ids_compressed)
315 .copied()
316 .collect();
317
318 (Arc::new(statediffnew), Arc::new(statediffremoved))
319 } else {
320 (state_ids_compressed, Arc::new(CompressedState::new()))
321 };
322
323 let save_statediff = |txn: &mut Txn, shortstatehash| {
324 self.services
325 .state_compressor
326 .save_state_from_diff(
327 txn,
328 shortstatehash,
329 statediffnew,
330 statediffremoved,
331 1_000_000, states_parents,
333 )
334 };
335
336 let (shortstatehash, _) = self
337 .services
338 .short
339 .get_or_create_shortstatehash(&state_hash, save_statediff)
340 .await?;
341
342 self.db
343 .shorteventid_shortstatehash
344 .aput::<KEY_LEN, VAL_LEN, _, _>(shorteventid, shortstatehash);
345
346 Ok(shortstatehash)
347}
348
349#[implement(Service)]
362#[tracing::instrument(
363 name = "set",
364 level = "debug",
365 skip(self, new_pdu),
366 fields(
367 count = ?self.services.globals.pending_count(),
368 )
369)]
370pub async fn append_to_state(&self, new_pdu: &PduEvent) -> Result<u64> {
371 const KEY_LEN: usize = size_of::<ShortEventId>();
372 const VAL_LEN: usize = size_of::<ShortStateHash>();
373
374 let shorteventid = self
375 .services
376 .short
377 .get_or_create_shorteventid(&new_pdu.event_id)
378 .await;
379
380 let previous_shortstatehash = self
381 .get_room_shortstatehash(&new_pdu.room_id)
382 .await;
383
384 if let Ok(p) = previous_shortstatehash {
385 self.db
386 .shorteventid_shortstatehash
387 .aput::<KEY_LEN, VAL_LEN, _, _>(shorteventid, p);
388 }
389
390 match &new_pdu.state_key {
391 | Some(state_key) => {
392 let states_parents = match previous_shortstatehash {
393 | Ok(p) =>
394 self.services
395 .state_compressor
396 .load_shortstatehash_info(p)
397 .await?,
398 | _ => Vec::new(),
399 };
400
401 let shortstatekey = self
402 .services
403 .short
404 .get_or_create_shortstatekey(&new_pdu.kind.to_string().into(), state_key)
405 .await;
406
407 let new = self
408 .services
409 .state_compressor
410 .compress_state_event(shortstatekey, &new_pdu.event_id)
411 .await;
412
413 let replaces = states_parents
414 .last()
415 .map(|info| {
416 info.full_state
417 .iter()
418 .find(|bytes| bytes.starts_with(&shortstatekey.to_be_bytes()))
419 })
420 .unwrap_or_default();
421
422 if Some(&new) == replaces {
423 return Ok(previous_shortstatehash.expect("must exist"));
424 }
425
426 let shortstatehash = self.services.globals.next_count();
428 let mut txn = self.services.db.txn();
429
430 let mut statediffnew = CompressedState::new();
431 statediffnew.insert(new);
432
433 let mut statediffremoved = CompressedState::new();
434 if let Some(replaces) = replaces {
435 statediffremoved.insert(*replaces);
436 }
437
438 self.services
439 .state_compressor
440 .save_state_from_diff(
441 &mut txn,
442 *shortstatehash,
443 Arc::new(statediffnew),
444 Arc::new(statediffremoved),
445 2,
446 states_parents,
447 )?;
448
449 txn.execute();
450
451 Ok(*shortstatehash)
452 },
453 | _ => Ok(previous_shortstatehash.expect("first event in room must be a state event")),
454 }
455}
456
457#[implement(Service)]
462#[tracing::instrument(skip(self, _mutex_lock), level = "debug")]
463pub fn set_room_state(
464 &self,
465 room_id: &RoomId,
466 shortstatehash: u64,
467 _mutex_lock: &RoomMutexGuard,
469) {
470 const BUFSIZE: usize = size_of::<u64>();
471
472 self.db
473 .roomid_shortstatehash
474 .raw_aput::<BUFSIZE, _, _>(room_id, shortstatehash);
475}
476
477#[implement(Service)]
483#[expect(clippy::too_many_arguments)]
484#[tracing::instrument(skip(self, content), level = "debug")]
485pub async fn get_auth_events(
486 &self,
487 room_id: &RoomId,
488 kind: &TimelineEventType,
489 sender: &UserId,
490 state_key: Option<&str>,
491 content: &serde_json::value::RawValue,
492 auth_rules: &AuthorizationRules,
493 include_create: bool,
494) -> Result<StateMap<PduEvent>>
495where
496 StateEventType: Send + Sync,
497 StateKey: Send + Sync,
498{
499 let Some(shortstatehash) = self
500 .get_room_shortstatehash(room_id)
501 .await
502 .optional()?
503 else {
504 return Ok(StateMap::new());
505 };
506
507 let sauthevents: HashMap<ShortStateKey, TypeStateKey> =
508 auth_types_for_event(kind, sender, state_key, content, auth_rules, include_create)?
509 .into_iter()
510 .try_stream()
511 .broad_and_then(async |(event_type, state_key): TypeStateKey| {
512 self.services
513 .short
514 .get_shortstatekey(&event_type, &state_key)
515 .await
516 .map(|sstatekey| (sstatekey, (event_type, state_key)))
517 .optional()
518 })
519 .ready_try_filter_map(Result::Ok)
520 .try_collect()
521 .await?;
522
523 let matching_state: Vec<_> = self
524 .services
525 .state_accessor
526 .state_full_shortids(shortstatehash)
527 .ready_try_filter_map(|(shortstatekey, shorteventid)| {
528 Ok(sauthevents
529 .get(&shortstatekey)
530 .map(move |(ty, sk)| ((ty, sk), shorteventid)))
531 })
532 .try_collect()
533 .await?;
534 let (state_keys, event_ids): (Vec<_>, Vec<_>) = matching_state.into_iter().unzip();
535
536 self.services
537 .short
538 .multi_get_eventid_from_short(event_ids.into_iter().stream())
539 .zip(state_keys.into_iter().stream())
540 .map(|(event_id, state_key)| {
541 event_id
542 .map(|event_id| (state_key, event_id))
543 .optional()
544 })
545 .ready_try_filter_map(Result::Ok)
546 .broad_and_then(async |((ty, sk), event_id): ((&_, &_), OwnedEventId)| {
547 self.services
548 .timeline
549 .get_pdu(&event_id)
550 .map_ok(|pdu| ((ty.clone(), sk.clone()), pdu))
551 .await
552 .optional()
553 })
554 .ready_try_filter_map(Result::Ok)
555 .try_collect()
556 .await
557}
558
559#[implement(Service)]
565#[tracing::instrument(skip_all, level = "debug")]
566pub async fn summary_stripped<Pdu: Event>(&self, event: &Pdu) -> Vec<Raw<AnyStrippedStateEvent>> {
567 let cells = [
568 (&StateEventType::RoomCreate, ""),
569 (&StateEventType::RoomJoinRules, ""),
570 (&StateEventType::RoomCanonicalAlias, ""),
571 (&StateEventType::RoomName, ""),
572 (&StateEventType::RoomAvatar, ""),
573 (&StateEventType::RoomMember, event.sender().as_str()), (&StateEventType::RoomEncryption, ""),
575 (&StateEventType::RoomTopic, ""),
576 ];
577
578 let fetches = cells.into_iter().map(|(event_type, state_key)| {
579 self.services
580 .state_accessor
581 .room_state_get(event.room_id(), event_type, state_key)
582 });
583
584 join_all(fetches)
585 .await
586 .into_iter()
587 .filter_map(Result::ok)
588 .map(Event::into_format)
589 .chain(once(event.to_format()))
590 .collect()
591}
592
593#[implement(Service)]
599#[tracing::instrument(skip_all, level = "debug")]
600pub async fn summary_pdus<Pdu: Event>(
601 &self,
602 event: &Pdu,
603 event_json: &CanonicalJsonObject,
604 room_version: &RoomVersionId,
605) -> Vec<Box<RawJsonValue>> {
606 let cells = [
607 (&StateEventType::RoomCreate, ""),
608 (&StateEventType::RoomJoinRules, ""),
609 (&StateEventType::RoomCanonicalAlias, ""),
610 (&StateEventType::RoomName, ""),
611 (&StateEventType::RoomAvatar, ""),
612 (&StateEventType::RoomMember, event.sender().as_str()),
613 (&StateEventType::RoomEncryption, ""),
614 (&StateEventType::RoomTopic, ""),
615 ];
616
617 let membership = self
618 .services
619 .federation
620 .format_pdu_into(event_json.clone(), Some(room_version))
621 .boxed() .await;
623
624 cells
625 .into_iter()
626 .stream()
627 .wide_filter_map(async |(event_type, state_key)| {
628 let pdu = self
629 .services
630 .state_accessor
631 .room_state_get(event.room_id(), event_type, state_key)
632 .await
633 .ok()?;
634
635 let pdu_json = self
636 .services
637 .timeline
638 .get_pdu_json(pdu.event_id())
639 .await
640 .ok()?;
641
642 Some(
643 self.services
644 .federation
645 .format_pdu_into(pdu_json, Some(room_version))
646 .await,
647 )
648 })
649 .chain(once(membership).stream())
650 .collect()
651 .await
652}
653
654#[implement(Service)]
658#[inline]
659pub async fn get_room_version_rules(&self, room_id: &RoomId) -> Result<RoomVersionRules> {
660 self.get_room_version(room_id)
661 .await
662 .and_then_ref(room_version::rules)
663}
664
665#[implement(Service)]
666#[tracing::instrument(
667 level = "debug"
668 skip(self),
669 ret(level = "trace"),
670)]
671pub async fn get_room_version(&self, room_id: &RoomId) -> Result<RoomVersionId> {
675 self.services
676 .state_accessor
677 .room_state_get_content(room_id, &StateEventType::RoomCreate, "")
678 .await
679 .as_ref()
680 .map(room_version::from_create_content)
681 .cloned()
682 .map_err(|e| err!(Request(NotFound("No create event found: {e:?}"))))
683}
684
685#[implement(Service)]
686#[tracing::instrument(
687 level = "debug"
688 skip(self),
689 ret(level = "trace"),
690)]
691pub async fn get_room_shortstatehash(&self, room_id: &RoomId) -> Result<ShortStateHash> {
696 self.db
697 .roomid_shortstatehash
698 .get(room_id)
699 .await
700 .deserialized()
701}
702
703#[implement(Service)]
708pub async fn pdu_shortstatehash(&self, event_id: &EventId) -> Result<ShortStateHash> {
709 self.services
710 .short
711 .get_shorteventid(event_id)
712 .and_then(|shorteventid| self.get_shortstatehash(shorteventid))
713 .await
714}
715
716#[implement(Service)]
717#[tracing::instrument(
718 level = "debug"
719 skip(self),
720 ret(level = "trace"),
721)]
722pub async fn get_shortstatehash(&self, shorteventid: ShortEventId) -> Result<ShortStateHash> {
726 const BUFSIZE: usize = size_of::<ShortEventId>();
727
728 self.db
729 .shorteventid_shortstatehash
730 .aqry::<BUFSIZE, _>(&shorteventid)
731 .await
732 .deserialized()
733}
734
735#[implement(Service)]
740pub(super) fn delete_room_shortstatehash(
741 &self,
742 room_id: &RoomId,
743 _mutex_lock: &Guard<OwnedRoomId, ()>,
744) -> Result {
745 self.db.roomid_shortstatehash.remove(room_id);
746
747 Ok(())
748}
749
750#[implement(Service)]
755#[tracing::instrument(
756 level = "debug"
757 skip_all,
758 fields(%room_id),
759)]
760pub async fn collapse_forward_extremities(
761 &self,
762 room_id: &RoomId,
763 state_lock: &RoomMutexGuard,
764) -> usize {
765 let extremities: ForwardExtremities = self
766 .get_forward_extremities(room_id)
767 .map(ToOwned::to_owned)
768 .collect()
769 .await;
770
771 if extremities.len() <= 1 {
772 return 0;
773 }
774
775 let survivor = join_all(extremities.iter().map(async |event_id| {
776 self.services
777 .timeline
778 .get_pdu_count(event_id)
779 .await
780 .ok()
781 .map(|count| (count, event_id))
782 }))
783 .await
784 .into_iter()
785 .flatten()
786 .max_by_key(|(count, _)| *count)
787 .map(|(_, event_id)| event_id);
788
789 let Some(survivor) = survivor else {
790 return 0;
791 };
792
793 self.set_forward_extremities(room_id, once(&**survivor), state_lock)
794 .await;
795
796 extremities.len().saturating_sub(1)
797}
798
799#[implement(Service)]
800#[tracing::instrument(
801 level = "trace"
802 skip(self),
803)]
804pub fn get_forward_extremities<'a>(
810 &'a self,
811 room_id: &'a RoomId,
812) -> impl Stream<Item = &EventId> + Send + '_ {
813 let prefix = (room_id, Interfix);
814
815 self.db
816 .roomid_pduleaves
817 .keys_prefix(&prefix)
818 .map_ok(|(_, event_id): (Ignore, &EventId)| event_id)
819 .ignore_err()
820}
821
822#[implement(Service)]
823#[tracing::instrument(
824 level = "debug"
825 skip_all,
826 fields(%room_id),
827)]
828pub async fn set_forward_extremities<'a, I>(
834 &'a self,
835 room_id: &'a RoomId,
836 event_ids: I,
837 _state_lock: &'a RoomMutexGuard,
838) where
839 I: Iterator<Item = &'a EventId> + Send + 'a,
840{
841 let prefix = (room_id, Interfix);
842 self.db
843 .roomid_pduleaves
844 .keys_prefix_raw(&prefix)
845 .ignore_err()
846 .ready_for_each(|key| self.db.roomid_pduleaves.remove(key))
847 .await;
848
849 for event_id in event_ids {
850 let key = (room_id, event_id);
851 self.db.roomid_pduleaves.put_raw(key, event_id);
852 }
853}
854
855#[implement(Service)]
860pub(super) async fn delete_all_rooms_forward_extremities(&self, room_id: &RoomId) -> Result {
861 let prefix = (room_id, Interfix);
862
863 self.db
864 .roomid_pduleaves
865 .keys_prefix_raw(&prefix)
866 .ignore_err()
867 .ready_for_each(|key| {
868 trace!("Removing key: {key:?}");
869 self.db.roomid_pduleaves.remove(key);
870 })
871 .await;
872
873 Ok(())
874}