use serde::Deserialize;
use std::io::{self, Read, Write};
use std::os::unix::process::{CommandExt, ExitStatusExt};
use std::process::{Command, Stdio};
use std::sync::OnceLock;
use std::sync::atomic::{AtomicBool, Ordering};
use std::thread;
use std::time::{Duration, Instant};
use thiserror::Error;
const CASCADE_DEADLINE: Duration = Duration::from_millis(2000);
const POLL_INTERVAL: Duration = Duration::from_millis(20);
#[derive(Debug, Deserialize)]
#[serde(deny_unknown_fields)]
struct Input {
command: String,
}
#[derive(Debug, Error)]
pub enum Error {
#[error("invalid input JSON: {0}")]
InvalidJson(#[source] serde_json::Error),
#[error("read input from stdin: {0}")]
StdinRead(#[source] io::Error),
#[error("spawn shell: {0}")]
Spawn(#[source] io::Error),
#[error("wait shell: {0}")]
Wait(#[source] io::Error),
#[error("write to stdout: {0}")]
Stdout(#[source] io::Error),
#[error("write to stderr: {0}")]
Stderr(#[source] io::Error),
}
#[rustfmt::skip]
pub fn run<R: Read, W: Write, E: Write>(
stdin: &mut R, stdout: &mut W, stderr: &mut E,
) -> Result<i32, Error> {
install_sigterm_handler();
run_with(stdin, stdout, stderr, "sh", sigterm_flag(), CASCADE_DEADLINE)
}
#[doc(hidden)]
pub(crate) fn run_with<R: Read, W: Write, E: Write>(
stdin: &mut R,
stdout: &mut W,
stderr: &mut E,
shell: &str,
stop: &AtomicBool,
deadline: Duration,
) -> Result<i32, Error> {
let mut buf = Vec::new();
stdin.read_to_end(&mut buf).map_err(Error::StdinRead)?;
let input: Input = serde_json::from_slice(&buf).map_err(Error::InvalidJson)?;
let mut cmd = Command::new(shell);
cmd.arg("-c")
.arg(&input.command)
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped());
unsafe {
cmd.pre_exec(enter_own_process_group);
}
let mut child = cmd.spawn().map_err(Error::Spawn)?;
let pgid = child.id() as i32;
let mut child_stdout = child.stdout.take().expect("piped");
let mut child_stderr = child.stderr.take().expect("piped");
let stdout_thread = thread::spawn(move || {
let mut buf = Vec::new();
let _ = child_stdout.read_to_end(&mut buf);
buf
});
let stderr_thread = thread::spawn(move || {
let mut buf = Vec::new();
let _ = child_stderr.read_to_end(&mut buf);
buf
});
let status = wait_with_cascade(&mut child, pgid, stop, deadline)?;
let captured_stdout = stdout_thread.join().expect("stdout reader did not panic");
let captured_stderr = stderr_thread.join().expect("stderr reader did not panic");
stdout.write_all(&captured_stdout).map_err(Error::Stdout)?;
stderr.write_all(&captured_stderr).map_err(Error::Stderr)?;
let exit_code = status
.code()
.or_else(|| status.signal().map(|sig| 128 + sig))
.unwrap_or(1);
Ok(exit_code)
}
fn enter_own_process_group() -> io::Result<()> {
unsafe {
libc::setpgid(0, 0);
}
Ok(())
}
fn wait_with_cascade(
child: &mut std::process::Child,
pgid: i32,
stop: &AtomicBool,
deadline: Duration,
) -> Result<std::process::ExitStatus, Error> {
loop {
if let Some(status) = child.try_wait().map_err(Error::Wait)? {
return Ok(status);
}
thread::sleep(POLL_INTERVAL);
if stop.load(Ordering::SeqCst) {
return cascade_terminate(child, pgid, deadline);
}
}
}
fn cascade_terminate(
child: &mut std::process::Child,
pgid: i32,
deadline: Duration,
) -> Result<std::process::ExitStatus, Error> {
unsafe {
libc::kill(-pgid, libc::SIGTERM);
}
let term_until = Instant::now() + deadline;
while Instant::now() < term_until {
if let Some(status) = child.try_wait().map_err(Error::Wait)? {
return Ok(status);
}
thread::sleep(POLL_INTERVAL);
}
unsafe {
libc::kill(-pgid, libc::SIGKILL);
}
child.wait().map_err(Error::Wait)
}
static SIGTERM_FLAG: AtomicBool = AtomicBool::new(false);
static HANDLER_INSTALLED: OnceLock<()> = OnceLock::new();
extern "C" fn on_sigterm(_signo: libc::c_int) {
SIGTERM_FLAG.store(true, Ordering::SeqCst);
}
fn install_sigterm_handler() {
HANDLER_INSTALLED.get_or_init(|| {
unsafe {
libc::signal(libc::SIGTERM, on_sigterm as *const () as libc::sighandler_t);
}
});
}
fn sigterm_flag() -> &'static AtomicBool {
&SIGTERM_FLAG
}
#[cfg(test)]
mod tests;