Luigit
repositories / smith

smith

There are many coding harnesses - but this one is fast

owned by admin

smith-core/src/bash.rs

Raw
//! Bounded, cancellable shell execution with owned-process-tree termination
//! (`SMH-SPEC-SPEC0001`, Tools).
//!
//! Pipes are drained by dedicated threads so the deadline loop is never
//! blocked by a silent child; per-stream caps keep the fitting prefix and
//! bound all collection, including post-exit draining.

use smith::error::{Result, SmithError};
use smith::tool::CancelHandle;
use std::io::Read;
use std::process::{Child, Command, ExitStatus, Stdio};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc;
use std::time::{Duration, Instant};

const READ_CHUNK: usize = 16 * 1024;
const GRACE_DRAIN: Duration = Duration::from_millis(500);
const POLL_SLICE: Duration = Duration::from_millis(25);

/// Result of a bounded shell execution.
#[derive(Clone, PartialEq, Eq, Debug, serde::Serialize, serde::Deserialize)]
pub struct BashResult {
    /// Exit status; `None` when killed by timeout, cancellation, or signal.
    pub exit_code: Option<i32>,
    /// Stdout, capped with the fitting prefix retained.
    pub stdout: String,
    /// Stderr, capped with the fitting prefix retained.
    pub stderr: String,
    /// Cancellation was requested.
    pub cancelled: bool,
    /// The deadline elapsed and the process tree was killed.
    pub timed_out: bool,
    /// Wall-clock duration in ms.
    pub duration_ms: u64,
}

/// Run `cmd` under `bash -c` in `workdir` with a hard deadline, per-stream
/// output caps, and a shared cancellation handle.
///
/// # Errors
///
/// Returns [`SmithError::Tool`] when the process cannot be spawned.
pub fn execute(
    cmd: &str,
    workdir: Option<&str>,
    timeout_ms: u64,
    max_output_bytes: usize,
    cancel: &CancelHandle,
) -> Result<BashResult> {
    let started = Instant::now();
    let deadline = started + Duration::from_millis(timeout_ms.max(1));
    let mut child = spawn(cmd, workdir)?;
    let (tx, rx) = mpsc::channel::<(Stream, Vec<u8>)>();
    let stdout_remaining = Arc::new(AtomicUsize::new(max_output_bytes));
    let stderr_remaining = Arc::new(AtomicUsize::new(max_output_bytes));
    let handles = drain_pipes(
        &mut child,
        &tx,
        Arc::clone(&stdout_remaining),
        Arc::clone(&stderr_remaining),
    );
    drop(tx); // reader threads hold the only senders

    let mut stdout: Vec<u8> = Vec::new();
    let mut stderr: Vec<u8> = Vec::new();
    let mut retain = |(stream, chunk): (Stream, Vec<u8>)| match stream {
        Stream::Stdout => stdout.extend_from_slice(&chunk),
        Stream::Stderr => stderr.extend_from_slice(&chunk),
    };
    let mut status: Option<ExitStatus> = None;
    let mut cancelled = false;
    let mut timed_out = false;

    while status.is_none() {
        if cancel.is_cancelled() {
            cancelled = true;
            kill_tree(&mut child);
            status = child.wait().ok();
            break;
        }
        let now = Instant::now();
        if now >= deadline {
            timed_out = true;
            kill_tree(&mut child);
            status = child.wait().ok();
            break;
        }
        let wait = POLL_SLICE.min(deadline - now);
        // A timeout or closed pipes fall through to polling the child, so
        // cancellation and the deadline stay live while it lingers. Closed
        // pipes return at once; pace the poll so the wait never spins.
        match rx.recv_timeout(wait) {
            Ok(item) => retain(item),
            Err(mpsc::RecvTimeoutError::Timeout) => {}
            Err(mpsc::RecvTimeoutError::Disconnected) => std::thread::sleep(wait),
        }
        if let Ok(Some(st)) = child.try_wait() {
            status = Some(st);
        }
    }

    // Bounded post-exit drain: retain late output until pipes close or the
    // grace window expires. Threads only ever send capped prefixes.
    let grace_end = Instant::now() + GRACE_DRAIN;
    while let Ok(item) = rx.recv_timeout(grace_end.saturating_duration_since(Instant::now())) {
        retain(item);
    }
    for handle in handles {
        let _ = handle.join();
    }

    Ok(BashResult {
        exit_code: status.as_ref().and_then(ExitStatus::code),
        stdout: String::from_utf8_lossy(&stdout).into_owned(),
        stderr: String::from_utf8_lossy(&stderr).into_owned(),
        cancelled,
        timed_out,
        duration_ms: u64::try_from(started.elapsed().as_millis()).unwrap_or(u64::MAX),
    })
}

