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)));
}
}