Skip to main content

tuwunel_service/rooms/event_handler/
state_at_incoming.rs

1use std::collections::HashMap;
2
3use futures::{FutureExt, StreamExt, TryFutureExt, TryStreamExt, future::try_join};
4use ruma::{EventId, OwnedEventId, RoomId, RoomVersionId};
5use tuwunel_core::{
6	Result, apply, debug, debug_warn, err, implement,
7	matrix::Event,
8	ref_at, trace,
9	utils::{
10		option::OptionExt,
11		stream::{
12			BroadbandExt, IterStream, ReadyExt, TryBroadbandExt, TryWidebandExt, WidebandExt,
13		},
14	},
15};
16
17use crate::rooms::{
18	short::{ShortStateHash, ShortStateKey},
19	state_res::{AuthSet, StateMap},
20};
21
22// TODO: if we know the prev_events of the incoming event we can avoid the
23#[implement(super::Service)]
24// request and build the state from a known point and resolve if > 1 prev_event
25#[tracing::instrument(name = "state1", level = "debug", skip_all)]
26pub(super) async fn state_at_incoming_degree_one<Pdu>(
27	&self,
28	incoming_pdu: &Pdu,
29) -> Result<Option<HashMap<u64, OwnedEventId>>>
30where
31	Pdu: Event,
32{
33	debug_assert!(
34		incoming_pdu.prev_events().count() == 1,
35		"Incoming PDU must have one prev_event to make this call"
36	);
37
38	let prev_event_id = incoming_pdu
39		.prev_events()
40		.next()
41		.expect("at least one prev_event");
42
43	let Ok(prev_event_sstatehash) = self
44		.services
45		.state
46		.pdu_shortstatehash(prev_event_id)
47		.inspect_err(|e| debug_warn!(?prev_event_id, "Missing state at prev_event: {e}"))
48		.await
49	else {
50		return Ok(None);
51	};
52
53	debug!(?prev_event_id, ?prev_event_sstatehash, "Resolving state at prev_event.");
54
55	let prev_event = self
56		.services
57		.timeline
58		.get_pdu(prev_event_id)
59		.map_err(|e| err!(Database("Could not find prev_event, but we know the state: {e:?}")));
60
61	let state = self
62		.services
63		.state_accessor
64		.state_full_ids(prev_event_sstatehash)
65		.collect::<HashMap<_, _>>()
66		.map(Ok);
67
68	let (prev_event, mut state) = try_join(prev_event, state).await?;
69
70	debug!(
71		?prev_event_id,
72		?prev_event_sstatehash,
73		state_ids = state.len(),
74		"Resolved state at prev_event.",
75	);
76
77	if let Some(state_key) = prev_event.state_key() {
78		let prev_event_type = prev_event.event_type().to_cow_str().into();
79
80		let shortstatekey = self
81			.services
82			.short
83			.get_or_create_shortstatekey(&prev_event_type, state_key)
84			.await;
85
86		state.insert(shortstatekey, prev_event.event_id().into());
87		// Now it's the state after the pdu
88		debug!(
89			?prev_event_id,
90			?prev_event_type,
91			?prev_event_sstatehash,
92			?shortstatekey,
93			state_ids = state.len(),
94			"Added prev_event to state.",
95		);
96	}
97
98	debug_assert!(!state.is_empty(), "should be returning None for empty HashMap result");
99
100	Ok(Some(state))
101}
102
103#[implement(super::Service)]
104#[tracing::instrument(name = "stateN", level = "debug", skip_all)]
105pub(super) async fn state_at_incoming_resolved<Pdu>(
106	&self,
107	incoming_pdu: &Pdu,
108	room_id: &RoomId,
109	room_version: &RoomVersionId,
110) -> Result<Option<HashMap<u64, OwnedEventId>>>
111where
112	Pdu: Event,
113{
114	debug_assert!(
115		incoming_pdu.prev_events().count() > 1,
116		"Incoming PDU should have more than one prev_event for this codepath"
117	);
118
119	trace!("Calculating extremity statehashes...");
120	let Ok(extremity_sstatehashes) = incoming_pdu
121		.prev_events()
122		.try_stream()
123		.broad_and_then(|prev_event_id| {
124			let sstatehash = self
125				.services
126				.state
127				.pdu_shortstatehash(prev_event_id);
128
129			let prev_event = self.services.timeline.get_pdu(prev_event_id);
130
131			try_join(sstatehash, prev_event).inspect_err(move |e| {
132				debug_warn!(?prev_event_id, "Missing state at prev_event: {e}");
133			})
134		})
135		.try_collect::<HashMap<_, _>>()
136		.await
137	else {
138		return Ok(None);
139	};
140
141	trace!("Calculating fork states...");
142	let (fork_states, auth_chain_sets) = extremity_sstatehashes
143		.into_iter()
144		.try_stream()
145		.wide_and_then(|(sstatehash, prev_event)| {
146			self.state_at_incoming_fork(room_id, room_version, sstatehash, prev_event)
147		})
148		.try_collect()
149		.map_ok(Vec::into_iter)
150		.map_ok(Iterator::unzip)
151		.map_ok(apply!(2, Vec::into_iter))
152		.map_ok(apply!(2, IterStream::stream))
153		.await?;
154
155	trace!("Resolving state");
156	let Ok(new_state) = self
157		.state_resolution(room_id, room_version, fork_states, auth_chain_sets)
158		.inspect_ok(|_| trace!("State resolution done."))
159		.await
160	else {
161		return Ok(None);
162	};
163
164	new_state
165		.into_iter()
166		.stream()
167		.broad_then(async |((event_type, state_key), event_id)| {
168			self.services
169				.short
170				.get_or_create_shortstatekey(&event_type, &state_key)
171				.map(move |shortstatekey| (shortstatekey, event_id))
172				.await
173		})
174		.collect::<HashMap<_, _>>()
175		.inspect(|state| trace!(state = state.len(), "Created shortstatekeys."))
176		.map(Some)
177		.map(Ok)
178		.await
179}
180
181#[implement(super::Service)]
182#[tracing::instrument(
183	name = "fork",
184	level = "debug",
185	skip_all,
186	fields(
187		?sstatehash,
188		prev_event = ?prev_event.event_id(),
189	)
190)]
191async fn state_at_incoming_fork<Pdu>(
192	&self,
193	room_id: &RoomId,
194	room_version: &RoomVersionId,
195	sstatehash: ShortStateHash,
196	prev_event: Pdu,
197) -> Result<(StateMap<OwnedEventId>, AuthSet<OwnedEventId>)>
198where
199	Pdu: Event,
200{
201	let leaf = prev_event
202		.state_key()
203		.map_stream(async |state_key| {
204			let event_id = prev_event.event_id();
205			let event_type = prev_event.kind().to_cow_str().into();
206			let shortstatekey = self
207				.services
208				.short
209				.get_or_create_shortstatekey(&event_type, state_key)
210				.await;
211
212			(shortstatekey, event_id.to_owned())
213		});
214
215	let leaf_state_after_event: Vec<_> = self
216		.services
217		.state_accessor
218		.state_full_ids(sstatehash)
219		.chain(leaf)
220		.collect()
221		.await;
222
223	trace!(
224		prev_event = ?prev_event.event_id(),
225		?sstatehash,
226		leaf_states = leaf_state_after_event.len(),
227		"leaf state after event"
228	);
229
230	let state = leaf_state_after_event
231		.iter()
232		.map(|(shortstatekey, event_id)| (*shortstatekey, event_id));
233
234	let starting_events = leaf_state_after_event
235		.iter()
236		.map(ref_at!(1))
237		.map(AsRef::as_ref);
238
239	try_join(
240		self.fork_state(state).map(Ok),
241		self.fork_chain(room_id, room_version, starting_events),
242	)
243	.await
244}
245
246/// Converts one fork branch from short state keys to typed state keys.
247///
248/// A later duplicate state key replaces an earlier entry. Short state key
249/// lookup failures are omitted.
250#[implement(super::Service)]
251pub(super) async fn fork_state<'a, State>(&'a self, state: State) -> StateMap<OwnedEventId>
252where
253	State: Iterator<Item = (ShortStateKey, &'a OwnedEventId)> + Send + 'a,
254{
255	state
256		.stream()
257		.wide_then(|(k, id)| {
258			self.services
259				.short
260				.get_statekey_from_short(k)
261				.map_ok(|(ty, sk)| ((ty, sk), id.clone()))
262		})
263		.ready_filter_map(Result::ok)
264		.collect()
265		.await
266}
267
268/// Collects the full auth chain for one fork branch.
269///
270/// The ids are distinct as [`AuthSet::from_distinct`] requires: the chain
271/// walk dedups short ids and the short-to-event-id mapping is injective.
272/// Iteration order is arbitrary.
273#[implement(super::Service)]
274pub(super) async fn fork_chain<'a, Events>(
275	&'a self,
276	room_id: &'a RoomId,
277	room_version: &'a RoomVersionId,
278	starting_events: Events,
279) -> Result<AuthSet<OwnedEventId>>
280where
281	Events: Iterator<Item = &'a EventId> + Clone + ExactSizeIterator + Send + 'a,
282{
283	self.services
284		.auth_chain
285		.event_ids_iter(room_id, room_version, starting_events)
286		.try_collect()
287		.map_ok(AuthSet::from_distinct)
288		.await
289}