use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, PoisonError};
use std::task::{Context, Poll, Waker};
use futures_core::Stream;
pub type MsgStream<Msg> = Pin<Box<dyn Stream<Item = Msg> + Send>>;
pub struct Task {
cancel: Option<Arc<Cancel>>,
}
impl Task {
pub fn detach(mut self) {
self.cancel = None;
}
}
impl Drop for Task {
fn drop(&mut self) {
if let Some(cancel) = &self.cancel {
cancel.cancel();
}
}
}
pub struct Cancel {
cancelled: AtomicBool,
waker: Mutex<Option<Waker>>,
}
impl Cancel {
pub(crate) fn new() -> Self {
Self {
cancelled: AtomicBool::new(false),
waker: Mutex::new(None),
}
}
fn waker(&self) -> std::sync::MutexGuard<'_, Option<Waker>> {
self.waker.lock().unwrap_or_else(PoisonError::into_inner)
}
pub(crate) fn cancel(&self) {
self.cancelled.store(true, Ordering::SeqCst);
if let Some(waker) = self.waker().take() {
waker.wake();
}
}
pub fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::SeqCst)
}
pub fn cancelled(self: &Arc<Self>) -> Cancelled {
Cancelled {
cancel: Arc::clone(self),
}
}
}
pub struct Cancelled {
cancel: Arc<Cancel>,
}
impl Future for Cancelled {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
if self.cancel.is_cancelled() {
return Poll::Ready(());
}
*self.cancel.waker() = Some(cx.waker().clone());
if self.cancel.is_cancelled() {
Poll::Ready(())
} else {
Poll::Pending
}
}
}
#[derive(Clone, Default)]
pub struct PersistTracker {
state: Arc<PersistState>,
}
#[derive(Default)]
struct PersistState {
count: AtomicUsize,
waker: Mutex<Option<Waker>>,
}
impl PersistTracker {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn guard(&self) -> PersistGuard {
self.state.count.fetch_add(1, Ordering::SeqCst);
PersistGuard {
state: Arc::clone(&self.state),
}
}
pub fn is_idle(&self) -> bool {
self.state.count.load(Ordering::SeqCst) == 0
}
pub fn wait(&self) -> PersistsDone {
PersistsDone {
state: Arc::clone(&self.state),
}
}
}
impl PersistState {
fn waker(&self) -> std::sync::MutexGuard<'_, Option<Waker>> {
self.waker.lock().unwrap_or_else(PoisonError::into_inner)
}
}
pub(crate) struct PersistGuard {
state: Arc<PersistState>,
}
impl Drop for PersistGuard {
fn drop(&mut self) {
if self.state.count.fetch_sub(1, Ordering::SeqCst) == 1
&& let Some(waker) = self.state.waker().take()
{
waker.wake();
}
}
}
pub struct PersistsDone {
state: Arc<PersistState>,
}
impl Future for PersistsDone {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
if self.state.count.load(Ordering::SeqCst) == 0 {
return Poll::Ready(());
}
*self.state.waker() = Some(cx.waker().clone());
if self.state.count.load(Ordering::SeqCst) == 0 {
Poll::Ready(())
} else {
Poll::Pending
}
}
}
pub enum Effect<Msg> {
Spawn {
stream: MsgStream<Msg>,
cancel: Arc<Cancel>,
},
}
pub(crate) fn spawn_effect<Msg>(
stream: impl Stream<Item = Msg> + Send + 'static,
) -> (Effect<Msg>, Task) {
let cancel = Arc::new(Cancel::new());
(
Effect::Spawn {
stream: Box::pin(stream),
cancel: Arc::clone(&cancel),
},
Task {
cancel: Some(cancel),
},
)
}
pub(crate) fn spawn_once_effect<Msg>(
future: impl Future<Output = Msg> + Send + 'static,
) -> (Effect<Msg>, Task) {
spawn_effect(FutureStream {
future: Some(Box::pin(future)),
})
}
struct FutureStream<F: Future> {
future: Option<Pin<Box<F>>>,
}
impl<F: Future> Stream for FutureStream<F> {
type Item = F::Output;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<F::Output>> {
let this = self.get_mut();
match &mut this.future {
Some(future) => match future.as_mut().poll(cx) {
Poll::Ready(value) => {
this.future = None;
Poll::Ready(Some(value))
}
Poll::Pending => Poll::Pending,
},
None => Poll::Ready(None),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Pending;
impl Stream for Pending {
type Item = ();
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<()>> {
Poll::Pending
}
}
#[test]
fn drop_cancels() {
let (effect, task) = spawn_effect(Pending);
let Effect::Spawn { cancel, .. } = effect;
assert!(!cancel.is_cancelled());
drop(task);
assert!(cancel.is_cancelled());
}
#[test]
fn detach_does_not_cancel() {
let (effect, task) = spawn_effect(Pending);
let Effect::Spawn { cancel, .. } = effect;
task.detach();
assert!(!cancel.is_cancelled());
}
fn poll_once<F: Future + Unpin>(future: &mut F) -> Poll<F::Output> {
let waker = Waker::noop();
Pin::new(future).poll(&mut Context::from_waker(waker))
}
#[test]
fn tracker_is_idle_until_guarded() {
let tracker = PersistTracker::new();
assert!(tracker.is_idle());
let guard = tracker.guard();
assert!(!tracker.is_idle());
drop(guard);
assert!(tracker.is_idle());
}
#[test]
fn wait_resolves_when_the_last_guard_drops() {
let tracker = PersistTracker::new();
let a = tracker.guard();
let b = tracker.guard();
let mut wait = tracker.wait();
assert_eq!(poll_once(&mut wait), Poll::Pending);
drop(a);
assert_eq!(poll_once(&mut wait), Poll::Pending);
drop(b);
assert_eq!(poll_once(&mut wait), Poll::Ready(()));
}
#[test]
fn wait_on_an_idle_tracker_is_immediate() {
let tracker = PersistTracker::new();
assert_eq!(poll_once(&mut tracker.wait()), Poll::Ready(()));
}
#[test]
fn the_final_guard_drop_wakes_the_waiter() {
use std::sync::atomic::AtomicUsize;
static WAKES: AtomicUsize = AtomicUsize::new(0);
fn count_waker() -> Waker {
use std::task::{RawWaker, RawWakerVTable};
fn wake(_: *const ()) {
WAKES.fetch_add(1, Ordering::SeqCst);
}
fn clone(p: *const ()) -> RawWaker {
RawWaker::new(p, &VTABLE)
}
fn drop_raw(_: *const ()) {}
static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, wake, wake, drop_raw);
unsafe { Waker::from_raw(RawWaker::new(std::ptr::null(), &VTABLE)) }
}
let tracker = PersistTracker::new();
let guard = tracker.guard();
let mut wait = tracker.wait();
let waker = count_waker();
assert!(
Pin::new(&mut wait)
.poll(&mut Context::from_waker(&waker))
.is_pending()
);
let before = WAKES.load(Ordering::SeqCst);
drop(guard);
assert_eq!(WAKES.load(Ordering::SeqCst), before + 1);
assert_eq!(poll_once(&mut wait), Poll::Ready(()));
}
}