use std::{
io::{Error, ErrorKind, Result},
os::windows::{
io::{AsRawHandle, BorrowedHandle},
process::CommandExt,
},
process::{Command, ExitStatus},
time::Duration,
};
#[cfg(feature = "tracing")]
use tracing::{debug, instrument};
use windows::Win32::{
Foundation::{CloseHandle, HANDLE},
System::Threading::PROCESS_CREATION_FLAGS,
};
use crate::{
ChildExitStatus,
windows::{
JobPort, job_creation_flags, make_job_object, resume_threads, terminate_job, wait_on_job,
},
};
#[cfg(feature = "creation-flags")]
use super::CreationFlags;
use super::{ChildWrapper, CommandWrap, CommandWrapper};
#[derive(Clone, Copy, Debug)]
pub struct JobObject;
fn user_creation_flags(core: &CommandWrap) -> PROCESS_CREATION_FLAGS {
#[cfg(feature = "creation-flags")]
{
core.get_wrap::<CreationFlags>()
.map_or(PROCESS_CREATION_FLAGS(0), |flags| flags.0)
}
#[cfg(not(feature = "creation-flags"))]
{
let _ = core;
PROCESS_CREATION_FLAGS(0)
}
}
fn terminate_child(child: &mut dyn ChildWrapper) {
if child.start_kill().is_ok() {
let _ = child.wait();
}
}
fn child_process_handle(child: &dyn ChildWrapper) -> Option<BorrowedHandle<'_>> {
child.process_handle().or_else(|| {
child
.try_inner_child()
.and_then(|child| child.process_handle())
})
}
impl CommandWrapper for JobObject {
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(self)))]
fn pre_spawn(&mut self, command: &mut Command, core: &CommandWrap) -> Result<()> {
let policy = job_creation_flags(user_creation_flags(core));
command.creation_flags(policy.flags.0);
Ok(())
}
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(self)))]
fn wrap_child(
&mut self,
mut inner: Box<dyn ChildWrapper>,
core: &CommandWrap,
) -> Result<Box<dyn ChildWrapper>> {
let policy = job_creation_flags(user_creation_flags(core));
#[cfg(feature = "tracing")]
debug!(
resume_after_assignment = policy.resume_after_assignment,
"options from other wrappers"
);
let handle = match child_process_handle(inner.as_ref()) {
Some(handle) => HANDLE(handle.as_raw_handle()),
None => {
terminate_child(&mut *inner);
return Err(Error::new(
ErrorKind::Unsupported,
"child wrapper does not expose a Windows process handle",
));
}
};
let job_port = match make_job_object(handle, false) {
Ok(job_port) => job_port,
Err(error) => {
terminate_child(&mut *inner);
return Err(error);
}
};
if policy.resume_after_assignment {
if let Err(error) = resume_threads(handle) {
let _ = terminate_job(job_port.job, 1);
terminate_child(&mut *inner);
return Err(error);
}
}
Ok(Box::new(JobObjectChild::new(inner, job_port)))
}
}
#[derive(Debug)]
pub struct JobObjectChild {
inner: Box<dyn ChildWrapper>,
exit_status: ChildExitStatus,
job_port: JobPort,
}
impl JobObjectChild {
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(job_port)))]
pub(crate) fn new(inner: Box<dyn ChildWrapper>, job_port: JobPort) -> Self {
Self {
inner,
exit_status: ChildExitStatus::Running,
job_port,
}
}
}
impl ChildWrapper for JobObjectChild {
fn inner(&self) -> &dyn ChildWrapper {
self.inner.as_ref()
}
fn inner_mut(&mut self) -> &mut dyn ChildWrapper {
self.inner.as_mut()
}
fn into_inner(self: Box<Self>) -> Box<dyn ChildWrapper> {
let its = std::mem::ManuallyDrop::new(self.job_port);
unsafe { CloseHandle(its.completion_port.0) }.ok();
self.inner
}
fn process_handle(&self) -> Option<BorrowedHandle<'_>> {
child_process_handle(self.inner.as_ref())
}
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(self)))]
fn start_kill(&mut self) -> Result<()> {
terminate_job(self.job_port.job, 1)
}
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(self)))]
fn wait(&mut self) -> Result<ExitStatus> {
if let ChildExitStatus::Exited(status) = &self.exit_status {
return Ok(*status);
}
let status = self.inner.wait()?;
self.exit_status = ChildExitStatus::Exited(status);
let JobPort {
completion_port, ..
} = self.job_port;
let _ = wait_on_job(completion_port, None)?;
Ok(status)
}
#[cfg_attr(feature = "tracing", instrument(level = "debug", skip(self)))]
fn try_wait(&mut self) -> Result<Option<ExitStatus>> {
let _ = wait_on_job(self.job_port.completion_port, Some(Duration::ZERO))?;
self.inner.try_wait()
}
}