tuwunel_service/rooms/event_handler/
state_at_incoming.rs1use 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#[implement(super::Service)]
30#[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 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 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 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#[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#[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#[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}