tuwunel_service/rooms/event_handler/
fetch_auth.rs1use 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#[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() .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 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 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 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 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}