use std::sync::atomic::AtomicU32;
use std::sync::atomic::Ordering;
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 CountdownState {
state: AtomicU32,
waiters: Mutex<WakerSet>,
}
impl CountdownState {
pub const fn new(count: u32) -> Self {
Self {
state: AtomicU32::new(count),
waiters: Mutex::new(WakerSet::new()),
}
}
pub fn state(&self) -> u32 {
self.state.load(Ordering::Acquire)
}
fn cas_state(&self, current: u32, new: u32) -> Result<(), u32> {
self.state
.compare_exchange_weak(current, new, Ordering::Release, Ordering::Relaxed)
.map(|_| ())
}
pub fn wake_all(&self) {
let wakers = {
let mut waiters = self.waiters.lock();
waiters.take_all()
};
wake_all(wakers);
}
pub fn poll_wait(&self, token: &mut Option<WakerToken>, cx: &mut Context<'_>) -> Poll<()> {
if self.try_wait().is_ok() {
*token = None;
return Poll::Ready(());
}
let mut waiters = self.waiters.lock();
if self.state() == 0 {
*token = None;
return Poll::Ready(());
}
let retired_waker = waiters.register(token, cx.waker());
drop(waiters);
drop(retired_waker);
Poll::Pending
}
#[inline]
pub fn unregister(&self, token: &mut Option<WakerToken>) {
if token.is_none() {
return;
}
let mut waiters = self.waiters.lock();
if self.state() == 0 {
*token = None;
return;
}
let removed_waker = waiters.unregister(token);
drop(waiters);
drop(removed_waker);
}
pub fn try_wait(&self) -> Result<(), u32> {
match self.state() {
0 => Ok(()),
s => Err(s),
}
}
pub fn decrement(&self, n: u32) -> bool {
let mut cnt = self.state();
loop {
if cnt == 0 {
return false;
}
let new_cnt = cnt.saturating_sub(n);
match self.cas_state(cnt, new_cnt) {
Ok(_) => return new_cnt == 0,
Err(x) => cnt = x,
}
}
}
}