Skip to main content

tuwunel_service/rooms/event_handler/
resolve_state.rs

1use std::{
2	borrow::Borrow,
3	collections::HashMap,
4	sync::{
5		Arc,
6		atomic::{AtomicBool, Ordering},
7	},
8};
9
10use futures::{FutureExt, Stream, StreamExt, TryFutureExt, TryStreamExt};
11use ruma::{EventId, OwnedEventId, RoomId, RoomVersionId};
12use serde::Deserialize;
13use tuwunel_core::{
14	Err, Result, err, error, implement,
15	matrix::room_version,
16	trace,
17	utils::stream::{IterStream, ReadyExt, TryWidebandExt, WidebandExt},
18};
19
20use crate::rooms::{
21	state_compressor::CompressedState,
22	state_res::{self, AuthSet, FetchEvent, StateMap},
23	timeline,
24};
25
26#[derive(Clone, Copy)]
27struct Strict<'a> {
28	timeline: &'a timeline::Service,
29	complete: Option<&'a AtomicBool>,
30}
31
32#[implement(super::Service)]
33#[tracing::instrument(
34	name = "state",
35	level = "debug",
36	skip_all,
37	fields(
38		incoming = ?incoming_state.len()
39	),
40)]
41pub async fn resolve_state(
42	&self,
43	room_id: &RoomId,
44	room_version: &RoomVersionId,
45	incoming_state: HashMap<u64, OwnedEventId>,
46) -> Result<Arc<CompressedState>> {
47	trace!("Loading current room state ids");
48	let current_sstatehash = self
49		.services
50		.state
51		.get_room_shortstatehash(room_id)
52		.map_err(|e| err!(Database(error!("No state for {room_id:?}: {e:?}"))))
53		.await?;
54
55	let current_state_ids: HashMap<_, _> = self
56		.services
57		.state_accessor
58		.state_full_ids(current_sstatehash)
59		.collect()
60		.await;
61
62	trace!("Loading fork states");
63	let fork_states = [current_state_ids, incoming_state];
64	let auth_chains = fork_states
65		.iter()
66		.try_stream()
67		.wide_and_then(|state| {
68			// The chain walk dedups short ids and maps them injectively, so
69			// the collected ids are distinct as `from_distinct` requires.
70			self.services
71				.auth_chain
72				.event_ids_iter(room_id, room_version, state.values().map(Borrow::borrow))
73				.try_collect()
74				.map_ok(AuthSet::from_distinct)
75		})
76		.ready_filter_map(Result::ok);
77
78	let fork_states = fork_states
79		.iter()
80		.stream()
81		.wide_then(|fork_state| {
82			let shortstatekeys = fork_state.keys().copied().stream();
83			let event_ids = fork_state.values().cloned().stream();
84			self.services
85				.short
86				.multi_get_statekey_from_short(shortstatekeys)
87				.zip(event_ids)
88				.ready_filter_map(|(ty_sk, id)| Some((ty_sk.ok()?, id)))
89				.collect::<StateMap<OwnedEventId>>()
90		});
91
92	trace!("Resolving state");
93	let state = self
94		.state_resolution(room_id, room_version, fork_states, auth_chains, None)
95		.await?;
96
97	trace!("State resolution done.");
98	let state_events: Vec<_> = state
99		.iter()
100		.stream()
101		.wide_then(|((event_type, state_key), event_id)| {
102			self.services
103				.short
104				.get_or_create_shortstatekey(event_type, state_key)
105				.map(move |shortstatekey| (shortstatekey, event_id))
106		})
107		.collect()
108		.await;
109
110	trace!("Compressing state...");
111	let new_room_state: CompressedState = self
112		.services
113		.state_compressor
114		.compress_state_events(
115			state_events
116				.iter()
117				.map(|(ssk, eid)| (ssk, (*eid).borrow())),
118		)
119		.collect()
120		.await;
121
122	Ok(Arc::new(new_room_state))
123}
124
125#[implement(super::Service)]
126#[tracing::instrument(name = "resolve", level = "debug", skip_all)]
127pub(super) async fn state_resolution<StateSets, AuthSets>(
128	&self,
129	_room_id: &RoomId,
130	room_version: &RoomVersionId,
131	state_sets: StateSets,
132	auth_chains: AuthSets,
133	complete: Option<&AtomicBool>,
134) -> Result<StateMap<OwnedEventId>>
135where
136	StateSets: Stream<Item = StateMap<OwnedEventId>> + Send,
137	AuthSets: Stream<Item = AuthSet<OwnedEventId>> + Send,
138{
139	let fetch = Strict {
140		timeline: &self.services.timeline,
141		complete,
142	};
143
144	state_res::resolve(
145		&room_version::rules(room_version)?,
146		state_sets,
147		auth_chains,
148		fetch,
149		self.services.server.config.hydra_backports,
150	)
151	.inspect_err(|error| {
152		if let Some(complete) = complete {
153			complete.store(false, Ordering::Relaxed);
154		}
155
156		error!(?error, "State resolution failed.");
157	})
158	.await
159}
160
161impl FetchEvent for Strict<'_> {
162	async fn get<T>(self, event_id: &EventId) -> Result<T>
163	where
164		T: for<'de> Deserialize<'de> + Send,
165	{
166		FetchEvent::get(self.timeline, event_id)
167			.map_err(|error| match self.complete.filter(|_| error.is_not_found()) {
168				| None => error,
169				| Some(complete) => {
170					complete.store(false, Ordering::Relaxed);
171
172					err!(Database("State resolution references missing event {event_id}."))
173				},
174			})
175			.await
176	}
177
178	async fn exists(self, event_id: &EventId) -> Result<bool> {
179		let found = FetchEvent::exists(self.timeline, event_id).await?;
180
181		match self.complete.filter(|_| !found) {
182			| None => Ok(found),
183			| Some(complete) => {
184				complete.store(false, Ordering::Relaxed);
185
186				Err!(Database("State resolution references missing event {event_id}."))
187			},
188		}
189	}
190}