use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use windows::Win32::Foundation::{CloseHandle, HANDLE};
use windows::Win32::Storage::FileSystem::{ReadFile, WriteFile};
use windows::Win32::System::Console::{
COORD, ClosePseudoConsole, CreatePseudoConsole, HPCON, ResizePseudoConsole,
};
use windows::Win32::System::Pipes::CreatePipe;
use crate::pal::error::{PalError, PalErrorKind};
use crate::pal::ids::PtyId;
use crate::pal::pseudoconsole::{Pseudoconsole, WindowSize};
use crate::pal::raw_handle::PipeHandle;
struct Pty {
hpcon: Option<HPCON>,
host_input: Arc<PipeHandle>,
host_output: Arc<PipeHandle>,
}
struct PtyTable {
ptys: HashMap<u64, Pty>,
}
fn table() -> &'static Mutex<PtyTable> {
static TABLE: OnceLock<Mutex<PtyTable>> = OnceLock::new();
TABLE.get_or_init(|| {
Mutex::new(PtyTable {
ptys: HashMap::new(),
})
})
}
fn next_id() -> u64 {
static NEXT: AtomicU64 = AtomicU64::new(1);
NEXT.fetch_add(1, Ordering::Relaxed)
}
fn close_handle(handle: HANDLE) {
if handle.is_invalid() {
return;
}
_ = unsafe { CloseHandle(handle) };
}
fn to_coord(size: WindowSize) -> Result<COORD, PalError> {
Ok(COORD {
X: i16::try_from(size.cols.max(1)).map_err(|_error| PalError::new(PalErrorKind::Other))?,
Y: i16::try_from(size.rows.max(1)).map_err(|_error| PalError::new(PalErrorKind::Other))?,
})
}
pub(crate) fn hpcon_for(pty: PtyId) -> Option<HPCON> {
table()
.lock()
.expect("pty table")
.ptys
.get(&pty.0)
.and_then(|pty| pty.hpcon)
}
#[derive(Debug, Default)]
pub(crate) struct BuildTargetPseudoconsole;
#[cfg_attr(coverage_nightly, coverage(off))]
#[cfg_attr(test, mutants::skip)]
impl Pseudoconsole for BuildTargetPseudoconsole {
fn create(&self, size: WindowSize) -> Result<PtyId, PalError> {
let mut input_read = HANDLE::default();
let mut input_write = HANDLE::default();
let mut output_read = HANDLE::default();
let mut output_write = HANDLE::default();
unsafe { CreatePipe(&raw mut input_read, &raw mut input_write, None, 0) }
.map_err(|_error| PalError::new(PalErrorKind::Other))?;
if unsafe { CreatePipe(&raw mut output_read, &raw mut output_write, None, 0) }.is_err() {
close_handle(input_read);
close_handle(input_write);
return Err(PalError::new(PalErrorKind::Other));
}
let coord = match to_coord(size) {
Ok(coord) => coord,
Err(error) => {
close_handle(input_read);
close_handle(input_write);
close_handle(output_read);
close_handle(output_write);
return Err(error);
}
};
let hpcon = unsafe { CreatePseudoConsole(coord, input_read, output_write, 0) };
let Ok(hpcon) = hpcon else {
close_handle(input_read);
close_handle(input_write);
close_handle(output_read);
close_handle(output_write);
return Err(PalError::new(PalErrorKind::Other));
};
close_handle(input_read);
close_handle(output_write);
let id = next_id();
table().lock().expect("pty table").ptys.insert(
id,
Pty {
hpcon: Some(hpcon),
host_input: PipeHandle::new(input_write),
host_output: PipeHandle::new(output_read),
},
);
Ok(PtyId(id))
}
fn resize(&self, pty: PtyId, size: WindowSize) -> Result<(), PalError> {
let coord = to_coord(size)?;
let table = table().lock().expect("pty table");
let hpcon = table
.ptys
.get(&pty.0)
.and_then(|pty| pty.hpcon)
.ok_or_else(|| PalError::new(PalErrorKind::NotFound))?;
unsafe { ResizePseudoConsole(hpcon, coord) }
.map_err(|_error| PalError::new(PalErrorKind::Other))
}
fn write_input(&self, pty: PtyId, data: &[u8]) -> Result<(), PalError> {
let handle = table()
.lock()
.expect("pty table")
.ptys
.get(&pty.0)
.map(|pty| Arc::clone(&pty.host_input))
.ok_or_else(|| PalError::new(PalErrorKind::NotFound))?;
let mut remaining = data;
while !remaining.is_empty() {
let mut transferred = 0_u32;
unsafe {
WriteFile(
handle.as_handle(),
Some(remaining),
Some(&raw mut transferred),
None,
)
}
.map_err(|_error| PalError::new(PalErrorKind::Other))?;
if transferred == 0 {
return Err(PalError::new(PalErrorKind::Other));
}
remaining = remaining
.get(transferred as usize..)
.ok_or_else(|| PalError::new(PalErrorKind::Other))?;
}
Ok(())
}
fn read_output(&self, pty: PtyId) -> Result<Vec<u8>, PalError> {
let handle = table()
.lock()
.expect("pty table")
.ptys
.get(&pty.0)
.map(|pty| Arc::clone(&pty.host_output))
.ok_or_else(|| PalError::new(PalErrorKind::NotFound))?;
let mut buf = vec![0_u8; 4096];
let mut transferred = 0_u32;
unsafe {
ReadFile(
handle.as_handle(),
Some(buf.as_mut_slice()),
Some(&raw mut transferred),
None,
)
}
.map_err(|_error| PalError::new(PalErrorKind::Other))?;
buf.truncate(transferred as usize);
Ok(buf)
}
fn finish(&self, pty: PtyId) {
let taken = table()
.lock()
.expect("pty table")
.ptys
.get_mut(&pty.0)
.and_then(|entry| {
entry
.hpcon
.take()
.map(|hpcon| (hpcon, Arc::clone(&entry.host_input)))
});
let Some((hpcon, host_input)) = taken else {
return;
};
unsafe {
ClosePseudoConsole(hpcon);
}
host_input.cancel();
}
fn close(&self, pty: PtyId) {
let Some(entry) = table().lock().expect("pty table").ptys.remove(&pty.0) else {
return;
};
if let Some(hpcon) = entry.hpcon {
unsafe {
ClosePseudoConsole(hpcon);
}
}
entry.host_input.cancel();
entry.host_output.cancel();
}
}