use std::io;
use std::process::{Child, Command, ExitStatus, Output, Stdio};
use std::sync::atomic::{AtomicI32, AtomicUsize, Ordering};
const SLOTS: usize = 512;
static CHILDREN: [AtomicI32; SLOTS] = [const { AtomicI32::new(0) }; SLOTS];
static GUARDS: AtomicUsize = AtomicUsize::new(0);
pub fn signal_registered_children(signal: i32) {
for slot in &CHILDREN {
let pid = slot.load(Ordering::SeqCst);
if pid > 0 {
#[cfg(unix)]
unsafe {
libc::kill(pid, signal);
}
#[cfg(not(unix))]
let _ = signal;
}
}
}
pub struct RegisteredChild {
slot: Option<usize>,
}
impl Drop for RegisteredChild {
fn drop(&mut self) {
if let Some(slot) = self.slot {
CHILDREN[slot].store(0, Ordering::SeqCst);
}
}
}
pub fn register(child: &Child) -> RegisteredChild {
let Ok(pid) = i32::try_from(child.id()) else {
return RegisteredChild { slot: None };
};
#[cfg(windows)]
windows_job::assign(child);
for (index, slot) in CHILDREN.iter().enumerate() {
if slot
.compare_exchange(0, pid, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
return RegisteredChild { slot: Some(index) };
}
}
RegisteredChild { slot: None }
}
#[cfg(windows)]
mod windows_job {
use std::process::Child;
use std::sync::Mutex;
use crate::process_supervision::JobHandle;
static JOB: Mutex<Option<JobHandle>> = Mutex::new(None);
fn job() -> std::sync::MutexGuard<'static, Option<JobHandle>> {
JOB.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn create() {
if let Ok(handle) = JobHandle::new() {
*job() = Some(handle);
}
}
pub fn release() {
*job() = None;
}
pub fn assign(child: &Child) {
if let Some(handle) = job().as_ref() {
let _ = handle.assign(child);
}
}
}
pub fn output(command: &mut Command) -> io::Result<Output> {
let child = command
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?;
let _registered = register(&child);
child.wait_with_output()
}
pub fn status(command: &mut Command) -> io::Result<ExitStatus> {
let mut child = command.spawn()?;
let _registered = register(&child);
child.wait()
}
#[cfg(unix)]
extern "C" fn on_signal(signal: libc::c_int) {
signal_registered_children(libc::SIGTERM);
unsafe {
libc::signal(signal, libc::SIG_DFL);
libc::raise(signal);
}
}
pub struct ChildSignalGuard {
#[cfg(unix)]
previous: Vec<(libc::c_int, libc::sighandler_t)>,
}
impl ChildSignalGuard {
pub fn install() -> io::Result<Self> {
#[cfg(unix)]
{
let mut previous = Vec::new();
if GUARDS.fetch_add(1, Ordering::SeqCst) == 0 {
for signal in [libc::SIGHUP, libc::SIGINT, libc::SIGTERM] {
let old = unsafe {
libc::signal(
signal,
on_signal as extern "C" fn(libc::c_int) as libc::sighandler_t,
)
};
if old == libc::SIG_ERR {
let error = io::Error::last_os_error();
for (installed, disposition) in previous.drain(..).rev() {
unsafe {
libc::signal(installed, disposition);
}
}
GUARDS.fetch_sub(1, Ordering::SeqCst);
return Err(error);
}
previous.push((signal, old));
}
}
Ok(Self { previous })
}
#[cfg(not(unix))]
{
if GUARDS.fetch_add(1, Ordering::SeqCst) == 0 {
#[cfg(windows)]
windows_job::create();
}
Ok(Self {})
}
}
}
impl Drop for ChildSignalGuard {
fn drop(&mut self) {
#[cfg(unix)]
for (signal, disposition) in self.previous.drain(..).rev() {
unsafe {
libc::signal(signal, disposition);
}
}
let outermost = GUARDS.fetch_sub(1, Ordering::SeqCst) == 1;
#[cfg(windows)]
if outermost {
windows_job::release();
}
#[cfg(not(windows))]
let _ = outermost;
}
}
#[cfg(all(test, unix))]
mod tests {
use std::process::{Command, Stdio};
use std::time::{Duration, Instant};
use super::*;
fn sleeper() -> Child {
Command::new("sleep")
.arg("30")
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.unwrap()
}
fn exits_within(child: &mut Child, limit: Duration) -> bool {
let started = Instant::now();
while started.elapsed() < limit {
if child.try_wait().unwrap().is_some() {
return true;
}
std::thread::sleep(Duration::from_millis(20));
}
false
}
#[test]
fn a_registered_child_receives_the_signal_and_a_released_one_does_not() {
let _guard = ChildSignalGuard::install().unwrap();
let mut tracked = sleeper();
let mut released = sleeper();
let registration = register(&tracked);
drop(register(&released));
signal_registered_children(libc::SIGTERM);
assert!(
exits_within(&mut tracked, Duration::from_secs(5)),
"the registered child was not signalled"
);
assert!(
released.try_wait().unwrap().is_none(),
"a child whose registration was dropped must not be signalled"
);
drop(registration);
released.kill().unwrap();
released.wait().unwrap();
}
}