use std::ffi::c_void;
#[cfg(feature = "tokio")]
use std::future::Future;
use std::io;
use std::os::windows::io::{AsHandle, AsRawHandle, BorrowedHandle, OwnedHandle};
use std::panic::{catch_unwind, AssertUnwindSafe};
#[cfg(feature = "tokio")]
use std::pin::Pin;
use std::ptr;
#[cfg(test)]
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex, PoisonError, Weak};
#[cfg(feature = "tokio")]
use std::task::{Context, Poll, Waker};
use std::thread;
use std::time::Duration;
use windows_sys::Win32::Foundation::ERROR_IO_PENDING;
use windows_sys::Win32::Foundation::{
HANDLE, INVALID_HANDLE_VALUE, WAIT_FAILED, WAIT_OBJECT_0, WAIT_TIMEOUT,
};
use windows_sys::Win32::System::Threading::{
GetExitCodeProcess, RegisterWaitForSingleObject, UnregisterWaitEx, WaitForSingleObject,
INFINITE, WT_EXECUTELONGFUNCTION, WT_EXECUTEONLYONCE,
};
use crate::core::job::{Job, KILL_EXIT_CODE};
use crate::core::pseudocon::ConsoleShared;
#[cfg(any(feature = "blocking", feature = "tokio"))]
pub(super) const LEGACY_CLOSE_GRACE: Duration = Duration::from_secs(1);
#[derive(Debug)]
pub(crate) struct ProcessWaiter {
process: OwnedHandle,
}
impl ProcessWaiter {
pub(crate) const fn new(process: OwnedHandle) -> Self {
Self { process }
}
#[cfg(any(feature = "blocking", test))]
pub(crate) fn wait(&self) -> io::Result<u32> {
let waited = unsafe { WaitForSingleObject(self.process.as_raw_handle(), INFINITE) };
match waited {
WAIT_OBJECT_0 => self.exit_code(),
WAIT_FAILED => Err(io::Error::last_os_error()),
other => Err(io::Error::other(format!(
"unexpected WaitForSingleObject result {other:#x} while waiting for child exit"
))),
}
}
pub(crate) fn try_wait(&self) -> io::Result<Option<u32>> {
let waited = unsafe { WaitForSingleObject(self.process.as_raw_handle(), 0) };
match waited {
WAIT_OBJECT_0 => self.exit_code().map(Some),
WAIT_TIMEOUT => Ok(None),
WAIT_FAILED => Err(io::Error::last_os_error()),
other => Err(io::Error::other(format!(
"unexpected WaitForSingleObject result {other:#x} while polling child exit"
))),
}
}
fn exit_code(&self) -> io::Result<u32> {
let mut code: u32 = 0;
let ok = unsafe { GetExitCodeProcess(self.process.as_raw_handle(), &mut code) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(code)
}
}
impl AsHandle for ProcessWaiter {
fn as_handle(&self) -> BorrowedHandle<'_> {
self.process.as_handle()
}
}
#[cfg(feature = "tokio")]
pub(crate) struct RegisteredWait {
wait_object: HANDLE,
context: *mut RegisteredWaitContext,
}
#[cfg(feature = "tokio")]
struct RegisteredWaitContext {
process: ProcessWaiter,
state: Mutex<RegisteredWaitState>,
}
#[cfg(feature = "tokio")]
#[derive(Default)]
struct RegisteredWaitState {
result: Option<io::Result<u32>>,
waker: Option<Waker>,
callback_active: bool,
owner_dropping: bool,
cleanup_after_callback: bool,
#[cfg(test)]
cleanup_observer: Option<Arc<AtomicBool>>,
}
#[cfg(all(feature = "tokio", test))]
impl Drop for RegisteredWaitContext {
fn drop(&mut self) {
let state = self.state.get_mut().unwrap_or_else(PoisonError::into_inner);
if let Some(observer) = &state.cleanup_observer {
observer.store(true, Ordering::Release);
}
}
}
#[cfg(feature = "tokio")]
unsafe impl Send for RegisteredWait {}
#[cfg(feature = "tokio")]
unsafe impl Sync for RegisteredWait {}
#[cfg(feature = "tokio")]
impl RegisteredWait {
pub(crate) fn new(process: BorrowedHandle<'_>) -> io::Result<Self> {
let process = ProcessWaiter::new(process.try_clone_to_owned()?);
let context = Box::into_raw(Box::new(RegisteredWaitContext {
process,
state: Mutex::new(RegisteredWaitState::default()),
}));
let mut wait_object: HANDLE = ptr::null_mut();
let process_handle = unsafe { (*context).process.as_handle().as_raw_handle() };
let registered = unsafe {
RegisterWaitForSingleObject(
&mut wait_object,
process_handle,
Some(registered_wait_callback),
context.cast(),
INFINITE,
WT_EXECUTEONLYONCE | WT_EXECUTELONGFUNCTION,
)
};
if registered == 0 {
let err = io::Error::last_os_error();
drop(unsafe { Box::from_raw(context) });
return Err(err);
}
Ok(Self {
wait_object,
context,
})
}
#[cfg(test)]
fn observe_cleanup(&self, observer: Arc<AtomicBool>) {
let context = unsafe { &*self.context };
context
.state
.lock()
.unwrap_or_else(PoisonError::into_inner)
.cleanup_observer = Some(observer);
}
}
#[cfg(feature = "tokio")]
impl Future for RegisteredWait {
type Output = io::Result<u32>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let context = unsafe { &*self.context };
let mut state = context.state.lock().unwrap_or_else(PoisonError::into_inner);
if let Some(result) = state.result.take() {
return Poll::Ready(result);
}
if should_replace_waker(state.waker.as_ref(), cx.waker()) {
state.waker = Some(cx.waker().clone());
}
Poll::Pending
}
}
#[cfg(feature = "tokio")]
fn should_replace_waker(registered: Option<&Waker>, candidate: &Waker) -> bool {
registered.map_or(true, |registered| !registered.will_wake(candidate))
}
fn callback_unregister_transfers_cleanup(result: i32, error_code: Option<i32>) -> bool {
result != 0 || error_code == i32::try_from(ERROR_IO_PENDING).ok()
}
#[cfg(feature = "tokio")]
impl Drop for RegisteredWait {
fn drop(&mut self) {
let context = unsafe { &*self.context };
let mut state = context.state.lock().unwrap_or_else(PoisonError::into_inner);
state.owner_dropping = true;
if state.callback_active {
let unregistered = unsafe { UnregisterWaitEx(self.wait_object, ptr::null_mut()) };
let err = io::Error::last_os_error();
if callback_unregister_transfers_cleanup(unregistered, err.raw_os_error()) {
state.cleanup_after_callback = true;
} else {
log_unregister_failure(&err);
}
return;
}
drop(state);
let unregistered = unsafe { UnregisterWaitEx(self.wait_object, INVALID_HANDLE_VALUE) };
if unregistered == 0 {
log_unregister_failure(&io::Error::last_os_error());
return;
}
drop(unsafe { Box::from_raw(self.context) });
}
}
#[cfg(feature = "tokio")]
unsafe extern "system" fn registered_wait_callback(raw: *mut c_void, _timed_out: bool) {
match catch_callback_unwind(|| {
unsafe { registered_wait_callback_inner(raw) }
}) {
Ok(true) => log_callback_panic_safely("registered-wait Waker"),
Ok(false) => {},
Err(()) => log_callback_panic_safely("registered-wait callback"),
}
}
#[cfg(feature = "tokio")]
unsafe fn registered_wait_callback_inner(raw: *mut c_void) -> bool {
let context = unsafe { &*raw.cast::<RegisteredWaitContext>() };
{
let mut state = context.state.lock().unwrap_or_else(PoisonError::into_inner);
if state.owner_dropping {
return false;
}
state.callback_active = true;
}
let result = context.process.exit_code();
let waker = {
let mut state = context.state.lock().unwrap_or_else(PoisonError::into_inner);
if state.owner_dropping {
None
} else {
state.result = Some(result);
state.waker.take()
}
};
let wake_panicked = waker.is_some_and(|waker| catch_callback_unwind(|| waker.wake()).is_err());
let cleanup = {
let mut state = context.state.lock().unwrap_or_else(PoisonError::into_inner);
state.callback_active = false;
state.cleanup_after_callback
};
if cleanup {
drop(unsafe { Box::from_raw(raw.cast::<RegisteredWaitContext>()) });
}
wake_panicked
}
fn catch_callback_unwind<T>(operation: impl FnOnce() -> T) -> Result<T, ()> {
match catch_unwind(AssertUnwindSafe(operation)) {
Ok(value) => Ok(value),
Err(payload) => {
std::mem::forget(payload);
Err(())
},
}
}
fn log_callback_panic_safely(callback: &'static str) {
let _contained = catch_callback_unwind(|| log_callback_panic(callback));
}
#[cfg(feature = "tracing")]
fn log_callback_panic(callback: &'static str) {
tracing::error!(
callback,
"contained a panic at a Windows thread-pool callback boundary"
);
}
#[cfg(not(feature = "tracing"))]
const fn log_callback_panic(_callback: &'static str) {}
fn log_unregister_failure(err: &io::Error) {
let _contained = catch_callback_unwind(|| log_unregister_failure_inner(err));
}
#[cfg(feature = "tracing")]
fn log_unregister_failure_inner(err: &io::Error) {
tracing::error!(
error = %err,
"failed to unregister process wait; retaining its callback context for safety"
);
}
#[cfg(not(feature = "tracing"))]
const fn log_unregister_failure_inner(_err: &io::Error) {}
pub(super) fn spawn_root_watcher(
process: OwnedHandle,
job: Weak<Job>,
shared: Arc<ConsoleShared>,
grace: Duration,
close_legacy: bool,
) -> io::Result<()> {
#[cfg(not(test))]
{
spawn_root_watcher_inner(process, job, shared, grace, close_legacy, true)
}
#[cfg(test)]
{
spawn_root_watcher_inner(process, job, shared, grace, close_legacy, true, None)
}
}
#[cfg(test)]
fn spawn_root_watcher_with_worker_spawn_failure(
process: OwnedHandle,
job: Weak<Job>,
shared: Arc<ConsoleShared>,
grace: Duration,
close_legacy: bool,
cleanup_observer: Arc<AtomicBool>,
) -> io::Result<()> {
spawn_root_watcher_inner(
process,
job,
shared,
grace,
close_legacy,
false,
Some(cleanup_observer),
)
}
fn spawn_root_watcher_inner(
process: OwnedHandle,
job: Weak<Job>,
shared: Arc<ConsoleShared>,
grace: Duration,
close_legacy: bool,
spawn_close_worker: bool,
#[cfg(test)] cleanup_observer: Option<Arc<AtomicBool>>,
) -> io::Result<()> {
let weak = Arc::downgrade(&shared);
drop(shared);
let context = Box::into_raw(Box::new(LegacyWaitContext {
process: ProcessWaiter::new(process),
job,
shared: weak,
grace,
close_legacy,
spawn_close_worker,
wait_object: Mutex::new(None),
wait_object_ready: Condvar::new(),
#[cfg(test)]
cleanup_observer,
}));
let mut wait_object: HANDLE = ptr::null_mut();
let process_handle = unsafe { (*context).process.as_handle().as_raw_handle() };
let registered = unsafe {
RegisterWaitForSingleObject(
&mut wait_object,
process_handle,
Some(legacy_wait_callback),
context.cast(),
INFINITE,
WT_EXECUTEONLYONCE | WT_EXECUTELONGFUNCTION,
)
};
if registered == 0 {
let err = io::Error::last_os_error();
drop(unsafe { Box::from_raw(context) });
return Err(err);
}
let context_ref = unsafe { &*context };
let mut slot = context_ref
.wait_object
.lock()
.unwrap_or_else(PoisonError::into_inner);
*slot = Some(wait_object);
drop(slot);
context_ref.wait_object_ready.notify_all();
Ok(())
}
struct LegacyWaitContext {
process: ProcessWaiter,
job: Weak<Job>,
shared: Weak<ConsoleShared>,
grace: Duration,
close_legacy: bool,
spawn_close_worker: bool,
wait_object: Mutex<Option<HANDLE>>,
wait_object_ready: Condvar,
#[cfg(test)]
cleanup_observer: Option<Arc<AtomicBool>>,
}
#[cfg(test)]
impl Drop for LegacyWaitContext {
fn drop(&mut self) {
if let Some(observer) = &self.cleanup_observer {
observer.store(true, Ordering::Release);
}
}
}
struct LegacyContextPtr(*mut LegacyWaitContext);
unsafe impl Send for LegacyContextPtr {}
unsafe extern "system" fn legacy_wait_callback(raw: *mut c_void, _timed_out: bool) {
if catch_callback_unwind(|| {
unsafe { legacy_wait_callback_inner(raw) }
})
.is_err()
{
log_callback_panic_safely("legacy-wait callback");
}
}
unsafe fn legacy_wait_callback_inner(raw: *mut c_void) {
let pointer = LegacyContextPtr(raw.cast());
let worker_pointer = LegacyContextPtr(pointer.0);
let context = unsafe { &*pointer.0 };
let spawned = if context.spawn_close_worker {
thread::Builder::new()
.name("conpty-oxide-legacy-close".into())
.spawn(move || {
unsafe { finish_root_wait(&worker_pointer) };
})
} else {
Err(io::Error::other(
"legacy close worker spawn failure injected by a crate-local test",
))
};
if let Err(err) = spawned {
finish_root_exit(context);
let wait_object = legacy_wait_object(context);
let unregistered = unsafe { UnregisterWaitEx(wait_object, ptr::null_mut()) };
let unregister_error = io::Error::last_os_error();
if callback_unregister_transfers_cleanup(unregistered, unregister_error.raw_os_error()) {
drop(unsafe { Box::from_raw(pointer.0) });
} else {
log_unregister_failure(&unregister_error);
}
log_legacy_worker_failure(&err);
}
}
unsafe fn finish_root_wait(pointer: &LegacyContextPtr) {
let context = unsafe { &*pointer.0 };
finish_root_exit(context);
let wait_object = legacy_wait_object(context);
let unregistered = unsafe { UnregisterWaitEx(wait_object, INVALID_HANDLE_VALUE) };
if unregistered != 0 {
drop(unsafe { Box::from_raw(pointer.0) });
} else {
log_unregister_failure(&io::Error::last_os_error());
}
}
fn finish_root_exit(context: &LegacyWaitContext) {
if let Some(job) = context.job.upgrade() {
if let Err(err) = job.terminate(KILL_EXIT_CODE) {
log_root_watcher_kill_failure(&err);
}
}
if !context.close_legacy {
return;
}
let Some(shared) = context.shared.upgrade() else {
return;
};
if should_wait_for_legacy_reader(shared.reader_finished()) {
thread::sleep(context.grace);
}
shared.request_close();
}
const fn should_wait_for_legacy_reader(reader_finished: bool) -> bool {
!reader_finished
}
fn legacy_wait_object(context: &LegacyWaitContext) -> HANDLE {
let mut slot = context
.wait_object
.lock()
.unwrap_or_else(PoisonError::into_inner);
loop {
if let Some(wait_object) = *slot {
return wait_object;
}
slot = context
.wait_object_ready
.wait(slot)
.unwrap_or_else(PoisonError::into_inner);
}
}
fn log_root_watcher_kill_failure(err: &io::Error) {
let _contained = catch_callback_unwind(|| {
log_root_watcher_kill_failure_inner(err);
});
}
#[cfg(feature = "tracing")]
fn log_root_watcher_kill_failure_inner(err: &io::Error) {
tracing::error!(
error = %err,
"failed to terminate the managed Job after root process exit"
);
}
#[cfg(not(feature = "tracing"))]
const fn log_root_watcher_kill_failure_inner(_err: &io::Error) {}
fn log_legacy_worker_failure(err: &io::Error) {
let _contained = catch_callback_unwind(|| {
log_legacy_worker_failure_inner(err);
});
}
#[cfg(feature = "tracing")]
fn log_legacy_worker_failure_inner(err: &io::Error) {
tracing::error!(
error = %err,
"failed to spawn legacy close worker; completed close in callback"
);
}
#[cfg(not(feature = "tracing"))]
const fn log_legacy_worker_failure_inner(_err: &io::Error) {}
#[cfg(test)]
#[path = "wait_tests.rs"]
mod tests;