#![allow(unsafe_code)]
use nix::fcntl::{FcntlArg, OFlag, fcntl};
use nix::pty::{OpenptyResult, Winsize, openpty};
use nix::sys::signal::{self, Signal};
use nix::sys::wait::{WaitPidFlag, WaitStatus, waitpid};
use nix::unistd::{ForkResult, Pid, execvp, fork, setsid};
use std::ffi::CString;
use std::os::unix::io::{AsRawFd, OwnedFd, RawFd};
use thiserror::Error;
#[derive(Debug, Error)]
pub enum PtyError {
#[error("failed to open PTY: {0}")]
OpenPty(#[source] nix::Error),
#[error("failed to fork: {0}")]
Fork(#[source] nix::Error),
#[error("failed to create session: {0}")]
Setsid(#[source] nix::Error),
#[error("failed to set controlling terminal: {0}")]
SetControllingTerminal(#[source] nix::Error),
#[error("failed to change directory: {0}")]
Chdir(#[source] std::io::Error),
#[error("failed to exec: {0}")]
Exec(#[source] nix::Error),
#[error("command is empty")]
EmptyCommand,
#[error("invalid command string: {0}")]
InvalidCommand(#[source] std::ffi::NulError),
#[error("failed to send signal: {0}")]
Signal(#[source] nix::Error),
#[error("failed to wait: {0}")]
Wait(#[source] nix::Error),
}
pub struct PtyProcess {
pub master: OwnedFd,
pub pid: Pid,
pub size: Winsize,
}
impl PtyProcess {
#[must_use]
pub fn master_fd(&self) -> RawFd {
self.master.as_raw_fd()
}
pub fn signal(&self, sig: Signal) -> Result<(), PtyError> {
signal::kill(self.pid, sig).map_err(PtyError::Signal)
}
pub fn try_wait(&self) -> Result<Option<i32>, PtyError> {
match waitpid(self.pid, Some(WaitPidFlag::WNOHANG)).map_err(PtyError::Wait)? {
WaitStatus::Exited(_, code) => Ok(Some(code)),
WaitStatus::Signaled(_, sig, _) => Ok(Some(128 + sig as i32)),
_ => Ok(None),
}
}
pub fn wait(&self) -> Result<i32, PtyError> {
match waitpid(self.pid, None).map_err(PtyError::Wait)? {
WaitStatus::Exited(_, code) => Ok(code),
WaitStatus::Signaled(_, sig, _) => Ok(128 + sig as i32),
status => {
tracing::warn!(?status, "unexpected wait status");
Ok(-1)
}
}
}
pub fn resize(&self, rows: u16, cols: u16) -> Result<(), PtyError> {
let winsize = Winsize {
ws_row: rows,
ws_col: cols,
ws_xpixel: 0,
ws_ypixel: 0,
};
unsafe {
let ret = libc::ioctl(self.master.as_raw_fd(), libc::TIOCSWINSZ, &winsize);
if ret < 0 {
return Err(PtyError::SetControllingTerminal(nix::Error::last()));
}
}
Ok(())
}
}
const ESSENTIAL_ENV_VARS: &[&str] = &[
"PATH", "HOME", "USER", "TERM", "SHELL", "LANG", "XDG_RUNTIME_DIR", "DBUS_SESSION_BUS_ADDRESS", ];
#[derive(Debug, Default)]
pub struct SpawnEnv {
pub vars: Vec<(String, String)>,
}
pub fn spawn(cmd: &[String], rows: u16, cols: u16) -> Result<PtyProcess, PtyError> {
spawn_with_env(cmd, rows, cols, &SpawnEnv::default(), None)
}
pub fn spawn_with_env(
cmd: &[String],
rows: u16,
cols: u16,
env: &SpawnEnv,
cwd: Option<&str>,
) -> Result<PtyProcess, PtyError> {
if cmd.is_empty() {
return Err(PtyError::EmptyCommand);
}
let winsize = Winsize {
ws_row: rows,
ws_col: cols,
ws_xpixel: 0,
ws_ypixel: 0,
};
let explicit_keys: std::collections::HashSet<&str> =
env.vars.iter().map(|(k, _)| k.as_str()).collect();
#[allow(unused_variables)]
let essential: Vec<(String, String)> = ESSENTIAL_ENV_VARS
.iter()
.filter(|k| !explicit_keys.contains(**k))
.filter_map(|k| std::env::var(k).ok().map(|v| (k.to_string(), v)))
.collect();
let OpenptyResult { master, slave } = openpty(&winsize, None).map_err(PtyError::OpenPty)?;
match unsafe { fork() }.map_err(PtyError::Fork)? {
ForkResult::Parent { child } => {
drop(slave);
let flags = fcntl(&master, FcntlArg::F_GETFL).map_err(PtyError::OpenPty)?;
let mut flags = OFlag::from_bits_retain(flags);
flags.insert(OFlag::O_NONBLOCK);
fcntl(&master, FcntlArg::F_SETFL(flags)).map_err(PtyError::OpenPty)?;
Ok(PtyProcess {
master,
pid: child,
size: winsize,
})
}
ForkResult::Child => {
drop(master);
if setsid().is_err() {
unsafe { libc::_exit(1) };
}
unsafe {
if libc::ioctl(slave.as_raw_fd(), libc::TIOCSCTTY as _, 0) < 0 {
libc::_exit(1);
}
}
let slave_fd = slave.as_raw_fd();
unsafe {
if libc::dup2(slave_fd, libc::STDIN_FILENO) < 0
|| libc::dup2(slave_fd, libc::STDOUT_FILENO) < 0
|| libc::dup2(slave_fd, libc::STDERR_FILENO) < 0
{
libc::_exit(1);
}
}
if slave_fd > 2 {
drop(slave);
}
unsafe {
for (key, _) in std::env::vars() {
std::env::remove_var(&key);
}
for (key, value) in &essential {
std::env::set_var(key, value);
}
for (key, value) in &env.vars {
std::env::set_var(key, value);
}
}
if let Some(dir) = cwd
&& std::env::set_current_dir(dir).is_err()
{
unsafe { libc::_exit(1) };
}
let Ok(prog) = CString::new(cmd[0].as_str()) else {
unsafe { libc::_exit(1) };
};
let Ok(args) = cmd
.iter()
.map(|s| CString::new(s.as_str()))
.collect::<Result<Vec<CString>, _>>()
else {
unsafe { libc::_exit(1) };
};
let _: Result<std::convert::Infallible, _> = execvp(&prog, &args);
unsafe { libc::_exit(127) };
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn test_spawn_echo() {
let pty = spawn(&["sh".into(), "-c".into(), "echo hello".into()], 24, 80).unwrap();
let exit_code = pty.wait().unwrap();
assert_eq!(exit_code, 0);
}
#[test]
fn test_spawn_exit_code() {
let pty = spawn(&["sh".into(), "-c".into(), "exit 42".into()], 24, 80).unwrap();
let exit_code = pty.wait().unwrap();
assert_eq!(exit_code, 42);
}
#[test]
fn test_spawn_empty_command() {
let result = spawn(&[], 24, 80);
assert!(matches!(result, Err(PtyError::EmptyCommand)));
}
#[test]
fn test_try_wait() {
let pty = spawn(&["sleep".into(), "0.1".into()], 24, 80).unwrap();
let result = pty.try_wait().unwrap();
assert!(result.is_none());
std::thread::sleep(Duration::from_millis(200));
let result = pty.try_wait().unwrap();
assert_eq!(result, Some(0));
}
}