use super::wait_queue::WaitQueue;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, Poll};
#[repr(align(64))]
pub struct Barrier {
n: usize,
count: AtomicUsize,
generation: AtomicUsize,
wait: WaitQueue,
}
impl Barrier {
#[must_use]
pub const fn new(n: usize) -> Self {
Self {
n: if n == 0 { 1 } else { n },
count: AtomicUsize::new(0),
generation: AtomicUsize::new(0),
wait: WaitQueue::new(),
}
}
#[inline(always)]
pub async fn wait(&self) -> BarrierWaitResult {
let observed_gen = self.generation.load(Ordering::Acquire);
let arrived = self.count.fetch_add(1, Ordering::AcqRel) + 1;
if arrived == self.n {
self.count.store(0, Ordering::Relaxed);
self.generation.fetch_add(1, Ordering::Release);
self.wait.wake_all();
return BarrierWaitResult { is_leader: true };
}
std::future::poll_fn(|cx| self.poll_generation_advanced(cx, observed_gen)).await;
BarrierWaitResult { is_leader: false }
}
#[inline(always)]
fn poll_generation_advanced(&self, cx: &Context<'_>, observed_gen: usize) -> Poll<()> {
if self.generation.load(Ordering::Acquire) != observed_gen {
return Poll::Ready(());
}
let token = self.wait.register(cx.waker());
if self.generation.load(Ordering::Acquire) != observed_gen {
self.wait.cancel(token);
return Poll::Ready(());
}
Poll::Pending
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(align(64))]
pub struct BarrierWaitResult {
is_leader: bool,
}
impl BarrierWaitResult {
#[must_use]
pub const fn is_leader(&self) -> bool {
self.is_leader
}
}