Skip to main content

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}