use std::sync::Mutex;
use tokio_util::sync::CancellationToken;
pub struct CancelSignal {
inner: Mutex<CancellationToken>,
}
impl CancelSignal {
#[must_use]
pub fn new() -> Self {
Self {
inner: Mutex::new(CancellationToken::new()),
}
}
pub fn cancel(&self) {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.cancel();
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.is_cancelled()
}
pub async fn notified(&self) {
let token = self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone();
token.cancelled().await;
}
pub fn reset(&self) {
*self
.inner
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = CancellationToken::new();
}
}
impl Default for CancelSignal {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn test_cancel_wakes_instantly() {
let signal = Arc::new(CancelSignal::new());
let signal_clone = signal.clone();
let handle = tokio::spawn(async move {
signal_clone.notified().await;
"woke up"
});
tokio::task::yield_now().await;
signal.cancel();
let result = handle.await.unwrap();
assert_eq!(result, "woke up");
}
#[tokio::test]
async fn test_notified_returns_immediately_if_already_cancelled() {
let signal = CancelSignal::new();
signal.cancel();
signal.notified().await;
}
#[test]
fn test_is_cancelled() {
let signal = CancelSignal::new();
assert!(!signal.is_cancelled());
signal.cancel();
assert!(signal.is_cancelled());
}
#[tokio::test]
async fn test_cancel_is_idempotent() {
let signal = Arc::new(CancelSignal::new());
signal.cancel();
signal.cancel();
signal.cancel();
assert!(signal.is_cancelled());
signal.notified().await;
}
#[tokio::test]
async fn test_multiple_waiters_all_wake() {
let signal = Arc::new(CancelSignal::new());
let mut handles = Vec::new();
for _ in 0..10 {
let s = signal.clone();
handles.push(tokio::spawn(async move {
s.notified().await;
true
}));
}
tokio::task::yield_now().await;
signal.cancel();
for handle in handles {
assert!(handle.await.unwrap());
}
}
#[test]
fn test_reset_clears_cancelled_state() {
let signal = CancelSignal::new();
signal.cancel();
assert!(signal.is_cancelled());
signal.reset();
assert!(
!signal.is_cancelled(),
"reset must re-arm the signal to non-cancelled"
);
}
#[test]
fn test_reset_is_visible_through_existing_arc_clone() {
let signal = Arc::new(CancelSignal::new());
let handle = Arc::clone(&signal);
signal.cancel();
assert!(handle.is_cancelled());
signal.reset();
assert!(!handle.is_cancelled());
}
#[tokio::test]
async fn test_notified_after_reset_waits_for_new_cancel() {
let signal = Arc::new(CancelSignal::new());
signal.cancel();
signal.reset();
let handle = {
let s = Arc::clone(&signal);
tokio::spawn(async move {
s.notified().await;
"woke"
})
};
tokio::task::yield_now().await;
let mut pending = tokio::time::interval(std::time::Duration::from_millis(5));
pending.tick().await;
assert!(
!handle.is_finished(),
"notified must not fire after reset until a new cancel arrives"
);
signal.cancel();
assert_eq!(handle.await.unwrap(), "woke");
}
}