use std::io::{self, Read, Write};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc::{sync_channel, SyncSender, TrySendError};
use std::time::Duration;
use portable_pty::{
native_pty_system, ChildKiller, CommandBuilder, ExitStatus, MasterPty, PtySize,
};
use tokio::sync::mpsc;
const READ_CHUNK: usize = 8192;
const OUTPUT_CHANNEL_DEPTH: usize = 512;
const WRITE_CHANNEL_DEPTH: usize = 1024;
fn default_shell() -> String {
resolve_shell(std::env::var_os("SHELL"))
}
fn build_command(command: &[String], fallback: impl FnOnce() -> String) -> CommandBuilder {
match command.split_first() {
Some((program, args)) => {
let mut cmd = CommandBuilder::new(program);
cmd.args(args);
cmd
}
None => CommandBuilder::new(fallback()),
}
}
fn scrub_koh_env(cmd: &mut CommandBuilder) {
for (key, _) in std::env::vars_os() {
if is_koh_env_key(&key) {
cmd.env_remove(&key);
}
}
}
pub(crate) fn is_koh_env_key(key: &std::ffi::OsStr) -> bool {
key.to_string_lossy().starts_with("KOH_")
}
fn resolve_shell(shell_env: Option<std::ffi::OsString>) -> String {
if let Some(sh) = shell_env {
if !sh.is_empty() {
return sh.to_string_lossy().into_owned();
}
}
if cfg!(target_os = "android") {
"/system/bin/sh".to_string()
} else {
"/bin/sh".to_string()
}
}
#[derive(Debug, thiserror::Error)]
pub enum PtyError {
#[error("opening pty: {0}")]
OpenPty(#[source] io::Error),
#[error("spawning shell: {0}")]
Spawn(#[source] io::Error),
#[error("starting pty reader: {0}")]
Reader(#[from] io::Error),
#[error("resizing pty: {0}")]
Resize(#[source] io::Error),
}
pub struct Pty {
master: Box<dyn MasterPty + Send>,
writer_tx: SyncSender<Vec<u8>>,
child: Box<dyn portable_pty::Child + Send + Sync>,
killer: Box<dyn ChildKiller + Send + Sync>,
reaped: AtomicBool,
reader_handle: Option<std::thread::JoinHandle<()>>,
writer_handle: Option<std::thread::JoinHandle<()>>,
}
impl Pty {
pub fn spawn(
rows: u16,
cols: u16,
command: &[String],
term: &str,
) -> Result<(Self, mpsc::Receiver<Vec<u8>>), PtyError> {
let pty_system = native_pty_system();
let pair = pty_system
.openpty(PtySize {
rows,
cols,
pixel_width: 0,
pixel_height: 0,
})
.map_err(|e| PtyError::OpenPty(io::Error::other(e)))?;
let mut cmd = build_command(command, default_shell);
cmd.env("TERM", term);
scrub_koh_env(&mut cmd);
let child = pair
.slave
.spawn_command(cmd)
.map_err(|e| PtyError::Spawn(io::Error::other(e)))?;
let killer = child.clone_killer();
drop(pair.slave);
let mut reader = pair
.master
.try_clone_reader()
.map_err(|e| PtyError::Reader(io::Error::other(e)))?;
let mut writer = pair
.master
.take_writer()
.map_err(|e| PtyError::Reader(io::Error::other(e)))?;
let (tx, rx) = mpsc::channel::<Vec<u8>>(OUTPUT_CHANNEL_DEPTH);
let reader_handle = std::thread::Builder::new()
.name("koh-pty-reader".into())
.spawn(move || {
let mut buf = [0u8; READ_CHUNK];
loop {
match reader.read(&mut buf) {
Ok(0) => break, Ok(n) => {
let Some(chunk) = buf.get(..n) else { break };
if tx.blocking_send(chunk.to_vec()).is_err() {
break; }
}
Err(e) => {
tracing::debug!(error = %e, "pty reader stopping");
break;
}
}
}
})?;
let (writer_tx, writer_rx) = sync_channel::<Vec<u8>>(WRITE_CHANNEL_DEPTH);
let writer_handle = std::thread::Builder::new()
.name("koh-pty-writer".into())
.spawn(move || {
while let Ok(chunk) = writer_rx.recv() {
if writer
.write_all(&chunk)
.and_then(|()| writer.flush())
.is_err()
{
break; }
}
})?;
Ok((
Self {
master: pair.master,
writer_tx,
child,
killer,
reaped: AtomicBool::new(false),
reader_handle: Some(reader_handle),
writer_handle: Some(writer_handle),
},
rx,
))
}
pub fn shutdown(mut self) {
if !self.reaped.load(Ordering::SeqCst) {
if let Err(e) = self.killer.kill() {
tracing::warn!(error = %e, "pty kill on shutdown failed; reader join may stall");
}
}
let reader = self.reader_handle.take();
let writer = self.writer_handle.take();
drop(self);
if let Some(h) = writer {
let _ = h.join();
}
if let Some(h) = reader {
let _ = h.join();
}
}
pub fn write_input(&self, data: &[u8]) -> io::Result<()> {
match self.writer_tx.try_send(data.to_vec()) {
Ok(()) => Ok(()),
Err(TrySendError::Full(_)) => Err(io::Error::new(
io::ErrorKind::WouldBlock,
"pty writer queue full (child not draining its input)",
)),
Err(TrySendError::Disconnected(_)) => Err(io::Error::from(io::ErrorKind::BrokenPipe)),
}
}
pub fn resize(&self, rows: u16, cols: u16) -> Result<(), PtyError> {
self.master
.resize(PtySize {
rows,
cols,
pixel_width: 0,
pixel_height: 0,
})
.map_err(|e| PtyError::Resize(io::Error::other(e)))
}
pub fn try_wait(&mut self) -> std::io::Result<Option<ExitStatus>> {
let r = self.child.try_wait();
if matches!(r, Ok(Some(_))) {
self.reaped.store(true, Ordering::SeqCst);
}
r
}
pub fn kill(&mut self) -> std::io::Result<()> {
if self.reaped.load(Ordering::SeqCst) {
return Ok(());
}
self.killer.kill()
}
pub fn terminate_process_group(&self, force: bool) -> std::io::Result<()> {
if self.reaped.load(Ordering::SeqCst) {
return Ok(());
}
#[cfg(unix)]
if let Some(pid) = self.process_id().and_then(|pid| i32::try_from(pid).ok()) {
use nix::sys::signal::{kill, Signal};
use nix::unistd::Pid;
return kill(
Pid::from_raw(-pid),
if force {
Signal::SIGKILL
} else {
Signal::SIGHUP
},
)
.map_err(std::io::Error::other);
}
Ok(())
}
pub fn shutdown_process_group(
&mut self,
grace: Duration,
) -> std::io::Result<Option<ExitStatus>> {
if self.reaped.load(Ordering::SeqCst) {
return Ok(None);
}
let direct = self.killer.kill();
let group = self.terminate_process_group(false);
let hup_error = match (direct, group) {
(Err(error), Err(_)) => Some(error),
_ => None,
};
std::thread::sleep(grace);
let force_error = match self.terminate_process_group(true) {
Ok(()) => None,
Err(error) if error.raw_os_error() == Some(nix::libc::ESRCH) => None,
Err(error) => Some(error),
};
let deadline = std::time::Instant::now() + Duration::from_secs(1);
loop {
if let Some(status) = self.try_wait()? {
return Ok(Some(status));
}
if std::time::Instant::now() >= deadline {
return match force_error.or(hup_error) {
Some(error) => Err(error),
None => Ok(None),
};
}
std::thread::sleep(Duration::from_millis(5));
}
}
pub fn kill_hard(&self) {
if self.reaped.load(Ordering::SeqCst) {
return;
}
#[cfg(unix)]
if let Some(pid) = self.process_id() {
use nix::sys::signal::{kill, Signal};
use nix::unistd::Pid;
if let Ok(pid) = i32::try_from(pid) {
let _ = kill(Pid::from_raw(pid), Signal::SIGKILL);
}
}
}
#[must_use]
pub fn process_id(&self) -> Option<u32> {
self.child.process_id()
}
}
impl Drop for Pty {
fn drop(&mut self) {
if self.reaped.load(Ordering::SeqCst) {
return;
}
let _ = self.killer.kill();
self.kill_hard();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_command_passes_argv_verbatim_and_falls_back_when_empty() {
let argv: Vec<String> = ["zellij", "attach", "-c", "my session"]
.into_iter()
.map(String::from)
.collect();
let cmd = build_command(&argv, || "FALLBACK-MUST-NOT-BE-USED".to_owned());
let got: Vec<String> = cmd
.get_argv()
.iter()
.map(|a| a.to_string_lossy().into_owned())
.collect();
assert_eq!(got, argv, "argv must reach the child exactly as given");
let cmd = build_command(&[], || "/custom/shell".to_owned());
let got: Vec<String> = cmd
.get_argv()
.iter()
.map(|a| a.to_string_lossy().into_owned())
.collect();
assert_eq!(got, ["/custom/shell"]);
}
#[test]
fn resolve_shell_prefers_env_then_platform_default() {
use std::ffi::OsString;
assert_eq!(
resolve_shell(Some(OsString::from("/usr/bin/fish"))),
"/usr/bin/fish"
);
let empty = resolve_shell(Some(OsString::new()));
let unset = resolve_shell(None);
assert_eq!(empty, unset, "empty SHELL falls through like unset");
assert!(
unset.starts_with('/') && !unset.is_empty(),
"an absolute fallback path"
);
if cfg!(target_os = "android") {
assert_eq!(unset, "/system/bin/sh");
} else {
assert_eq!(unset, "/bin/sh");
}
}
#[test]
fn scrub_removes_koh_key_passphrase_even_when_inherited() {
std::env::set_var("KOH_KEY_PASSPHRASE", "topsecret-unit");
std::env::set_var("KOH_DNS", "1.1.1.1");
let mut cmd = CommandBuilder::new("/bin/sh");
assert!(
cmd.get_env("KOH_KEY_PASSPHRASE").is_some(),
"the builder seeds the parent env, so the var is present before scrubbing"
);
scrub_koh_env(&mut cmd);
assert!(
cmd.get_env("KOH_KEY_PASSPHRASE").is_none(),
"the identity-key passphrase must be scrubbed from the child env"
);
assert!(
cmd.get_env("KOH_DNS").is_none(),
"operational KOH_* vars are scrubbed too"
);
std::env::remove_var("KOH_KEY_PASSPHRASE");
std::env::remove_var("KOH_DNS");
}
#[test]
#[allow(
clippy::items_after_statements,
reason = "`_assert_typed` is a deliberate compile-time signature assertion kept beside the runtime checks it documents"
)]
fn pty_error_variants_are_constructible_and_reachable() {
let mk = || io::Error::other("boom");
for e in [
PtyError::OpenPty(mk()),
PtyError::Spawn(mk()),
PtyError::Reader(mk()),
PtyError::Resize(mk()),
] {
assert!(!e.to_string().is_empty(), "variant must Display");
}
let from_io: PtyError = mk().into();
assert!(matches!(from_io, PtyError::Reader(_)));
let absorbed: anyhow::Error = PtyError::OpenPty(mk()).into();
assert!(absorbed.to_string().contains("opening pty"));
fn _assert_typed(r: Result<(), PtyError>) -> Result<(), PtyError> {
r
}
}
}