Skip to main content

tuwunel_service/rooms/event_handler/
fetch_auth.rs

1use std::{
2	collections::{HashSet, VecDeque},
3	time::Duration,
4};
5
6use futures::{FutureExt, StreamExt, TryFutureExt};
7use ruma::{
8	CanonicalJsonObject, CanonicalJsonValue, EventId, OwnedEventId, RoomId, RoomVersionId,
9	ServerName,
10};
11use tuwunel_core::{
12	debug, debug_error, debug_warn, expected, implement,
13	matrix::{PduEvent, pdu::MAX_AUTH_EVENTS},
14	trace,
15	utils::stream::{BroadbandExt, IterStream},
16	warn,
17};
18
19use super::backoff::{Context, Disposition};
20use crate::fetcher::{Op, Opts};
21
22/// Find the event and auth it. Once the event is validated (steps 1 - 8)
23/// it is appended to the outliers Tree.
24///
25/// Returns pdu and if we fetched it over federation the raw json.
26///
27/// a. Look in the main timeline (pduid_pdu tree)
28/// b. Look at outlier pdu tree
29/// c. Ask origin server over federation
30/// d. TODO: Ask other servers over federation?
31#[implement(super::Service)]
32#[tracing::instrument(
33	level = "debug",
34	skip_all,
35	fields(
36		%origin,
37		events = %events.clone().count(),
38		lev = %recursion_level,
39	),
40)]
41pub(super) async fn fetch_auth<'a, Events>(
42	&self,
43	origin: &ServerName,
44	room_id: &RoomId,
45	events: Events,
46	room_version: &RoomVersionId,
47	recursion_level: usize,
48) -> Vec<(PduEvent, Option<CanonicalJsonObject>)>
49where
50	Events: Iterator<Item = &'a EventId> + Clone + Send,
51{
52	let events_with_auth_events: Vec<_> = events
53		.stream()
54		.broad_then(|event_id| self.fetch_auth_chain(origin, room_id, event_id, room_version))
55		.collect()
56		.boxed() // size firewall
57		.await;
58
59	events_with_auth_events
60		.into_iter()
61		.stream()
62		.fold(Vec::new(), async |mut pdus, (id, local_pdu, events_in_reverse_order)| {
63			if self.services.server.check_running().is_err() {
64				return pdus;
65			}
66
67			// a. Look in the main timeline (pduid_pdu tree)
68			// b. Look at outlier pdu tree
69			// (get_pdu_json checks both)
70			if let Some(local_pdu) = local_pdu {
71				pdus.push((local_pdu, None));
72			}
73
74			events_in_reverse_order
75				.into_iter()
76				.rev()
77				.stream()
78				.fold(pdus, async |mut pdus, (next_id, value)| {
79					if self
80						.is_suppressed(
81							Context::Auth,
82							&next_id,
83							Duration::from_mins(5)..Duration::from_hours(24),
84						)
85						.await
86						.is_deny()
87					{
88						return pdus;
89					}
90
91					// recursion cycle
92					let outlier = Box::pin(self.handle_outlier_pdu(
93						origin,
94						room_id,
95						&next_id,
96						value.clone(),
97						room_version,
98						expected!(recursion_level + 1),
99						true,
100					));
101
102					if let Ok((pdu, json)) = outlier
103						.await
104						.inspect_err(|e| warn!("Authentication of event {next_id} failed: {e:?}"))
105					{
106						if next_id == id {
107							pdus.push((pdu, Some(json)));
108						}
109						self.record_success(Context::Auth, &next_id).await;
110					} else {
111						self.record_outcome(Context::Auth, &next_id, Disposition::Transient);
112					}
113
114					pdus
115				})
116				.await
117		})
118		.await
119}
120
121#[implement(super::Service)]
122#[tracing::instrument(
123	name = "chain",
124	level = "trace",
125	skip_all,
126	fields(%event_id),
127)]
128async fn fetch_auth_chain(
129	&self,
130	origin: &ServerName,
131	room_id: &RoomId,
132	event_id: &EventId,
133	room_version: &RoomVersionId,
134) -> (OwnedEventId, Option<PduEvent>, Vec<(OwnedEventId, CanonicalJsonObject)>) {
135	// a. Look in the main timeline (pduid_pdu tree)
136	// b. Look at outlier pdu tree
137	// (get_pdu_json checks both)
138	if let Ok(local_pdu) = self.services.timeline.get_pdu(event_id).await {
139		trace!(?event_id, "Found in database");
140		return (event_id.to_owned(), Some(local_pdu), vec![]);
141	}
142
143	// c. Ask origin server over federation
144	// We also handle its auth chain here so we don't get a stack overflow in
145	// handle_outlier_pdu.
146	let mut events_all = HashSet::new();
147	let mut events_in_reverse_order = Vec::new();
148	let mut todo_auth_events: VecDeque<_> = [event_id.to_owned()].into();
149	while let Some(next_id) = todo_auth_events.pop_front() {
150		if events_all.contains(&next_id) {
151			continue;
152		}
153
154		if self
155			.is_suppressed(
156				Context::Fetch,
157				&next_id,
158				Duration::from_mins(2)..Duration::from_hours(8),
159			)
160			.await
161			.is_deny()
162		{
163			debug_warn!("Backed off from {next_id}");
164			continue;
165		}
166
167		if self.services.timeline.pdu_exists(&next_id).await {
168			trace!(?next_id, "Found in database");
169			continue;
170		}
171
172		if self.services.server.check_running().is_err() {
173			debug_warn!(?next_id, "Server shutting down");
174			break;
175		}
176
177		debug!("Fetching {next_id} over federation.");
178		let opts = Opts::new(Op::AuthEvent, room_id.to_owned())
179			.event_id(next_id.clone())
180			.hint(origin.to_owned())
181			.room_version(room_version.to_owned())
182			.attempt_limit(super::EVENT_FETCH_ATTEMPT_LIMIT)
183			.fanout_for_op();
184
185		let Ok(outcome) = self
186			.services
187			.fetcher
188			.fetch(opts)
189			.inspect_err(|e| debug_error!(?next_id, "Failed to fetch event: {e}"))
190			.await
191		else {
192			debug_warn!("Backing off from {next_id}");
193			self.record_outcome(Context::Fetch, &next_id, Disposition::Transient);
194			continue;
195		};
196
197		let Ok(value) = serde_json::from_slice::<CanonicalJsonObject>(&outcome.bytes) else {
198			self.record_outcome(Context::Fetch, &next_id, Disposition::Transient);
199			continue;
200		};
201
202		debug!("Got {next_id} over federation");
203		self.record_success(Context::Fetch, &next_id)
204			.await;
205		value
206			.get("auth_events")
207			.and_then(CanonicalJsonValue::as_array)
208			.into_iter()
209			.flatten()
210			.filter_map(|auth_event| auth_event.try_into().ok())
211			.take(MAX_AUTH_EVENTS)
212			.for_each(|auth_event: &EventId| {
213				todo_auth_events.push_back(auth_event.to_owned());
214			});
215
216		events_in_reverse_order.push((next_id.clone(), value));
217		events_all.insert(next_id);
218	}
219
220	(event_id.to_owned(), None, events_in_reverse_order)
221}