use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, OnceLock};
use std::task::{Context, Poll};
use lgwks_deps::tokio::sync::watch;
struct Inner {
cancelled: AtomicBool,
signal: OnceLock<watch::Sender<bool>>,
parent: Option<Arc<Inner>>,
}
impl Inner {
fn root() -> Arc<Self> {
Arc::new(Self {
cancelled: AtomicBool::new(false),
signal: OnceLock::new(),
parent: None,
})
}
fn child_of(parent: &Arc<Self>) -> Arc<Self> {
Arc::new(Self {
cancelled: AtomicBool::new(false),
signal: OnceLock::new(),
parent: Some(Arc::clone(parent)),
})
}
fn subscribe(&self) -> watch::Receiver<bool> {
let signal = self.signal.get_or_init(|| watch::channel(false).0);
if self.cancelled.load(Ordering::SeqCst) {
signal.send_replace(true);
}
signal.subscribe()
}
fn is_cancelled(&self) -> bool {
let mut node = self;
loop {
if node.cancelled.load(Ordering::SeqCst) {
return true;
}
match node.parent.as_ref() {
Some(parent) => node = parent.as_ref(),
None => return false,
}
}
}
fn cancelled(&self) -> Pin<Box<dyn Future<Output = ()> + Send + '_>> {
let mut waiters: Vec<Pin<Box<dyn Future<Output = ()> + Send + '_>>> = Vec::new();
let mut node = Some(self);
while let Some(current) = node {
let mut receiver = current.subscribe();
if *receiver.borrow_and_update() {
return Box::pin(std::future::ready(()));
}
waiters.push(Box::pin(async move { wait_for_signal(receiver).await }));
node = current.parent.as_deref();
}
Box::pin(std::future::poll_fn(move |context: &mut Context<'_>| {
let mut ready = false;
for waiter in &mut waiters {
if waiter.as_mut().poll(context).is_ready() {
ready = true;
}
}
if ready {
Poll::Ready(())
} else {
Poll::Pending
}
}))
}
fn cancel(&self) {
self.cancelled.store(true, Ordering::SeqCst);
if let Some(signal) = self.signal.get() {
signal.send_replace(true);
}
}
}
impl Drop for Inner {
fn drop(&mut self) {
let mut next = self.parent.take();
while let Some(node) = next {
match Arc::into_inner(node) {
Some(mut inner) => next = inner.parent.take(),
None => return,
}
}
}
}
impl fmt::Debug for Inner {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("Inner")
.field("cancelled", &self.cancelled.load(Ordering::SeqCst))
.field("has_parent", &self.parent.is_some())
.finish_non_exhaustive()
}
}
#[derive(Clone, Debug)]
pub struct CancellationToken {
inner: Arc<Inner>,
}
impl CancellationToken {
#[must_use]
pub fn new() -> Self {
Self {
inner: Inner::root(),
}
}
pub fn cancel(&self) {
self.inner.cancel();
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.inner.is_cancelled()
}
pub async fn cancelled(&self) {
self.inner.cancelled().await;
}
pub fn cancelled_owned(&self) -> impl Future<Output = ()> + Send + 'static {
let token = self.clone();
async move { token.cancelled().await }
}
#[must_use]
pub fn child_token(&self) -> Self {
Self {
inner: Inner::child_of(&self.inner),
}
}
#[must_use]
pub fn drop_guard(self) -> DropGuard {
DropGuard { token: self }
}
pub async fn run_until_cancelled<F: Future>(&self, future: F) -> Option<F::Output> {
let mut inner = Box::pin(future);
let mut cancelled = Box::pin(self.cancelled());
std::future::poll_fn(move |context: &mut Context<'_>| {
if let Poll::Ready(output) = inner.as_mut().poll(context) {
return Poll::Ready(Some(output));
}
if cancelled.as_mut().poll(context).is_ready() {
return Poll::Ready(None);
}
Poll::Pending
})
.await
}
}
impl Default for CancellationToken {
fn default() -> Self {
Self::new()
}
}
async fn wait_for_signal(mut receiver: watch::Receiver<bool>) {
while receiver.changed().await.is_ok() {
if *receiver.borrow_and_update() {
return;
}
}
}
#[derive(Debug)]
pub struct DropGuard {
token: CancellationToken,
}
impl Drop for DropGuard {
fn drop(&mut self) {
self.token.cancel();
}
}
#[cfg(test)]
mod tests {
use super::CancellationToken;
use crate::rt::runtime::block_on;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::task::{Context, Waker};
use std::time::Duration;
const DEPTHS: [usize; 3] = [1_024, 8_192, 50_000];
const SMALL_STACK: usize = 256 * 1024;
const JOURNEY_LIMIT: Duration = Duration::from_secs(30);
fn chain_of(depth: usize) -> (CancellationToken, CancellationToken) {
let root = CancellationToken::new();
let mut leaf = root.clone();
for _ in 0..depth {
leaf = leaf.child_token();
}
(root, leaf)
}
fn middle_link(depth: usize) -> usize {
depth.saturating_sub(depth.div_ceil(2))
}
fn polls_ready<F: Future + ?Sized>(future: &mut Pin<Box<F>>) -> bool {
let mut context = Context::from_waker(Waker::noop());
future.as_mut().poll(&mut context).is_ready()
}
#[derive(Debug)]
struct Doorbell {
rung: Mutex<bool>,
bell: Condvar,
}
impl Doorbell {
fn new() -> Self {
Self {
rung: Mutex::new(false),
bell: Condvar::new(),
}
}
fn ring(&self) {
let mut rung = crate::journal::owner::lock(&self.rung);
*rung = true;
self.bell.notify_all();
}
fn wait(&self, timeout: Duration) -> bool {
let rung = crate::journal::owner::lock(&self.rung);
let (rung, _timeout) =
crate::journal::owner::wait_timeout_while(&self.bell, rung, timeout, |rung| !*rung);
*rung
}
}
struct RingOnDrop {
bell: Arc<Doorbell>,
}
impl Drop for RingOnDrop {
fn drop(&mut self) {
self.bell.ring();
}
}
#[derive(Debug, PartialEq, Eq)]
enum StackFailure {
Finished,
Panicked(String),
TimedOut,
Unspawnable(String),
}
fn panic_message(payload: Box<dyn std::any::Any + Send>) -> String {
match payload.downcast::<String>() {
Ok(text) => *text,
Err(payload) => match payload.downcast::<&'static str>() {
Ok(text) => String::from(*text),
Err(_other) => String::from("<panic payload was not a string>"),
},
}
}
fn on_a_stack(stack_size: usize, journey: impl FnOnce() + Send + 'static) -> StackFailure {
let bell = Arc::new(Doorbell::new());
let ring = Arc::clone(&bell);
let started = std::thread::Builder::new()
.name(String::from("cancellation-depth"))
.stack_size(stack_size)
.spawn(move || {
let _ring = RingOnDrop { bell: ring };
journey();
});
let handle = match started {
Ok(handle) => handle,
Err(error) => return StackFailure::Unspawnable(error.to_string()),
};
if !bell.wait(JOURNEY_LIMIT) {
return StackFailure::TimedOut;
}
match handle.join() {
Ok(()) => StackFailure::Finished,
Err(payload) => StackFailure::Panicked(panic_message(payload)),
}
}
#[test]
fn a_fresh_wait_on_a_deep_chain_is_pending_then_resolves() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let (root, leaf) = chain_of(depth);
let mut wait = Box::pin(leaf.cancelled());
assert!(
!polls_ready(&mut wait),
"a fresh wait on an uncancelled {depth}-deep chain must be pending"
);
root.cancel();
assert!(
polls_ready(&mut wait),
"the wait must resolve once the root of a {depth}-deep chain is cancelled"
);
assert!(
leaf.is_cancelled(),
"the leaf of a {depth}-deep chain must observe the root cancel"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn cancelling_an_intermediate_node_reaches_a_deep_leaf() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let root = CancellationToken::new();
let mut leaf = root.clone();
let mut middle = root.clone();
let mut reached_middle = false;
for index in 0..depth {
leaf = leaf.child_token();
if index == middle_link(depth) {
middle = leaf.clone();
reached_middle = true;
}
}
assert!(
reached_middle,
"a {depth}-deep chain contains a middle node"
);
middle.cancel();
assert!(
leaf.is_cancelled(),
"a cancel at the middle of a {depth}-deep chain must reach the leaf"
);
assert!(
!root.is_cancelled(),
"cancelling a descendant must not cancel the root"
);
let mut wait = Box::pin(leaf.cancelled());
assert!(
polls_ready(&mut wait),
"a wait on a chain already cancelled mid-way must be ready on its first poll"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn dropping_a_pending_deep_wait_is_stack_bounded() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let (root, leaf) = chain_of(depth);
let mut wait = Box::pin(leaf.cancelled());
assert!(
!polls_ready(&mut wait),
"a fresh wait must be pending before it is dropped"
);
drop(wait);
assert!(
!root.is_cancelled(),
"dropping a pending wait must not cancel the chain"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn a_deep_chain_dropped_while_other_owners_release_concurrently_flattens_without_nesting() {
const CHURN_THREADS: usize = 4;
const CHURN_STACK: usize = 256 * 1024;
let outcome = on_a_stack(SMALL_STACK, move || {
let root = CancellationToken::new();
let weak_root = Arc::downgrade(&root.inner);
let mut leaf = root.clone();
for _ in 0..2048 {
leaf = leaf.child_token();
}
let weak_leaf = Arc::downgrade(&leaf.inner);
let door = Arc::new(AtomicBool::new(false));
let refused = Arc::new(AtomicUsize::new(0));
let mut churn = Vec::new();
let mut churn_rounds = Vec::new();
for _ in 0..CHURN_THREADS {
let root = root.clone();
let door = Arc::clone(&door);
let rounds = Arc::new(AtomicUsize::new(0));
churn_rounds.push(Arc::clone(&rounds));
let refused = Arc::clone(&refused);
let thread = std::thread::Builder::new()
.name(String::from("cancellation-churn"))
.stack_size(CHURN_STACK)
.spawn(move || {
while !door.load(Ordering::Relaxed) {
let child = root.child_token();
drop(child);
rounds.fetch_add(1, Ordering::Relaxed);
}
});
match thread {
Ok(handle) => churn.push(handle),
Err(_refused) => {
refused.fetch_add(1, Ordering::Relaxed);
}
}
}
assert_eq!(
refused.load(Ordering::Relaxed),
0,
"the OS refused a churn thread; the hammer needs all of them"
);
let mut spins = 0u32;
while churn_rounds
.iter()
.any(|rounds| rounds.load(Ordering::Relaxed) == 0)
{
std::thread::yield_now();
spins += 1;
assert!(
spins < 10_000_000,
"churn threads never ran; the hammer is invalid"
);
}
drop(leaf);
assert!(
weak_leaf.upgrade().is_none(),
"the flattened leaf must be freed even under concurrent release"
);
assert!(
weak_root.upgrade().is_some(),
"the root is still held by this frame and by the churn threads"
);
door.store(true, Ordering::Relaxed);
for handle in churn {
assert!(
handle.join().is_ok(),
"a churn thread panicked; the hammer is invalid"
);
}
for (index, rounds) in churn_rounds.iter().enumerate() {
assert!(
rounds.load(Ordering::Relaxed) > 0,
"churn thread {index} ran no rounds"
);
}
drop(root);
assert!(
weak_root.upgrade().is_none(),
"the root and its ancestry must be freed once the last handle goes"
);
});
assert_eq!(outcome, StackFailure::Finished);
}
#[test]
fn destroying_a_deep_chain_is_stack_bounded_and_leaks_nothing() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let root = CancellationToken::new();
let weak_root = Arc::downgrade(&root.inner);
let mut leaf = root.clone();
let mut weak_middle = Arc::downgrade(&root.inner);
let mut reached_middle = false;
for index in 0..depth {
leaf = leaf.child_token();
if index == middle_link(depth) {
weak_middle = Arc::downgrade(&leaf.inner);
reached_middle = true;
}
}
assert!(
reached_middle,
"a {depth}-deep chain contains a middle node"
);
let weak_leaf = Arc::downgrade(&leaf.inner);
assert!(
weak_root.upgrade().is_some(),
"the root must be alive while a handle to it is held"
);
drop(leaf);
assert!(
weak_leaf.upgrade().is_none(),
"the leaf must be freed when its last handle goes"
);
assert!(
weak_middle.upgrade().is_none(),
"freeing the leaf of a {depth}-deep chain must free its ancestry, not retain it"
);
assert!(
weak_root.upgrade().is_some(),
"the root is still held, so the chain must not have been freed from under it"
);
drop(root);
assert!(
weak_root.upgrade().is_none(),
"the root must be freed once its last handle goes"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn a_fresh_token_is_not_cancelled() {
let token = CancellationToken::new();
assert!(!token.is_cancelled(), "a new token must start uncancelled");
}
#[test]
fn cancel_reaches_every_clone_once() {
let token = CancellationToken::new();
let clone = token.clone();
token.cancel();
token.cancel();
assert!(
token.is_cancelled(),
"the original must observe its own cancel"
);
assert!(clone.is_cancelled(), "a clone must observe the cancel");
}
#[test]
fn cancelling_a_child_leaves_the_parent_running() {
let parent = CancellationToken::new();
let child = parent.child_token();
child.cancel();
assert!(child.is_cancelled(), "the child must be cancelled");
assert!(
!parent.is_cancelled(),
"cancelling a child must not cancel its parent"
);
}
#[test]
fn cancelling_a_parent_reaches_a_grandchild() {
let parent = CancellationToken::new();
let child = parent.child_token();
let grandchild = child.child_token();
parent.cancel();
assert!(child.is_cancelled(), "the child must follow its parent");
assert!(
grandchild.is_cancelled(),
"cancellation must reach a grandchild"
);
}
#[test]
fn a_child_created_after_cancellation_starts_cancelled() {
let parent = CancellationToken::new();
parent.cancel();
let child = parent.child_token();
assert!(
child.is_cancelled(),
"a child created after cancel must not start live"
);
}
#[test]
fn a_descendant_survives_its_intermediate_parent_being_dropped() {
let root = CancellationToken::new();
let leaves: Vec<CancellationToken> = {
let branch = root.child_token();
(0..8).map(|_| branch.child_token()).collect()
};
root.cancel();
assert!(
leaves.iter().all(CancellationToken::is_cancelled),
"a leaf must stay cancellable after its intermediate parent is dropped"
);
}
#[test]
fn drop_guard_cancels_on_scope_exit() {
let token = CancellationToken::new();
{
let _guard = token.clone().drop_guard();
assert!(
!token.is_cancelled(),
"the guard must not cancel before it is dropped"
);
}
assert!(
token.is_cancelled(),
"dropping the guard must cancel the token"
);
}
#[test]
fn run_until_cancelled_returns_the_value_when_the_future_wins() {
let token = CancellationToken::new();
let output = block_on(token.run_until_cancelled(async { 7u8 }));
assert_eq!(
output,
Some(7),
"a completed future must yield its value, not None"
);
}
#[test]
fn run_until_cancelled_returns_none_when_already_cancelled() {
let token = CancellationToken::new();
token.cancel();
let output = block_on(token.run_until_cancelled(std::future::pending::<u8>()));
assert_eq!(
output, None,
"an already-cancelled token must abandon a pending future"
);
}
#[test]
fn cancelled_terminates_for_a_deep_already_cancelled_chain() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let (root, leaf) = chain_of(depth);
root.cancel();
assert!(
leaf.is_cancelled(),
"a {depth}-deep descendant must observe the root cancel"
);
let resolved = block_on(crate::rt::time::timeout(JOURNEY_LIMIT, leaf.cancelled()));
assert!(
resolved.is_ok(),
"waiting on an already-cancelled {depth}-deep chain must resolve"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn a_cancel_that_finds_no_subscriber_is_observed_by_the_next_one() {
for depth in DEPTHS {
let (root, leaf) = chain_of(depth);
let intermediate = leaf.child_token();
root.cancel();
assert!(
intermediate.is_cancelled(),
"depth {depth}: a child of an already-cancelled node must observe the cancel \
without any channel having been built for either"
);
let resolved = block_on(crate::rt::time::timeout(JOURNEY_LIMIT, leaf.cancelled()));
assert!(
resolved.is_ok(),
"depth {depth}: a cancel published to no channel must still wake the next \
waiter; it did not resolve within the journey limit"
);
}
}
#[test]
fn a_child_cancelled_before_its_own_channel_exists_still_wakes_its_own_waiter() {
let root = CancellationToken::new();
let child = root.child_token();
child.cancel();
assert!(
child.is_cancelled(),
"a child cancelled on its own account must observe its own cancel"
);
assert!(
!root.is_cancelled(),
"cancelling a child must still leave its parent running"
);
let resolved = block_on(crate::rt::time::timeout(JOURNEY_LIMIT, child.cancelled()));
assert!(
resolved.is_ok(),
"a child cancelled before any waiter existed must still resolve the next waiter"
);
}
#[test]
fn a_deep_chain_propagates_to_every_level() {
const EVERY_LEVEL: usize = 2_048;
let root = CancellationToken::new();
let mut chain = vec![root.clone()];
for _ in 0..EVERY_LEVEL {
let next = chain.last().map(CancellationToken::child_token);
match next {
Some(token) => chain.push(token),
None => break,
}
}
root.cancel();
assert!(
chain.iter().all(CancellationToken::is_cancelled),
"every level of a {EVERY_LEVEL}-deep chain must be cancelled"
);
}
#[test]
fn dropping_a_deep_intermediate_handle_keeps_the_leaf_cancellable() {
for depth in DEPTHS {
let outcome = on_a_stack(SMALL_STACK, move || {
let root = CancellationToken::new();
let mut leaf = root.clone();
let mut weak_middle = Arc::downgrade(&root.inner);
let mut reached_middle = false;
for index in 0..depth {
let next = leaf.child_token();
leaf = next;
if index == middle_link(depth) {
weak_middle = Arc::downgrade(&leaf.inner);
reached_middle = true;
}
}
assert!(
reached_middle,
"a {depth}-deep chain contains a middle node"
);
assert!(
weak_middle.upgrade().is_some(),
"the middle node must be held by its descendant after its own handle goes"
);
root.cancel();
assert!(
leaf.is_cancelled(),
"the leaf of a {depth}-deep chain must stay cancellable"
);
});
assert_eq!(outcome, StackFailure::Finished, "depth {depth}");
}
}
#[test]
fn a_dropped_child_does_not_keep_its_parent_cancellable_state_alive() {
let parent = CancellationToken::new();
{
let child = parent.child_token();
assert!(!child.is_cancelled(), "the child starts live");
}
parent.cancel();
assert!(
parent.is_cancelled(),
"the parent must still cancel with no live children"
);
}
}