use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
use crate::internal::mutex::Mutex;
use crate::internal::wake_all;
use crate::internal::wakerset::WakerSet;
use crate::internal::wakerset::WakerToken;
#[derive(Debug)]
pub struct Barrier {
n: u32,
state: Mutex<BarrierState>,
}
struct BarrierState {
arrived: u32,
generation: usize,
waiters: WakerSet,
}
impl fmt::Debug for BarrierState {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierState")
.field("arrived", &self.arrived)
.field("generation", &self.generation)
.finish_non_exhaustive()
}
}
pub struct BarrierWaitResult(bool);
impl fmt::Debug for BarrierWaitResult {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierWaitResult")
.field("is_leader", &self.is_leader())
.finish()
}
}
impl BarrierWaitResult {
#[must_use]
pub fn is_leader(&self) -> bool {
self.0
}
}
impl Barrier {
pub fn new(n: u32) -> Self {
let n = if n > 0 { n } else { 1 };
Self {
n,
state: Mutex::new(BarrierState {
arrived: 0,
generation: 0,
waiters: WakerSet::with_capacity((n - 1) as usize),
}),
}
}
pub async fn wait(&self) -> BarrierWaitResult {
let generation = {
let mut state = self.state.lock();
let generation = state.generation;
state.arrived += 1;
if state.arrived == self.n {
state.arrived = 0;
state.generation += 1;
let wakers = state.waiters.drain();
drop(state);
wake_all(wakers);
return BarrierWaitResult(true);
}
generation
};
let fut = BarrierWait {
token: None,
generation,
barrier: self,
};
fut.await;
BarrierWaitResult(false)
}
}
#[must_use = "futures do nothing unless you `.await` or poll them"]
struct BarrierWait<'a> {
token: Option<WakerToken>,
generation: usize,
barrier: &'a Barrier,
}
impl fmt::Debug for BarrierWait<'_> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("BarrierWait")
.field("generation", &self.generation)
.finish_non_exhaustive()
}
}
impl Future for BarrierWait<'_> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self {
token,
generation,
barrier,
} = self.get_mut();
let mut state = barrier.state.lock();
if *generation < state.generation {
*token = None;
return Poll::Ready(());
}
let retired_waker = state.waiters.register(token, cx.waker());
drop(state);
drop(retired_waker);
Poll::Pending
}
}
impl Drop for BarrierWait<'_> {
fn drop(&mut self) {
if self.token.is_none() {
return;
}
let mut state = self.barrier.state.lock();
if self.generation != state.generation {
self.token = None;
return;
}
let removed_waker = state.waiters.unregister(&mut self.token);
drop(state);
drop(removed_waker);
}
}