use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::Notify;
#[derive(Default)]
struct Inner {
cancelled: AtomicBool,
notify: Notify,
}
#[derive(Clone, Default)]
pub struct CancelToken {
inner: Arc<Inner>,
}
impl CancelToken {
pub fn new() -> Self {
Self::default()
}
pub fn cancel(&self) {
self.inner.cancelled.store(true, Ordering::SeqCst);
self.inner.notify.notify_waiters();
}
pub fn is_cancelled(&self) -> bool {
self.inner.cancelled.load(Ordering::SeqCst)
}
pub async fn cancelled(&self) {
let notified = self.inner.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.is_cancelled() {
return;
}
notified.await;
}
}
impl std::fmt::Debug for CancelToken {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CancelToken")
.field("cancelled", &self.is_cancelled())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[tokio::test]
async fn cancelled_resolves_when_the_token_fires() {
let token = CancelToken::new();
assert!(!token.is_cancelled());
let waiter = tokio::spawn({
let token = token.clone();
async move { token.cancelled().await }
});
tokio::task::yield_now().await;
token.cancel();
tokio::time::timeout(Duration::from_secs(5), waiter)
.await
.expect("a fired token wakes its waiter")
.unwrap();
assert!(token.is_cancelled());
}
#[tokio::test]
async fn cancelled_returns_immediately_for_an_already_fired_token() {
let token = CancelToken::new();
token.cancel();
tokio::time::timeout(Duration::from_secs(5), token.cancelled())
.await
.expect("no wait for a token that already fired");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn a_cancel_concurrent_with_the_wait_is_observed() {
for _ in 0..500 {
let token = CancelToken::new();
let firing = tokio::spawn({
let token = token.clone();
async move { token.cancel() }
});
tokio::time::timeout(Duration::from_secs(5), token.cancelled())
.await
.expect("the cancel was observed");
firing.await.unwrap();
}
}
#[tokio::test]
async fn cancel_is_idempotent_and_wakes_every_clone() {
let token = CancelToken::new();
let waiters: Vec<_> = (0..3)
.map(|_| {
let token = token.clone();
tokio::spawn(async move { token.cancelled().await })
})
.collect();
tokio::task::yield_now().await;
token.cancel();
token.cancel();
for waiter in waiters {
tokio::time::timeout(Duration::from_secs(5), waiter)
.await
.expect("every clone observes the cancel")
.unwrap();
}
}
#[test]
fn debug_reports_the_state() {
let token = CancelToken::new();
assert!(format!("{token:?}").contains("cancelled: false"));
token.cancel();
assert!(format!("{token:?}").contains("cancelled: true"));
}
}