Skip to main content

tuwunel_service/rooms/event_handler/
state_at_incoming.rs

1use std::{collections::HashMap, sync::atomic::AtomicBool};
2
3use futures::{
4	FutureExt, StreamExt, TryFutureExt, TryStreamExt,
5	future::{join, try_join},
6};
7use ruma::{EventId, OwnedEventId, RoomId, RoomVersionId};
8use tuwunel_core::{
9	Error, Result, debug, debug_warn, err, implement,
10	matrix::Event,
11	ref_at, trace,
12	utils::{
13		option::OptionExt,
14		stream::{BroadbandExt, IterStream, TryBroadbandExt, WidebandExt},
15	},
16};
17
18use crate::rooms::{
19	short::{ShortStateHash, ShortStateKey},
20	state_res::{AuthSet, StateMap},
21};
22
23enum ForkError {
24	State(Error),
25	Chain(Error),
26}
27
28// TODO: if we know the prev_events of the incoming event we can avoid the
29#[implement(super::Service)]
30// request and build the state from a known point and resolve if > 1 prev_event
31#[tracing::instrument(name = "state1", level = "debug", skip_all)]
32pub(super) async fn state_at_incoming_degree_one<Pdu>(
33	&self,
34	incoming_pdu: &Pdu,
35) -> Result<Option<HashMap<u64, OwnedEventId>>>
36where
37	Pdu: Event,
38{
39	debug_assert!(
40		incoming_pdu.prev_events().count() == 1,
41		"Incoming PDU must have one prev_event to make this call"
42	);
43
44	let prev_event_id = incoming_pdu
45		.prev_events()
46		.next()
47		.expect("at least one prev_event");
48
49	let Ok(prev_event_sstatehash) = self
50		.services
51		.state
52		.pdu_shortstatehash(prev_event_id)
53		.inspect_err(|e| debug_warn!(?prev_event_id, "Missing state at prev_event: {e}"))
54		.await
55	else {
56		return Ok(None);
57	};
58
59	debug!(?prev_event_id, ?prev_event_sstatehash, "Resolving state at prev_event.");
60
61	let prev_event = self
62		.services
63		.timeline
64		.get_pdu(prev_event_id)
65		.map_err(|e| err!(Database("Could not find prev_event, but we know the state: {e:?}")));
66
67	let state = self
68		.services
69		.state_accessor
70		.state_full_ids_strict(prev_event_sstatehash)
71		.try_collect::<HashMap<_, _>>()
72		.map(|state| {
73			state
74				.inspect_err(|error| {
75					debug_warn!(
76						%prev_event_id,
77						?prev_event_sstatehash,
78						%error,
79						"Failed to materialize state at prev_event.",
80					);
81				})
82				.ok()
83		})
84		.map(Ok);
85
86	let (prev_event, state) = try_join(prev_event, state).await?;
87	let Some(mut state) = state else {
88		return Ok(None);
89	};
90
91	debug!(
92		?prev_event_id,
93		?prev_event_sstatehash,
94		state_ids = state.len(),
95		"Resolved state at prev_event.",
96	);
97
98	// Every event holding a `shorteventid_shortstatehash` row passed spec check 5
99	// (auth against the state at its own position) as a hard reject; soft failure
100	// (spec check 6) still writes the row, so soft-failed events are valid fold
101	// inputs while positionally rejected events never gain a row.
102	if let Some(state_key) = prev_event.state_key() {
103		let prev_event_type = prev_event.event_type().to_cow_str().into();
104
105		let shortstatekey = self
106			.services
107			.short
108			.get_or_create_shortstatekey(&prev_event_type, state_key)
109			.await;
110
111		state.insert(shortstatekey, prev_event.event_id().into());
112		// Now it's the state after the pdu
113		debug!(
114			?prev_event_id,
115			?prev_event_type,
116			?prev_event_sstatehash,
117			?shortstatekey,
118			state_ids = state.len(),
119			"Added prev_event to state.",
120		);
121	}
122
123	debug_assert!(!state.is_empty(), "should be returning None for empty HashMap result");
124
125	Ok(Some(state))
126}
127
128#[implement(super::Service)]
129#[tracing::instrument(name = "stateN", level = "debug", skip_all)]
130pub(super) async fn state_at_incoming_resolved<Pdu>(
131	&self,
132	incoming_pdu: &Pdu,
133	room_id: &RoomId,
134	room_version: &RoomVersionId,
135) -> Result<Option<HashMap<u64, OwnedEventId>>>
136where
137	Pdu: Event,
138{
139	debug_assert!(
140		incoming_pdu.prev_events().count() > 1,
141		"Incoming PDU should have more than one prev_event for this codepath"
142	);
143
144	trace!("Calculating extremity statehashes...");
145	let Ok(extremity_sstatehashes) = incoming_pdu
146		.prev_events()
147		.try_stream()
148		.broad_and_then(|prev_event_id| {
149			let sstatehash = self
150				.services
151				.state
152				.pdu_shortstatehash(prev_event_id);
153
154			let prev_event = self.services.timeline.get_pdu(prev_event_id);
155
156			try_join(sstatehash, prev_event).inspect_err(move |e| {
157				debug_warn!(?prev_event_id, "Missing state at prev_event: {e}");
158			})
159		})
160		.try_collect::<HashMap<_, _>>()
161		.await
162	else {
163		return Ok(None);
164	};
165
166	trace!("Calculating fork states...");
167	let forks = extremity_sstatehashes
168		.into_iter()
169		.stream()
170		.wide_then(|(sstatehash, prev_event)| {
171			self.state_at_incoming_fork(room_id, room_version, sstatehash, prev_event)
172		})
173		.map(|fork| match fork {
174			| Ok(fork) => Ok(Ok(fork)),
175			| Err(ForkError::State(error)) => Ok(Err(error)),
176			| Err(ForkError::Chain(error)) => Err(error),
177		})
178		.try_collect::<Vec<_>>()
179		.await?;
180
181	let mut state_error = None;
182	let mut complete_forks = Vec::with_capacity(forks.len());
183	for fork in forks {
184		match fork {
185			| Ok(fork) => complete_forks.push(fork),
186			| Err(error) => {
187				state_error.get_or_insert(error);
188			},
189		}
190	}
191
192	if let Some(error) = state_error {
193		debug_warn!(%error, "Failed to materialize sibling fork state.");
194		return Ok(None);
195	}
196
197	let (fork_states, auth_chain_sets): (Vec<_>, Vec<_>) = complete_forks.into_iter().unzip();
198	let fork_states = fork_states.into_iter().stream();
199	let auth_chain_sets = auth_chain_sets.into_iter().stream();
200
201	trace!("Resolving state");
202	let Ok(new_state) = self
203		.state_resolution(room_id, room_version, fork_states, auth_chain_sets, None)
204		.inspect_ok(|_| trace!("State resolution done."))
205		.await
206	else {
207		return Ok(None);
208	};
209
210	new_state
211		.into_iter()
212		.stream()
213		.broad_then(async |((event_type, state_key), event_id)| {
214			self.services
215				.short
216				.get_or_create_shortstatekey(&event_type, &state_key)
217				.map(move |shortstatekey| (shortstatekey, event_id))
218				.await
219		})
220		.collect::<HashMap<_, _>>()
221		.inspect(|state| trace!(state = state.len(), "Created shortstatekeys."))
222		.map(Some)
223		.map(Ok)
224		.await
225}
226
227#[implement(super::Service)]
228#[tracing::instrument(
229	name = "fork",
230	level = "debug",
231	skip_all,
232	fields(
233		?sstatehash,
234		prev_event = ?prev_event.event_id(),
235	)
236)]
237async fn state_at_incoming_fork<Pdu>(
238	&self,
239	room_id: &RoomId,
240	room_version: &RoomVersionId,
241	sstatehash: ShortStateHash,
242	prev_event: Pdu,
243) -> Result<(StateMap<OwnedEventId>, AuthSet<OwnedEventId>), ForkError>
244where
245	Pdu: Event,
246{
247	// Every event holding a `shorteventid_shortstatehash` row passed spec check 5
248	// (auth against the state at its own position) as a hard reject; soft failure
249	// (spec check 6) still writes the row, so soft-failed events are valid fold
250	// inputs while positionally rejected events never gain a row.
251	let leaf = prev_event
252		.state_key()
253		.map_stream(async |state_key| {
254			let event_id = prev_event.event_id();
255			let event_type = prev_event.kind().to_cow_str().into();
256			let shortstatekey = self
257				.services
258				.short
259				.get_or_create_shortstatekey(&event_type, state_key)
260				.await;
261
262			(shortstatekey, event_id.to_owned())
263		});
264
265	let leaf_state_after_event: Vec<_> = self
266		.services
267		.state_accessor
268		.state_full_ids_strict(sstatehash)
269		.chain(leaf.map(Ok))
270		.try_collect()
271		.await
272		.map_err(ForkError::State)?;
273
274	trace!(
275		prev_event = ?prev_event.event_id(),
276		?sstatehash,
277		leaf_states = leaf_state_after_event.len(),
278		"leaf state after event"
279	);
280
281	let state = leaf_state_after_event
282		.iter()
283		.map(|(shortstatekey, event_id)| (*shortstatekey, event_id));
284
285	let starting_events = leaf_state_after_event
286		.iter()
287		.map(ref_at!(1))
288		.map(AsRef::as_ref);
289
290	let (state, chain) =
291		join(self.fork_state(state), self.fork_chain(room_id, room_version, starting_events))
292			.await;
293
294	let chain = chain.map_err(ForkError::Chain)?;
295	let state = state.map_err(ForkError::State)?;
296
297	Ok((state, chain))
298}
299
300/// Converts one fork branch from short state keys to typed state keys.
301///
302/// A later duplicate state key replaces an earlier entry. Short state key
303/// lookup failures are returned without constructing a partial map.
304#[implement(super::Service)]
305pub(super) async fn fork_state<'a, State>(
306	&'a self,
307	state: State,
308) -> Result<StateMap<OwnedEventId>>
309where
310	State: Iterator<Item = (ShortStateKey, &'a OwnedEventId)> + Send + 'a,
311{
312	state
313		.stream()
314		.wide_then(|(k, id)| {
315			self.services
316				.short
317				.get_statekey_from_short(k)
318				.map_ok(|(ty, sk)| ((ty, sk), id.clone()))
319		})
320		.try_collect()
321		.await
322}
323
324/// Collects the full auth chain for one fork branch.
325///
326/// The ids are distinct as [`AuthSet::from_distinct`] requires: the chain
327/// walk dedups short ids and the short-to-event-id mapping is injective.
328/// Iteration order is arbitrary.
329#[implement(super::Service)]
330pub(super) async fn fork_chain<'a, Events>(
331	&'a self,
332	room_id: &'a RoomId,
333	room_version: &'a RoomVersionId,
334	starting_events: Events,
335) -> Result<AuthSet<OwnedEventId>>
336where
337	Events: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
338{
339	self.services
340		.auth_chain
341		.event_ids_iter(room_id, room_version, starting_events)
342		.try_collect()
343		.map_ok(AuthSet::from_distinct)
344		.await
345}
346
347/// Collects the strict full auth chain for one fork branch.
348///
349/// The shared completeness flag is cleared when a polled chain source cannot
350/// be loaded.
351#[implement(super::Service)]
352pub(super) async fn fork_chain_strict<'a, Events>(
353	&'a self,
354	room_id: &'a RoomId,
355	room_version: &'a RoomVersionId,
356	starting_events: Events,
357	complete: &'a AtomicBool,
358) -> Result<AuthSet<OwnedEventId>>
359where
360	Events: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
361{
362	self.services
363		.auth_chain
364		.event_ids_iter_strict(room_id, room_version, starting_events, complete)
365		.try_collect()
366		.map_ok(AuthSet::from_distinct)
367		.await
368}