Skip to main content

arcsec_core/
cancel.rs

1//! Cooperative cancellation of a running solve, and progress reports from it.
2//!
3//! A solve can take seconds (a wide spiral search, a blind index pass), and a host
4//! application — an imaging suite with a Stop button — needs to end one early
5//! without killing a thread. Cancellation here is cooperative: the solver polls a
6//! [`CancelToken`] at its natural checkpoints (each spiral position, each blind
7//! pass, each index hypothesis) and returns [`ArcsecError::Cancelled`] at the next
8//! one after the token fires. Polling is one relaxed atomic load, so the cost is
9//! nothing measurable.
10//!
11//! The token is *ambient* rather than a parameter: [`with_token`] installs it for
12//! the duration of a closure on the calling thread, and the solvers read it with
13//! [`current`] when they start, then hand it to any worker threads they spawn. That
14//! keeps every existing signature (and every `SolveParams` literal) unchanged, and a
15//! call made without a token behaves exactly as before.
16//!
17//! A token can also carry a progress observer ([`CancelToken::with_progress`]),
18//! told which [`stage`] the solve is in and, for the spiral search, how far through
19//! it is. Progress messages for people go through the `log` crate as before; this
20//! is for a progress bar.
21//!
22//! ```
23//! use arcsec_core::cancel::{CancelToken, with_token};
24//!
25//! let token = CancelToken::new();
26//! let stop = token.clone(); // hand this to the UI thread
27//! stop.cancel();
28//! let cancelled = with_token(&token, || arcsec_core::cancel::is_cancelled());
29//! assert!(cancelled);
30//! ```
31//!
32//! [`ArcsecError::Cancelled`]: crate::ArcsecError::Cancelled
33
34use alloc::sync::Arc;
35use core::cell::RefCell;
36use core::fmt;
37use core::sync::atomic::{AtomicBool, Ordering};
38
39/// A poll function: returns `true` once the solve should stop.
40type Poll = dyn Fn() -> bool + Send + Sync;
41/// A progress observer: a [`stage`] name and a fraction in `[0, 1]`, or a negative
42/// value when the stage's extent is unknown.
43type Progress = dyn Fn(&'static str, f64) + Send + Sync;
44
45/// Names of the stages a solve reports to a progress observer.
46pub mod stage {
47    /// Finding and measuring the image's stars. Fraction unknown.
48    pub const DETECTING: &str = "detecting stars";
49    /// The spiral search around the hint; the fraction is of the positions within
50    /// the search radius. A solve usually ends well before 1.
51    pub const SEARCHING: &str = "searching";
52    /// A blind index looking for the field. Fraction unknown.
53    pub const BLIND_INDEX: &str = "blind index";
54}
55
56/// A shared, thread-safe "please stop" flag, with an optional progress observer.
57///
58/// Clones share the flag: cancel any clone and every holder sees it. Optionally it
59/// also polls a caller-supplied function (see [`CancelToken::with_poll`]), so a
60/// host that already keeps its own stop flag need not mirror it into this one.
61#[derive(Clone)]
62pub struct CancelToken {
63    flag: Arc<AtomicBool>,
64    poll: Option<Arc<Poll>>,
65    progress: Option<Arc<Progress>>,
66}
67
68impl fmt::Debug for CancelToken {
69    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
70        f.debug_struct("CancelToken")
71            .field("cancelled", &self.flag.load(Ordering::Relaxed))
72            .field("poll", &self.poll.is_some())
73            .field("progress", &self.progress.is_some())
74            .finish()
75    }
76}
77
78impl Default for CancelToken {
79    fn default() -> Self {
80        Self::new()
81    }
82}
83
84impl CancelToken {
85    /// A token that has not been cancelled.
86    #[must_use]
87    pub fn new() -> Self {
88        Self {
89            flag: Arc::new(AtomicBool::new(false)),
90            poll: None,
91            progress: None,
92        }
93    }
94
95    /// A token that is also cancelled once `poll` returns `true`.
96    ///
97    /// `poll` is called from whichever thread reaches a checkpoint, worker threads
98    /// included, possibly from several at once, so it must be cheap and
99    /// thread-safe. Once it has returned `true` the token stays cancelled and
100    /// `poll` is not called again.
101    #[must_use]
102    pub fn with_poll(poll: impl Fn() -> bool + Send + Sync + 'static) -> Self {
103        Self {
104            poll: Some(Arc::new(poll)),
105            ..Self::new()
106        }
107    }
108
109    /// This token (still sharing its flag and poll function with its clones),
110    /// with `progress` told of the solve's progress: a [`stage`] name and a
111    /// fraction in `[0, 1]`, or `-1` when the stage's extent is unknown.
112    ///
113    /// Called from the solver's worker threads, possibly several at once, so it
114    /// must be thread-safe. The search reports at most a few hundred times.
115    #[must_use]
116    pub fn with_progress(
117        self,
118        progress: impl Fn(&'static str, f64) + Send + Sync + 'static,
119    ) -> Self {
120        Self {
121            progress: Some(Arc::new(progress)),
122            ..self
123        }
124    }
125
126    /// Ask every solve using this token to stop at its next checkpoint.
127    pub fn cancel(&self) {
128        self.flag.store(true, Ordering::Relaxed);
129    }
130
131    /// Whether the token has been cancelled (or its poll function says so).
132    #[must_use]
133    pub fn is_cancelled(&self) -> bool {
134        if self.flag.load(Ordering::Relaxed) {
135            return true;
136        }
137        if let Some(poll) = &self.poll
138            && poll()
139        {
140            self.flag.store(true, Ordering::Relaxed);
141            return true;
142        }
143        false
144    }
145
146    /// Report progress to the observer, if there is one.
147    pub fn progress(&self, stage: &'static str, fraction: f64) {
148        if let Some(p) = &self.progress {
149            p(stage, fraction);
150        }
151    }
152}
153
154std::thread_local! {
155    static CURRENT: RefCell<Option<CancelToken>> = const { RefCell::new(None) };
156}
157
158/// Restores the previous ambient token when dropped, panics included.
159struct Restore(Option<CancelToken>);
160
161impl Drop for Restore {
162    fn drop(&mut self) {
163        let prev = self.0.take();
164        CURRENT.with(|c| *c.borrow_mut() = prev);
165    }
166}
167
168/// Run `f` with `token` as this thread's ambient cancellation token.
169///
170/// Solves started inside `f` on this thread poll `token`, and pass it on to the
171/// worker threads they spawn. Calls nest; the previous token is restored when `f`
172/// returns or unwinds.
173pub fn with_token<R>(token: &CancelToken, f: impl FnOnce() -> R) -> R {
174    let prev = CURRENT.with(|c| c.borrow_mut().replace(token.clone()));
175    let _restore = Restore(prev);
176    f()
177}
178
179/// Like [`with_token`], but with no token at all when `token` is `None`.
180pub fn with_optional<R>(token: Option<&CancelToken>, f: impl FnOnce() -> R) -> R {
181    match token {
182        Some(t) => with_token(t, f),
183        None => f(),
184    }
185}
186
187/// This thread's ambient token, if one is installed.
188#[must_use]
189pub fn current() -> Option<CancelToken> {
190    CURRENT.with(|c| c.borrow().clone())
191}
192
193/// Whether this thread's ambient token (if any) has been cancelled.
194#[must_use]
195pub fn is_cancelled() -> bool {
196    CURRENT.with(|c| c.borrow().as_ref().is_some_and(CancelToken::is_cancelled))
197}
198
199/// Report progress to this thread's ambient token's observer, if any.
200pub fn progress(stage: &'static str, fraction: f64) {
201    CURRENT.with(|c| {
202        if let Some(t) = c.borrow().as_ref() {
203            t.progress(stage, fraction);
204        }
205    });
206}
207
208/// `Some(token)` cancelled, as a predicate the hot loops can capture by reference.
209pub(crate) fn fired(token: Option<&CancelToken>) -> bool {
210    token.is_some_and(CancelToken::is_cancelled)
211}
212
213#[cfg(test)]
214mod tests {
215    use super::*;
216    use core::sync::atomic::AtomicUsize;
217
218    #[test]
219    fn clones_share_the_flag() {
220        let a = CancelToken::new();
221        let b = a.clone();
222        assert!(!a.is_cancelled());
223        b.cancel();
224        assert!(a.is_cancelled());
225    }
226
227    #[test]
228    fn a_poll_function_cancels_and_sticks() {
229        let calls = Arc::new(AtomicUsize::new(0));
230        let c = Arc::clone(&calls);
231        let t = CancelToken::with_poll(move || c.fetch_add(1, Ordering::Relaxed) >= 2);
232        assert!(!t.is_cancelled());
233        assert!(!t.is_cancelled());
234        assert!(t.is_cancelled());
235        assert!(t.is_cancelled());
236        assert_eq!(calls.load(Ordering::Relaxed), 3, "not polled once fired");
237    }
238
239    #[test]
240    fn progress_reaches_the_observer() {
241        let seen = Arc::new(std::sync::Mutex::new(Vec::new()));
242        let s = Arc::clone(&seen);
243        let t = CancelToken::with_poll(|| false)
244            .with_progress(move |stage, f| s.lock().unwrap().push((stage, f)));
245        with_token(&t, || progress(stage::SEARCHING, 0.5));
246        t.progress(stage::DETECTING, -1.0);
247        progress(stage::SEARCHING, 0.9); // no ambient token: nowhere to go
248        assert_eq!(
249            *seen.lock().unwrap(),
250            [(stage::SEARCHING, 0.5), (stage::DETECTING, -1.0)]
251        );
252        // The poll function survived.
253        assert!(!t.is_cancelled());
254    }
255
256    #[test]
257    fn the_ambient_token_is_scoped_and_nests() {
258        assert!(current().is_none());
259        let outer = CancelToken::new();
260        let inner = CancelToken::new();
261        inner.cancel();
262        with_token(&outer, || {
263            assert!(!is_cancelled());
264            with_token(&inner, || assert!(is_cancelled()));
265            assert!(!is_cancelled(), "outer restored");
266        });
267        assert!(current().is_none());
268        assert!(!is_cancelled());
269    }
270
271    #[test]
272    fn the_ambient_token_is_restored_after_a_panic() {
273        let t = CancelToken::new();
274        let r = std::panic::catch_unwind(core::panic::AssertUnwindSafe(|| {
275            with_token(&t, || panic!("boom"));
276        }));
277        assert!(r.is_err());
278        assert!(current().is_none());
279    }
280
281    #[test]
282    fn other_threads_do_not_see_it() {
283        let t = CancelToken::new();
284        t.cancel();
285        with_token(&t, || {
286            let seen = std::thread::spawn(is_cancelled).join().unwrap();
287            assert!(!seen);
288        });
289    }
290}