use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use thiserror::Error;
#[derive(Debug, Clone, Copy, Error, PartialEq, Eq)]
#[error("canceling statement due to user request")]
pub struct QueryCancelled;
pub const SQLSTATE_QUERY_CANCELED: &str = "57014";
#[derive(Debug, Clone, Default)]
pub struct CancellationToken {
flag: Arc<AtomicBool>,
}
impl CancellationToken {
pub fn new() -> Self {
Self {
flag: Arc::new(AtomicBool::new(false)),
}
}
pub fn cancel(&self) {
self.flag.store(true, Ordering::Release);
}
pub fn reset(&self) {
self.flag.store(false, Ordering::Release);
}
pub fn is_cancelled(&self) -> bool {
self.flag.load(Ordering::Acquire)
}
pub fn check(&self) -> Result<(), QueryCancelled> {
if self.is_cancelled() {
Err(QueryCancelled)
} else {
Ok(())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn fresh_token_is_not_cancelled() {
let tok = CancellationToken::new();
assert!(!tok.is_cancelled());
assert!(tok.check().is_ok());
}
#[test]
fn cancel_propagates_through_clone() {
let tok = CancellationToken::new();
let observer = tok.clone();
tok.cancel();
assert!(observer.is_cancelled());
assert_eq!(observer.check(), Err(QueryCancelled));
}
#[test]
fn reset_clears_signal() {
let tok = CancellationToken::new();
tok.cancel();
tok.reset();
assert!(!tok.is_cancelled());
assert!(tok.check().is_ok());
}
#[test]
fn cancel_visible_across_threads() {
let tok = CancellationToken::new();
let worker = tok.clone();
let handle = thread::spawn(move || {
while !worker.is_cancelled() {
std::hint::spin_loop();
}
worker.check()
});
tok.cancel();
let res = handle.join().unwrap();
assert_eq!(res, Err(QueryCancelled));
}
}