use std::io;
use tokio::process::{Child, Command};
pub(crate) fn spawn(command: &mut Command) -> io::Result<(Child, Option<ProcessTree>)> {
#[cfg(windows)]
{
platform::spawn(command).map(|(child, tree)| (child, Some(ProcessTree(Some(tree)))))
}
#[cfg(not(windows))]
{
command.spawn().map(|child| (child, None))
}
}
pub(crate) struct ProcessTree(Option<platform::Tree>);
impl ProcessTree {
pub(crate) fn terminate(mut self) -> io::Result<()> {
self.0.take().expect("process tree is armed").terminate()
}
}
impl Drop for ProcessTree {
fn drop(&mut self) {
if let Some(tree) = self.0.take() {
let _ = tree.terminate();
}
}
}
#[cfg(not(windows))]
mod platform {
pub(super) struct Tree;
impl Tree {
pub(super) fn terminate(&self) -> std::io::Result<()> {
unreachable!("process-tree ownership is Windows-only")
}
}
}
#[cfg(windows)]
mod platform {
use std::mem::size_of;
use std::os::windows::process::CommandExt;
use std::{io, ptr};
use tokio::process::{Child, Command};
use windows_sys::Win32::Foundation::{CloseHandle, HANDLE, INVALID_HANDLE_VALUE};
use windows_sys::Win32::System::Diagnostics::ToolHelp::{
CreateToolhelp32Snapshot, TH32CS_SNAPTHREAD, THREADENTRY32, Thread32First, Thread32Next,
};
use windows_sys::Win32::System::JobObjects::{
AssignProcessToJobObject, CreateJobObjectW, JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE,
JOBOBJECT_EXTENDED_LIMIT_INFORMATION, JobObjectExtendedLimitInformation,
SetInformationJobObject, TerminateJobObject,
};
use windows_sys::Win32::System::Threading::{
CREATE_NO_WINDOW, CREATE_SUSPENDED, OpenThread, ResumeThread, THREAD_SUSPEND_RESUME,
};
struct OwnedHandle(HANDLE);
unsafe impl Send for OwnedHandle {}
unsafe impl Sync for OwnedHandle {}
impl Drop for OwnedHandle {
fn drop(&mut self) {
unsafe {
CloseHandle(self.0);
}
}
}
pub(super) struct Tree {
job: OwnedHandle,
}
pub(super) fn spawn(command: &mut Command) -> io::Result<(Child, Tree)> {
command
.as_std_mut()
.creation_flags(CREATE_NO_WINDOW | CREATE_SUSPENDED);
let mut child = command.spawn()?;
match attach_and_resume(&child) {
Ok(tree) => Ok((child, tree)),
Err(error) => {
let _ = child.start_kill();
Err(error)
}
}
}
fn attach_and_resume(child: &Child) -> io::Result<Tree> {
let raw_job = unsafe { CreateJobObjectW(ptr::null(), ptr::null()) };
if raw_job.is_null() {
return Err(io::Error::last_os_error());
}
let job = OwnedHandle(raw_job);
let mut limits = JOBOBJECT_EXTENDED_LIMIT_INFORMATION::default();
limits.BasicLimitInformation.LimitFlags = JOB_OBJECT_LIMIT_KILL_ON_JOB_CLOSE;
if unsafe {
SetInformationJobObject(
job.0,
JobObjectExtendedLimitInformation,
ptr::from_ref(&limits).cast(),
size_of::<JOBOBJECT_EXTENDED_LIMIT_INFORMATION>() as u32,
)
} == 0
{
return Err(io::Error::last_os_error());
}
let process = child.raw_handle().ok_or_else(|| {
io::Error::new(
io::ErrorKind::NotFound,
"CLI exited before Job Object assignment",
)
})?;
if unsafe { AssignProcessToJobObject(job.0, process.cast()) } == 0 {
return Err(io::Error::last_os_error());
}
resume_initial_thread(child.id().ok_or_else(|| {
io::Error::new(io::ErrorKind::NotFound, "CLI exited before thread resume")
})?)?;
Ok(Tree { job })
}
fn resume_initial_thread(pid: u32) -> io::Result<()> {
let snapshot = unsafe { CreateToolhelp32Snapshot(TH32CS_SNAPTHREAD, 0) };
if snapshot == INVALID_HANDLE_VALUE {
return Err(io::Error::last_os_error());
}
let snapshot = OwnedHandle(snapshot);
let mut entry = THREADENTRY32 {
dwSize: size_of::<THREADENTRY32>() as u32,
..Default::default()
};
let mut found = unsafe { Thread32First(snapshot.0, &mut entry) } != 0;
while found {
if entry.th32OwnerProcessID == pid {
let raw_thread =
unsafe { OpenThread(THREAD_SUSPEND_RESUME, 0, entry.th32ThreadID) };
if raw_thread.is_null() {
return Err(io::Error::last_os_error());
}
let thread = OwnedHandle(raw_thread);
if unsafe { ResumeThread(thread.0) } == u32::MAX {
return Err(io::Error::last_os_error());
}
return Ok(());
}
found = unsafe { Thread32Next(snapshot.0, &mut entry) } != 0;
}
Err(io::Error::new(
io::ErrorKind::NotFound,
"CLI initial thread was not found",
))
}
impl Tree {
pub(super) fn terminate(&self) -> io::Result<()> {
if unsafe { TerminateJobObject(self.job.0, 1) } != 0 {
Ok(())
} else {
Err(io::Error::last_os_error())
}
}
}
}