use std::collections::HashMap;
use std::future::Future;
use std::io::{Read, Write};
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::Mutex;
use std::task::{Context, Poll};
use std::thread;
use crate::backend::{
BoxFuture, TerminalParams, TtyControl, TtyControlHandle, TtyError, TtyHandle,
};
use crate::control::signal_from_name;
use bytes::Bytes;
use futures_core::Stream;
use portable_pty::{native_pty_system, ChildKiller, CommandBuilder, MasterPty, PtySize};
use tokio::sync::{mpsc, oneshot};
use tokio_stream::wrappers::ReceiverStream;
use tracing::{debug, warn};
pub enum StdinCmd {
Bytes(Vec<u8>),
Eof,
}
#[derive(Clone)]
pub struct PtyControl {
master: Arc<Mutex<Box<dyn MasterPty + Send>>>,
killer: Arc<Mutex<Box<dyn ChildKiller + Send + Sync>>>,
pid: Option<u32>,
}
impl PtyControl {
pub fn new(
master: Arc<Mutex<Box<dyn MasterPty + Send>>>,
killer: Arc<Mutex<Box<dyn ChildKiller + Send + Sync>>>,
pid: Option<u32>,
) -> Self {
Self {
master,
killer,
pid,
}
}
}
impl TtyControl for PtyControl {
fn resize(&self, cols: u16, rows: u16, pixel_width: u16, pixel_height: u16) {
let size = PtySize {
cols,
rows,
pixel_width,
pixel_height,
};
let master = self.master.lock().unwrap_or_else(|e| e.into_inner());
if let Err(e) = master.resize(size) {
warn!("pty resize failed: {e}");
}
}
fn signal(&self, name: &str) {
#[cfg(unix)]
{
if let Some(pid) = self.pid {
if let Some(sig) = signal_from_name(name) {
let pgid = pid as i32;
let r = unsafe { libc::kill(-pgid, sig) };
if r == 0 {
return;
}
let err = std::io::Error::last_os_error();
let r2 = unsafe { libc::kill(pgid, sig) };
if r2 == 0 {
return;
}
warn!(
"pty signal `{name}` (group {pgid}) failed: {err}; \
direct kill also failed: {}",
std::io::Error::last_os_error()
);
return;
}
}
let mut killer = self.killer.lock().unwrap_or_else(|e| e.into_inner());
if let Err(e) = killer.kill() {
warn!("pty fallback ChildKiller::kill failed: {e}");
}
}
#[cfg(not(unix))]
{
let _ = name;
let mut killer = self.killer.lock().unwrap_or_else(|e| e.into_inner());
if let Err(e) = killer.kill() {
warn!("pty ChildKiller::kill failed: {e}");
}
}
}
}
pub struct LocalExitFuture {
rx: oneshot::Receiver<i32>,
killer: Option<Box<dyn ChildKiller + Send + Sync>>,
}
impl LocalExitFuture {
fn new(rx: oneshot::Receiver<i32>, killer: Box<dyn ChildKiller + Send + Sync>) -> Self {
Self {
rx,
killer: Some(killer),
}
}
}
impl Future for LocalExitFuture {
type Output = Result<i32, TtyError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match Pin::new(&mut self.rx).poll(cx) {
Poll::Pending => Poll::Pending,
Poll::Ready(Ok(code)) => {
self.killer.take();
Poll::Ready(Ok(code))
}
Poll::Ready(Err(_)) => {
self.killer.take();
Poll::Ready(Err(TtyError::WaitFailed {
message: "waiter thread exited without sending exit code".to_string(),
}))
}
}
}
}
impl Drop for LocalExitFuture {
fn drop(&mut self) {
if let Some(killer) = self.killer.take() {
let mut killer = killer;
if let Err(e) = killer.kill() {
debug!("LocalExitFuture drop: ChildKiller::kill failed: {e}");
}
}
}
}
struct StdinSink {
tx: mpsc::Sender<StdinCmd>,
inflight: Option<InflightSend>,
inflight_len: usize,
inflight_close: Option<InflightSend>,
close_sent: bool,
}
type InflightSend = Pin<Box<dyn Future<Output = Result<(), mpsc::error::SendError<()>>> + Send>>;
impl StdinSink {
fn new(tx: mpsc::Sender<StdinCmd>) -> Self {
Self {
tx,
inflight: None,
inflight_len: 0,
inflight_close: None,
close_sent: false,
}
}
}
impl tokio::io::AsyncWrite for StdinSink {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<Result<usize, std::io::Error>> {
if let Some(fut) = self.inflight.as_mut() {
match fut.as_mut().poll(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Ok(())) => {
let n = self.inflight_len;
self.inflight = None;
self.inflight_len = 0;
return Poll::Ready(Ok(n));
}
Poll::Ready(Err(_)) => {
self.inflight = None;
self.inflight_len = 0;
return Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"stdin channel closed",
)));
}
}
}
match self.tx.try_send(StdinCmd::Bytes(buf.to_vec())) {
Ok(()) => Poll::Ready(Ok(buf.len())),
Err(mpsc::error::TrySendError::Full(_)) => {
let tx = self.tx.clone();
let bytes = buf.to_vec();
let len = bytes.len();
self.inflight_len = len;
self.inflight = Some(Box::pin(async move {
let permit = tx.reserve().await?;
permit.send(StdinCmd::Bytes(bytes));
Ok(())
}));
self.poll_write(cx, buf)
}
Err(mpsc::error::TrySendError::Closed(_)) => Poll::Ready(Err(std::io::Error::new(
std::io::ErrorKind::BrokenPipe,
"stdin channel closed",
))),
}
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), std::io::Error>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Result<(), std::io::Error>> {
if self.close_sent {
return Poll::Ready(Ok(()));
}
if let Some(fut) = self.inflight_close.as_mut() {
match fut.as_mut().poll(cx) {
Poll::Pending => return Poll::Pending,
Poll::Ready(Ok(())) => {
self.inflight_close = None;
self.close_sent = true;
return Poll::Ready(Ok(()));
}
Poll::Ready(Err(_)) => {
self.inflight_close = None;
self.close_sent = true;
return Poll::Ready(Ok(()));
}
}
}
match self.tx.try_send(StdinCmd::Eof) {
Ok(()) => {
self.close_sent = true;
Poll::Ready(Ok(()))
}
Err(mpsc::error::TrySendError::Full(_)) => {
let tx = self.tx.clone();
self.inflight_close = Some(Box::pin(async move {
let permit = tx.reserve().await?;
permit.send(StdinCmd::Eof);
Ok(())
}));
self.poll_shutdown(cx)
}
Err(mpsc::error::TrySendError::Closed(_)) => {
self.close_sent = true;
Poll::Ready(Ok(()))
}
}
}
}
pub fn allocate_pty(
terminal: TerminalParams,
cmd: Vec<String>,
cwd: Option<PathBuf>,
env: HashMap<String, String>,
) -> Result<TtyHandle, TtyError> {
if cmd.is_empty() {
return Err(TtyError::AllocFailed {
message: "cmd must be non-empty".to_string(),
});
}
let pty_system = native_pty_system();
let size = PtySize {
cols: terminal.cols,
rows: terminal.rows,
pixel_width: terminal.pixel_width,
pixel_height: terminal.pixel_height,
};
let pair = pty_system
.openpty(size)
.map_err(|e| TtyError::AllocFailed {
message: format!("openpty: {e}"),
})?;
let mut builder = CommandBuilder::new(&cmd[0]);
for arg in &cmd[1..] {
builder.arg(arg);
}
if let Some(cwd) = cwd {
builder.cwd(cwd);
}
for (k, v) in env {
builder.env(k, v);
}
builder.set_controlling_tty(true);
let mut child = pair
.slave
.spawn_command(builder)
.map_err(|e| TtyError::AllocFailed {
message: format!("spawn_command: {e}"),
})?;
drop(pair.slave);
let pid = child.process_id();
let killer = child.clone_killer();
let killer_for_control = killer.clone_killer();
let master: Arc<Mutex<Box<dyn MasterPty + Send>>> = Arc::new(Mutex::new(pair.master));
let reader_master = master.clone();
let (stdout_tx, stdout_rx) = mpsc::channel::<Bytes>(64);
thread::Builder::new()
.name("pty-reader".into())
.spawn(move || {
let reader = {
let m = reader_master.lock().unwrap_or_else(|e| e.into_inner());
match m.try_clone_reader() {
Ok(r) => r,
Err(e) => {
warn!("pty-reader: try_clone_reader failed: {e}");
let _ = stdout_tx.blocking_send(Bytes::new());
return;
}
}
};
let mut reader = reader;
let mut buf = vec![0u8; 8192];
loop {
match reader.read(&mut buf) {
Ok(0) => break,
Ok(n) => {
let chunk = Bytes::copy_from_slice(&buf[..n]);
if stdout_tx.blocking_send(chunk).is_err() {
break;
}
}
Err(e) if e.kind() == std::io::ErrorKind::Interrupted => continue,
Err(e) => {
warn!("pty-reader: read error: {e}");
break;
}
}
}
let _ = stdout_tx.blocking_send(Bytes::new());
debug!("pty-reader thread done");
})
.map_err(|e| TtyError::AllocFailed {
message: format!("spawn pty-reader thread: {e}"),
})?;
let writer_master = master.clone();
let (stdin_tx, mut stdin_rx) = mpsc::channel::<StdinCmd>(64);
thread::Builder::new()
.name("pty-writer".into())
.spawn(move || {
let writer = {
let m = writer_master.lock().unwrap_or_else(|e| e.into_inner());
match m.take_writer() {
Ok(w) => w,
Err(e) => {
warn!("pty-writer: take_writer failed: {e}");
return;
}
}
};
let mut writer = writer;
while let Some(cmd) = stdin_rx.blocking_recv() {
match cmd {
StdinCmd::Bytes(bytes) => {
if let Err(e) = writer.write_all(&bytes) {
warn!("pty-writer: write_all failed: {e}");
break;
}
if let Err(e) = writer.flush() {
warn!("pty-writer: flush failed: {e}");
break;
}
}
StdinCmd::Eof => {
drop(writer);
break;
}
}
}
debug!("pty-writer thread done");
})
.map_err(|e| TtyError::AllocFailed {
message: format!("spawn pty-writer thread: {e}"),
})?;
let (exit_tx, exit_rx) = oneshot::channel::<i32>();
thread::Builder::new()
.name("pty-waiter".into())
.spawn(move || {
let status = match child.wait() {
Ok(s) => s,
Err(e) => {
warn!("pty-waiter: wait failed: {e}");
let _ = exit_tx.send(-1);
return;
}
};
let code = status.exit_code() as i32;
debug!(exit_code = code, "pty-waiter: child reaped");
let _ = exit_tx.send(code);
})
.map_err(|e| TtyError::AllocFailed {
message: format!("spawn pty-waiter thread: {e}"),
})?;
let stdout: Pin<Box<dyn Stream<Item = Bytes> + Send>> =
Box::pin(ReceiverStream::new(stdout_rx));
let stdin: Box<dyn tokio::io::AsyncWrite + Send + Unpin> = Box::new(StdinSink::new(stdin_tx));
let exit_code: BoxFuture<Result<i32, TtyError>> =
Box::pin(LocalExitFuture::new(exit_rx, killer));
let control = Some(TtyControlHandle::new(Arc::new(PtyControl::new(
master,
Arc::new(Mutex::new(killer_for_control)),
pid,
))));
Ok(TtyHandle {
stdin,
stdout,
stderr: None,
exit_code,
control,
})
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::AsyncWriteExt;
use tokio_stream::StreamExt;
fn term() -> TerminalParams {
TerminalParams {
term: None,
cols: 80,
rows: 24,
pixel_width: 0,
pixel_height: 0,
modes: serde_json::Value::Null,
}
}
fn env_default() -> HashMap<String, String> {
let mut env = HashMap::new();
env.insert("TERM".to_string(), "dumb".to_string());
env
}
async fn wait_marker(marker: &std::path::Path) {
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5);
loop {
if marker.exists() {
let _ = std::fs::remove_file(marker);
return;
}
if tokio::time::Instant::now() >= deadline {
panic!("child never became ready (marker {:?} missing)", marker);
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
}
fn marker_path(name: &str) -> std::path::PathBuf {
std::env::temp_dir().join(format!(
"alktty_pty_{name}_{}_{}.txt",
std::process::id(),
nanos_seed()
))
}
fn nanos_seed() -> u64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("clock went backwards")
.as_nanos() as u64
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn happy_path_echo_exits_zero() {
let handle = allocate_pty(
term(),
vec!["echo".to_string(), "hello".to_string()],
None,
env_default(),
)
.expect("allocate");
let mut stdout = handle.stdout;
let mut collected = Vec::new();
while let Some(chunk) = stdout.next().await {
if chunk.is_empty() {
break;
}
collected.extend_from_slice(&chunk);
}
let code = handle.exit_code.await.expect("exit_code");
assert_eq!(code, 0);
let s = String::from_utf8_lossy(&collected);
assert!(s.contains("hello"), "stdout should contain hello: {s:?}");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn stdin_round_trip_cat() {
let handle =
allocate_pty(term(), vec!["cat".to_string()], None, env_default()).expect("allocate");
let mut stdin = handle.stdin;
let mut stdout = handle.stdout;
stdin.write_all(b"ping\n").await.expect("write");
stdin.shutdown().await.expect("shutdown (eof)");
let mut collected = Vec::new();
while let Some(chunk) = stdout.next().await {
if chunk.is_empty() {
break;
}
collected.extend_from_slice(&chunk);
}
let code = handle.exit_code.await.expect("exit_code");
let s = String::from_utf8_lossy(&collected);
assert!(s.contains("ping"), "stdout should contain ping: {s:?}");
assert_eq!(code, 0);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn resize_does_not_error() {
let handle = allocate_pty(
term(),
vec!["sleep".to_string(), "1".to_string()],
None,
env_default(),
)
.expect("allocate");
let control = handle.control.as_ref().expect("control");
control.resize(120, 40, 0, 0);
let _ = handle.exit_code.await.expect("exit_code");
}
#[cfg(unix)]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn signal_int_kills_child() {
let marker = marker_path("sigint_ready");
let cmd = format!("echo ready > '{}'; exec sleep 60", marker.display());
let handle = allocate_pty(
term(),
vec!["bash".to_string(), "-c".to_string(), cmd],
None,
env_default(),
)
.expect("allocate");
let control = handle.control.as_ref().expect("control");
wait_marker(&marker).await;
control.signal("INT");
let code = tokio::time::timeout(std::time::Duration::from_secs(5), handle.exit_code)
.await
.expect("exit timed out")
.expect("exit_code");
assert_ne!(
code, 0,
"signal-terminated child should report non-zero exit: {code}"
);
}
#[cfg(unix)]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn signal_reaches_process_group_child() {
let marker = marker_path("pgroup_ready");
let cmd = format!("echo ready > '{}'; sleep 60", marker.display());
let handle = allocate_pty(
term(),
vec!["bash".to_string(), "-c".to_string(), cmd],
None,
env_default(),
)
.expect("allocate");
let control = handle.control.as_ref().expect("control");
wait_marker(&marker).await;
control.signal("INT");
let code = tokio::time::timeout(std::time::Duration::from_secs(5), handle.exit_code)
.await
.expect("exit timed out")
.expect("exit_code");
assert_ne!(code, 0, "process group should have been killed: {code}");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancel_cleanup_kills_child_on_drop() {
let handle = allocate_pty(
term(),
vec!["sleep".to_string(), "60".to_string()],
None,
env_default(),
)
.expect("allocate");
drop(handle);
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
let probe = allocate_pty(
term(),
vec!["echo".to_string(), "ok".to_string()],
None,
env_default(),
)
.expect("allocate");
let code = tokio::time::timeout(std::time::Duration::from_secs(5), probe.exit_code)
.await
.expect("probe timed out")
.expect("probe exit_code");
assert_eq!(code, 0);
}
#[cfg(unix)]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn unknown_signal_falls_back_to_child_killer() {
let marker = marker_path("unknown_sig_ready");
let cmd = format!("echo ready > '{}'; exec sleep 60", marker.display());
let handle = allocate_pty(
term(),
vec!["bash".to_string(), "-c".to_string(), cmd],
None,
env_default(),
)
.expect("allocate");
let control = handle.control.as_ref().expect("control");
wait_marker(&marker).await;
control.signal("NOSUCH");
let code = tokio::time::timeout(std::time::Duration::from_secs(5), handle.exit_code)
.await
.expect("exit timed out")
.expect("exit_code");
assert_ne!(code, 0, "fallback kill should terminate the child: {code}");
}
#[cfg(unix)]
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn signal_after_child_exit_takes_both_kill_fallbacks() {
let handle = allocate_pty(
term(),
vec!["echo".to_string(), "done".to_string()],
None,
env_default(),
)
.expect("allocate");
let control = handle.control.clone().expect("control");
let code = handle.exit_code.await.expect("exit_code");
assert_eq!(code, 0);
control.signal("INT");
control.signal("NOSUCH");
}
}