use futures_util::lock::Mutex;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
#[derive(Clone)]
pub struct WaitGroup {
inner: Arc<Inner>,
}
struct Inner {
count: Mutex<isize>,
waker: Mutex<Option<Waker>>,
}
impl WaitGroup {
pub fn new() -> WaitGroup {
WaitGroup {
inner: Arc::new(Inner {
count: Mutex::new(0),
waker: Mutex::new(None),
}),
}
}
pub async fn add(&self, delta: isize) {
if delta <= 0 {
panic!("The argument `delta` of wait group `add` must be a positive number");
}
let mut count = self.inner.count.lock().await;
*count += delta;
if *count >= isize::max_value() / 2 {
panic!("wait group count is too large");
}
}
pub async fn done(&self) {
let mut count = self.inner.count.lock().await;
*count -= 1;
if *count <= 0 {
if let Some(waker) = &*self.inner.waker.lock().await {
waker.clone().wake();
}
}
}
pub async fn count(&self) -> isize {
*self.inner.count.lock().await
}
}
impl Future for WaitGroup {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let mut count = self.inner.count.lock();
let pin_count = Pin::new(&mut count);
if let Poll::Ready(count) = pin_count.poll(cx) {
if *count <= 0 {
return Poll::Ready(());
}
}
drop(count);
let mut waker = self.inner.waker.lock();
let pin_waker = Pin::new(&mut waker);
if let Poll::Ready(mut waker) = pin_waker.poll(cx) {
*waker = Some(cx.waker().clone());
}
Poll::Pending
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
#[should_panic]
async fn add_zero() {
let wg = WaitGroup::new();
wg.add(0).await;
}
#[tokio::test]
#[should_panic]
async fn add_neg_one() {
let wg = WaitGroup::new();
wg.add(-1).await;
}
#[tokio::test]
#[should_panic]
async fn add_very_max() {
let wg = WaitGroup::new();
wg.add(isize::max_value()).await;
}
#[tokio::test]
async fn add() {
let wg = WaitGroup::new();
wg.add(1).await;
wg.add(10).await;
assert_eq!(*wg.inner.count.lock().await, 11);
}
#[tokio::test]
async fn done() {
let wg = WaitGroup::new();
wg.done().await;
wg.done().await;
assert_eq!(*wg.inner.count.lock().await, -2);
}
#[tokio::test]
async fn count() {
let wg = WaitGroup::new();
assert_eq!(wg.count().await, 0);
wg.add(10).await;
assert_eq!(wg.count().await, 10);
wg.done().await;
assert_eq!(wg.count().await, 9);
}
}