use std::ffi::{OsStr, OsString};
use std::io;
use std::os::windows::ffi::OsStrExt;
use std::os::windows::io::FromRawHandle;
use std::pin::Pin;
use std::sync::{
Arc,
atomic::{AtomicU64, Ordering},
};
use std::task::{Context, Poll};
use std::time::Duration;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use tokio::process::Command;
use tokio::sync::Notify;
use windows_sys::Win32::Foundation::{CloseHandle, HANDLE, WAIT_OBJECT_0};
use windows_sys::Win32::Security::SECURITY_ATTRIBUTES;
use windows_sys::Win32::System::Console::{
COORD, ClosePseudoConsole, CreatePseudoConsole, GetConsoleWindow, GetStdHandle, HPCON,
ResizePseudoConsole, STD_ERROR_HANDLE, STD_HANDLE, STD_INPUT_HANDLE, STD_OUTPUT_HANDLE,
SetStdHandle,
};
use windows_sys::Win32::System::JobObjects::AssignProcessToJobObject;
use windows_sys::Win32::System::Pipes::CreatePipe;
use windows_sys::Win32::System::Threading::{
CREATE_NEW_PROCESS_GROUP, CREATE_SUSPENDED, CREATE_UNICODE_ENVIRONMENT, CreateProcessW,
DeleteProcThreadAttributeList, EXTENDED_STARTUPINFO_PRESENT, GetExitCodeProcess, INFINITE,
InitializeProcThreadAttributeList, LPPROC_THREAD_ATTRIBUTE_LIST,
PROC_THREAD_ATTRIBUTE_PSEUDOCONSOLE, PROCESS_INFORMATION, ResumeThread, STARTF_USESTDHANDLES,
STARTUPINFOEXW, TerminateProcess, UpdateProcThreadAttribute, WaitForSingleObject,
};
use crate::sys::SpawnOptions;
use crate::sys::pid_gate::PidGate;
use super::{PtyExitStatus, PtyReader, PtySpawn, PtyWriter};
const STILL_ACTIVE: u32 = 259;
fn coord(cols: u16, rows: u16) -> COORD {
let clamp = |v: u16| v.min(i16::MAX as u16) as i16;
COORD {
X: clamp(cols),
Y: clamp(rows),
}
}
const CONPTY_OUTPUT_QUIESCENCE: Duration = Duration::from_millis(100);
#[derive(Default)]
struct OutputActivity {
sequence: AtomicU64,
changed: Notify,
}
impl OutputActivity {
fn record_chunk(&self) {
self.sequence.fetch_add(1, Ordering::Release);
self.changed.notify_waiters();
}
async fn wait_for_quiescence(&self) {
let mut sequence = self.sequence.load(Ordering::Acquire);
loop {
let changed = self.changed.notified();
tokio::pin!(changed);
changed.as_mut().enable();
let current = self.sequence.load(Ordering::Acquire);
if current != sequence {
sequence = current;
continue;
}
if tokio::time::timeout(CONPTY_OUTPUT_QUIESCENCE, changed)
.await
.is_err()
{
return;
}
sequence = self.sequence.load(Ordering::Acquire);
}
}
}
struct OwnedProcess(HANDLE);
unsafe impl Send for OwnedProcess {}
unsafe impl Sync for OwnedProcess {}
impl Drop for OwnedProcess {
fn drop(&mut self) {
if !self.0.is_null() {
unsafe { CloseHandle(self.0) };
}
}
}
pub(crate) struct PtyChild {
process: Arc<OwnedProcess>,
thread: HANDLE,
hpc: HPCON,
input_keepalive: Option<tokio::sync::mpsc::UnboundedSender<Vec<u8>>>,
hpc_closed: bool,
output_activity: Arc<OutputActivity>,
pid: u32,
}
unsafe impl Send for PtyChild {}
unsafe impl Sync for PtyChild {}
impl PtyChild {
#[allow(dead_code)]
pub(crate) fn id(&self) -> Option<u32> {
Some(self.pid)
}
pub(crate) async fn reap(&mut self, _gate: &PidGate) -> io::Result<PtyExitStatus> {
self.wait().await
}
pub(crate) async fn wait(&mut self) -> io::Result<PtyExitStatus> {
let process = Arc::clone(&self.process);
let code = tokio::task::spawn_blocking(move || {
unsafe { WaitForSingleObject(process.0, INFINITE) };
let mut code: u32 = 0;
let ok = unsafe { GetExitCodeProcess(process.0, &mut code) };
(ok != 0).then_some(code)
})
.await
.map_err(io::Error::other)?;
self.input_keepalive.take();
self.output_activity.wait_for_quiescence().await;
self.close_pty();
Ok(PtyExitStatus::from_code(code.map(|c| c as i32)))
}
fn close_pty(&mut self) {
if !self.hpc_closed {
self.hpc_closed = true;
unsafe { ClosePseudoConsole(self.hpc) };
}
}
pub(crate) fn resize(&self, cols: u16, rows: u16) -> io::Result<()> {
let hr = unsafe { ResizePseudoConsole(self.hpc, coord(cols, rows)) };
if hr != 0 {
return Err(io::Error::from_raw_os_error(hr));
}
Ok(())
}
pub(crate) fn try_wait(&mut self) -> io::Result<Option<PtyExitStatus>> {
let waited = unsafe { WaitForSingleObject(self.process.0, 0) };
if waited != WAIT_OBJECT_0 {
return Ok(None);
}
let mut code: u32 = 0;
let ok = unsafe { GetExitCodeProcess(self.process.0, &mut code) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
if code == STILL_ACTIVE {
return Ok(None);
}
Ok(Some(PtyExitStatus::from_code(Some(code as i32))))
}
pub(crate) fn start_kill(&mut self) -> io::Result<()> {
let ok = unsafe { TerminateProcess(self.process.0, 1) };
if ok == 0 {
return Ok(());
}
Ok(())
}
}
impl Drop for PtyChild {
fn drop(&mut self) {
self.close_pty();
self.input_keepalive.take();
unsafe {
if !self.thread.is_null() {
CloseHandle(self.thread);
}
}
}
}
unsafe fn close(handle: HANDLE) {
if !handle.is_null() {
unsafe { CloseHandle(handle) };
}
}
fn append_arg(out: &mut Vec<u16>, arg: &OsStr, force_quote: bool) {
let wide: Vec<u16> = arg.encode_wide().collect();
let quote = force_quote
|| wide.is_empty()
|| wide
.iter()
.any(|&c| c == u16::from(b' ') || c == u16::from(b'\t'));
if quote {
out.push(u16::from(b'"'));
}
let mut backslashes = 0usize;
for &c in &wide {
if c == u16::from(b'\\') {
backslashes += 1;
} else {
if c == u16::from(b'"') {
for _ in 0..=backslashes {
out.push(u16::from(b'\\'));
}
}
backslashes = 0;
}
out.push(c);
}
if quote {
for _ in 0..backslashes {
out.push(u16::from(b'\\'));
}
out.push(u16::from(b'"'));
}
}
fn build_command_line(cmd: &Command) -> Vec<u16> {
let std_cmd = cmd.as_std();
let mut line: Vec<u16> = Vec::new();
append_arg(&mut line, std_cmd.get_program(), true);
for arg in std_cmd.get_args() {
line.push(u16::from(b' '));
append_arg(&mut line, arg, false);
}
line.push(0);
line
}
fn build_env_block(env: Option<Vec<(OsString, OsString)>>) -> Option<Vec<u16>> {
let pairs = env?;
let mut block: Vec<u16> = Vec::new();
for (k, v) in pairs {
block.extend(k.encode_wide());
block.push(u16::from(b'='));
block.extend(v.encode_wide());
block.push(0);
}
block.push(0);
if block.len() == 1 {
block.push(0);
}
Some(block)
}
fn to_wide_nul(s: &OsStr) -> Vec<u16> {
let mut v: Vec<u16> = s.encode_wide().collect();
v.push(0);
v
}
fn launcher_has_console() -> bool {
!unsafe { GetConsoleWindow() }.is_null()
}
fn conpty_startup_info(
attr_list: LPPROC_THREAD_ATTRIBUTE_LIST,
sever_console_std_handles: bool,
) -> STARTUPINFOEXW {
let mut si = STARTUPINFOEXW::default();
si.StartupInfo.cb = std::mem::size_of::<STARTUPINFOEXW>() as u32;
if sever_console_std_handles {
si.StartupInfo.dwFlags = STARTF_USESTDHANDLES;
si.StartupInfo.hStdInput = std::ptr::null_mut();
si.StartupInfo.hStdOutput = std::ptr::null_mut();
si.StartupInfo.hStdError = std::ptr::null_mut();
}
si.lpAttributeList = attr_list;
si
}
struct NulledLauncherStdio {
saved: [(STD_HANDLE, HANDLE); 3],
_lock: std::sync::MutexGuard<'static, ()>,
}
impl NulledLauncherStdio {
fn install() -> Self {
static LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
let lock = LOCK.lock().unwrap_or_else(|poisoned| poisoned.into_inner());
let mut saved = [
(STD_INPUT_HANDLE, std::ptr::null_mut()),
(STD_OUTPUT_HANDLE, std::ptr::null_mut()),
(STD_ERROR_HANDLE, std::ptr::null_mut()),
];
for (which, prev) in &mut saved {
unsafe {
*prev = GetStdHandle(*which);
SetStdHandle(*which, std::ptr::null_mut());
}
}
Self { saved, _lock: lock }
}
}
impl Drop for NulledLauncherStdio {
fn drop(&mut self) {
for (which, prev) in self.saved {
unsafe { SetStdHandle(which, prev) };
}
}
}
pub(crate) fn spawn_pty(
cmd: &mut Command,
opts: &SpawnOptions,
env: Option<Vec<(OsString, OsString)>>,
job: HANDLE,
skip_drop_kill: &crate::sys::SkipDropKill,
) -> io::Result<PtySpawn> {
let mut input_read: HANDLE = std::ptr::null_mut();
let mut input_write: HANDLE = std::ptr::null_mut();
let mut output_read: HANDLE = std::ptr::null_mut();
let mut output_write: HANDLE = std::ptr::null_mut();
let sec: *const SECURITY_ATTRIBUTES = std::ptr::null();
if unsafe { CreatePipe(&mut input_read, &mut input_write, sec, 0) } == 0 {
return Err(io::Error::last_os_error());
}
if unsafe { CreatePipe(&mut output_read, &mut output_write, sec, 0) } == 0 {
let e = io::Error::last_os_error();
unsafe {
close(input_read);
close(input_write);
}
return Err(e);
}
let (cols, rows) = opts.pty_size.unwrap_or(super::DEFAULT_PTY_SIZE);
let size = coord(cols, rows);
let mut hpc: HPCON = 0;
let hr = unsafe { CreatePseudoConsole(size, input_read, output_write, 0, &mut hpc) };
if hr != 0 {
unsafe {
close(input_read);
close(input_write);
close(output_read);
close(output_write);
}
return Err(io::Error::from_raw_os_error(hr));
}
let cleanup_all = || unsafe {
ClosePseudoConsole(hpc);
close(input_read);
close(input_write);
close(output_read);
close(output_write);
};
let mut attr_size: usize = 0;
unsafe { InitializeProcThreadAttributeList(std::ptr::null_mut(), 1, 0, &mut attr_size) };
let mut attr_buf = vec![0u8; attr_size];
let attr_list = attr_buf.as_mut_ptr().cast();
if unsafe { InitializeProcThreadAttributeList(attr_list, 1, 0, &mut attr_size) } == 0 {
let e = io::Error::last_os_error();
cleanup_all();
return Err(e);
}
let updated = unsafe {
UpdateProcThreadAttribute(
attr_list,
0,
PROC_THREAD_ATTRIBUTE_PSEUDOCONSOLE as usize,
hpc as *const std::ffi::c_void,
std::mem::size_of::<HPCON>(),
std::ptr::null_mut(),
std::ptr::null_mut(),
)
};
if updated == 0 {
let e = io::Error::last_os_error();
unsafe { DeleteProcThreadAttributeList(attr_list) };
cleanup_all();
return Err(e);
}
let has_console = launcher_has_console();
let si = conpty_startup_info(attr_list, has_console);
let mut command_line = build_command_line(cmd);
let env_block = build_env_block(env);
let cwd_wide = cmd
.as_std()
.get_current_dir()
.map(|p| to_wide_nul(p.as_os_str()));
let mut flags = CREATE_SUSPENDED | EXTENDED_STARTUPINFO_PRESENT | opts.creation_flags;
if opts.windows_new_process_group {
flags |= CREATE_NEW_PROCESS_GROUP;
}
let (env_ptr, env_flag): (*const std::ffi::c_void, u32) = match &env_block {
Some(block) => (block.as_ptr().cast(), CREATE_UNICODE_ENVIRONMENT),
None => (std::ptr::null(), 0),
};
flags |= env_flag;
let cwd_ptr = cwd_wide.as_ref().map_or(std::ptr::null(), |w| w.as_ptr());
let mut pi: PROCESS_INFORMATION = unsafe { std::mem::zeroed() };
let created = {
let _nulled_stdio = (!has_console).then(NulledLauncherStdio::install);
unsafe {
CreateProcessW(
std::ptr::null(),
command_line.as_mut_ptr(),
std::ptr::null(),
std::ptr::null(),
0, flags,
env_ptr,
cwd_ptr,
std::ptr::from_ref(&si.StartupInfo),
&mut pi,
)
}
};
unsafe { DeleteProcThreadAttributeList(attr_list) };
if created == 0 {
let e = io::Error::last_os_error();
cleanup_all();
return Err(e);
}
unsafe {
close(input_read);
close(output_write);
}
if unsafe { AssignProcessToJobObject(job, pi.hProcess) } == 0 {
let e = io::Error::last_os_error();
unsafe {
TerminateProcess(pi.hProcess, 1);
close(pi.hProcess);
close(pi.hThread);
ClosePseudoConsole(hpc);
close(input_write);
close(output_read);
}
return Err(e);
}
skip_drop_kill.clear();
unsafe { ResumeThread(pi.hThread) };
let pid = pi.dwProcessId;
let output_activity = Arc::new(OutputActivity::default());
let reader: PtyReader = Box::new(bridge_reader(output_read, Arc::clone(&output_activity)));
let (writer, input_keepalive) = bridge_writer(input_write);
let writer: PtyWriter = Box::new(writer);
Ok(PtySpawn {
child: PtyChild {
process: Arc::new(OwnedProcess(pi.hProcess)),
thread: pi.hThread,
hpc,
input_keepalive: Some(input_keepalive),
hpc_closed: false,
output_activity,
pid,
},
reader,
writer,
pid: Some(pid),
})
}
struct SendHandle(HANDLE);
unsafe impl Send for SendHandle {}
impl SendHandle {
unsafe fn into_file(self) -> std::fs::File {
unsafe { std::fs::File::from_raw_handle(self.0 as _) }
}
}
fn bridge_reader(handle: HANDLE, output_activity: Arc<OutputActivity>) -> ChannelReader {
let (tx, rx) = tokio::sync::mpsc::channel::<Vec<u8>>(64);
let h = SendHandle(handle);
std::thread::spawn(move || {
use std::io::Read;
let mut file = unsafe { h.into_file() };
let mut buf = [0u8; 8192];
loop {
match file.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
output_activity.record_chunk();
if tx.blocking_send(buf[..n].to_vec()).is_err() {
break;
}
}
}
}
});
ChannelReader {
rx: std::sync::Mutex::new(rx),
leftover: Vec::new(),
pos: 0,
}
}
fn bridge_writer(handle: HANDLE) -> (ChannelWriter, tokio::sync::mpsc::UnboundedSender<Vec<u8>>) {
let (tx, mut rx) = tokio::sync::mpsc::unbounded_channel::<Vec<u8>>();
let h = SendHandle(handle);
std::thread::spawn(move || {
use std::io::Write;
let mut file = unsafe { h.into_file() };
while let Some(chunk) = rx.blocking_recv() {
if file.write_all(&chunk).is_err() {
break;
}
let _ = file.flush();
}
});
let keepalive = tx.clone();
(
ChannelWriter {
tx,
shutdown: false,
},
keepalive,
)
}
struct ChannelReader {
rx: std::sync::Mutex<tokio::sync::mpsc::Receiver<Vec<u8>>>,
leftover: Vec<u8>,
pos: usize,
}
impl AsyncRead for ChannelReader {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
if this.pos < this.leftover.len() {
let start = this.pos;
let n = (this.leftover.len() - start).min(buf.remaining());
buf.put_slice(&this.leftover[start..start + n]);
this.pos += n;
return Poll::Ready(Ok(()));
}
let poll = this
.rx
.get_mut()
.expect("pty reader mutex poisoned")
.poll_recv(cx);
match poll {
Poll::Ready(Some(chunk)) => {
let n = chunk.len().min(buf.remaining());
buf.put_slice(&chunk[..n]);
if n < chunk.len() {
this.leftover = chunk;
this.pos = n;
} else {
this.leftover.clear();
this.pos = 0;
}
Poll::Ready(Ok(()))
}
Poll::Ready(None) => Poll::Ready(Ok(())),
Poll::Pending => Poll::Pending,
}
}
}
struct ChannelWriter {
tx: tokio::sync::mpsc::UnboundedSender<Vec<u8>>,
shutdown: bool,
}
impl ChannelWriter {
fn send_eof(&mut self) -> io::Result<()> {
if self.shutdown {
return Ok(());
}
self.shutdown = true;
self.tx
.send(vec![0x1a, b'\r'])
.map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "pty stdin writer closed"))
}
}
impl AsyncWrite for ChannelWriter {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
if self.shutdown {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"pty stdin writer closed",
)));
}
match self.tx.send(buf.to_vec()) {
Ok(()) => Poll::Ready(Ok(buf.len())),
Err(_) => Poll::Ready(Err(io::Error::new(
io::ErrorKind::BrokenPipe,
"pty stdin writer closed",
))),
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(self.send_eof())
}
}
impl Drop for ChannelWriter {
fn drop(&mut self) {
let _ = self.send_eof();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn conpty_startup_info_severs_std_handles_only_when_requested() {
let severed = conpty_startup_info(std::ptr::null_mut(), true);
assert_eq!(
severed.StartupInfo.dwFlags & STARTF_USESTDHANDLES,
STARTF_USESTDHANDLES
);
assert!(severed.StartupInfo.hStdInput.is_null());
assert!(severed.StartupInfo.hStdOutput.is_null());
assert!(severed.StartupInfo.hStdError.is_null());
let inherited = conpty_startup_info(std::ptr::null_mut(), false);
assert_eq!(inherited.StartupInfo.dwFlags & STARTF_USESTDHANDLES, 0);
}
#[test]
fn coord_maps_a_window_size_and_clamps_beyond_i16() {
let c = coord(120, 40);
assert_eq!((c.X, c.Y), (120, 40));
let d = coord(
super::super::DEFAULT_PTY_SIZE.0,
super::super::DEFAULT_PTY_SIZE.1,
);
assert_eq!((d.X, d.Y), (80, 24));
let big = coord(u16::MAX, 40_000);
assert_eq!((big.X, big.Y), (i16::MAX, i16::MAX));
}
}