1use std::collections::{BTreeSet, VecDeque};
2
3use axum::extract::State;
4use futures::{
5 StreamExt, TryFutureExt, TryStreamExt,
6 future::try_join,
7 stream::{FuturesOrdered, unfold},
8};
9use ruma::{
10 CanonicalJsonObject, CanonicalJsonValue, EventId, OwnedEventId,
11 api::federation::event::get_missing_events, canonical_json::redact_in_place,
12};
13use tuwunel_core::{
14 Error, Result, debug, err,
15 matrix::room_version::rules as room_version_rules,
16 utils::stream::{TryWidebandExt, automatic_width},
17};
18
19use super::AccessCheck;
20use crate::Ruma;
21
22type Seen = BTreeSet<OwnedEventId>;
23type Pending = VecDeque<(OwnedEventId, bool)>;
24
25const LIMIT_MAX: usize = 50;
27
28const LIMIT_DEFAULT: usize = 10;
30
31const WALK_MAX: usize = 256;
33
34const EARLIEST_MAX: usize = 4096;
36
37pub(crate) async fn get_missing_events_route(
41 State(services): State<crate::State>,
42 body: Ruma<get_missing_events::v1::Request>,
43) -> Result<get_missing_events::v1::Response> {
44 let access_check = AccessCheck {
45 services: &services,
46 origin: body.origin(),
47 room_id: &body.room_id,
48 event_id: None,
49 };
50
51 let room_version = services.state.get_room_version(&body.room_id);
52
53 let (room_version, ()) = try_join(room_version, access_check.check()).await?;
54
55 let rules = room_version_rules(&room_version)?;
56
57 let fetch = async |(event_id, is_latest): (OwnedEventId, bool)| {
58 let event = services.timeline.get_pdu_json(&event_id).await;
59
60 (event_id, is_latest, event)
61 };
62
63 let limit = body
65 .limit
66 .try_into()
67 .unwrap_or(LIMIT_DEFAULT)
68 .min(LIMIT_MAX);
69
70 let (seen, pending) = walk_seed(&body);
71 let seen_max = seen.len().saturating_add(WALK_MAX);
72
73 let fetches = FuturesOrdered::new();
74
75 let events =
76 unfold((seen, pending, fetches), async |(mut seen, mut pending, mut fetches)| {
77 let event = next_missing_event(
78 &body,
79 &fetch,
80 seen_max,
81 &mut seen,
82 &mut pending,
83 &mut fetches,
84 )
85 .await?;
86
87 Some((event, (seen, pending, fetches)))
88 })
89 .take(limit)
90 .map(Ok::<_, Error>)
91 .wide_and_then(async |(event_id, mut event)| {
92 let visible = services
93 .state_accessor
94 .server_can_see_event(body.origin(), &body.room_id, &event_id)
95 .await;
96
97 let event = if visible {
98 services
99 .state_accessor
100 .erased_for_server(body.origin(), event)
101 .await
102 } else {
103 redact_in_place(&mut event, &rules.redaction, None)
104 .map_err(|error| err!(Database("Failed to redact event: {error}")))?;
105
106 event
107 };
108
109 let event = services
110 .federation
111 .format_pdu_into(event, Some(&room_version))
112 .await;
113
114 Ok(event)
115 })
116 .try_collect::<Vec<_>>()
117 .map_ok(|mut vec| {
118 vec.reverse();
119 vec
120 })
121 .await?;
122
123 Ok(get_missing_events::v1::Response { events })
124}
125
126fn walk_seed(body: &get_missing_events::v1::Request) -> (Seen, Pending) {
132 let mut seen: Seen = body
133 .earliest_events
134 .iter()
135 .take(EARLIEST_MAX)
136 .cloned()
137 .collect();
138
139 let pending = body
140 .latest_events
141 .iter()
142 .take(WALK_MAX)
143 .filter(|event_id| seen.insert((*event_id).clone()))
144 .cloned()
145 .map(|event_id| (event_id, true))
146 .collect();
147
148 (seen, pending)
149}
150
151async fn next_missing_event<Fetch, Fut>(
152 body: &Ruma<get_missing_events::v1::Request>,
153 fetch: &Fetch,
154 seen_max: usize,
155 seen: &mut Seen,
156 pending: &mut Pending,
157 fetches: &mut FuturesOrdered<Fut>,
158) -> Option<(OwnedEventId, CanonicalJsonObject)>
159where
160 Fetch: Fn((OwnedEventId, bool)) -> Fut + Sync,
161 Fut: Future<Output = (OwnedEventId, bool, Result<CanonicalJsonObject>)> + Send,
162{
163 loop {
164 let width = automatic_width();
165
166 while fetches.len() < width
167 && let Some(input) = pending.pop_front()
168 {
169 fetches.push_back(fetch(input));
170 }
171
172 let (event_id, is_latest, event) = fetches.next().await?;
173 let Ok(event) = event else {
174 debug!(
175 ?body.origin,
176 %event_id,
177 "Event does not exist locally, skipping"
178 );
179
180 continue;
181 };
182
183 if event
184 .get("room_id")
185 .and_then(CanonicalJsonValue::as_str)
186 != Some(body.room_id.as_str())
187 {
188 continue;
189 }
190
191 event
192 .get("prev_events")
193 .and_then(CanonicalJsonValue::as_array)
194 .into_iter()
195 .flatten()
196 .filter_map(CanonicalJsonValue::as_str)
197 .filter_map(|event_id| EventId::parse(event_id).ok())
198 .filter(|event_id| seen.len() < seen_max && seen.insert(event_id.clone()))
199 .for_each(|event_id| pending.push_back((event_id, false)));
200
201 if !is_latest {
202 return Some((event_id, event));
203 }
204 }
205}
206
207#[cfg(test)]
208mod tests {
209 use std::sync::atomic::{AtomicUsize, Ordering::Relaxed};
210
211 use axum_extra::extract::cookie::CookieJar;
212 use futures::stream::FuturesOrdered;
213 use ruma::{
214 CanonicalJsonObject, EventId, OwnedEventId, RoomId,
215 api::federation::event::get_missing_events::v1::Request, room_id, server_name,
216 };
217 use serde_json::json;
218 use tuwunel_core::matrix::pdu::MAX_PREV_EVENTS;
219
220 use super::{EARLIEST_MAX, Ruma, WALK_MAX, err, next_missing_event, walk_seed};
221
222 fn event_ids(prefix: &str, len: usize) -> Vec<OwnedEventId> {
223 (0..len)
224 .map(|index| {
225 EventId::parse(format!("${prefix}{index}:example.com")).expect("valid event id")
226 })
227 .collect()
228 }
229
230 fn request(earliest: Vec<OwnedEventId>, latest: Vec<OwnedEventId>) -> Ruma<Request> {
231 let body = Request::new(room_id!("!room:example.com").to_owned(), earliest, latest);
232
233 Ruma {
234 body,
235 cookie: CookieJar::new(),
236 origin: Some(server_name!("example.com").to_owned()),
237 sender_user: None,
238 sender_device: None,
239 appservice_info: None,
240 json_body: None,
241 }
242 }
243
244 fn event_with_prevs(room_id: &RoomId, index: usize) -> CanonicalJsonObject {
245 let prev_events: Vec<_> = (0..MAX_PREV_EVENTS)
246 .map(|prev| format!("$p{index}a{prev}:example.com"))
247 .collect();
248 let value = json!({
249 "room_id": room_id,
250 "prev_events": prev_events,
251 });
252
253 serde_json::from_value(value).expect("valid canonical json")
254 }
255
256 #[test]
257 fn seed_bounded_by_cap_not_request() {
258 let body = request(Vec::new(), event_ids("latest", 5_000));
259 let (seen, pending) = walk_seed(&body);
260
261 assert_eq!(seen.len(), WALK_MAX);
262 assert_eq!(pending.len(), WALK_MAX);
263
264 let body = request(event_ids("earliest", 10_000), event_ids("latest", 5_000));
265 let (seen, pending) = walk_seed(&body);
266
267 assert_eq!(seen.len(), EARLIEST_MAX + WALK_MAX);
268 assert_eq!(pending.len(), WALK_MAX);
269 }
270
271 #[tokio::test]
272 async fn walk_lookups_bounded_by_cap_not_request() {
273 let body = request(Vec::new(), event_ids("latest", 5_000));
274 let (mut seen, mut pending) = walk_seed(&body);
275 let seen_max = seen.len().saturating_add(WALK_MAX);
276 let lookups = AtomicUsize::new(0);
277 let room_id = body.room_id.clone();
278 let fetch = async |(event_id, is_latest): (OwnedEventId, bool)| {
279 let index = lookups.fetch_add(1, Relaxed);
280 let event = is_latest
281 .then(|| event_with_prevs(&room_id, index))
282 .ok_or_else(|| err!(Request(NotFound("Event not found."))));
283
284 (event_id, is_latest, event)
285 };
286
287 let mut fetches = FuturesOrdered::new();
288
289 let result =
290 next_missing_event(&body, &fetch, seen_max, &mut seen, &mut pending, &mut fetches)
291 .await;
292
293 assert!(result.is_none());
294 assert_eq!(lookups.load(Relaxed), 2 * WALK_MAX);
295 }
296}