#[derive(Clone, Debug)]
pub struct CancellationToken {
rx: Option<tokio::sync::watch::Receiver<bool>>,
}
#[derive(Clone, Debug)]
pub struct CancellationHandle {
tx: tokio::sync::watch::Sender<bool>,
}
impl CancellationToken {
#[must_use]
pub fn new() -> (CancellationHandle, CancellationToken) {
let (tx, rx) = tokio::sync::watch::channel(false);
(
CancellationHandle { tx },
CancellationToken { rx: Some(rx) },
)
}
#[must_use]
pub fn disabled() -> Self {
Self { rx: None }
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.rx.as_ref().is_some_and(|rx| *rx.borrow())
}
pub(crate) fn can_cancel(&self) -> bool {
self.rx.is_some()
}
pub async fn cancelled(&self) {
match &self.rx {
None => std::future::pending::<()>().await,
Some(rx) => {
let mut rx = rx.clone();
if *rx.borrow() {
return;
}
while rx.changed().await.is_ok() {
if *rx.borrow() {
return;
}
}
std::future::pending::<()>().await;
}
}
}
}
impl CancellationHandle {
pub fn cancel(&self) {
let _ = self.tx.send(true);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[tokio::test]
async fn cancel_propagates_to_receivers() {
let (handle, token) = CancellationToken::new();
assert!(!token.is_cancelled());
handle.cancel();
tokio::time::timeout(Duration::from_secs(1), token.cancelled())
.await
.expect("cancelled() should resolve after cancel");
assert!(token.is_cancelled());
}
#[tokio::test]
async fn clone_after_cancel_still_observes() {
let (handle, token) = CancellationToken::new();
handle.cancel();
let cloned = token.clone();
assert!(cloned.is_cancelled());
tokio::time::timeout(Duration::from_secs(1), cloned.cancelled())
.await
.expect("a clone made after cancel still observes it");
}
#[test]
fn is_cancelled_polls_without_await() {
let (handle, token) = CancellationToken::new();
assert!(!token.is_cancelled());
handle.cancel();
assert!(token.is_cancelled());
}
#[tokio::test]
async fn cancel_is_idempotent() {
let (handle, token) = CancellationToken::new();
handle.cancel();
handle.cancel(); assert!(token.is_cancelled());
}
#[tokio::test]
async fn dropping_handle_without_cancel_keeps_token_pending() {
let (handle, token) = CancellationToken::new();
drop(handle);
assert!(!token.is_cancelled());
let res = tokio::time::timeout(Duration::from_millis(50), token.cancelled()).await;
assert!(
res.is_err(),
"cancelled() should stay pending after handle drop"
);
}
#[tokio::test]
async fn disabled_is_never_cancelled_and_can_cancel_false() {
let token = CancellationToken::disabled();
assert!(!token.is_cancelled());
assert!(!token.can_cancel());
let res = tokio::time::timeout(Duration::from_millis(50), token.cancelled()).await;
assert!(res.is_err(), "disabled token should stay pending");
}
#[test]
fn can_cancel_true_for_new_token() {
let (_handle, token) = CancellationToken::new();
assert!(token.can_cancel());
}
#[tokio::test]
async fn cancelled_resolves_when_cancel_fires_while_awaiting() {
let (handle, token) = CancellationToken::new();
let waiter = tokio::spawn(async move { token.cancelled().await });
tokio::time::sleep(Duration::from_millis(20)).await;
handle.cancel();
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.expect("waiter should wake on cancel")
.expect("waiter task should not panic");
}
#[tokio::test]
async fn clone_before_cancel_is_observed() {
let (handle, token) = CancellationToken::new();
let cloned = token.clone();
assert!(!cloned.is_cancelled());
handle.cancel();
assert!(cloned.is_cancelled());
assert!(token.is_cancelled());
tokio::time::timeout(Duration::from_secs(1), cloned.cancelled())
.await
.expect("a clone made before cancel still resolves");
}
#[tokio::test]
async fn cloned_handle_triggers_cancellation() {
let (handle, token) = CancellationToken::new();
let handle2 = handle.clone();
drop(handle); assert!(!token.is_cancelled());
handle2.cancel();
assert!(token.is_cancelled());
}
#[tokio::test]
async fn multiple_awaiters_all_wake_on_cancel() {
let (handle, token) = CancellationToken::new();
let waiters: Vec<_> = (0..5)
.map(|_| {
let t = token.clone();
tokio::spawn(async move { t.cancelled().await })
})
.collect();
tokio::time::sleep(Duration::from_millis(20)).await;
handle.cancel();
for w in waiters {
tokio::time::timeout(Duration::from_secs(1), w)
.await
.expect("every awaiter should wake")
.expect("awaiter task should not panic");
}
}
}