use alloc::sync::Arc;
use core::cell::RefCell;
use core::fmt;
use core::sync::atomic::{AtomicBool, Ordering};
type Poll = dyn Fn() -> bool + Send + Sync;
type Progress = dyn Fn(&'static str, f64) + Send + Sync;
pub mod stage {
pub const DETECTING: &str = "detecting stars";
pub const SEARCHING: &str = "searching";
pub const BLIND_INDEX: &str = "blind index";
}
#[derive(Clone)]
pub struct CancelToken {
flag: Arc<AtomicBool>,
poll: Option<Arc<Poll>>,
progress: Option<Arc<Progress>>,
}
impl fmt::Debug for CancelToken {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CancelToken")
.field("cancelled", &self.flag.load(Ordering::Relaxed))
.field("poll", &self.poll.is_some())
.field("progress", &self.progress.is_some())
.finish()
}
}
impl Default for CancelToken {
fn default() -> Self {
Self::new()
}
}
impl CancelToken {
#[must_use]
pub fn new() -> Self {
Self {
flag: Arc::new(AtomicBool::new(false)),
poll: None,
progress: None,
}
}
#[must_use]
pub fn with_poll(poll: impl Fn() -> bool + Send + Sync + 'static) -> Self {
Self {
poll: Some(Arc::new(poll)),
..Self::new()
}
}
#[must_use]
pub fn with_progress(
self,
progress: impl Fn(&'static str, f64) + Send + Sync + 'static,
) -> Self {
Self {
progress: Some(Arc::new(progress)),
..self
}
}
pub fn cancel(&self) {
self.flag.store(true, Ordering::Relaxed);
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
if self.flag.load(Ordering::Relaxed) {
return true;
}
if let Some(poll) = &self.poll
&& poll()
{
self.flag.store(true, Ordering::Relaxed);
return true;
}
false
}
pub fn progress(&self, stage: &'static str, fraction: f64) {
if let Some(p) = &self.progress {
p(stage, fraction);
}
}
}
std::thread_local! {
static CURRENT: RefCell<Option<CancelToken>> = const { RefCell::new(None) };
}
struct Restore(Option<CancelToken>);
impl Drop for Restore {
fn drop(&mut self) {
let prev = self.0.take();
CURRENT.with(|c| *c.borrow_mut() = prev);
}
}
pub fn with_token<R>(token: &CancelToken, f: impl FnOnce() -> R) -> R {
let prev = CURRENT.with(|c| c.borrow_mut().replace(token.clone()));
let _restore = Restore(prev);
f()
}
pub fn with_optional<R>(token: Option<&CancelToken>, f: impl FnOnce() -> R) -> R {
match token {
Some(t) => with_token(t, f),
None => f(),
}
}
#[must_use]
pub fn current() -> Option<CancelToken> {
CURRENT.with(|c| c.borrow().clone())
}
#[must_use]
pub fn is_cancelled() -> bool {
CURRENT.with(|c| c.borrow().as_ref().is_some_and(CancelToken::is_cancelled))
}
pub fn progress(stage: &'static str, fraction: f64) {
CURRENT.with(|c| {
if let Some(t) = c.borrow().as_ref() {
t.progress(stage, fraction);
}
});
}
pub(crate) fn fired(token: Option<&CancelToken>) -> bool {
token.is_some_and(CancelToken::is_cancelled)
}
#[cfg(test)]
mod tests {
use super::*;
use core::sync::atomic::AtomicUsize;
#[test]
fn clones_share_the_flag() {
let a = CancelToken::new();
let b = a.clone();
assert!(!a.is_cancelled());
b.cancel();
assert!(a.is_cancelled());
}
#[test]
fn a_poll_function_cancels_and_sticks() {
let calls = Arc::new(AtomicUsize::new(0));
let c = Arc::clone(&calls);
let t = CancelToken::with_poll(move || c.fetch_add(1, Ordering::Relaxed) >= 2);
assert!(!t.is_cancelled());
assert!(!t.is_cancelled());
assert!(t.is_cancelled());
assert!(t.is_cancelled());
assert_eq!(calls.load(Ordering::Relaxed), 3, "not polled once fired");
}
#[test]
fn progress_reaches_the_observer() {
let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
let s = Arc::clone(&seen);
let t = CancelToken::with_poll(|| false)
.with_progress(move |stage, f| s.lock().unwrap().push((stage, f)));
with_token(&t, || progress(stage::SEARCHING, 0.5));
t.progress(stage::DETECTING, -1.0);
progress(stage::SEARCHING, 0.9); assert_eq!(
*seen.lock().unwrap(),
[(stage::SEARCHING, 0.5), (stage::DETECTING, -1.0)]
);
assert!(!t.is_cancelled());
}
#[test]
fn the_ambient_token_is_scoped_and_nests() {
assert!(current().is_none());
let outer = CancelToken::new();
let inner = CancelToken::new();
inner.cancel();
with_token(&outer, || {
assert!(!is_cancelled());
with_token(&inner, || assert!(is_cancelled()));
assert!(!is_cancelled(), "outer restored");
});
assert!(current().is_none());
assert!(!is_cancelled());
}
#[test]
fn the_ambient_token_is_restored_after_a_panic() {
let t = CancelToken::new();
let r = std::panic::catch_unwind(core::panic::AssertUnwindSafe(|| {
with_token(&t, || panic!("boom"));
}));
assert!(r.is_err());
assert!(current().is_none());
}
#[test]
fn other_threads_do_not_see_it() {
let t = CancelToken::new();
t.cancel();
with_token(&t, || {
let seen = std::thread::spawn(is_cancelled).join().unwrap();
assert!(!seen);
});
}
}