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