tuwunel_service/rooms/event_handler/
fetch_prev.rs1use std::{collections::HashMap, iter::once, time::Duration};
2
3use futures::{
4 FutureExt, StreamExt,
5 stream::{FuturesOrdered, FuturesUnordered},
6};
7use ruma::{
8 CanonicalJsonObject, CanonicalJsonValue, EventId, MilliSecondsSinceUnixEpoch, OwnedEventId,
9 RoomId, RoomVersionId, ServerName, int, uint,
10};
11use serde_json::value::RawValue as RawJsonValue;
12use tokio::time::{Instant, timeout_at};
13use tuwunel_core::{
14 Result, err, implement,
15 matrix::{
16 Event, PduEvent,
17 event::gen_event_id,
18 pdu::{MAX_PREV_EVENTS, check_room_id},
19 },
20 utils::{
21 BoolExt,
22 stream::{IterStream, automatic_width},
23 },
24};
25
26use super::handle_prev_pdu::PrevUpgrade;
27use crate::{
28 fetcher::{EventWindow, Op, Opts},
29 rooms::state_res::topological_sort,
30};
31
32pub(super) type Pdus = HashMap<OwnedEventId, (PduEvent, CanonicalJsonObject)>;
33
34pub(super) struct PrevFetch {
39 pub(super) sorted: Vec<OwnedEventId>,
42
43 pub(super) pdus: Pdus,
45
46 pub(super) capped: bool,
48}
49
50#[implement(super::Service)]
56#[tracing::instrument(
57 level = "debug",
58 skip_all,
59 fields(
60 %origin,
61 events = %initial_set.clone().count(),
62 ),
63)]
64pub(super) async fn fetch_prev<'a, Events>(
65 &self,
66 PrevUpgrade {
67 origin,
68 room_id,
69 event_id: incoming_event_id,
70 room_version,
71 recursion_level,
72 first_ts_in_room,
73 ..
74 }: PrevUpgrade<'_>,
75 initial_set: Events,
76) -> Result<PrevFetch>
77where
78 Events: Iterator<Item = &'a EventId> + Clone + Send,
79{
80 let has_gap = initial_set
81 .clone()
82 .stream()
83 .any(async |event_id| !self.services.timeline.pdu_exists(event_id).await)
84 .await;
85
86 let wait_ms = self.services.server.config.fetch_prev_wait_ms;
87 let has_gap = (has_gap && wait_ms > 0)
88 .then_async(|| self.await_prev_gap(initial_set.clone(), Duration::from_millis(wait_ms)))
89 .await
90 .unwrap_or(has_gap);
91
92 has_gap
93 .then_async(|| {
94 self.prefetch_missing_events(
95 origin,
96 room_id,
97 incoming_event_id,
98 room_version,
99 recursion_level,
100 )
101 })
102 .await;
103
104 let mut todo_outlier_stack: FuturesOrdered<_> = initial_set
105 .stream()
106 .map(ToOwned::to_owned)
107 .filter_map(async |event_id| {
108 self.services
109 .timeline
110 .non_outlier_pdu_exists(&event_id)
111 .await
112 .is_err()
113 .then_some(event_id)
114 })
115 .map(async |event_id| {
116 let events = once(event_id.as_ref());
117 let auth = self
118 .fetch_auth(origin, room_id, events, room_version, recursion_level)
119 .await;
120
121 (event_id, auth)
122 })
123 .map(FutureExt::boxed) .collect()
125 .await;
126
127 let limit = usize::from(self.services.server.config.max_fetch_prev_events);
128 let mut amount = 0_usize;
129 let mut pdus = HashMap::new();
130 let mut graph: HashMap<OwnedEventId, _> = HashMap::with_capacity(todo_outlier_stack.len());
131
132 while let Some((prev_event_id, mut outlier)) = todo_outlier_stack.next().await {
133 self.services.server.check_running()?;
134
135 let Some((pdu, mut json_opt)) = outlier.pop() else {
136 graph.insert(prev_event_id.clone(), Default::default());
138 continue;
139 };
140
141 check_room_id(&pdu, room_id)?;
142
143 if amount > limit {
144 amount = amount.saturating_add(1);
146 graph.insert(prev_event_id.clone(), Default::default());
147 continue;
148 }
149
150 if json_opt.is_none() {
151 json_opt = self
152 .services
153 .timeline
154 .get_outlier_pdu_json(&prev_event_id)
155 .await
156 .ok();
157 }
158
159 let Some(json) = json_opt else {
160 graph.insert(prev_event_id.clone(), Default::default());
162 continue;
163 };
164
165 if pdu.origin_server_ts() > first_ts_in_room {
166 amount = amount.saturating_add(1);
167 debug_assert!(
168 pdu.prev_events().count() <= MAX_PREV_EVENTS,
169 "PduEvent {prev_event_id} has too many prev_events"
170 );
171
172 for prev_prev in pdu.prev_events() {
173 if graph.contains_key(prev_prev) {
174 continue;
175 }
176
177 let prev_prev = prev_prev.to_owned();
178 let fetch = async move {
179 let fetch = self
180 .fetch_auth(
181 origin,
182 room_id,
183 once(prev_prev.as_ref()),
184 room_version,
185 recursion_level,
186 )
187 .await;
188
189 (prev_prev, fetch)
190 };
191
192 todo_outlier_stack.push_back(fetch.boxed()); }
194
195 graph.insert(
196 prev_event_id.clone(),
197 pdu.prev_events().map(ToOwned::to_owned).collect(),
198 );
199 } else {
200 graph.insert(prev_event_id.clone(), Default::default());
202 }
203
204 pdus.insert(prev_event_id.clone(), (pdu, json));
205 }
206
207 let event_fetch = async |event_id: OwnedEventId| {
208 let origin_server_ts = pdus
209 .get(&event_id)
210 .map_or_else(|| uint!(0), |info| info.0.origin_server_ts().get());
211
212 Ok((int!(0).into(), MilliSecondsSinceUnixEpoch(origin_server_ts)))
216 };
217
218 let graph_len = graph.len();
219 let sorted = topological_sort(graph, &event_fetch)
220 .await
221 .map_err(|e| err!(Database(error!("Error sorting prev events: {e}"))))?;
222
223 debug_assert_eq!(
224 sorted.len(),
225 graph_len,
226 "topological sort returned a different number of outputs than inputs"
227 );
228
229 debug_assert!(
230 sorted.len() >= pdus.len(),
231 "returned topologically sorted events differ from pdus"
232 );
233
234 let capped = amount > limit.saturating_add(1);
236
237 Ok(PrevFetch { sorted, pdus, capped })
238}
239
240#[implement(super::Service)]
241async fn await_prev_gap<'a, Events>(&self, initial_set: Events, wait: Duration) -> bool
242where
243 Events: Iterator<Item = &'a EventId> + Send,
244{
245 let deadline = Instant::now()
246 .checked_add(wait)
247 .expect("wait deadline overflows");
248
249 let pending: FuturesUnordered<_> = initial_set
252 .map(|event_id| (event_id, self.services.timeline.watch_event(event_id)))
253 .stream()
254 .filter_map(async |(event_id, watcher)| {
255 (!self.services.timeline.pdu_exists(event_id).await).then_some(watcher)
256 })
257 .collect()
258 .await;
259
260 if pending.is_empty() {
261 return false;
262 }
263
264 timeout_at(deadline, pending.count())
265 .await
266 .is_err()
267}
268
269#[implement(super::Service)]
276#[tracing::instrument(name = "missing", level = "debug", skip_all)]
277async fn prefetch_missing_events(
278 &self,
279 origin: &ServerName,
280 room_id: &RoomId,
281 incoming_event_id: &EventId,
282 room_version: &RoomVersionId,
283 recursion_level: usize,
284) {
285 let boundary: EventWindow = self
286 .services
287 .state
288 .get_forward_extremities(room_id)
289 .map(ToOwned::to_owned)
290 .collect()
291 .await;
292
293 let opts = Opts::new(Op::MissingEvents, room_id.to_owned())
294 .latest_events([incoming_event_id.to_owned()])
295 .earliest_events(boundary)
296 .hint(origin.to_owned())
297 .room_version(room_version.to_owned())
298 .attempt_limit(super::EVENT_FETCH_ATTEMPT_LIMIT)
299 .fanout_for_op();
300
301 let Ok(outcome) = self.services.fetcher.fetch(opts).await else {
302 return;
303 };
304
305 let Ok(events) = serde_json::from_slice::<Vec<Box<RawJsonValue>>>(&outcome.bytes) else {
306 return;
307 };
308
309 events
310 .into_iter()
311 .stream()
312 .for_each_concurrent(automatic_width(), async |pdu| {
313 self.land_missing_event(origin, room_id, &pdu, room_version, recursion_level)
314 .await
315 .ok();
316 })
317 .await;
318}
319
320#[implement(super::Service)]
323#[tracing::instrument(name = "land", level = "trace", skip_all)]
324async fn land_missing_event(
325 &self,
326 origin: &ServerName,
327 room_id: &RoomId,
328 pdu: &RawJsonValue,
329 room_version: &RoomVersionId,
330 recursion_level: usize,
331) -> Result {
332 let value: CanonicalJsonObject = serde_json::from_str(pdu.get())
333 .map_err(|e| err!(BadServerResponse("missing-events pdu is not canonical json: {e}")))?;
334
335 value
336 .get("room_id")
337 .and_then(CanonicalJsonValue::as_str)
338 .is_some_and(|id| id == room_id.as_str())
339 .then_some(())
340 .ok_or_else(|| {
341 err!(Request(InvalidParam("missing-events pdu is for a different room")))
342 })?;
343
344 let event_id = gen_event_id(&value, room_version)?;
345
346 Box::pin(self.handle_outlier_pdu(
348 origin,
349 room_id,
350 &event_id,
351 value,
352 room_version,
353 recursion_level,
354 false,
355 ))
356 .await
357 .map(|_| ())
358}