#[derive(Clone, Copy)]
enum Stream {
    Stdout,
    Stderr,
}

fn spawn(cmd: &str, workdir: Option<&str>) -> Result<Child> {
    let mut command = Command::new("bash");
    command
        .arg("-c")
        .arg(cmd)
        .stdout(Stdio::piped())
        .stderr(Stdio::piped());
    if let Some(dir) = workdir {
        command.current_dir(dir);
    }
    #[cfg(unix)]
    {
        use std::os::unix::process::CommandExt;
        // Own the process tree: the child becomes a group leader so the
        // deadline path can signal every descendant.
        command.process_group(0);
    }
    command.spawn().map_err(|e| SmithError::Tool {
        code: "BASH_SPAWN".to_string(),
        message: e.to_string(),
    })
}

/// Kill the whole owned process group where the platform supports it,
/// falling back to the direct child.
fn kill_tree(child: &mut Child) {
    #[cfg(unix)]
    {
        use nix::sys::signal::{Signal, killpg};
        use nix::unistd::Pid;
        let pgid = Pid::from_raw(i32::try_from(child.id()).unwrap_or(i32::MAX));
        let _ = killpg(pgid, Signal::SIGKILL);
    }
    let _ = child.kill();
}

/// Spawn one reader thread per pipe. Each thread forwards capped prefixes
/// (never more than `remaining` permits) and keeps draining to EOF so a
/// chatty descendant can never deadlock the parent.
fn drain_pipes(
    child: &mut Child,
    tx: &mpsc::Sender<(Stream, Vec<u8>)>,
    stdout_remaining: Arc<AtomicUsize>,
    stderr_remaining: Arc<AtomicUsize>,
) -> Vec<std::thread::JoinHandle<()>> {
    let mut handles = Vec::new();
    if let Some(out) = child.stdout.take() {
        handles.push(spawn_reader(
            out,
            tx.clone(),
            Stream::Stdout,
            stdout_remaining,
        ));
    }
    if let Some(err) = child.stderr.take() {
        handles.push(spawn_reader(
            err,
            tx.clone(),
            Stream::Stderr,
            stderr_remaining,
        ));
    }
    handles
}

/// Forward capped prefixes from one pipe until it closes.
fn spawn_reader<R: Read + Send + 'static>(
    mut pipe: R,
    tx: mpsc::Sender<(Stream, Vec<u8>)>,
    stream: Stream,
    remaining: Arc<AtomicUsize>,
) -> std::thread::JoinHandle<()> {
    std::thread::spawn(move || {
        let mut buf = vec![0u8; READ_CHUNK];
        loop {
            match pipe.read(&mut buf) {
                // End of pipe or a read failure both end this reader.
                Ok(0) | Err(_) => break,
                Ok(n) => {
                    let keep = n.min(remaining.load(Ordering::Relaxed));
                    if keep > 0 {
                        if tx.send((stream, buf[..keep].to_vec())).is_err() {
                            break;
                        }
                        remaining.fetch_sub(keep, Ordering::Relaxed);
                    }
                    // Beyond the cap we keep reading (without retaining) so
                    // the pipe never back-pressures the child indefinitely.
                }
            }
        }
    })
}

#[cfg(test)]
#[cfg(unix)]
#[expect(
    clippy::unwrap_used,
    reason = "tests may panic on invariant violations"
)]
mod tests {
    use super::*;
    use crate::process_signal::{ProcessSignal, closed};
    use std::net::TcpStream;
    use std::thread;

    fn run(cmd: &str, timeout_ms: u64, cap: usize) -> BashResult {
        execute(cmd, None, timeout_ms, cap, &CancelHandle::new()).unwrap()
    }

    #[test]
    fn completes_and_captures_output() {
        let res = run("printf 'hello'", 5_000, 1024);
        assert_eq!(res.exit_code, Some(0));
        assert_eq!(res.stdout, "hello");
        assert!(!res.timed_out && !res.cancelled);
    }

    #[test]
    fn timeout_kills_silent_child_promptly() {
        let res = run("sleep 30", 400, 1024);
        assert!(res.timed_out);
        assert_eq!(res.exit_code, None);
        // keep-with-deadline: time is the contract; wide margin over the 400ms timeout.
        assert!(res.duration_ms < 2_000, "took {}ms", res.duration_ms);
    }

    #[test]
    fn stderr_heavy_child_cannot_deadlock_the_deadline() {
        let res = run("while true; do echo noise >&2; done", 400, 4096);
        assert!(res.timed_out);
        // keep-with-deadline: time is the contract; wide margin over the 400ms timeout.
        assert!(res.duration_ms < 2_000, "took {}ms", res.duration_ms);
    }

