1use std::{
2 collections::{HashMap, VecDeque},
3 iter::once,
4 ops::Deref,
5};
6
7use futures::{
8 Stream, StreamExt, TryStreamExt,
9 stream::{FuturesUnordered, try_unfold},
10};
11use ruma::{EventId, OwnedEventId};
12use tuwunel_core::{
13 Result, implement, is_equal_to,
14 matrix::{event_id::RandomState, pdu::AuthEvents},
15 smallvec::SmallVec,
16 utils::{
17 BoolExt,
18 math::expect_into,
19 stream::{IterStream, automatic_width},
20 },
21};
22
23use super::super::{FetchEvent, fetch_event::AuthRefs};
24
25struct Global<Fut: Future + Send> {
26 subgraph: Subgraph,
27 todo: Todo<Fut>,
28 locals: Locals,
29 waiters: Waiters,
30 ready: Ready,
31 deferred: Deferred,
32 parked: usize,
33}
34
35struct Context<'a> {
36 subgraph: &'a mut Subgraph,
37 waiters: &'a mut Waiters,
38 ready: &'a mut Ready,
39 parked: &'a mut usize,
40 outputs: &'a mut Path,
41}
42
43#[derive(Debug, Default)]
44struct Local {
45 path: Path,
46 stack: Stack,
47 marked: usize,
48}
49
50#[derive(Debug)]
51struct Wake {
52 event_id: OwnedEventId,
53 locals: Waiting,
54 result: Resolution,
55}
56
57#[derive(Debug)]
58enum Evaluation {
59 Continue,
60 Fetch(OwnedEventId),
61 Park,
62}
63
64#[derive(Clone, Copy, Debug)]
65enum Resolution {
66 Dead,
67 Subgraph,
68}
69
70#[derive(Clone, Copy, Debug)]
71enum Substate {
72 Conflicted,
73 Pending(LocalId),
76 Dead,
77 Subgraph,
78}
79
80type Todo<Fut> = FuturesUnordered<Fut>;
81type Subgraph = HashMap<OwnedEventId, Substate, RandomState>;
82type Locals = Vec<Local>;
83type LocalId = u32;
84type Waiters = HashMap<OwnedEventId, Waiting, RandomState>;
85type Waiting = SmallVec<[usize; WAITING_INLINE]>;
86type Ready = Vec<Wake>;
87type Deferred = VecDeque<(usize, OwnedEventId)>;
88type Path = SmallVec<[OwnedEventId; PATH_INLINE]>;
89type Stack = SmallVec<[Frame; STACK_INLINE]>;
90type Frame = AuthEvents;
91
92const PATH_INLINE: usize = 4;
93const STACK_INLINE: usize = 4;
94const WAITING_INLINE: usize = 1;
95const CAPACITY_MULTIPLIER: usize = 4;
96
97#[tracing::instrument(
98 name = "subgraph_dfs",
99 level = "debug",
100 skip_all,
101 fields(
102 starting_events = %conflicted_set.len(),
103 )
104)]
105pub(super) fn conflicted_subgraph_dfs(
106 conflicted_set: &Vec<&OwnedEventId>,
107 fetch: impl FetchEvent,
108) -> impl Stream<Item = Result<OwnedEventId>> + Send {
109 let initial_capacity = conflicted_set
110 .len()
111 .saturating_mul(CAPACITY_MULTIPLIER);
112
113 let seeds = || conflicted_set.iter().map(Deref::deref).cloned();
114
115 let mut subgraph = Subgraph::with_capacity_and_hasher(initial_capacity, RandomState);
118
119 subgraph.extend(seeds().map(|event_id| (event_id, Substate::Conflicted)));
120
121 let state = Global {
122 subgraph,
123 todo: Todo::new(),
124 locals: Locals::with_capacity(conflicted_set.len()),
125 waiters: Waiters::with_hasher(RandomState),
126 ready: Ready::new(),
127 deferred: Deferred::new(),
128 parked: 0,
129 };
130
131 #[expect(closure_returning_async_block)]
134 try_unfold((seeds(), state), move |(mut inputs, mut state)| async move {
135 let width = automatic_width();
136
137 debug_assert!(
138 state.todo.len() <= width,
139 "Excessive in-flight conflicted-subgraph fetches"
140 );
141
142 while state.todo.len() < width {
143 if let Some((id, event_id)) = state.deferred.pop_front() {
144 state.todo.push(fetch_auth(id, event_id, fetch));
145 continue;
146 }
147
148 let Some(seed) = inputs.next() else {
149 break;
150 };
151
152 let id = state.locals.len();
153
154 state.locals.push(Local::default());
155 state.todo.push(fetch_auth(id, seed, fetch));
156 }
157
158 let Some((id, event_id, event)) = state.todo.next().await else {
159 debug_assert!(state.waiters.is_empty(), "Unresolved conflicted-subgraph waiters");
160 debug_assert!(state.ready.is_empty(), "Undrained conflicted-subgraph wakes");
161 debug_assert!(state.deferred.is_empty(), "Deferred conflicted-subgraph fetches");
162 debug_assert_eq!(state.parked, 0, "Parked conflicted-subgraph walkers");
163
164 let output: Result<_> = Ok(None);
165
166 return output;
167 };
168
169 while state.todo.len() < width
170 && let Some((deferred_id, deferred_event_id)) = state.deferred.pop_front()
171 {
172 state
173 .todo
174 .push(fetch_auth(deferred_id, deferred_event_id, fetch));
175 }
176
177 let mut outputs = Path::new();
178
179 if let Some(next_id) = process_fetch(&mut state, id, event_id, event, &mut outputs)? {
180 if state.todo.len() < width {
181 state.todo.push(fetch_auth(id, next_id, fetch));
182 } else {
183 state.deferred.push_back((id, next_id));
184 }
185 }
186
187 while let Some(Wake { event_id, locals, result }) = state.ready.pop() {
188 for id in locals {
189 if let Some(next_id) = resume(&mut state, id, &event_id, result, &mut outputs) {
190 if state.todo.len() < width {
191 state.todo.push(fetch_auth(id, next_id, fetch));
192 } else {
193 state.deferred.push_back((id, next_id));
194 }
195 }
196 }
197 }
198
199 let output: Result<_> = Ok(Some((outputs, (inputs, state))));
200
201 output
202 })
203 .map_ok(|outputs| outputs.into_iter().map(Ok).stream())
204 .try_flatten()
205}
206
207async fn fetch_auth(
208 id: usize,
209 event_id: OwnedEventId,
210 fetch: impl FetchEvent,
211) -> (usize, OwnedEventId, Result<AuthRefs>) {
212 let event = fetch.get::<AuthRefs>(&event_id).await;
213
214 (id, event_id, event)
215}
216
217fn process_fetch<Fut>(
218 state: &mut Global<Fut>,
219 id: usize,
220 event_id: OwnedEventId,
221 event: Result<AuthRefs>,
222 outputs: &mut Path,
223) -> Result<Option<OwnedEventId>>
224where
225 Fut: Future + Send,
226{
227 match event {
228 | Ok(event) => {
229 let local = &mut state.locals[id];
230
231 local.path.push(event_id);
232 local.stack.push(event.auth_events);
233 },
234 | Err(error) if error.is_not_found() => {
235 let Global { subgraph, waiters, ready, parked, .. } = state;
236 let mut context = Context {
237 subgraph,
238 waiters,
239 ready,
240 parked,
241 outputs,
242 };
243
244 complete_pending(&mut context, event_id, Resolution::Dead);
245 },
246 | Err(error) => return Err(error),
247 }
248
249 Ok(advance(state, id, outputs))
250}
251
252fn resume<Fut: Future + Send>(
253 state: &mut Global<Fut>,
254 id: usize,
255 event_id: &EventId,
256 result: Resolution,
257 outputs: &mut Path,
258) -> Option<OwnedEventId> {
259 if matches!(result, Resolution::Subgraph) {
260 let Global {
261 subgraph, locals, waiters, ready, parked, ..
262 } = state;
263
264 let mut context = Context {
265 subgraph,
266 waiters,
267 ready,
268 parked,
269 outputs,
270 };
271
272 locals[id].insert_path(&mut context, event_id);
273 }
274
275 advance(state, id, outputs)
276}
277
278fn advance<Fut: Future + Send>(
279 state: &mut Global<Fut>,
280 id: usize,
281 outputs: &mut Path,
282) -> Option<OwnedEventId> {
283 let Global {
284 subgraph, locals, waiters, ready, parked, ..
285 } = state;
286
287 let local = &mut locals[id];
288 let mut context = Context {
289 subgraph,
290 waiters,
291 ready,
292 parked,
293 outputs,
294 };
295
296 while let Some(event_id) = local.pop(&mut context) {
297 match local.eval(id, &mut context, event_id) {
298 | Evaluation::Continue => {},
299 | Evaluation::Fetch(event_id) => return Some(event_id),
300 | Evaluation::Park => return None,
301 }
302 }
303
304 if local.stack.is_empty() {
305 *local = Local::default();
306 }
307
308 None
309}
310
311#[implement(Local)]
312fn pop(&mut self, context: &mut Context<'_>) -> Option<OwnedEventId> {
313 while self.stack.last().is_some_and(Frame::is_empty) {
314 self.stack.pop();
315
316 if let Some(event_id) = self.path.pop() {
317 complete_pending(context, event_id, Resolution::Dead);
318 }
319 }
320
321 self.marked = self.marked.min(self.path.len());
322 self.stack.last_mut().and_then(Frame::pop)
323}
324
325#[implement(Local)]
326#[tracing::instrument(
327 name = "descent",
328 level = "trace",
329 skip_all,
330 fields(
331 s = ?context
332 .subgraph
333 .values()
334 .fold((0_u64, 0_u64, 0_u64, 0_u64), |(pending, dead, conflicted, subgraph), state| {
335 match state {
336 | Substate::Pending(_) =>
337 (pending.saturating_add(1), dead, conflicted, subgraph),
338 | Substate::Dead => (pending, dead.saturating_add(1), conflicted, subgraph),
339 | Substate::Conflicted => {
340 (pending, dead, conflicted.saturating_add(1), subgraph)
341 },
342 | Substate::Subgraph => {
343 (pending, dead, conflicted, subgraph.saturating_add(1))
344 },
345 }
346 }),
347
348 %event_id,
349 path = self.path.len(),
350 stack = self.stack.iter().flatten().count(),
351 )
352)]
353fn eval(&mut self, id: usize, context: &mut Context<'_>, event_id: OwnedEventId) -> Evaluation {
354 match context.subgraph.get(&event_id).copied() {
355 | Some(Substate::Subgraph) => {
356 self.insert_path(context, &event_id);
357 Evaluation::Continue
358 },
359 | Some(Substate::Dead) => Evaluation::Continue,
360 | Some(Substate::Pending(owner)) => {
361 if expect_into::<usize, _>(owner) == id {
362 return Evaluation::Continue;
363 }
364
365 context
366 .waiters
367 .entry(event_id)
368 .or_default()
369 .push(id);
370
371 *context.parked = context.parked.saturating_add(1);
372 Evaluation::Park
373 },
374 | Some(Substate::Conflicted) => {
375 self.insert_path(context, &event_id);
376
377 self.path
378 .first()
379 .is_some_and(is_equal_to!(&event_id))
380 .is_false()
381 .then_some(event_id)
382 .map_or(Evaluation::Continue, Evaluation::Fetch)
383 },
384 | None => {
385 context
386 .subgraph
387 .insert(event_id.clone(), Substate::Pending(expect_into(id)));
388
389 Evaluation::Fetch(event_id)
390 },
391 }
392}
393
394#[implement(Local)]
395fn insert_path(&mut self, context: &mut Context<'_>, event_id: &EventId) {
396 let Context {
397 subgraph,
398 waiters,
399 ready,
400 parked,
401 outputs,
402 } = context;
403
404 let inserted = self.path[self.marked..]
405 .iter()
406 .map(AsRef::as_ref)
407 .chain(once(event_id))
408 .filter(|event_id| insert_path_filter(subgraph, waiters, ready, parked, event_id))
409 .map(ToOwned::to_owned);
410
411 outputs.extend(inserted);
412 self.marked = self.path.len();
413}
414
415fn insert_path_filter(
416 subgraph: &mut Subgraph,
417 waiters: &mut Waiters,
418 ready: &mut Ready,
419 parked: &mut usize,
420 event_id: &EventId,
421) -> bool {
422 let Some(state) = subgraph.get_mut(event_id) else {
423 subgraph.insert(event_id.to_owned(), Substate::Subgraph);
424 return true;
425 };
426
427 if matches!(*state, Substate::Subgraph) {
428 return false;
429 }
430
431 let pending = matches!(*state, Substate::Pending(_));
432
433 debug_assert!(
434 !matches!(*state, Substate::Dead),
435 "Dead node inserted into conflicted subgraph"
436 );
437
438 *state = Substate::Subgraph;
439
440 if pending
441 && !waiters.is_empty()
442 && let Some(locals) = waiters.remove(event_id)
443 {
444 debug_assert!(*parked >= locals.len(), "Invalid parked walker count");
445 *parked = parked.saturating_sub(locals.len());
446 ready.push(Wake {
447 event_id: event_id.to_owned(),
448 locals,
449 result: Resolution::Subgraph,
450 });
451 }
452
453 true
454}
455
456fn complete_pending(context: &mut Context<'_>, event_id: OwnedEventId, result: Resolution) {
457 let Some(state) = context.subgraph.get_mut(&event_id) else {
458 return;
459 };
460
461 if !matches!(*state, Substate::Pending(_)) {
462 return;
463 }
464
465 *state = match result {
466 | Resolution::Dead => Substate::Dead,
467 | Resolution::Subgraph => Substate::Subgraph,
468 };
469
470 if context.waiters.is_empty() {
471 return;
472 }
473
474 let Some(locals) = context.waiters.remove(&event_id) else {
475 return;
476 };
477
478 debug_assert!(*context.parked >= locals.len(), "Invalid parked walker count");
479 *context.parked = context.parked.saturating_sub(locals.len());
480 context
481 .ready
482 .push(Wake { event_id, locals, result });
483}