use std::{
cell::{Cell, RefCell},
error::Error,
fmt::{Debug, Display},
future::Future,
pin::Pin,
rc::{Rc, Weak},
task::{Context, Poll, Waker},
};
use slab::Slab;
#[derive(Default)]
struct Inner {
cancelled: Cell<bool>,
wakers: RefCell<Vec<Waker>>,
callbacks: RefCell<Slab<Box<dyn FnOnce()>>>,
}
impl Debug for Inner {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Inner")
.field("cancelled", &self.cancelled)
.field("wakers.len", &self.wakers.borrow().len())
.field("callbacks.len", &self.callbacks.borrow().len())
.finish()
}
}
#[derive(Debug, Default, Clone)]
pub struct CancellationTokenSource {
inner: Rc<Inner>,
}
#[derive(Debug, Clone)]
pub struct CancellationToken {
inner: Rc<Inner>,
}
#[derive(Debug, Copy, Clone, Default, Eq, Ord, PartialEq, PartialOrd, Hash)]
pub struct Cancelled;
impl Display for Cancelled {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("cancelled by CancellationTokenSource")
}
}
impl Error for Cancelled {}
impl CancellationTokenSource {
pub fn new() -> Self {
Default::default()
}
pub fn token(&self) -> CancellationToken {
CancellationToken {
inner: self.inner.clone(),
}
}
pub fn cancel(&self) {
if !self.inner.cancelled.replace(true) {
for cb in self.inner.callbacks.borrow_mut().drain() {
cb();
}
for w in self.inner.wakers.borrow_mut().drain(..) {
w.wake();
}
}
}
pub fn is_cancelled(&self) -> bool {
self.inner.cancelled.get()
}
}
impl CancellationToken {
pub fn is_cancelled(&self) -> bool {
self.inner.cancelled.get()
}
pub fn check_cancelled(&self) -> Result<(), Cancelled> {
if self.is_cancelled() {
Err(Cancelled)
} else {
Ok(())
}
}
pub fn cancelled(&self) -> CancelledFuture {
CancelledFuture {
token: self.clone(),
}
}
pub fn register(&self, f: impl FnOnce() + 'static) -> Option<CancellationTokenRegistration> {
if self.is_cancelled() {
f();
None
} else {
CancellationTokenRegistration {
inner: Rc::downgrade(&self.inner),
key: self.inner.callbacks.borrow_mut().insert(Box::new(f)),
}
.into()
}
}
}
#[derive(Debug)]
pub struct CancellationTokenRegistration {
inner: Weak<Inner>,
key: usize,
}
impl Drop for CancellationTokenRegistration {
fn drop(&mut self) {
if let Some(inner) = self.inner.upgrade() {
if inner.cancelled.get() {
return;
}
let _ = inner.callbacks.borrow_mut().remove(self.key);
}
}
}
#[derive(Debug)]
pub struct CancelledFuture {
token: CancellationToken,
}
impl Future for CancelledFuture {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
if self.token.is_cancelled() {
Poll::Ready(())
} else {
let mut wakers = self.token.inner.wakers.borrow_mut();
if !wakers.iter().any(|w| w.will_wake(cx.waker())) {
wakers.push(cx.waker().clone());
}
Poll::Pending
}
}
}
#[cfg(test)]
mod tests {
use std::cell::Cell;
use std::rc::Rc;
use std::time::Duration;
use futures::{FutureExt, executor::LocalPool, pin_mut, select, task::LocalSpawnExt};
use futures_timer::Delay;
use super::*;
#[test]
fn cancel_two_tasks() {
let cancelled_a = Rc::new(Cell::new(false));
let cancelled_b = Rc::new(Cell::new(false));
let task_a = |token: CancellationToken| {
let cancelled_a = Rc::clone(&cancelled_a);
async move {
for _ in 1..=5 {
let delay = Delay::new(Duration::from_millis(50)).fuse();
let cancelled = token.cancelled().fuse();
pin_mut!(delay, cancelled);
select! {
_ = delay => {},
_ = cancelled => {
cancelled_a.set(true);
break;
},
}
}
}
};
let task_b = |token: CancellationToken| {
let cancelled_b = Rc::clone(&cancelled_b);
async move {
for _ in 1..=5 {
Delay::new(Duration::from_millis(80)).await;
if token.check_cancelled().is_err() {
cancelled_b.set(true);
break;
}
}
}
};
let cts = CancellationTokenSource::new();
let mut pool = LocalPool::new();
let spawner = pool.spawner();
spawner
.spawn_local(task_a(cts.token()).map(|_| ()))
.unwrap();
spawner
.spawn_local(task_b(cts.token()).map(|_| ()))
.unwrap();
{
let cts_clone = cts.clone();
spawner
.spawn_local(
async move {
Delay::new(Duration::from_millis(200)).await;
cts_clone.cancel();
}
.map(|_| ()),
)
.unwrap();
}
pool.run();
assert!(cts.is_cancelled());
assert!(cancelled_a.get());
assert!(cancelled_b.get());
cts.cancel();
assert!(cts.is_cancelled());
}
#[test]
fn cancellation_register_callbacks() {
let cts = CancellationTokenSource::new();
let token = cts.token();
let flag_before = Rc::new(Cell::new(false));
let flag_after = Rc::new(Cell::new(false));
let flag_drop = Rc::new(Cell::new(false));
let reg_before = {
let flag = Rc::clone(&flag_before);
token
.register(move || {
flag.set(true);
})
.unwrap()
};
cts.cancel();
assert!(flag_before.get());
drop(reg_before);
{
let flag = Rc::clone(&flag_after);
token.register(move || {
flag.set(true);
});
}
assert!(flag_after.get());
let token2 = CancellationTokenSource::new().token();
let reg_drop = {
let flag = Rc::clone(&flag_drop);
token2
.register(move || {
flag.set(true);
})
.unwrap()
};
drop(reg_drop); token2.inner.cancelled.set(true); assert!(!flag_drop.get());
}
#[test]
fn cancelled_future_poll_ready() {
let cts = CancellationTokenSource::new();
let token = cts.token();
let mut pool = LocalPool::new();
let spawner = pool.spawner();
let finished = Rc::new(Cell::new(false));
let finished_clone = Rc::clone(&finished);
spawner
.spawn_local(
async move {
token.cancelled().await;
finished_clone.set(true);
}
.map(|_| ()),
)
.unwrap();
cts.cancel();
pool.run();
assert!(finished.get());
}
#[test]
fn multiple_callbacks_and_idempotent_cancel() {
let cts = CancellationTokenSource::new();
let token = cts.token();
let flags: Vec<_> = (0..3).map(|_| Rc::new(Cell::new(false))).collect();
let regs: Vec<_> = flags
.iter()
.map(|flag| {
let f = Rc::clone(flag);
token
.register(move || {
f.set(true);
})
.unwrap()
})
.collect();
cts.cancel();
for flag in &flags {
assert!(flag.get());
}
cts.cancel();
for flag in &flags {
assert!(flag.get());
}
drop(regs); }
}