unb-runtime 2.0.0

unb session runtime: transport codec, session/writer engine, Wire handle, cancellation
Documentation
use std::collections::HashMap;
use std::fmt::Debug;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, Weak};
use std::task::{Context, Poll, Waker};
use std::time::Duration;

use n0_future::task::{spawn, JoinHandle};

#[derive(Debug, Default)]
struct CancelState {
    cancelled: AtomicBool,
    next_waiter: AtomicU64,
    wakers: Mutex<HashMap<u64, Waker>>,
    children: Mutex<Vec<Weak<CancelState>>>,
}

impl CancelState {
    fn wake_all(&self) {
        let drained: Vec<Waker> = self
            .wakers
            .lock()
            .unwrap()
            .drain()
            .map(|(_, waker)| waker)
            .collect();
        for waker in drained {
            waker.wake();
        }
    }
}

pub struct CancellationToken {
    state: Arc<CancelState>,
    timeout_handle: Option<JoinHandle<()>>,
}

impl Clone for CancellationToken {
    fn clone(&self) -> Self {
        Self {
            state: self.state.clone(),
            timeout_handle: None,
        }
    }
}

impl Debug for CancellationToken {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("CancellationToken")
            .field("cancelled", &self.is_cancelled())
            .finish()
    }
}

impl Default for CancellationToken {
    fn default() -> Self {
        Self::new()
    }
}

impl Drop for CancellationToken {
    fn drop(&mut self) {
        if let Some(handle) = self.timeout_handle.take() {
            handle.abort();
        }
    }
}

impl CancellationToken {
    pub fn new() -> Self {
        Self {
            state: Arc::new(CancelState::default()),
            timeout_handle: None,
        }
    }

    pub fn timeout(duration: Duration) -> Self {
        let mut token = CancellationToken::new();
        let child = token.clone();
        token.timeout_handle = Some(spawn(async move {
            n0_future::time::sleep(duration).await;
            child.cancel();
        }));
        token
    }

    pub fn is_cancelled(&self) -> bool {
        self.state.cancelled.load(Ordering::Acquire)
    }

    pub fn cancel_after(&self, duration: Duration) {
        let token = self.clone();
        spawn(async move {
            n0_future::time::sleep(duration).await;
            token.cancel();
        });
    }

    pub fn cancel(&self) {
        if self.state.cancelled.swap(true, Ordering::AcqRel) {
            return;
        }
        self.state.wake_all();
        let mut stack: Vec<Arc<CancelState>> = Self::collect_children(&self.state);
        while let Some(node) = stack.pop() {
            if !node.cancelled.swap(true, Ordering::AcqRel) {
                node.wake_all();
                stack.extend(Self::collect_children(&node));
            }
        }
    }

    fn collect_children(state: &Arc<CancelState>) -> Vec<Arc<CancelState>> {
        let mut children = state.children.lock().unwrap();
        let mut alive = Vec::new();
        children.retain(|weak| match weak.upgrade() {
            Some(child) => {
                alive.push(child);
                true
            }
            None => false,
        });
        alive
    }

    pub fn child_token(&self) -> Self {
        let child = CancellationToken::new();
        {
            let mut children = self.state.children.lock().unwrap();
            children.retain(|weak| weak.strong_count() > 0);
            children.push(Arc::downgrade(&child.state));
        }
        if self.is_cancelled() {
            child.cancel();
        }
        child
    }

    pub fn cancelled(&self) -> Cancelled {
        Cancelled {
            state: Arc::downgrade(&self.state),
            id: self.state.next_waiter.fetch_add(1, Ordering::Relaxed),
        }
    }

    pub fn drop_guard(&self) -> DropGuard {
        DropGuard::new(self.clone())
    }
}

pub struct DropGuard {
    token: Option<CancellationToken>,
}

impl DropGuard {
    pub fn new(token: CancellationToken) -> Self {
        Self { token: Some(token) }
    }

    pub fn disarm(&mut self) {
        self.token = None;
    }
}

impl Drop for DropGuard {
    fn drop(&mut self) {
        if let Some(token) = &self.token {
            token.cancel();
        }
    }
}

