use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use futures_util::task::{waker_ref, ArcWake, AtomicWaker};
struct WeakWakerInner {
waker: AtomicWaker,
}
impl ArcWake for WeakWakerInner {
fn wake_by_ref(arc_self: &Arc<Self>) {
arc_self.waker.wake();
}
}
struct WeakWaker {
inner: Arc<WeakWakerInner>,
}
impl Drop for WeakWaker {
fn drop(&mut self) {
self.inner.waker.take();
}
}
pub struct WeakWakerFuture<F: Future> {
fut: F,
weak_waker: WeakWaker,
}
impl<F: Future> WeakWakerFuture<F> {
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub fn new(fut: F) -> WeakWakerFuture<F> {
WeakWakerFuture {
fut,
weak_waker: WeakWaker {
inner: Arc::new(WeakWakerInner {
waker: AtomicWaker::new(),
}),
},
}
}
}
impl<F: Future> Future for WeakWakerFuture<F> {
type Output = F::Output;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> std::task::Poll<Self::Output> {
let m = unsafe { self.get_unchecked_mut() };
m.weak_waker.inner.waker.register(cx.waker());
unsafe {
Pin::new_unchecked(&mut m.fut)
.poll(&mut Context::from_waker(&waker_ref(&m.weak_waker.inner)))
}
}
}
#[cfg(test)]
mod tests {
use std::future::Future;
use std::pin::pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::Context;
use super::*;
use futures::task::waker;
use futures_util::task::ArcWake;
struct TestWaker {
waked: AtomicBool,
}
impl ArcWake for TestWaker {
fn wake_by_ref(s: &Arc<Self>) {
s.waked.store(true, Ordering::SeqCst);
}
}
#[test]
fn test_mpsc_queue_weak_waker_drop_correctly() {
let (tx, mut rx) = futures::channel::mpsc::unbounded::<()>();
let mut fut = WeakWakerFuture::new(futures::StreamExt::next(&mut rx));
let pinned_fut = Pin::new(&mut fut);
let base_waker = Arc::new(TestWaker {
waked: AtomicBool::new(false),
});
assert!(pinned_fut
.poll(&mut Context::from_waker(&waker(base_waker.clone())))
.is_pending());
assert_eq!(Arc::strong_count(&base_waker), 2);
drop(fut);
assert!(!base_waker.waked.load(Ordering::SeqCst));
assert_eq!(Arc::strong_count(&base_waker), 1);
tx.unbounded_send(()).unwrap();
}
#[test]
fn test_mpsc_queue_weak_waker_smoke() {
let (tx, mut rx) = futures::channel::mpsc::unbounded::<()>();
let mut pinned_fut = pin!(WeakWakerFuture::new(futures::StreamExt::next(&mut rx)));
let base_waker = Arc::new(TestWaker {
waked: AtomicBool::new(false),
});
assert!(pinned_fut
.as_mut()
.poll(&mut Context::from_waker(&waker(base_waker.clone())))
.is_pending());
assert_eq!(Arc::strong_count(&base_waker), 2);
tx.unbounded_send(()).unwrap();
assert!(base_waker.waked.load(Ordering::SeqCst));
assert!(pinned_fut
.poll(&mut Context::from_waker(&waker(base_waker.clone())))
.is_ready());
}
}