Skip to main content

tuwunel_service/rooms/event_handler/
fetch_prev.rs

1use 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
34/// An incoming event's missing previous events, walked backwards.
35///
36/// The incoming event's handler upgrades the walked events and records
37/// `capped` in the prev-walk counters.
38pub(super) struct PrevFetch {
39	/// Walked event ids in topological order, including placeholders for events
40	/// that failed to fetch or fell past the cap.
41	pub(super) sorted: Vec<OwnedEventId>,
42
43	/// Fetched PDUs of the walked events.
44	pub(super) pdus: Pdus,
45
46	/// Whether the `max_fetch_prev_events` cap cut the walk.
47	pub(super) capped: bool,
48}
49
50/// Walk an incoming event's missing previous events backwards.
51///
52/// Each fetched event queues its own missing previous events in turn, until
53/// they reach the timeline, predate the room, or exceed the
54/// `max_fetch_prev_events` cap.
55#[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) // heterogeneous FuturesOrdered
124		.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			// Fetch and handle failed
137			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			// keep counting so the capped check below sees the cut
145			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			// Get json failed, so this was not fetched over federation
161			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()); // heterogeneous FuturesOrdered
193			}
194
195			graph.insert(
196				prev_event_id.clone(),
197				pdu.prev_events().map(ToOwned::to_owned).collect(),
198			);
199		} else {
200			// Time based check failed
201			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		// This return value is the key used for sorting events,
213		// events are then sorted by power level, time,
214		// and lexically by event_id.
215		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	// at most limit + 1 events are admitted, so more means one was cut
235	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	// Each watcher registers before its existence recheck, so a prev that
250	// arrives during the recheck still wakes us.
251	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/// Fill the prev gap below `incoming_event_id` with one `/get_missing_events`
270/// batch, landing each returned event as a local outlier so the per-event walk
271/// resolves it without a federation fetch. `latest_events` is the held event
272/// the server walks back from, bounded by our forward extremities so it returns
273/// only the gap; best effort, so a failed batch or rejected event just leaves
274/// that id for the walk.
275#[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/// Authenticate and persist one event from the missing-events batch as an
321/// outlier, deriving its id from content rather than trusting a requested id.
322#[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	// cold arm: missing events
347	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}