pub struct Cancelled {
    state: Weak<CancelState>,
    id: u64,
}

impl Future for Cancelled {
    type Output = ();

    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
        let Some(state) = self.state.upgrade() else {
            return Poll::Ready(());
        };
        if state.cancelled.load(Ordering::Acquire) {
            return Poll::Ready(());
        }
        state
            .wakers
            .lock()
            .unwrap()
            .insert(self.id, cx.waker().clone());
        if state.cancelled.load(Ordering::Acquire) {
            Poll::Ready(())
        } else {
            Poll::Pending
        }
    }
}

impl Drop for Cancelled {
    fn drop(&mut self) {
        if let Some(state) = self.state.upgrade() {
            state.wakers.lock().unwrap().remove(&self.id);
        }
    }
}

#[derive(thiserror::Error, Debug)]
pub enum TaskErrors {
    #[error("task cancelled")]
    Cancelled,
}

pub trait FutureExtension: Future + Sized {
    fn with_cancel(
        self,
        cancellation: &CancellationToken,
    ) -> impl Future<Output = Result<Self::Output, TaskErrors>>;
}

impl<T: Future> FutureExtension for T {
    async fn with_cancel(
        self,
        cancellation: &CancellationToken,
    ) -> Result<Self::Output, TaskErrors> {
        if cancellation.is_cancelled() {
            return Err(TaskErrors::Cancelled);
        }
        let this = std::pin::pin!(self);
        tokio::select! {
            output = this => Ok(output),
            () = cancellation.cancelled() => Err(TaskErrors::Cancelled),
        }
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn a_fresh_token_is_not_cancelled_until_cancelled() {
        let token = CancellationToken::new();
        assert!(!token.is_cancelled());
        token.cancel();
        assert!(token.is_cancelled());
    }

    #[test]
    fn cancelling_a_parent_cascades_to_the_whole_subtree() {
        let parent = CancellationToken::new();
        let child = parent.child_token();
        let grandchild = child.child_token();
        parent.cancel();
        assert!(child.is_cancelled());
        assert!(grandchild.is_cancelled());
    }

    #[test]
    fn cascade_survives_a_dropped_intermediate_that_is_kept_alive() {
        let root = CancellationToken::new();
        let intermediate = root.child_token();
        let leaf = intermediate.child_token();
        root.cancel();
        assert!(
            leaf.is_cancelled(),
            "an alive intermediate carries the cascade"
        );
    }

    #[test]
    fn a_child_of_an_already_cancelled_parent_is_born_cancelled() {
        let parent = CancellationToken::new();
        parent.cancel();
        assert!(parent.child_token().is_cancelled());
    }

    #[test]
    fn dropped_children_are_pruned_so_the_parent_does_not_grow_unbounded() {
        let parent = CancellationToken::new();
        for _ in 0..1000 {
            let _ = parent.child_token();
        }
        assert!(
            parent.state.children.lock().unwrap().len() <= 1,
            "dead child weaks are reclaimed"
        );
    }

    #[tokio::test]
    async fn every_concurrent_waiter_on_one_token_wakes_on_cancel() {
        let token = CancellationToken::new();
        let waiters: Vec<_> = (0..8)
            .map(|_| {
                let token = token.clone();
                tokio::spawn(async move { token.cancelled().await })
            })
            .collect();
        token.cancel();
        for waiter in waiters {
            tokio::time::timeout(std::time::Duration::from_secs(5), waiter)
                .await
                .expect("a single AtomicWaker would have starved all but one waiter")
                .unwrap();
        }
    }

    #[tokio::test]
    async fn a_dropped_waiter_leaves_no_registration_behind() {
        let token = CancellationToken::new();
        {
            let fut = token.cancelled();
            let _ = futures_util::poll!(std::pin::pin!(fut));
        }
        assert!(
            token.state.wakers.lock().unwrap().is_empty(),
            "drop deregisters the waker"
        );
    }

    #[tokio::test]
    async fn with_cancel_short_circuits_a_pending_future() {
        let token = CancellationToken::new();
        token.cancel();
        let result = std::future::pending::<()>().with_cancel(&token).await;
        assert!(matches!(result, Err(TaskErrors::Cancelled)));
    }
}