use parking_lot::Mutex;
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
struct CancelInner {
cancelled: AtomicBool,
wakers: Mutex<Vec<Option<Waker>>>,
}
pub struct CancelToken(Arc<CancelInner>);
impl CancelToken {
pub fn new() -> Self {
CancelToken(Arc::new(CancelInner {
cancelled: AtomicBool::new(false),
wakers: Mutex::new(Vec::new()),
}))
}
pub fn cancel(&self) {
if self.0.cancelled.swap(true, Ordering::SeqCst) {
return;
}
let waiters = std::mem::take(&mut *self.0.wakers.lock());
for waker in waiters.into_iter().flatten() {
waker.wake();
}
}
pub fn is_cancelled(&self) -> bool {
self.0.cancelled.load(Ordering::SeqCst)
}
pub fn cancelled(&self) -> Cancelled {
Cancelled {
token: self.clone(),
slot: None,
}
}
}
impl Clone for CancelToken {
fn clone(&self) -> Self {
CancelToken(Arc::clone(&self.0))
}
}
impl Default for CancelToken {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for CancelToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CancelToken")
.field("cancelled", &self.is_cancelled())
.finish()
}
}
pub struct Cancelled {
token: CancelToken,
slot: Option<usize>,
}
impl Future for Cancelled {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> {
let this = self.get_mut();
let mut wakers = this.token.0.wakers.lock();
if this.token.is_cancelled() {
return Poll::Ready(());
}
match this.slot {
Some(i) => {
if !wakers[i].as_ref().is_some_and(|w| w.will_wake(cx.waker())) {
wakers[i] = Some(cx.waker().clone());
}
}
None => {
let i = match wakers.iter().position(|w| w.is_none()) {
Some(free) => free,
None => {
wakers.push(None);
wakers.len() - 1
}
};
wakers[i] = Some(cx.waker().clone());
this.slot = Some(i);
}
}
Poll::Pending
}
}
impl Drop for Cancelled {
fn drop(&mut self) {
let Some(i) = self.slot else { return };
if let Some(entry) = self.token.0.wakers.lock().get_mut(i) {
*entry = None;
}
}
}
impl fmt::Debug for Cancelled {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Cancelled").field("token", &self.token).finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::task::{RawWaker, RawWakerVTable, Waker};
fn noop_waker() -> Waker {
fn clone(_: *const ()) -> RawWaker {
RawWaker::new(std::ptr::null(), &VTABLE)
}
fn noop(_: *const ()) {}
static VTABLE: RawWakerVTable = RawWakerVTable::new(clone, noop, noop, noop);
unsafe { Waker::from_raw(RawWaker::new(std::ptr::null(), &VTABLE)) }
}
fn poll_once(fut: &mut Cancelled) -> Poll<()> {
let waker = noop_waker();
Pin::new(fut).poll(&mut Context::from_waker(&waker))
}
#[test]
fn repolling_one_waiter_registers_a_single_slot() {
let token = CancelToken::new();
let mut waiter = token.cancelled();
for _ in 0..100 {
assert!(poll_once(&mut waiter).is_pending());
}
assert_eq!(token.0.wakers.lock().len(), 1);
}
#[test]
fn a_dropped_waiter_frees_its_slot_for_reuse() {
let token = CancelToken::new();
for _ in 0..100 {
let mut waiter = token.cancelled();
assert!(poll_once(&mut waiter).is_pending());
assert_eq!(token.0.wakers.lock().len(), 1);
}
assert_eq!(token.0.wakers.lock().iter().filter(|w| w.is_some()).count(), 0);
}
#[test]
fn concurrent_waiters_each_get_their_own_slot() {
let token = CancelToken::new();
let mut a = token.cancelled();
let mut b = token.cancelled();
assert!(poll_once(&mut a).is_pending());
assert!(poll_once(&mut b).is_pending());
assert_eq!(token.0.wakers.lock().len(), 2);
}
#[test]
fn a_cancelled_token_resolves_without_registering() {
let token = CancelToken::new();
token.cancel();
let mut waiter = token.cancelled();
assert!(poll_once(&mut waiter).is_ready());
assert!(token.0.wakers.lock().is_empty());
}
}