use std::thread;
use std::time::{Duration, Instant};
pub trait Signaler {
fn term(&self, pgid: i32);
fn kill(&self, pgid: i32);
fn alive(&self, pid: i32) -> bool;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct RealSignaler;
impl Signaler for RealSignaler {
fn term(&self, pgid: i32) {
unsafe {
libc::kill(-pgid, libc::SIGTERM);
}
}
fn kill(&self, pgid: i32) {
unsafe {
libc::kill(-pgid, libc::SIGKILL);
}
}
fn alive(&self, pid: i32) -> bool {
let r = unsafe { libc::kill(pid, 0) };
if r == 0 {
return true;
}
let errno = std::io::Error::last_os_error()
.raw_os_error()
.unwrap_or(libc::ESRCH);
errno != libc::ESRCH
}
}
pub fn cascade(pgids: &[i32], signaler: &dyn Signaler, deadline: Duration, poll: Duration) {
for pgid in pgids {
signaler.term(*pgid);
}
let term_until = Instant::now() + deadline;
while Instant::now() < term_until {
if pgids.iter().all(|pgid| !signaler.alive(*pgid)) {
return;
}
thread::sleep(poll);
}
for pgid in pgids {
if signaler.alive(*pgid) {
signaler.kill(*pgid);
}
}
}
#[cfg(test)]
#[derive(Default)]
pub(crate) struct RecordingSignaler {
pub(crate) invocations: std::sync::Mutex<Vec<(&'static str, i32)>>,
pub(crate) alive_polls_remaining: std::sync::atomic::AtomicI32,
}
#[cfg(test)]
impl RecordingSignaler {
pub(crate) fn new(alive_polls: i32) -> Self {
Self {
invocations: std::sync::Mutex::new(Vec::new()),
alive_polls_remaining: std::sync::atomic::AtomicI32::new(alive_polls),
}
}
pub(crate) fn took(&self) -> Vec<(&'static str, i32)> {
self.invocations.lock().unwrap().clone()
}
}
#[cfg(test)]
impl Signaler for RecordingSignaler {
fn term(&self, pgid: i32) {
self.invocations.lock().unwrap().push(("term", pgid));
}
fn kill(&self, pgid: i32) {
self.invocations.lock().unwrap().push(("kill", pgid));
}
fn alive(&self, _pid: i32) -> bool {
let prev = self
.alive_polls_remaining
.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
prev > 0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cascade_sigterm_drains_within_deadline_no_sigkill() {
let s = RecordingSignaler::new(0);
cascade(&[42], &s, Duration::from_secs(60), Duration::from_millis(1));
let invocations = s.took();
assert_eq!(invocations, vec![("term", 42)]);
}
#[test]
fn cascade_sigkill_after_deadline_when_still_alive() {
let s = RecordingSignaler::new(i32::MAX);
cascade(
&[42],
&s,
Duration::from_millis(20),
Duration::from_millis(1),
);
let invocations = s.took();
assert_eq!(invocations.first(), Some(&("term", 42)));
assert!(invocations.iter().any(|&(sig, _)| sig == "kill"));
let kill_count = invocations
.iter()
.filter(|&&(sig, _)| sig == "kill")
.count();
assert_eq!(kill_count, 1);
}
#[test]
fn cascade_signals_every_pgid_in_term_phase() {
let s = RecordingSignaler::new(0);
cascade(
&[1, 2, 3],
&s,
Duration::from_secs(60),
Duration::from_millis(1),
);
let invocations = s.took();
let term_targets: Vec<i32> = invocations
.iter()
.filter(|(sig, _)| *sig == "term")
.map(|(_, pgid)| *pgid)
.collect();
assert_eq!(term_targets, vec![1, 2, 3]);
}
fn spawn_in_own_pgid(args: &[&str], stdout: std::process::Stdio) -> std::process::Child {
use std::os::unix::process::CommandExt as _;
use std::process::{Command, Stdio};
let (program, rest) = args.split_first().expect("non-empty argv");
unsafe {
Command::new(program)
.args(rest)
.stdin(Stdio::null())
.stdout(stdout)
.stderr(Stdio::null())
.pre_exec(|| {
libc::setpgid(0, 0);
Ok(())
})
.spawn()
.expect("spawn child")
}
}
#[test]
fn real_signaler_term_kills_then_alive_reports_dead_after_reap() {
let mut child = spawn_in_own_pgid(&["sleep", "60"], std::process::Stdio::null());
let pid = child.id() as i32;
let s = RealSignaler;
assert!(
s.alive(pid),
"child should be alive immediately after spawn"
);
s.term(pid);
child.wait().unwrap();
assert!(!s.alive(pid), "alive should report dead after reap");
}
#[test]
fn real_signaler_kill_uncatchable_takes_down_sigterm_handler() {
use std::io::BufRead as _;
let mut child = spawn_in_own_pgid(
&[
"sh",
"-c",
"trap '' TERM; echo ready; while :; do sleep 1; done",
],
std::process::Stdio::piped(),
);
let pid = child.id() as i32;
let mut out = std::io::BufReader::new(child.stdout.take().expect("stdout is piped"));
let mut line = String::new();
out.read_line(&mut line).expect("read fixture handshake");
assert_eq!(line.trim(), "ready", "fixture reached its post-trap echo");
let s = RealSignaler;
s.term(pid);
assert!(s.alive(pid), "TERM is trapped; child should still be alive");
s.kill(pid);
child.wait().unwrap();
assert!(!s.alive(pid), "alive should report dead after reap");
}
#[test]
fn real_signaler_alive_returns_false_for_nonexistent_pid() {
let s = RealSignaler;
assert!(!s.alive(i32::MAX), "synthetic pid should be reported dead");
}
#[test]
fn cascade_only_kills_still_alive_pgids() {
let s = RecordingSignaler::new(2);
cascade(&[42], &s, Duration::from_secs(60), Duration::from_millis(1));
let invocations = s.took();
let kills: Vec<&(&'static str, i32)> = invocations
.iter()
.filter(|(sig, _)| *sig == "kill")
.collect();
assert!(kills.is_empty());
}
}