use std::{
ffi::OsStr,
io,
num::NonZeroU32,
process::{Child, ChildStderr, ChildStdin, ChildStdout, Command, ExitStatus},
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
thread,
time::{Duration, Instant},
};
use process_wrap::std::{StdChildWrapper, StdCommandWrap, StdCommandWrapper};
#[derive(Debug)]
pub struct OwnedCommand {
command: Command,
windows_hidden: bool,
caller_job: CallerJob,
}
#[cfg(windows)]
type CallerJob = Option<std::os::windows::io::OwnedHandle>;
#[cfg(not(windows))]
type CallerJob = Option<std::convert::Infallible>;
impl OwnedCommand {
pub fn new(program: impl AsRef<OsStr>) -> Self {
Self::from_command(Command::new(program))
}
pub fn from_command(command: Command) -> Self {
Self {
command,
windows_hidden: true,
caller_job: Default::default(),
}
}
pub fn command_mut(&mut self) -> &mut Command {
&mut self.command
}
pub fn env_allowlist<I, S>(&mut self, keys: I) -> &mut Self
where
I: IntoIterator<Item = S>,
S: AsRef<OsStr>,
{
let kept: Vec<_> = keys
.into_iter()
.filter_map(|k| std::env::var_os(k.as_ref()).map(|v| (k.as_ref().to_os_string(), v)))
.collect();
self.command.env_clear();
for (k, v) in kept {
self.command.env(k, v);
}
self
}
pub fn env_strip<I, S>(&mut self, keys: I) -> &mut Self
where
I: IntoIterator<Item = S>,
S: AsRef<OsStr>,
{
for k in keys {
self.command.env_remove(k);
}
self
}
pub fn windows_hide(&mut self) -> &mut Self {
self.windows_hidden = true;
self
}
#[cfg(windows)]
pub fn windows_job(
&mut self,
job: std::os::windows::io::BorrowedHandle<'_>,
) -> io::Result<&mut Self> {
self.caller_job = Some(job.try_clone_to_owned()?);
Ok(self)
}
pub fn spawn(self) -> io::Result<OwnedChild> {
spawn_owned(
self.command,
self.windows_hidden,
false,
None,
self.caller_job,
)
}
}
pub const DEFAULT_ENV_ALLOWLIST: &[&str] = &[
"PATH",
"HOME",
"USER",
"LOGNAME",
"LANG",
"LC_ALL",
"TMPDIR",
"TEMP",
"TMP",
"SystemRoot",
"SYSTEMROOT",
"USERPROFILE",
"APPDATA",
"LOCALAPPDATA",
"ProgramData",
"COMSPEC",
"PATHEXT",
"XDG_DATA_HOME",
"XDG_RUNTIME_DIR",
];
#[derive(Debug)]
pub struct OwnedChild {
child: Option<Box<dyn StdChildWrapper>>,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ReapOutcome {
Exited(ExitStatus),
TimedOut,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum WaitOutcome {
Exited(ExitStatus),
Terminated(ExitStatus),
}
impl OwnedChild {
pub fn id(&self) -> u32 {
self.child().id()
}
pub fn take_stdin(&mut self) -> Option<ChildStdin> {
self.child_mut().stdin().take()
}
pub fn take_stdout(&mut self) -> Option<ChildStdout> {
self.child_mut().stdout().take()
}
pub fn take_stderr(&mut self) -> Option<ChildStderr> {
self.child_mut().stderr().take()
}
pub fn try_wait(&mut self) -> io::Result<Option<ExitStatus>> {
self.child_mut().try_wait()
}
pub fn wait(&mut self) -> io::Result<ExitStatus> {
self.child_mut().wait()
}
pub fn wait_timeout(&mut self, timeout: Duration) -> io::Result<Option<ExitStatus>> {
let started = Instant::now();
loop {
if let Some(status) = self.try_wait()? {
return Ok(Some(status));
}
if started.elapsed() >= timeout {
return Ok(None);
}
thread::sleep(Duration::from_millis(5).min(timeout.saturating_sub(started.elapsed())));
}
}
pub fn wait_or_kill(&mut self, grace: Duration) -> io::Result<WaitOutcome> {
#[cfg(unix)]
self.child().signal(nix::libc::SIGTERM)?;
if let Some(status) = self.wait_timeout(grace)? {
return Ok(WaitOutcome::Exited(status));
}
self.terminate_tree().map(WaitOutcome::Terminated)
}
pub fn terminate_tree(&mut self) -> io::Result<ExitStatus> {
if let Some(status) = self.try_wait()? {
return Ok(status);
}
self.child_mut().start_kill()?;
self.child_mut().wait()
}
pub fn terminate_tree_bounded(&mut self, timeout: Duration) -> io::Result<ReapOutcome> {
if let Some(status) = self.try_wait()? {
return Ok(ReapOutcome::Exited(status));
}
self.child_mut().start_kill()?;
Ok(match self.wait_timeout(timeout)? {
Some(status) => ReapOutcome::Exited(status),
None => ReapOutcome::TimedOut,
})
}
fn child(&self) -> &dyn StdChildWrapper {
self.child
.as_deref()
.expect("owned child is always present")
}
fn child_mut(&mut self) -> &mut dyn StdChildWrapper {
self.child
.as_deref_mut()
.expect("owned child is always present")
}
}
impl Drop for OwnedChild {
fn drop(&mut self) {
let _ = self.terminate_tree();
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct AdoptedProcess {
pid: NonZeroU32,
}
impl AdoptedProcess {
pub fn new(pid: NonZeroU32) -> Self {
Self { pid }
}
pub fn id(&self) -> u32 {
self.pid.get()
}
pub fn is_running(&self) -> io::Result<bool> {
process_is_running(self.pid)
}
}
#[derive(Debug)]
struct CleanupWrapper {
captured_pid: Option<Arc<AtomicU32>>,
caller_job: CallerJob,
}
impl StdCommandWrapper for CleanupWrapper {
fn post_spawn(&mut self, child: &mut Child, _core: &StdCommandWrap) -> io::Result<()> {
if let Some(pid) = &self.captured_pid {
pid.store(child.id(), Ordering::SeqCst);
}
Ok(())
}
fn wrap_child(
&mut self,
#[cfg_attr(not(windows), allow(unused_mut))] mut child: Box<dyn StdChildWrapper>,
_core: &StdCommandWrap,
) -> io::Result<Box<dyn StdChildWrapper>> {
#[cfg(not(windows))]
let _ = self.caller_job.take();
#[cfg(windows)]
let kill_on_close = {
let bound = match self.caller_job.take() {
Some(job) => assign_to_caller_job(&job, child.inner()),
None => Ok(()),
}
.and_then(|()| KillOnCloseJob::assign(child.inner()));
match bound {
Ok(job) => Some(job),
Err(error) => {
let _ = child.start_kill();
let _ = child.wait();
return Err(error);
}
}
};
Ok(Box::new(CleanupChild {
child: Some(child),
#[cfg(windows)]
_kill_on_close: kill_on_close,
}))
}
}
#[cfg(windows)]
fn assign_to_caller_job(job: &std::os::windows::io::OwnedHandle, child: &Child) -> io::Result<()> {
use std::os::windows::io::AsRawHandle;
use windows::Win32::{Foundation::HANDLE, System::JobObjects::AssignProcessToJobObject};
unsafe {
AssignProcessToJobObject(
HANDLE(job.as_raw_handle() as _),
HANDLE(child.as_raw_handle() as _),
)
}
.map_err(|e| io::Error::other(format!("bind child to the caller's job object: {e}")))
}
#[cfg(windows)]
#[derive(Debug)]
struct KillOnCloseJob(windows::Win32::Foundation::HANDLE);
#[cfg(windows)]
impl KillOnCloseJob {
fn assign(child: &Child) -> io::Result<Self> {
use std::os::windows::io::AsRawHandle;
use windows::Win32::{
Foundation::HANDLE,
System::JobObjects::{
AssignProcessToJobObject, CreateJobObjectW, JobObjectExtendedLimitInformation,
SetInformationJobObject, JOBOBJECT_EXTENDED_LIMIT_INFORMATION,
JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE,
},
};
let job = Self(unsafe { CreateJobObjectW(None, None) }.map_err(io::Error::other)?);
let mut info = JOBOBJECT_EXTENDED_LIMIT_INFORMATION::default();
info.BasicLimitInformation.LimitFlags = JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE;
unsafe {
SetInformationJobObject(
job.0,
JobObjectExtendedLimitInformation,
&info as *const _ as *const _,
std::mem::size_of_val(&info) as u32,
)
}
.map_err(io::Error::other)?;
unsafe { AssignProcessToJobObject(job.0, HANDLE(child.as_raw_handle() as _)) }
.map_err(io::Error::other)?;
Ok(job)
}
}
#[cfg(windows)]
impl Drop for KillOnCloseJob {
fn drop(&mut self) {
let _ = unsafe { windows::Win32::Foundation::CloseHandle(self.0) };
}
}
#[cfg(windows)]
unsafe impl Send for KillOnCloseJob {}
#[cfg(windows)]
unsafe impl Sync for KillOnCloseJob {}
#[derive(Debug)]
struct CleanupChild {
child: Option<Box<dyn StdChildWrapper>>,
#[cfg(windows)]
_kill_on_close: Option<KillOnCloseJob>,
}
impl CleanupChild {
fn child(&self) -> &dyn StdChildWrapper {
self.child.as_deref().expect("cleanup child is present")
}
fn child_mut(&mut self) -> &mut dyn StdChildWrapper {
self.child.as_deref_mut().expect("cleanup child is present")
}
}
impl StdChildWrapper for CleanupChild {
fn inner(&self) -> &Child {
self.child().inner()
}
fn inner_mut(&mut self) -> &mut Child {
self.child_mut().inner_mut()
}
fn into_inner(mut self: Box<Self>) -> Child {
self.child
.take()
.expect("cleanup child is present")
.into_inner()
}
fn stdin(&mut self) -> &mut Option<ChildStdin> {
self.child_mut().stdin()
}
fn stdout(&mut self) -> &mut Option<ChildStdout> {
self.child_mut().stdout()
}
fn stderr(&mut self) -> &mut Option<ChildStderr> {
self.child_mut().stderr()
}
fn id(&self) -> u32 {
self.child().id()
}
fn start_kill(&mut self) -> io::Result<()> {
self.child_mut().start_kill()
}
fn try_wait(&mut self) -> io::Result<Option<ExitStatus>> {
self.child_mut().try_wait()
}
fn wait(&mut self) -> io::Result<ExitStatus> {
self.child_mut().wait()
}
#[cfg(unix)]
fn signal(&self, signal: i32) -> io::Result<()> {
self.child().signal(signal)
}
}
impl Drop for CleanupChild {
fn drop(&mut self) {
let Some(child) = self.child.as_deref_mut() else {
return;
};
if !matches!(child.try_wait(), Ok(Some(_))) {
let _ = child.start_kill();
let _ = child.wait();
}
}
}
#[derive(Debug)]
struct FailAfterSpawn;
impl StdCommandWrapper for FailAfterSpawn {
fn wrap_child(
&mut self,
_child: Box<dyn StdChildWrapper>,
_core: &StdCommandWrap,
) -> io::Result<Box<dyn StdChildWrapper>> {
Err(io::Error::other("injected post-spawn wrapping failure"))
}
}
fn spawn_owned(
command: Command,
windows_hidden: bool,
fail_after_spawn: bool,
captured_pid: Option<Arc<AtomicU32>>,
caller_job: CallerJob,
) -> io::Result<OwnedChild> {
let mut command = StdCommandWrap::from(command);
command.wrap(CleanupWrapper {
captured_pid,
caller_job,
});
#[cfg(windows)]
{
use process_wrap::std::{CreationFlags, JobObject};
use windows::Win32::System::Threading::CREATE_NO_WINDOW;
if windows_hidden {
command.wrap(CreationFlags(CREATE_NO_WINDOW));
}
command.wrap(JobObject);
}
#[cfg(unix)]
{
use process_wrap::std::ProcessGroup;
let _ = windows_hidden;
command.wrap(ProcessGroup::leader());
}
if fail_after_spawn {
command.wrap(FailAfterSpawn);
}
command
.spawn()
.map(|child| OwnedChild { child: Some(child) })
}
#[cfg(test)]
fn spawn_for_test(command: Command, captured_pid: Arc<AtomicU32>) -> io::Result<OwnedChild> {
spawn_owned(command, true, true, Some(captured_pid), Default::default())
}
#[cfg(windows)]
fn process_is_running(pid: NonZeroU32) -> io::Result<bool> {
use windows::Win32::{
Foundation::{CloseHandle, ERROR_INVALID_PARAMETER, WAIT_TIMEOUT},
System::Threading::{
OpenProcess, WaitForSingleObject, PROCESS_QUERY_LIMITED_INFORMATION,
PROCESS_SYNCHRONIZE,
},
};
let handle = unsafe {
OpenProcess(
PROCESS_QUERY_LIMITED_INFORMATION | PROCESS_SYNCHRONIZE,
false,
pid.get(),
)
};
let handle = match handle {
Ok(handle) => handle,
Err(error) => {
if error.code() == ERROR_INVALID_PARAMETER.to_hresult() {
return Ok(false);
}
return Err(io::Error::other(error.to_string()));
}
};
let wait = unsafe { WaitForSingleObject(handle, 0) };
let close = unsafe { CloseHandle(handle) };
close.map_err(|error| io::Error::other(error.to_string()))?;
match wait.0 {
0 => Ok(false),
value if value == WAIT_TIMEOUT.0 => Ok(true),
_ => Err(io::Error::last_os_error()),
}
}
#[cfg(unix)]
fn process_is_running(pid: NonZeroU32) -> io::Result<bool> {
use nix::{errno::Errno, sys::signal::kill, unistd::Pid};
let pid = i32::try_from(pid.get())
.map(Pid::from_raw)
.map_err(io::Error::other)?;
match kill(pid, None) {
Ok(()) | Err(Errno::EPERM) => Ok(true),
Err(Errno::ESRCH) => Ok(false),
Err(error) => Err(io::Error::from(error)),
}
}
#[cfg(test)]
mod tests {
#[test]
fn owned_commands_hide_their_console_window_by_default() {
assert!(super::OwnedCommand::new("x").windows_hidden);
assert!(super::OwnedCommand::from_command(std::process::Command::new("x")).windows_hidden);
}
#[cfg(windows)]
fn host_job() -> std::os::windows::io::OwnedHandle {
use std::os::windows::io::FromRawHandle;
use windows::Win32::System::JobObjects::{
CreateJobObjectW, JobObjectExtendedLimitInformation, SetInformationJobObject,
JOBOBJECT_EXTENDED_LIMIT_INFORMATION, JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE,
};
let job = unsafe { CreateJobObjectW(None, None) }.unwrap();
let job_handle = unsafe { std::os::windows::io::OwnedHandle::from_raw_handle(job.0 as _) };
let mut info = JOBOBJECT_EXTENDED_LIMIT_INFORMATION::default();
info.BasicLimitInformation.LimitFlags = JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE;
unsafe {
SetInformationJobObject(
job,
JobObjectExtendedLimitInformation,
&info as *const _ as *const _,
std::mem::size_of_val(&info) as u32,
)
}
.unwrap();
job_handle
}
#[cfg(windows)]
fn in_job(pid: u32, job: &std::os::windows::io::OwnedHandle) -> bool {
use std::os::windows::io::AsRawHandle;
#[link(name = "kernel32")]
unsafe extern "system" {
fn OpenProcess(access: u32, inherit: i32, pid: u32) -> *mut std::ffi::c_void;
fn IsProcessInJob(
process: *mut std::ffi::c_void,
job: *mut std::ffi::c_void,
result: *mut i32,
) -> i32;
fn CloseHandle(handle: *mut std::ffi::c_void) -> i32;
}
unsafe {
let process = OpenProcess(0x1000, 0, pid); assert!(!process.is_null());
let mut inside = 0;
assert_ne!(
IsProcessInJob(process, job.as_raw_handle() as _, &mut inside),
0
);
CloseHandle(process);
inside != 0
}
}
#[cfg(windows)]
#[test]
fn children_join_the_callers_job_and_die_when_the_caller_closes_it() {
use std::os::windows::io::AsHandle;
use std::time::{Duration, Instant};
let job = host_job();
let sleeper = || {
let mut command = super::OwnedCommand::new("powershell.exe");
command.command_mut().args([
"-NoLogo",
"-NoProfile",
"-NonInteractive",
"-Command",
"Start-Sleep -Seconds 60",
]);
command.windows_job(job.as_handle()).unwrap();
command.spawn().unwrap()
};
let (first, second) = (sleeper(), sleeper());
let pids = [first.id(), second.id()];
assert!(
pids.iter().all(|pid| in_job(*pid, &job)),
"children must be in the caller's job"
);
drop(job);
let deadline = Instant::now() + Duration::from_secs(5);
for pid in pids {
let process = super::AdoptedProcess::new(std::num::NonZeroU32::new(pid).unwrap());
while process.is_running().unwrap_or(false) {
assert!(
Instant::now() < deadline,
"child {pid} survived its caller's job"
);
std::thread::sleep(Duration::from_millis(20));
}
}
drop((first, second));
}
#[test]
fn a_failure_after_spawn_cleans_up_the_partially_wrapped_child() {
use std::{
process::Command,
sync::{
atomic::{AtomicU32, Ordering},
Arc,
},
thread,
time::{Duration, Instant},
};
let command = if cfg!(windows) {
let mut command = Command::new("powershell.exe");
command.args([
"-NoLogo",
"-NoProfile",
"-NonInteractive",
"-Command",
"Start-Sleep -Seconds 60",
]);
command
} else {
let mut command = Command::new("sleep");
command.arg("60");
command
};
let spawned_pid = Arc::new(AtomicU32::new(0));
assert!(super::spawn_for_test(command, Arc::clone(&spawned_pid)).is_err());
let pid = spawned_pid.load(Ordering::SeqCst);
assert_ne!(pid, 0, "test must fail after the OS process was spawned");
let process = super::AdoptedProcess::new(std::num::NonZeroU32::new(pid).unwrap());
let deadline = Instant::now() + Duration::from_secs(5);
while process.is_running().unwrap_or(false) {
assert!(
Instant::now() < deadline,
"partially spawned process leaked"
);
thread::sleep(Duration::from_millis(20));
}
}
}