1use std::{cmp::Ordering, collections::BTreeMap};
8
9use futures::StreamExt;
10use ruma::{EventId, OwnedEventId, OwnedServerName, RoomId};
11use tuwunel_core::{
12 implement,
13 matrix::{Event, PduCount},
14 smallvec::SmallVec,
15 utils::{IterStream, stream::BroadbandExt},
16};
17
18use crate::federation::ShouldAttempt;
19
20type Servers = SmallVec<[OwnedServerName; 1]>;
21
22#[derive(Clone, Copy, Debug, Default)]
29pub struct PruneSummary {
30 pub before: usize,
32
33 pub after: usize,
35
36 pub dangling: usize,
38
39 pub referenced: usize,
41
42 pub message: usize,
44
45 pub state: usize,
47}
48
49#[derive(Clone, Copy, Debug)]
54pub enum Trigger {
55 Receive,
57
58 Admin,
60}
61
62struct Candidate<Id> {
63 id: Id,
64 class: Class,
65}
66
67enum Class {
68 Dangling,
69 Referenced,
70 Own,
71 Live(Live),
72}
73
74struct Live {
75 state: bool,
76 redacted: bool,
77 owner_online: bool,
78 count: PduCount,
79}
80
81enum Partial {
85 Dangling,
86 Referenced,
87 Own,
88 Live {
89 state: bool,
90 redacted: bool,
91 count: PduCount,
92 server: OwnedServerName,
93 },
94}
95
96#[implement(super::Service)]
103#[tracing::instrument(
104 level = "debug"
105 skip_all,
106 fields(%room_id),
107)]
108pub async fn prune_forward_extremities(
109 &self,
110 room_id: &RoomId,
111 extremities: &mut Vec<OwnedEventId>,
112 goal: usize,
113 trigger: Trigger,
114) -> PruneSummary {
115 let candidates = self
116 .classify_extremities(room_id, extremities, trigger)
117 .await;
118
119 let (survivors, summary) = select(candidates, goal);
120 *extremities = survivors;
121
122 summary
123}
124
125#[implement(super::Service)]
126async fn classify_extremities(
127 &self,
128 room_id: &RoomId,
129 extremities: &[OwnedEventId],
130 trigger: Trigger,
131) -> Vec<Candidate<OwnedEventId>> {
132 let partials: Vec<(OwnedEventId, Partial)> = extremities
133 .iter()
134 .stream()
135 .broad_then(async |event_id| {
136 let partial = self
137 .classify_leaf(room_id, event_id, trigger)
138 .await;
139
140 (event_id.clone(), partial)
141 })
142 .collect()
143 .await;
144
145 let mut servers: Servers = partials
146 .iter()
147 .filter_map(|(_, partial)| match partial {
148 | Partial::Live { server, .. } => Some(server.clone()),
149 | _ => None,
150 })
151 .collect();
152
153 servers.sort_unstable();
154 servers.dedup();
155
156 let online: BTreeMap<OwnedServerName, bool> = servers
157 .into_iter()
158 .stream()
159 .broad_then(async |server| {
160 let verdict = self
161 .services
162 .federation
163 .should_attempt(&server)
164 .await;
165
166 (server, !matches!(verdict, ShouldAttempt::No { .. }))
167 })
168 .collect()
169 .await;
170
171 partials
172 .into_iter()
173 .map(|(id, partial)| {
174 let class = match partial {
175 | Partial::Dangling => Class::Dangling,
176 | Partial::Referenced => Class::Referenced,
177 | Partial::Own => Class::Own,
178 | Partial::Live { state, redacted, count, server } => Class::Live(Live {
179 state,
180 redacted,
181 owner_online: online.get(&server).copied().unwrap_or(false),
182 count,
183 }),
184 };
185
186 Candidate { id, class }
187 })
188 .collect()
189}
190
191#[implement(super::Service)]
192async fn classify_leaf(&self, room_id: &RoomId, event_id: &EventId, trigger: Trigger) -> Partial {
193 let Ok(count) = self
194 .services
195 .timeline
196 .get_pdu_count(event_id)
197 .await
198 else {
199 return Partial::Dangling;
200 };
201
202 let Ok(pdu) = self.services.timeline.get_pdu(event_id).await else {
203 return Partial::Dangling;
204 };
205
206 if matches!(trigger, Trigger::Admin)
207 && self
208 .services
209 .pdu_metadata
210 .is_event_referenced(room_id, event_id)
211 .await
212 {
213 return Partial::Referenced;
214 }
215
216 let server = pdu.sender().server_name();
217 if server == self.services.globals.server_name() {
218 return Partial::Own;
219 }
220
221 Partial::Live {
222 state: pdu.state_key().is_some(),
223 redacted: pdu.is_redacted(),
224 count,
225 server: server.to_owned(),
226 }
227}
228
229fn select<Id>(candidates: Vec<Candidate<Id>>, goal: usize) -> (Vec<Id>, PruneSummary) {
234 let before = candidates.len();
235
236 let mut own = Vec::new();
237 let mut live = Vec::new();
238 let mut dangling = Vec::new();
239 let mut referenced = Vec::new();
240
241 for Candidate { id, class } in candidates {
242 match class {
243 | Class::Own => own.push(id),
244 | Class::Dangling => dangling.push(id),
245 | Class::Referenced => referenced.push(id),
246 | Class::Live(scored) => live.push((scored, id)),
247 }
248 }
249
250 live.sort_unstable_by(|(a, _), (b, _)| drop_order(a, b));
251
252 let floor = usize::from(own.is_empty());
255 let drop_live = goal.min(live.len().saturating_sub(floor));
256
257 let message = live[..drop_live]
258 .iter()
259 .filter(|(scored, _)| !scored.state)
260 .count();
261
262 let state = drop_live.saturating_sub(message);
263
264 let mut survivors: Vec<Id> = own
265 .into_iter()
266 .chain(live.into_iter().skip(drop_live).map(|(_, id)| id))
267 .collect();
268
269 if survivors.is_empty()
272 && let Some(id) = referenced.pop().or_else(|| dangling.pop())
273 {
274 survivors.push(id);
275 }
276
277 let summary = PruneSummary {
278 before,
279 after: survivors.len(),
280 dangling: dangling.len(),
281 referenced: referenced.len(),
282 message,
283 state,
284 };
285
286 (survivors, summary)
287}
288
289fn drop_order(a: &Live, b: &Live) -> Ordering {
292 a.state
293 .cmp(&b.state)
294 .then(b.redacted.cmp(&a.redacted))
295 .then(b.owner_online.cmp(&a.owner_online))
296 .then(a.count.cmp(&b.count))
297}
298
299pub(crate) fn prune_goal(len: usize, max: usize, emergency: usize, batch: usize) -> usize {
305 let emergency = emergency.max(max);
307
308 len.saturating_sub(emergency)
309 .max(len.saturating_sub(max).min(batch))
310}
311
312#[cfg(test)]
313mod tests {
314 use super::*;
315
316 fn live(
317 id: u32,
318 state: bool,
319 redacted: bool,
320 owner_online: bool,
321 count: u64,
322 ) -> Candidate<u32> {
323 Candidate {
324 id,
325 class: Class::Live(Live {
326 state,
327 redacted,
328 owner_online,
329 count: PduCount::Normal(count),
330 }),
331 }
332 }
333
334 fn own(id: u32) -> Candidate<u32> { Candidate { id, class: Class::Own } }
335
336 fn dangling(id: u32) -> Candidate<u32> { Candidate { id, class: Class::Dangling } }
337
338 fn referenced(id: u32) -> Candidate<u32> { Candidate { id, class: Class::Referenced } }
339
340 #[test]
341 fn message_leaves_drop_before_state() {
342 let candidates = vec![
343 live(1, true, false, false, 1),
344 live(2, false, false, false, 5),
345 live(3, false, false, false, 6),
346 ];
347
348 let (survivors, summary) = select(candidates, 2);
349
350 assert_eq!(survivors, vec![1]);
351 assert_eq!(summary.message, 2);
352 assert_eq!(summary.state, 0);
353 assert_eq!(summary.after, 1);
354 assert_eq!(summary.before, 3);
355 }
356
357 #[test]
358 fn state_leaf_outlives_redacted_message() {
359 let candidates = vec![live(1, true, true, true, 1), live(2, false, false, false, 9)];
362
363 let (survivors, _) = select(candidates, 1);
364
365 assert_eq!(survivors, vec![1]);
366 }
367
368 #[test]
369 fn redacted_drops_before_intact() {
370 let candidates = vec![live(1, false, false, false, 5), live(2, false, true, false, 5)];
371
372 let (survivors, _) = select(candidates, 1);
373
374 assert_eq!(survivors, vec![1]);
375 }
376
377 #[test]
378 fn online_owner_drops_before_offline() {
379 let candidates = vec![live(1, false, false, false, 5), live(2, false, false, true, 5)];
380
381 let (survivors, _) = select(candidates, 1);
382
383 assert_eq!(survivors, vec![1]);
384 }
385
386 #[test]
387 fn oldest_drops_first() {
388 let candidates = vec![live(1, false, false, false, 9), live(2, false, false, false, 3)];
389
390 let (survivors, _) = select(candidates, 1);
391
392 assert_eq!(survivors, vec![1]);
393 }
394
395 #[test]
396 fn own_leaf_never_dropped() {
397 let candidates =
398 vec![own(1), live(2, false, false, false, 5), live(3, false, false, false, 6)];
399
400 let (survivors, summary) = select(candidates, 10);
401
402 assert_eq!(survivors, vec![1]);
403 assert_eq!(summary.after, 1);
404 assert_eq!(summary.message, 2);
405 }
406
407 #[test]
408 fn dangling_swept_free_and_uncounted() {
409 let candidates = vec![
410 dangling(1),
411 dangling(2),
412 dangling(3),
413 live(4, false, false, false, 5),
414 live(5, false, false, false, 6),
415 ];
416
417 let (survivors, summary) = select(candidates, 1);
418
419 assert_eq!(survivors, vec![5]);
420 assert_eq!(summary.dangling, 3);
421 assert_eq!(summary.message, 1);
422 assert_eq!(summary.after, 1);
423 }
424
425 #[test]
426 fn referenced_swept_free_on_admin() {
427 let candidates = vec![referenced(1), referenced(2), live(3, false, false, false, 5)];
428
429 let (survivors, summary) = select(candidates, 0);
430
431 assert_eq!(survivors, vec![3]);
432 assert_eq!(summary.referenced, 2);
433 assert_eq!(summary.after, 1);
434 }
435
436 #[test]
437 fn floor_keeps_one_live_when_all_would_drop() {
438 let candidates = vec![
439 live(1, false, false, false, 5),
440 live(2, false, false, false, 6),
441 live(3, false, false, false, 7),
442 ];
443
444 let (survivors, summary) = select(candidates, 10);
445
446 assert_eq!(survivors, vec![3]);
447 assert_eq!(summary.message, 2);
448 }
449
450 #[test]
451 fn floor_keeps_one_dangling_when_nothing_else() {
452 let candidates = vec![dangling(1), dangling(2), dangling(3)];
453
454 let (survivors, summary) = select(candidates, 10);
455
456 assert_eq!(survivors.len(), 1);
457 assert_eq!(summary.dangling, 2);
458 }
459
460 #[test]
461 fn floor_prefers_referenced_over_dangling() {
462 let candidates = vec![dangling(1), referenced(2)];
463
464 let (survivors, summary) = select(candidates, 10);
465
466 assert_eq!(survivors, vec![2]);
467 assert_eq!(summary.referenced, 0);
468 assert_eq!(summary.dangling, 1);
469 }
470
471 #[test]
472 fn pace_emergency_cut_above_bound() {
473 assert_eq!(prune_goal(1060, 60, 256, 32), 804);
474 }
475
476 #[test]
477 fn pace_batch_between_cap_and_bound() {
478 assert_eq!(prune_goal(257, 60, 256, 32), 32);
479 }
480
481 #[test]
482 fn pace_one_over_cap_drops_one() {
483 assert_eq!(prune_goal(61, 60, 256, 32), 1);
484 }
485
486 #[test]
487 fn pace_at_cap_drops_none() {
488 assert_eq!(prune_goal(60, 60, 256, 32), 0);
489 }
490
491 #[test]
492 fn pace_emergency_at_or_below_cap_is_unpaced() {
493 assert_eq!(prune_goal(100, 60, 0, 32), 40);
495 assert_eq!(prune_goal(100, 60, 60, 32), 40);
496 }
497
498 #[test]
499 fn pace_zero_batch_parks_at_emergency() {
500 assert_eq!(prune_goal(300, 60, 256, 0), 44);
501 assert_eq!(prune_goal(256, 60, 256, 0), 0);
502 }
503
504 #[test]
505 fn pace_monotone_across_emergency_bound() {
506 assert_eq!(prune_goal(256, 60, 256, 32), 32);
509 }
510}