    /// Run `script` (built around the signal's raise command) with a 30 s
    /// timeout, cancel it once it raises the signal, and return its result
    /// within the signal deadline together with the raising connection.
    fn cancel_once_raised(script: impl FnOnce(&str) -> String) -> (BashResult, TcpStream) {
        let cancel = CancelHandle::new();
        let signal = ProcessSignal::new();
        let script = script(&signal.raise());
        let (done, finished) = mpsc::channel();
        {
            let cancel = cancel.clone();
            thread::spawn(move || {
                let _ = done.send(execute(&script, None, 30_000, 1024, &cancel).unwrap());
            });
        }
        // synchronize: cancel once the command reports it runs, never after a guessed delay.
        let socket = signal.wait();
        cancel.cancel();
        let res = finished
            .recv_timeout(Duration::from_secs(10))
            .expect("cancelled command returned within the deadline");
        (res, socket)
    }

    #[test]
    fn cancellation_stops_execution_with_pipes_open_or_closed() {
        for redirect in ["", "exec >&- 2>&-; "] {
            let (res, mut socket) =
                cancel_once_raised(|raise| format!("{redirect}{raise}; sleep 30"));
            assert!(res.cancelled && !res.timed_out, "{redirect:?}");
            assert!(
                closed(&mut socket),
                "cancelled command kept running: {redirect:?}"
            );
        }
    }

    /// User plus system CPU clock ticks spent by the calling thread.
    #[cfg(target_os = "linux")]
    fn thread_cpu_ticks() -> u64 {
        let stat = std::fs::read_to_string("/proc/thread-self/stat").unwrap();
        // Fields after the parenthesised command name; utime and stime are
        // the 14th and 15th fields of the whole line.
        let fields: Vec<&str> = stat
            .rsplit_once(')')
            .unwrap()
            .1
            .split_whitespace()
            .collect();
        fields[11].parse::<u64>().unwrap() + fields[12].parse::<u64>().unwrap()
    }

    #[cfg(target_os = "linux")]
    #[test]
    fn waiting_on_a_command_that_closed_its_pipes_does_not_spin() {
        let before = thread_cpu_ticks();
        let res = run("exec >&- 2>&-; sleep 30", 500, 1024);
        let spent = thread_cpu_ticks() - before;
        assert!(res.timed_out);
        // A spinning wait burns the whole 500 ms (≈50 ticks at 100 Hz).
        assert!(spent < 10, "waiting thread burned {spent} CPU ticks");
    }

    #[test]
    fn output_written_after_exit_is_drained_within_the_grace_window() {
        // The background writer waits until the exited shell is reaped, so
        // its output can only arrive after the main loop has ended.
        let res = run(
            "(while kill -0 $$ 2>/dev/null; do :; done; printf late) & exit 0",
            5_000,
            1024,
        );
        assert_eq!(res.exit_code, Some(0));
        assert_eq!(res.stdout, "late");
    }

    #[test]
    fn reader_forwards_whole_chunks_up_to_the_cap_and_nothing_beyond() {
        const CHUNK: usize = 16 * 1024;
        let (tx, rx) = mpsc::channel();
        let remaining = Arc::new(AtomicUsize::new(CHUNK + 1));
        spawn_reader(
            std::io::Cursor::new(vec![b'x'; 3 * CHUNK]),
            tx,
            Stream::Stdout,
            Arc::clone(&remaining),
        )
        .join()
        .unwrap();
        let forwarded: Vec<usize> = rx.iter().map(|(_, chunk)| chunk.len()).collect();
        assert_eq!(forwarded, vec![CHUNK, 1]);
        assert_eq!(remaining.load(Ordering::Relaxed), 0);
    }

    #[test]
    fn output_cap_keeps_fitting_prefix() {
        let res = run("printf 'abcdefghij'", 5_000, 5);
        assert_eq!(res.stdout, "abcde");
    }

    #[test]
    fn high_volume_output_is_not_throttled() {
        // keep-with-deadline: the command's own timeout is the throughput bound.
        let res = run("head -c 300000 /dev/zero | tr '\\0' 'x'", 3_000, 1 << 20);
        assert!(
            !res.timed_out,
            "throughput regression: {}ms",
            res.duration_ms
        );
        assert_eq!(res.stdout.len(), 300_000);
    }

    #[test]
    fn descendant_processes_are_killed_with_the_tree() {
        // Only the background descendant holds the connection.
        let (res, mut socket) =
            cancel_once_raised(|raise| format!("({raise}; sleep 30) & sleep 30"));
        assert!(res.cancelled);
        assert!(closed(&mut socket), "descendant survived the tree kill");
    }
}