dtact_util/sync/barrier.rs
1//! Async rendezvous point for a fixed number of tasks.
2
3use super::wait_queue::WaitQueue;
4use std::sync::atomic::{AtomicUsize, Ordering};
5use std::task::{Context, Poll};
6
7/// A barrier that releases all `n` participants together once every one
8/// of them has called [`Barrier::wait`].
9#[repr(align(64))]
10pub struct Barrier {
11 n: usize,
12 /// Arrivals for the current generation; reset to 0 by whichever
13 /// arrival completes the generation.
14 count: AtomicUsize,
15 /// Bumped every time the barrier releases, so waiters parked in an
16 /// old generation know to stop waiting (and so a barrier can be
17 /// reused indefinitely, matching `tokio::sync::Barrier`).
18 generation: AtomicUsize,
19 wait: WaitQueue,
20}
21
22impl Barrier {
23 /// Create a barrier requiring `n` participants per generation. `n ==
24 /// 0` behaves like `n == 1` (a single `wait()` call releases
25 /// immediately as the leader) rather than a barrier no `wait()` call
26 /// could ever complete.
27 #[must_use]
28 pub const fn new(n: usize) -> Self {
29 Self {
30 n: if n == 0 { 1 } else { n },
31 count: AtomicUsize::new(0),
32 generation: AtomicUsize::new(0),
33 wait: WaitQueue::new(),
34 }
35 }
36
37 /// Wait for every one of this barrier's `n` participants to arrive.
38 /// Exactly one of the `n` calls in each generation gets back a
39 /// [`BarrierWaitResult`] with [`is_leader`](BarrierWaitResult::is_leader)
40 /// `true`; the rest get `false`. All `n` calls return together.
41 #[inline(always)]
42 pub async fn wait(&self) -> BarrierWaitResult {
43 let observed_gen = self.generation.load(Ordering::Acquire);
44 let arrived = self.count.fetch_add(1, Ordering::AcqRel) + 1;
45
46 if arrived == self.n {
47 self.count.store(0, Ordering::Relaxed);
48 self.generation.fetch_add(1, Ordering::Release);
49 self.wait.wake_all();
50 return BarrierWaitResult { is_leader: true };
51 }
52
53 std::future::poll_fn(|cx| self.poll_generation_advanced(cx, observed_gen)).await;
54 BarrierWaitResult { is_leader: false }
55 }
56
57 #[inline(always)]
58 fn poll_generation_advanced(&self, cx: &Context<'_>, observed_gen: usize) -> Poll<()> {
59 if self.generation.load(Ordering::Acquire) != observed_gen {
60 return Poll::Ready(());
61 }
62 let token = self.wait.register(cx.waker());
63 if self.generation.load(Ordering::Acquire) != observed_gen {
64 self.wait.cancel(token);
65 return Poll::Ready(());
66 }
67 Poll::Pending
68 }
69}
70
71/// Returned by [`Barrier::wait`]; tells the caller whether it was the one
72/// call (per generation) whose arrival released everyone else. Purely
73/// informational — every participant proceeds regardless.
74#[derive(Debug, Clone, Copy, PartialEq, Eq)]
75#[repr(align(64))]
76pub struct BarrierWaitResult {
77 is_leader: bool,
78}
79
80impl BarrierWaitResult {
81 /// `true` for exactly one of the `n` [`Barrier::wait`] calls per
82 /// generation — the one whose arrival completed it.
83 #[must_use]
84 pub const fn is_leader(&self) -> bool {
85 self.is_leader
86 }
87}