use std::sync::atomic::{AtomicU8, Ordering};
use std::sync::{Arc, LazyLock, Mutex};
static CTRLC_HANDLERS_NEW: Mutex<Vec<Box<dyn FnMut() + Send>>> = Mutex::new(Vec::new());
static CTRLC_SIGNAL_STACK: Mutex<Vec<CtrlcFrame>> = Mutex::new(Vec::new());
static INIT_ONCE: LazyLock<Result<(), String>> = LazyLock::new(|| {
let mut handlers = vec![];
let set_result = ctrlc::try_set_handler(move || {
if let Ok(mut new_handlers) = CTRLC_HANDLERS_NEW.lock() {
handlers.extend(new_handlers.drain(..));
}
let mut signalled = false;
if let Ok(stack) = CTRLC_SIGNAL_STACK.lock() {
if let Some(frame) = stack.last() {
signalled = true;
frame.signal.signal();
if let Some(f) = &frame.on_signal {
f(frame.signal.clone())
}
}
}
if !signalled && handlers.is_empty() {
std::process::exit(1);
}
for handler in handlers.iter_mut().rev() {
handler();
}
});
match set_result {
Err(ctrlc::Error::MultipleHandlers) => {
Err("failed to set ctrl-c handler: a handler is already set using the `ctrlc` crate. please set with cu::cli instead (see documentation for more information)".to_string())
}
Err(other_error) => {
Err(format!("failed to set ctrl-c handler: {other_error}"))
}
Ok(_) => Ok(())
}
});
#[cfg(feature = "cli")]
pub fn add_global_ctrlc_handler<F: FnMut() + Send + 'static>(handler: F) -> cu::Result<()> {
{
let Ok(mut handlers) = CTRLC_HANDLERS_NEW.lock() else {
cu::bail!("global ctrl-c handler vector is poisoned");
};
handlers.push(Box::new(handler));
}
if let Err(e) = &*INIT_ONCE {
cu::bail!("{e}");
}
Ok(())
}
#[cfg_attr(not(feature = "coroutine"), doc = "```rust,ignore")]
#[cfg_attr(feature = "coroutine", doc = "```rust,no_run")]
#[inline(always)]
pub fn ctrlc_frame() -> CtrlcBuilder {
CtrlcBuilder::default()
}
pub struct CtrlcBuilder {
abort_threshold: u8,
on_signal: Option<OnSignalFn>,
}
impl Default for CtrlcBuilder {
#[inline(always)]
fn default() -> Self {
Self {
abort_threshold: 1,
on_signal: None,
}
}
}
impl CtrlcBuilder {
#[inline(always)]
pub fn abort_threshold(mut self, threshold: u8) -> Self {
self.abort_threshold = threshold;
self
}
#[inline(always)]
pub fn on_signal<F: Fn(CtrlcSignal) + Send + 'static>(mut self, f: F) -> Self {
self.on_signal = Some(Box::new(f));
self
}
pub fn execute<T, F>(self, f: F) -> cu::Result<Option<T>>
where
F: FnOnce(CtrlcSignal) -> cu::Result<T>,
{
let signal = CtrlcSignal::new(self.abort_threshold);
let Some(ctrlc_frame_scope) = CtrlcFrame::push_scope(signal.clone(), self.on_signal) else {
return f(signal).map(Some);
};
if let Err(e) = &*INIT_ONCE {
cu::bail!("{e}");
}
let result = f(signal.clone());
if signal.should_abort() {
return Ok(None);
}
drop(ctrlc_frame_scope); result.map(Some)
}
#[cfg(feature = "coroutine")]
pub async fn co_execute<T, TFuture, F>(self, f: F) -> cu::Result<Option<T>>
where
TFuture: Future<Output = cu::Result<T>>,
F: FnOnce(CtrlcSignal) -> TFuture,
{
let signal = CtrlcSignal::new(self.abort_threshold);
let Some(ctrlc_frame_scope) = CtrlcFrame::push_scope(signal.clone(), self.on_signal) else {
return f(signal).await.map(Some);
};
if let Err(e) = &*INIT_ONCE {
cu::bail!("{e}");
}
let result = f(signal.clone()).await;
if signal.should_abort() {
return Ok(None);
}
drop(ctrlc_frame_scope); result.map(Some)
}
}
type OnSignalFn = Box<dyn Fn(CtrlcSignal) + Send>;
struct CtrlcFrame {
id: usize,
signal: CtrlcSignal,
on_signal: Option<OnSignalFn>,
}
struct CtrlcScope(usize);
#[derive(Clone)]
pub struct CtrlcSignal {
signaled_times: Arc<AtomicU8>,
abort_threshold: u8,
}
impl CtrlcFrame {
pub fn push_scope(signal: CtrlcSignal, on_signal: Option<OnSignalFn>) -> Option<CtrlcScope> {
let Ok(mut signal_stack) = CTRLC_SIGNAL_STACK.lock() else {
cu::trace!("failed to register new ctrl-c frame");
return None;
};
let id = crate::next_atomic_usize();
signal_stack.push(Self {
id,
signal,
on_signal,
});
Some(CtrlcScope(id))
}
}
impl Drop for CtrlcScope {
fn drop(&mut self) {
if let Ok(mut signal_stack) = CTRLC_SIGNAL_STACK.lock() {
signal_stack.retain(|x| x.id != self.0);
}
}
}
impl CtrlcSignal {
fn new(abort_threshold: u8) -> Self {
Self {
signaled_times: Arc::new(AtomicU8::new(0)),
abort_threshold,
}
}
pub fn check(&self) -> cu::Result<()> {
if self.should_abort() {
cu::bail!("interrupted")
}
Ok(())
}
pub fn should_abort(&self) -> bool {
self.signaled_times() >= self.abort_threshold
}
pub fn signaled(&self) -> bool {
self.signaled_times() > 0
}
pub fn signaled_times(&self) -> u8 {
self.signaled_times.load(Ordering::Acquire)
}
pub fn signal(&self) {
self.signaled_times.fetch_add(1, Ordering::SeqCst);
}
}