Skip to main content

tuwunel_service/rooms/event_handler/
resolve_state.rs

1use std::{borrow::Borrow, collections::HashMap, sync::Arc};
2
3use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
4use ruma::{OwnedEventId, RoomId, RoomVersionId};
5use tuwunel_core::{
6	Result, err, implement,
7	matrix::room_version,
8	trace,
9	utils::stream::{IterStream, ReadyExt, TryWidebandExt, WidebandExt},
10};
11
12use crate::rooms::{
13	state_compressor::CompressedState,
14	state_res::{self, AuthSet, StateMap},
15};
16
17#[implement(super::Service)]
18#[tracing::instrument(
19	name = "state",
20	level = "debug",
21	skip_all,
22	fields(
23		incoming = ?incoming_state.len()
24	),
25)]
26pub async fn resolve_state(
27	&self,
28	room_id: &RoomId,
29	room_version: &RoomVersionId,
30	incoming_state: HashMap<u64, OwnedEventId>,
31) -> Result<Arc<CompressedState>> {
32	trace!("Loading current room state ids");
33	let current_sstatehash = self
34		.services
35		.state
36		.get_room_shortstatehash(room_id)
37		.map_err(|e| err!(Database(error!("No state for {room_id:?}: {e:?}"))))
38		.await?;
39
40	let current_state_ids: HashMap<_, _> = self
41		.services
42		.state_accessor
43		.state_full_ids(current_sstatehash)
44		.collect()
45		.await;
46
47	trace!("Loading fork states");
48	let fork_states = [current_state_ids, incoming_state];
49	let auth_chains = fork_states
50		.iter()
51		.try_stream()
52		.wide_and_then(|state| {
53			// The chain walk dedups short ids and maps them injectively, so
54			// the collected ids are distinct as `from_distinct` requires.
55			self.services
56				.auth_chain
57				.event_ids_iter(room_id, room_version, state.values().map(Borrow::borrow))
58				.try_collect()
59				.map_ok(AuthSet::from_distinct)
60		})
61		.ready_filter_map(Result::ok);
62
63	let fork_states = fork_states
64		.iter()
65		.stream()
66		.wide_then(|fork_state| {
67			let shortstatekeys = fork_state.keys().copied().stream();
68			let event_ids = fork_state.values().cloned().stream();
69			self.services
70				.short
71				.multi_get_statekey_from_short(shortstatekeys)
72				.zip(event_ids)
73				.ready_filter_map(|(ty_sk, id)| Some((ty_sk.ok()?, id)))
74				.collect::<StateMap<OwnedEventId>>()
75		});
76
77	trace!("Resolving state");
78	let state = self
79		.state_resolution(room_id, room_version, fork_states, auth_chains)
80		.await?;
81
82	trace!("State resolution done.");
83	let state_events: Vec<_> = state
84		.iter()
85		.stream()
86		.wide_then(|((event_type, state_key), event_id)| {
87			self.services
88				.short
89				.get_or_create_shortstatekey(event_type, state_key)
90				.map(move |shortstatekey| (shortstatekey, event_id))
91		})
92		.collect()
93		.await;
94
95	trace!("Compressing state...");
96	let new_room_state: CompressedState = self
97		.services
98		.state_compressor
99		.compress_state_events(
100			state_events
101				.iter()
102				.map(|(ssk, eid)| (ssk, (*eid).borrow())),
103		)
104		.collect()
105		.await;
106
107	Ok(Arc::new(new_room_state))
108}
109
110#[implement(super::Service)]
111#[tracing::instrument(name = "resolve", level = "debug", skip_all)]
112pub(super) async fn state_resolution<StateSets, AuthSets>(
113	&self,
114	_room_id: &RoomId,
115	room_version: &RoomVersionId,
116	state_sets: StateSets,
117	auth_chains: AuthSets,
118) -> Result<StateMap<OwnedEventId>>
119where
120	StateSets: Stream<Item = StateMap<OwnedEventId>> + Send,
121	AuthSets: Stream<Item = AuthSet<OwnedEventId>> + Send,
122{
123	state_res::resolve(
124		&room_version::rules(room_version)?,
125		state_sets,
126		auth_chains,
127		&async |event_id: OwnedEventId| self.event_fetch(&event_id).await,
128		&async |event_id: OwnedEventId| self.event_exists(&event_id).await,
129		self.services.server.config.hydra_backports,
130	)
131	.map_err(|e| err!(error!("State resolution failed: {e:?}")))
132	.await
133}