use std::io::{self, Read};
use std::process::{Command, Output, Stdio};
use std::time::{Duration, Instant};
pub const COMMAND_OUTPUT_CAP_BYTES: usize = 16 * 1024 * 1024;
pub const OUTPUT_TRUNCATED_MARKER: &[u8] = b"\n[truncated by all-smi: output exceeded cap]\n";
fn read_capped<R: Read>(mut reader: R, buf: &mut Vec<u8>, cap: usize) {
const CHUNK: usize = 8 * 1024;
let mut chunk = [0u8; CHUNK];
let mut truncated = false;
loop {
match reader.read(&mut chunk) {
Ok(0) => break, Ok(n) => {
if truncated {
continue;
}
let remaining = cap.saturating_sub(buf.len());
if remaining == 0 {
truncated = true;
continue;
}
let take = n.min(remaining);
buf.extend_from_slice(&chunk[..take]);
if buf.len() >= cap {
truncated = true;
}
}
Err(_) => break,
}
}
if truncated {
buf.extend_from_slice(OUTPUT_TRUNCATED_MARKER);
}
}
pub fn run_command_with_timeout(
command: &str,
args: &[&str],
timeout: Duration,
) -> io::Result<Output> {
let mut child = Command::new(command)
.args(args)
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()?;
let start = Instant::now();
let poll_interval = Duration::from_millis(10);
loop {
match child.try_wait() {
Ok(Some(status)) => {
let mut stdout = Vec::new();
let mut stderr = Vec::new();
if let Some(out) = child.stdout.take() {
read_capped(out, &mut stdout, COMMAND_OUTPUT_CAP_BYTES);
}
if let Some(err) = child.stderr.take() {
read_capped(err, &mut stderr, COMMAND_OUTPUT_CAP_BYTES);
}
return Ok(Output {
status,
stdout,
stderr,
});
}
Ok(None) => {
if start.elapsed() >= timeout {
let _ = child.kill();
let _ = child.wait(); return Err(io::Error::new(
io::ErrorKind::TimedOut,
format!("Command '{command}' timed out after {timeout:?}"),
));
}
std::thread::sleep(poll_interval);
}
Err(e) => {
let _ = child.kill();
let _ = child.wait();
return Err(e);
}
}
}
}
pub fn run_command_fast_fail(command: &str, args: &[&str]) -> io::Result<Output> {
let timeout = if is_container_environment() {
Duration::from_millis(500) } else {
Duration::from_secs(2) };
run_command_with_timeout(command, args, timeout)
}
fn is_container_environment() -> bool {
std::path::Path::new("/.dockerenv").exists()
|| std::path::Path::new("/run/.containerenv").exists()
|| std::env::var("KUBERNETES_SERVICE_HOST").is_ok()
|| std::env::var("CONTAINER_RUNTIME").is_ok()
|| check_cgroup_container()
}
fn check_cgroup_container() -> bool {
if let Ok(contents) = std::fs::read_to_string("/proc/self/cgroup") {
contents.contains("/docker/")
|| contents.contains("/lxc/")
|| contents.contains("/kubepods/")
|| contents.contains("/containerd/")
} else {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[cfg(unix)]
fn output_cap_truncates_hostile_stdout() {
let output = run_command_with_timeout("yes", &["hello"], Duration::from_millis(150));
let output = match output {
Ok(o) => o,
Err(e) if e.kind() == io::ErrorKind::NotFound => return,
Err(e) if e.kind() == io::ErrorKind::TimedOut => return,
Err(e) => panic!("unexpected error: {e}"),
};
let max_allowed = COMMAND_OUTPUT_CAP_BYTES + OUTPUT_TRUNCATED_MARKER.len() + 4096;
assert!(
output.stdout.len() <= max_allowed,
"stdout uncapped: len={}",
output.stdout.len()
);
}
#[test]
fn output_cap_leaves_small_outputs_alone() {
#[cfg(unix)]
{
let out = run_command_with_timeout("printf", &["hello"], Duration::from_secs(2));
if let Ok(out) = out {
assert_eq!(out.stdout, b"hello");
let marker = String::from_utf8_lossy(OUTPUT_TRUNCATED_MARKER);
assert!(
!String::from_utf8_lossy(&out.stdout).contains(marker.as_ref()),
"small output should not carry the truncation marker"
);
}
}
}
}