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