moirai-pal 0.7.0

Platform Abstraction Layer for Moirai async I/O operations
Documentation
//! Browser-local future cancellation state.

use std::cell::{Cell, RefCell};
use std::future::Future;
use std::pin::Pin;
use std::rc::Rc;
use std::task::{Context, Poll, Waker};

struct LocalTaskState {
    cancelled: Cell<bool>,
    waker: RefCell<Option<Waker>>,
}

impl LocalTaskState {
    fn cancel(&self) {
        if self.cancelled.replace(true) {
            return;
        }
        // The borrow ends before the wake: a waker may poll the task again on
        // this thread, and that poll borrows the same cell.
        let waker = self.waker.borrow_mut().take();
        if let Some(waker) = waker {
            waker.wake();
        }
    }
}

/// Owns cancellation for one future scheduled on the browser event loop.
///
/// The handle is single-owner. Calling [`Self::cancel`] or dropping the handle
/// wakes the task, which then drops its child future and any PAL resources it
/// owns. A task that has already completed is unaffected.
#[must_use = "retain the handle to cancel the browser task"]
pub struct LocalTaskHandle {
    state: Rc<LocalTaskState>,
}

impl LocalTaskHandle {
    /// Requests cancellation of the task.
    pub fn cancel(&self) {
        self.state.cancel();
    }

    /// Returns whether cancellation has been requested.
    #[must_use]
    pub fn is_cancelled(&self) -> bool {
        self.state.cancelled.get()
    }
}

impl Drop for LocalTaskHandle {
    fn drop(&mut self) {
        self.state.cancel();
    }
}

pub(crate) struct CancellableFuture<F> {
    future: Option<Pin<Box<F>>>,
    state: Rc<LocalTaskState>,
}

pub(crate) fn cancellable<F>(future: F) -> (LocalTaskHandle, CancellableFuture<F>)
where
    F: Future<Output = ()> + 'static,
{
    let state = Rc::new(LocalTaskState {
        cancelled: Cell::new(false),
        waker: RefCell::new(None),
    });
    let handle = LocalTaskHandle {
        state: Rc::clone(&state),
    };
    let future = CancellableFuture {
        future: Some(Box::pin(future)),
        state,
    };
    (handle, future)
}

impl<F: Future<Output = ()>> Future for CancellableFuture<F> {
    type Output = ();

    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
        // `future` is boxed and pinned on its own, so the wrapper is `Unpin`.
        let this = self.get_mut();
        if this.state.cancelled.get() {
            this.future.take();
            this.state.waker.borrow_mut().take();
            return Poll::Ready(());
        }

        this.state.waker.borrow_mut().replace(cx.waker().clone());
        let Some(future) = this.future.as_mut() else {
            this.state.waker.borrow_mut().take();
            return Poll::Ready(());
        };
        let result = future.as_mut().poll(cx);
        if result.is_ready() {
            this.future.take();
            this.state.waker.borrow_mut().take();
        }
        result
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::sync::Arc;
    use std::sync::atomic::{AtomicUsize, Ordering};
    use std::task::{Wake, Waker};

    struct PendingFuture {
        polls: Rc<Cell<usize>>,
        dropped: Rc<Cell<bool>>,
    }

    impl Future for PendingFuture {
        type Output = ();

        fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
            let this = self.get_mut();
            this.polls.set(this.polls.get() + 1);
            Poll::Pending
        }
    }

    impl Drop for PendingFuture {
        fn drop(&mut self) {
            self.dropped.set(true);
        }
    }

    struct CountingFuture {
        polls: Rc<Cell<usize>>,
    }

    impl Future for CountingFuture {
        type Output = ();

        fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Self::Output> {
            let this = self.get_mut();
            this.polls.set(this.polls.get() + 1);
            Poll::Ready(())
        }
    }

    #[derive(Default)]
    struct WakeCounter(AtomicUsize);

    impl Wake for WakeCounter {
        fn wake(self: Arc<Self>) {
            self.0.fetch_add(1, Ordering::Relaxed);
        }

        fn wake_by_ref(self: &Arc<Self>) {
            self.0.fetch_add(1, Ordering::Relaxed);
        }
    }

    #[test]
    fn cancellation_before_poll_drops_the_child_without_polling_it() {
        let polls = Rc::new(Cell::new(0));
        let dropped = Rc::new(Cell::new(false));
        let (handle, mut future) = cancellable(PendingFuture {
            polls: Rc::clone(&polls),
            dropped: Rc::clone(&dropped),
        });

        handle.cancel();
        assert!(handle.is_cancelled());
        let waker = Waker::noop();
        let mut context = Context::from_waker(waker);
        assert!(Pin::new(&mut future).poll(&mut context).is_ready());
        assert_eq!(polls.get(), 0);
        assert!(dropped.get());
    }

    #[test]
    fn cancellation_wakes_a_pending_task_and_drops_the_child() {
        let dropped = Rc::new(Cell::new(false));
        let (handle, mut future) = cancellable(PendingFuture {
            polls: Rc::new(Cell::new(0)),
            dropped: Rc::clone(&dropped),
        });
        let signal = Arc::new(WakeCounter::default());
        let waker = Waker::from(Arc::clone(&signal));
        let mut context = Context::from_waker(&waker);

        assert!(Pin::new(&mut future).poll(&mut context).is_pending());
        handle.cancel();
        assert_eq!(signal.0.load(Ordering::Relaxed), 1);
        assert!(Pin::new(&mut future).poll(&mut context).is_ready());
        assert!(dropped.get());
    }

    thread_local! {
        static REENTRANT: RefCell<Option<Rc<LocalTaskState>>> = const { RefCell::new(None) };
        static WAKER_CELL_FREE_DURING_WAKE: Cell<Option<bool>> = const { Cell::new(None) };
    }

    struct ReentrantWake;

    impl Wake for ReentrantWake {
        fn wake(self: Arc<Self>) {
            self.wake_by_ref();
        }

        fn wake_by_ref(self: &Arc<Self>) {
            let free = REENTRANT.with(|state| {
                state
                    .borrow()
                    .as_ref()
                    .is_some_and(|state| state.waker.try_borrow_mut().is_ok())
            });
            WAKER_CELL_FREE_DURING_WAKE.with(|cell| cell.set(Some(free)));
        }
    }

    #[test]
    fn a_waker_that_re_enters_the_task_finds_its_cell_free() {
        let (handle, mut future) = cancellable(PendingFuture {
            polls: Rc::new(Cell::new(0)),
            dropped: Rc::new(Cell::new(false)),
        });
        REENTRANT.with(|state| *state.borrow_mut() = Some(Rc::clone(&handle.state)));
        let waker = Waker::from(Arc::new(ReentrantWake));
        let mut context = Context::from_waker(&waker);
        assert!(Pin::new(&mut future).poll(&mut context).is_pending());

        handle.cancel();

        assert_eq!(WAKER_CELL_FREE_DURING_WAKE.with(Cell::get), Some(true));
        REENTRANT.with(|state| state.borrow_mut().take());
    }

    #[test]
    fn dropping_the_handle_requests_cancellation() {
        let dropped = Rc::new(Cell::new(false));
        let (handle, mut future) = cancellable(PendingFuture {
            polls: Rc::new(Cell::new(0)),
            dropped: Rc::clone(&dropped),
        });
        drop(handle);

        let waker = Waker::noop();
        let mut context = Context::from_waker(waker);
        assert!(Pin::new(&mut future).poll(&mut context).is_ready());
        assert!(dropped.get());
    }

    #[test]
    fn completed_task_can_be_cancelled_without_repolling() {
        let polls = Rc::new(Cell::new(0));
        let (handle, mut future) = cancellable(CountingFuture {
            polls: Rc::clone(&polls),
        });
        let waker = Waker::noop();
        let mut context = Context::from_waker(waker);
        assert!(Pin::new(&mut future).poll(&mut context).is_ready());
        assert_eq!(polls.get(), 1);
        assert!(!handle.is_cancelled());
        handle.cancel();
        assert!(handle.is_cancelled());
        assert!(Pin::new(&mut future).poll(&mut context).is_ready());
        assert_eq!(polls.get(), 1);
    }
}