use anyhow::Result;
use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
use std::sync::Once;
static REGISTER_ONCE: Once = Once::new();
static CANCEL_FLAG: AtomicBool = AtomicBool::new(false);
static FLAG_SIGTERM: AtomicBool = AtomicBool::new(false);
static FORCE_EXIT: AtomicBool = AtomicBool::new(false);
static SIGNAL_HITS: AtomicU8 = AtomicU8::new(0);
pub fn register_handler() -> Result<()> {
let mut register_result: Result<()> = Ok(());
REGISTER_ONCE.call_once(|| {
register_result = register_handler_inner();
});
register_result
}
fn register_handler_inner() -> Result<()> {
ctrlc::set_handler(|| {
note_cooperative_signal(&CANCEL_FLAG, &FORCE_EXIT, "SIGINT");
})?;
tracing::debug!("Ctrl+C (SIGINT) handler registered successfully");
#[cfg(unix)]
{
unsafe {
signal_hook::low_level::register(signal_hook::consts::SIGTERM, || {
FLAG_SIGTERM.store(true, Ordering::Release);
let prev = record_signal_hit();
if prev >= 1 {
FORCE_EXIT.store(true, Ordering::Release);
}
})?;
}
tracing::debug!("SIGTERM handler registered");
}
#[cfg(not(unix))]
{
}
Ok(())
}
#[inline]
fn record_signal_hit() -> u8 {
SIGNAL_HITS.fetch_add(1, Ordering::Relaxed)
}
fn note_cooperative_signal(cancel: &AtomicBool, force: &AtomicBool, kind: &str) {
let prev = record_signal_hit();
if prev == 0 {
cancel.store(true, Ordering::Release);
tracing::debug!(signal = kind, "cancellation signal received");
} else {
force.store(true, Ordering::Release);
cancel.store(true, Ordering::Release);
tracing::debug!(signal = kind, "force-exit signal (second hit)");
}
}
#[must_use]
pub fn is_cancelled() -> bool {
CANCEL_FLAG.load(Ordering::Acquire)
}
#[must_use]
pub fn is_terminated() -> bool {
FLAG_SIGTERM.load(Ordering::Acquire)
}
#[must_use]
pub fn should_stop() -> bool {
is_cancelled() || is_terminated()
}
#[must_use]
pub fn is_force_exit() -> bool {
FORCE_EXIT.load(Ordering::Acquire)
}
#[must_use]
pub fn cancellation_flag() -> &'static AtomicBool {
&CANCEL_FLAG
}
#[must_use]
pub fn sigterm_flag() -> &'static AtomicBool {
&FLAG_SIGTERM
}
#[must_use]
pub fn force_exit_flag() -> &'static AtomicBool {
&FORCE_EXIT
}
#[must_use]
pub fn signal_exit_code() -> Option<i32> {
if is_terminated() {
Some(crate::errors::exit_codes::EX_SIGTERM)
} else if is_cancelled() {
Some(crate::errors::exit_codes::EX_SIGINT)
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
use serial_test::serial;
fn reset_flags_for_test() {
cancellation_flag().store(false, Ordering::Release);
sigterm_flag().store(false, Ordering::Release);
force_exit_flag().store(false, Ordering::Release);
SIGNAL_HITS.store(0, Ordering::Relaxed);
}
#[test]
#[serial]
fn is_cancelled_false_before_signal() {
reset_flags_for_test();
assert!(!is_cancelled());
}
#[test]
#[serial]
fn cancellation_flag_is_stable_static() {
let flag_a = cancellation_flag();
let flag_b = cancellation_flag();
assert!(std::ptr::eq(flag_a, flag_b));
}
#[test]
#[serial]
fn flag_can_be_set_and_read() {
let flag = cancellation_flag();
let previous_value = flag.load(Ordering::Acquire);
flag.store(previous_value, Ordering::Release);
assert_eq!(flag.load(Ordering::Acquire), previous_value);
}
#[test]
#[serial]
fn is_terminated_false_by_default() {
reset_flags_for_test();
assert!(!is_terminated());
}
#[test]
#[serial]
fn sigterm_flag_is_stable_static() {
let a = sigterm_flag();
let b = sigterm_flag();
assert!(std::ptr::eq(a, b));
}
#[test]
#[serial]
fn is_terminated_true_after_set() {
let flag = sigterm_flag();
flag.store(true, Ordering::Release);
assert!(is_terminated());
flag.store(false, Ordering::Release);
}
#[test]
#[serial]
fn is_cancelled_false_after_reset() {
let flag = cancellation_flag();
flag.store(true, Ordering::Release);
assert!(is_cancelled());
flag.store(false, Ordering::Release);
assert!(!is_cancelled());
}
#[test]
#[serial]
fn should_stop_true_when_either_flag_set() {
reset_flags_for_test();
assert!(!should_stop());
cancellation_flag().store(true, Ordering::Release);
assert!(should_stop());
cancellation_flag().store(false, Ordering::Release);
sigterm_flag().store(true, Ordering::Release);
assert!(should_stop());
sigterm_flag().store(false, Ordering::Release);
assert!(!should_stop());
}
#[test]
#[serial]
fn note_cooperative_signal_sets_force_on_second_hit() {
reset_flags_for_test();
let cancel = cancellation_flag();
let force = force_exit_flag();
note_cooperative_signal(cancel, force, "SIGINT");
assert!(is_cancelled());
assert!(!is_force_exit());
note_cooperative_signal(cancel, force, "SIGINT");
assert!(is_force_exit());
reset_flags_for_test();
}
#[test]
#[serial]
fn signal_exit_code_prefers_sigterm() {
reset_flags_for_test();
assert_eq!(signal_exit_code(), None);
cancellation_flag().store(true, Ordering::Release);
assert_eq!(
signal_exit_code(),
Some(crate::errors::exit_codes::EX_SIGINT)
);
sigterm_flag().store(true, Ordering::Release);
assert_eq!(
signal_exit_code(),
Some(crate::errors::exit_codes::EX_SIGTERM)
);
reset_flags_for_test();
}
#[test]
#[serial]
fn register_handler_is_idempotent() {
let _ = register_handler();
assert!(register_handler().is_ok());
}
#[test]
#[serial]
fn record_signal_hit_escalates_on_second() {
reset_flags_for_test();
assert_eq!(record_signal_hit(), 0);
assert_eq!(record_signal_hit(), 1);
reset_flags_for_test();
}
}