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