polars_async/primitives/
wait_group.rs1use std::future::Future;
2use std::pin::Pin;
3use std::sync::Arc;
4use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
5use std::task::{Context, Poll, Waker};
6
7use parking_lot::Mutex;
8
9#[derive(Default, Debug)]
10struct WaitGroupInner {
11 waker: Mutex<Option<Waker>>,
12 token_count: AtomicUsize,
13 is_waiting: AtomicBool,
14}
15
16#[derive(Default)]
17pub struct WaitGroup {
18 inner: Arc<WaitGroupInner>,
19}
20
21impl WaitGroup {
22 pub fn token(&self) -> WaitToken {
24 self.inner.token_count.fetch_add(1, Ordering::Relaxed);
25 WaitToken {
26 inner: Arc::clone(&self.inner),
27 }
28 }
29
30 pub async fn wait(&self) {
35 let was_waiting = self.inner.is_waiting.swap(true, Ordering::Relaxed);
36 assert!(!was_waiting);
37 WaitGroupFuture { inner: &self.inner }.await
38 }
39}
40
41struct WaitGroupFuture<'a> {
42 inner: &'a Arc<WaitGroupInner>,
43}
44
45impl Future for WaitGroupFuture<'_> {
46 type Output = ();
47
48 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
49 if self.inner.token_count.load(Ordering::Acquire) == 0 {
50 return Poll::Ready(());
51 }
52
53 let mut waker_lock = self.inner.waker.lock();
55 if self.inner.token_count.load(Ordering::Acquire) == 0 {
56 return Poll::Ready(());
57 }
58
59 let waker = cx.waker().clone();
60 *waker_lock = Some(waker);
61 Poll::Pending
62 }
63}
64
65impl Drop for WaitGroupFuture<'_> {
66 fn drop(&mut self) {
67 self.inner.is_waiting.store(false, Ordering::Relaxed);
68 }
69}
70
71#[derive(Debug)]
72pub struct WaitToken {
73 inner: Arc<WaitGroupInner>,
74}
75
76impl Clone for WaitToken {
77 fn clone(&self) -> Self {
78 self.inner.token_count.fetch_add(1, Ordering::Relaxed);
79 Self {
80 inner: self.inner.clone(),
81 }
82 }
83}
84
85impl Drop for WaitToken {
86 fn drop(&mut self) {
87 if self.inner.token_count.fetch_sub(1, Ordering::Release) == 1 {
89 if let Some(w) = self.inner.waker.lock().take() {
90 w.wake();
91 }
92 }
93 }
94}