use std::fmt;
use std::future::Future;
use std::future::IntoFuture;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::AtomicUsize;
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)]
struct State {
handles: AtomicUsize,
waiters: Mutex<WakerSet>,
}
impl State {
fn new() -> Self {
Self {
handles: AtomicUsize::new(1),
waiters: Mutex::new(WakerSet::new()),
}
}
fn register_handle(&self) {
self.handles.fetch_add(1, Ordering::Relaxed);
}
fn release_handle(&self) {
let previous = self.handles.fetch_sub(1, Ordering::Release);
debug_assert!(previous > 0, "a live handle must own one count");
if previous != 1 {
return;
}
let wakers = {
let mut waiters = self.waiters.lock();
waiters.take_all()
};
wake_all(wakers);
}
fn poll_wait(&self, token: &mut Option<WakerToken>, cx: &mut Context<'_>) -> Poll<()> {
if self.handles.load(Ordering::Acquire) == 0 {
*token = None;
return Poll::Ready(());
}
let mut waiters = self.waiters.lock();
if self.handles.load(Ordering::Acquire) == 0 {
*token = None;
return Poll::Ready(());
}
let retired_waker = waiters.register(token, cx.waker());
drop(waiters);
drop(retired_waker);
Poll::Pending
}
fn unregister(&self, token: &mut Option<WakerToken>) {
if token.is_none() {
return;
}
let mut waiters = self.waiters.lock();
if self.handles.load(Ordering::Acquire) == 0 {
*token = None;
return;
}
let removed_waker = waiters.unregister(token);
drop(waiters);
drop(removed_waker);
}
}
pub struct WaitGroup {
state: Option<Arc<State>>,
}
impl fmt::Debug for WaitGroup {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WaitGroup").finish_non_exhaustive()
}
}
impl Default for WaitGroup {
fn default() -> Self {
Self::new()
}
}
impl WaitGroup {
pub fn new() -> Self {
Self {
state: Some(Arc::new(State::new())),
}
}
}
impl Clone for WaitGroup {
fn clone(&self) -> Self {
let state = self
.state
.as_ref()
.expect("a live WaitGroup owns its state")
.clone();
state.register_handle();
Self { state: Some(state) }
}
}
impl Drop for WaitGroup {
fn drop(&mut self) {
if let Some(state) = self.state.take() {
state.release_handle();
}
}
}
impl IntoFuture for WaitGroup {
type Output = ();
type IntoFuture = Wait;
fn into_future(mut self) -> Self::IntoFuture {
let state = self.state.take().expect("a live WaitGroup owns its state");
state.release_handle();
Wait { token: None, state }
}
}
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct Wait {
token: Option<WakerToken>,
state: Arc<State>,
}
impl Clone for Wait {
fn clone(&self) -> Self {
Wait {
token: None,
state: self.state.clone(),
}
}
}
impl fmt::Debug for Wait {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Wait").finish_non_exhaustive()
}
}
impl Future for Wait {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let Self { token, state } = self.get_mut();
state.poll_wait(token, cx)
}
}
impl Drop for Wait {
fn drop(&mut self) {
self.state.unregister(&mut self.token);
}
}