#![allow(unsafe_code)]
use std::io::{self, Write as _};
use std::os::fd::{FromRawFd as _, OwnedFd, RawFd};
use std::sync::OnceLock;
use std::time::{Duration, Instant};
use ferroday_cage::{ExitStatus, Pty, Running, Terminal};
use rustix::event::{PollFd, PollFlags, Timespec};
use rustix::fs::OFlags;
use rustix::io::Errno;
use rustix::process::Signal;
use rustix::runtime::{How, KernelSigSet};
use rustix::termios::{self, OptionalActions, SpecialCodeIndex, Termios};
use syscalls::Sysno;
const RELAYED: [Signal; 8] = [
Signal::WINCH,
Signal::INT,
Signal::TERM,
Signal::HUP,
Signal::QUIT,
Signal::TSTP,
Signal::CONT,
Signal::TTOU,
];
const SIGINFO_LEN: usize = 128;
const RELAY_BUF: usize = 8192;
static ORIGINAL: OnceLock<Termios> = OnceLock::new();
static PANIC_HOOK: OnceLock<()> = OnceLock::new();
#[derive(Debug)]
pub enum RelayError {
Cage(ferroday_cage::Error),
Io(io::Error),
}
impl From<ferroday_cage::Error> for RelayError {
fn from(error: ferroday_cage::Error) -> RelayError {
RelayError::Cage(error)
}
}
impl From<Errno> for RelayError {
fn from(errno: Errno) -> RelayError {
RelayError::Io(errno.into())
}
}
impl From<io::Error> for RelayError {
fn from(error: io::Error) -> RelayError {
RelayError::Io(error)
}
}
pub struct Signals {
fd: OwnedFd,
}
impl Signals {
pub fn install() -> Result<Signals, RelayError> {
let mut blocked = KernelSigSet::empty();
let mut mask = 0u64;
for signal in RELAYED {
blocked.insert(signal);
mask |= 1u64 << (signal.as_raw() - 1);
}
unsafe { rustix::runtime::kernel_sigprocmask(How::BLOCK, Some(&blocked)) }?;
let fd = unsafe {
syscalls::syscall!(
Sysno::signalfd4,
-1i32,
&mask as *const u64,
size_of::<u64>(),
SFD_CLOEXEC | SFD_NONBLOCK
)
}
.map_err(|err| io::Error::from_raw_os_error(err.into_raw()))?;
let fd = unsafe { OwnedFd::from_raw_fd(fd as RawFd) };
Ok(Signals { fd })
}
fn next(&self) -> Result<Option<Signal>, RelayError> {
let mut record = [0u8; SIGINFO_LEN];
loop {
match rustix::io::read(&self.fd, &mut record) {
Ok(0) => return Ok(None),
Ok(_) => {
let number = i32::from_ne_bytes(
record[0..4]
.try_into()
.expect("a four-byte slice converts to an array"),
);
if let Some(signal) = Signal::from_named_raw(number) {
return Ok(Some(signal));
}
}
Err(Errno::AGAIN | Errno::INTR) => return Ok(None),
Err(errno) => return Err(errno.into()),
}
}
}
}
const SFD_CLOEXEC: usize = OFlags::CLOEXEC.bits() as usize;
const SFD_NONBLOCK: usize = OFlags::NONBLOCK.bits() as usize;
fn caller_is_terminal() -> bool {
termios::tcgetattr(rustix::stdio::stdin()).is_ok()
}
pub fn caller_size() -> Option<(u16, u16)> {
let size = termios::tcgetwinsize(rustix::stdio::stdin()).ok()?;
(size.ws_row != 0 && size.ws_col != 0).then_some((size.ws_row, size.ws_col))
}
pub fn terminal_for_caller() -> Terminal {
match caller_size() {
Some((rows, cols)) => Terminal::new().size(rows, cols),
None => Terminal::new(),
}
}
fn restore() {
if let Some(original) = ORIGINAL.get() {
let _ = termios::tcsetattr(rustix::stdio::stdin(), OptionalActions::Flush, original);
}
}
fn enter_raw() -> Result<(), RelayError> {
let current = termios::tcgetattr(rustix::stdio::stdin())?;
let original = ORIGINAL.get_or_init(|| current.clone());
let mut raw = original.clone();
raw.make_raw();
termios::tcsetattr(rustix::stdio::stdin(), OptionalActions::Now, &raw)?;
Ok(())
}
fn install_panic_hook() {
PANIC_HOOK.get_or_init(|| {
let default = std::panic::take_hook();
std::panic::set_hook(Box::new(move |info| {
restore();
default(info);
}));
});
}
struct RawMode;
impl Drop for RawMode {
fn drop(&mut self) {
restore();
}
}
const PRIMARY: usize = 0;
const SIGNALS: usize = 1;
const STDIN: usize = 2;
#[derive(Clone, Copy, PartialEq, Eq)]
enum Escalation {
Running,
Terminated,
Killed,
}
pub fn relay(
mut running: Running<'_>,
pty: Pty,
signals: Signals,
timeout: Option<Duration>,
kill_after: Duration,
) -> Result<Option<ExitStatus>, RelayError> {
let flags = rustix::fs::fcntl_getfl(&pty)?;
rustix::fs::fcntl_setfl(&pty, flags | OFlags::NONBLOCK)?;
let interactive = caller_is_terminal();
if interactive {
install_panic_hook();
enter_raw()?;
}
let _raw = RawMode;
let outcome = pump(
&mut running,
&pty,
&signals,
interactive,
timeout,
kill_after,
);
drop(_raw);
let expired = outcome?;
let status = running.wait()?;
Ok(if expired { None } else { Some(status) })
}
fn pump(
running: &mut Running<'_>,
pty: &Pty,
signals: &Signals,
interactive: bool,
timeout: Option<Duration>,
kill_after: Duration,
) -> Result<bool, RelayError> {
let stdin = rustix::stdio::stdin();
let mut stdin_open = true;
let mut pending: Vec<u8> = Vec::new();
let mut buf = [0u8; RELAY_BUF];
let mut escalation = Escalation::Running;
let mut expired = false;
let mut deadline = timeout.map(|timeout| Instant::now() + timeout);
let mut interrupted = false;
loop {
let watch_stdin = stdin_open && pending.is_empty();
let mut primary_events = PollFlags::IN;
if !pending.is_empty() {
primary_events |= PollFlags::OUT;
}
let mut slots = [
PollFd::new(pty, primary_events),
PollFd::new(&signals.fd, PollFlags::IN),
PollFd::new(&stdin, PollFlags::IN),
];
let watched = if watch_stdin {
&mut slots[..]
} else {
&mut slots[..STDIN]
};
let wait = deadline.map(timespec_until);
match rustix::event::poll(watched, wait.as_ref()) {
Ok(0) => {
expired = true;
match escalation {
Escalation::Running => {
running.terminate()?;
escalation = Escalation::Terminated;
deadline = Some(Instant::now() + kill_after);
}
Escalation::Terminated => {
running.kill()?;
escalation = Escalation::Killed;
deadline = None;
}
Escalation::Killed => deadline = None,
}
continue;
}
Ok(_) => {}
Err(Errno::INTR) => continue,
Err(errno) => return Err(errno.into()),
}
let ready: [PollFlags; 3] = std::array::from_fn(|index| slots[index].revents());
if ready[SIGNALS].intersects(PollFlags::IN) {
while let Some(signal) = signals.next()? {
match signal {
Signal::WINCH => {
if let Some((rows, cols)) = caller_size() {
pty.resize(rows, cols)?;
}
}
Signal::INT | Signal::TERM | Signal::HUP | Signal::QUIT => {
if interrupted {
running.kill()?;
} else {
running.terminate()?;
interrupted = true;
}
}
Signal::TSTP => {
restore();
rustix::process::kill_process(rustix::process::getpid(), Signal::STOP)
.map_err(RelayError::from)?;
}
Signal::CONT if interactive => {
enter_raw()?;
if let Some((rows, cols)) = caller_size() {
pty.resize(rows, cols)?;
}
}
_ => {}
}
}
}
if ready[PRIMARY].intersects(PollFlags::OUT) {
flush_pending(pty, &mut pending)?;
}
if ready[STDIN]
.intersects(PollFlags::IN | PollFlags::HUP | PollFlags::ERR | PollFlags::NVAL)
{
match rustix::io::read(stdin, &mut buf) {
Ok(0) | Err(Errno::BADF) => {
let veof = replica_veof(pty);
pending.push(veof);
flush_pending(pty, &mut pending)?;
stdin_open = false;
}
Ok(read) => {
pending.extend_from_slice(&buf[..read]);
flush_pending(pty, &mut pending)?;
}
Err(Errno::INTR | Errno::AGAIN) => {}
Err(errno) => return Err(errno.into()),
}
}
if ready[PRIMARY].intersects(PollFlags::IN | PollFlags::HUP) {
match rustix::io::read(pty, &mut buf) {
Ok(0) => return Ok(expired),
Ok(read) => {
let mut out = io::stdout();
out.write_all(&buf[..read])?;
out.flush()?;
}
Err(Errno::INTR | Errno::AGAIN) => {}
Err(Errno::IO) => return Ok(expired),
Err(errno) => return Err(errno.into()),
}
}
}
}
fn flush_pending(pty: &Pty, pending: &mut Vec<u8>) -> Result<(), RelayError> {
while !pending.is_empty() {
match rustix::io::write(pty, pending) {
Ok(0) => return Ok(()),
Ok(written) => {
pending.drain(..written);
}
Err(Errno::AGAIN | Errno::INTR) => return Ok(()),
Err(Errno::IO) => {
pending.clear();
return Ok(());
}
Err(errno) => return Err(errno.into()),
}
}
Ok(())
}
fn replica_veof(pty: &Pty) -> u8 {
termios::tcgetattr(pty)
.map(|settings| settings.special_codes[SpecialCodeIndex::VEOF])
.unwrap_or(0o004)
}
fn timespec_until(deadline: Instant) -> Timespec {
let remaining = deadline.saturating_duration_since(Instant::now());
Timespec {
tv_sec: i64::try_from(remaining.as_secs()).unwrap_or(i64::MAX),
tv_nsec: i64::from(remaining.subsec_nanos()),
}
}