1use std::collections::BTreeSet;
39use std::io::{self, Read, Write};
40use std::path::PathBuf;
41use std::process::{Child, Command, ExitStatus, Stdio};
42use std::time::{Duration, Instant};
43
44#[derive(Debug, Clone, PartialEq, Eq)]
46pub enum EnvPolicy {
47 InheritExcept { deny: Vec<String> },
56 ClearExcept { allow: Vec<String> },
62}
63
64impl EnvPolicy {
65 pub fn minimal_allow() -> Vec<String> {
69 let base: &[&str] = if cfg!(windows) {
70 &["PATH", "PATHEXT", "SYSTEMROOT", "SYSTEMDRIVE", "COMSPEC", "TEMP", "TMP", "USERPROFILE"]
71 } else {
72 &["PATH", "HOME", "TMPDIR", "LANG", "LC_ALL", "TZ"]
73 };
74 base.iter().map(|s| s.to_string()).collect()
75 }
76}
77
78impl Default for EnvPolicy {
79 fn default() -> Self {
80 EnvPolicy::InheritExcept { deny: secret_env_vars() }
81 }
82}
83
84static SECRET_ENV: std::sync::Mutex<Option<BTreeSet<String>>> = std::sync::Mutex::new(None);
92
93pub fn deny_env_var(name: &str) {
99 if name.trim().is_empty() {
100 return;
101 }
102 let mut guard = SECRET_ENV.lock().unwrap_or_else(|e| e.into_inner());
103 guard.get_or_insert_with(BTreeSet::new).insert(name.to_string());
104}
105
106pub fn secret_env_vars() -> Vec<String> {
108 let guard = SECRET_ENV.lock().unwrap_or_else(|e| e.into_inner());
109 guard.as_ref().map(|s| s.iter().cloned().collect()).unwrap_or_default()
110}
111
112#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
114pub enum StderrMode {
115 #[default]
118 Pipe,
119 Inherit,
122}
123
124#[derive(Debug, Clone)]
126pub struct SpawnPolicy {
127 pub timeout: Option<Duration>,
131 pub max_output_bytes: usize,
135 pub env: EnvPolicy,
136 pub stderr: StderrMode,
137 pub current_dir: Option<PathBuf>,
139}
140
141pub const DEFAULT_TIMEOUT: Duration = Duration::from_secs(300);
145
146pub const DEFAULT_MAX_OUTPUT: usize = 64 * 1024 * 1024;
150
151impl Default for SpawnPolicy {
152 fn default() -> Self {
153 SpawnPolicy {
154 timeout: Some(DEFAULT_TIMEOUT),
155 max_output_bytes: DEFAULT_MAX_OUTPUT,
156 env: EnvPolicy::default(),
157 stderr: StderrMode::default(),
158 current_dir: None,
159 }
160 }
161}
162
163impl SpawnPolicy {
164 pub fn deny_vars<I: IntoIterator<Item = String>>(mut self, vars: I) -> Self {
168 let mut denied: BTreeSet<String> = match self.env {
174 EnvPolicy::InheritExcept { deny } => deny.into_iter().collect(),
175 EnvPolicy::ClearExcept { allow } => {
176 self.env = EnvPolicy::ClearExcept { allow };
177 return self;
178 }
179 };
180 denied.extend(vars);
181 self.env = EnvPolicy::InheritExcept { deny: denied.into_iter().collect() };
182 self
183 }
184
185 pub fn timeout(mut self, timeout: Option<Duration>) -> Self {
186 self.timeout = timeout;
187 self
188 }
189
190 pub fn stderr(mut self, mode: StderrMode) -> Self {
191 self.stderr = mode;
192 self
193 }
194}
195
196#[derive(Debug)]
198pub struct SpawnOutput {
199 pub status: ExitStatus,
200 pub stdout: Vec<u8>,
201 pub stderr: Vec<u8>,
202 pub timed_out: bool,
205 pub stdout_truncated: bool,
206 pub stderr_truncated: bool,
207}
208
209impl SpawnOutput {
210 pub fn stderr_text(&self) -> String {
212 String::from_utf8_lossy(&self.stderr).trim().to_string()
213 }
214
215 pub fn failure(&self, what: &str) -> Option<String> {
218 if self.timed_out {
219 return Some(format!("{what} timed out and was killed"));
220 }
221 if !self.status.success() {
222 let err = self.stderr_text();
223 return Some(if err.is_empty() {
224 format!("{what} exited with {}", self.status)
225 } else {
226 format!("{what} exited with {}: {err}", self.status)
227 });
228 }
229 None
230 }
231}
232
233pub fn run(
243 mut cmd: Command,
244 stdin: Option<&[u8]>,
245 extra_env: &[(&str, &str)],
246 policy: &SpawnPolicy,
247) -> io::Result<SpawnOutput> {
248 match &policy.env {
249 EnvPolicy::InheritExcept { deny } => {
250 for var in deny {
251 cmd.env_remove(var);
252 }
253 }
254 EnvPolicy::ClearExcept { allow } => {
255 cmd.env_clear();
256 for var in allow {
257 if let Ok(val) = std::env::var(var) {
258 cmd.env(var, val);
259 }
260 }
261 }
262 }
263 for (k, v) in extra_env {
264 cmd.env(k, v);
265 }
266 if let Some(dir) = &policy.current_dir {
267 cmd.current_dir(dir);
268 }
269
270 cmd.stdin(if stdin.is_some() { Stdio::piped() } else { Stdio::null() })
271 .stdout(Stdio::piped())
272 .stderr(match policy.stderr {
273 StderrMode::Pipe => Stdio::piped(),
274 StderrMode::Inherit => Stdio::inherit(),
275 });
276
277 let mut child = cmd.spawn()?;
278
279 let stdin_thread = child.stdin.take().map(|mut pipe| {
283 let payload = stdin.unwrap_or_default().to_vec();
284 std::thread::spawn(move || {
285 let _ = pipe.write_all(&payload);
286 })
288 });
289
290 let cap = policy.max_output_bytes;
291 let out_thread = child.stdout.take().map(|pipe| std::thread::spawn(move || drain(pipe, cap)));
292 let err_thread = child.stderr.take().map(|pipe| std::thread::spawn(move || drain(pipe, cap)));
293
294 let (status, timed_out) = wait_bounded(&mut child, policy.timeout)?;
295
296 if let Some(t) = stdin_thread {
297 let _ = t.join();
298 }
299 let (stdout, stdout_truncated) = out_thread.and_then(|t| t.join().ok()).unwrap_or((Vec::new(), false));
300 let (stderr, stderr_truncated) = err_thread.and_then(|t| t.join().ok()).unwrap_or((Vec::new(), false));
301
302 Ok(SpawnOutput { status, stdout, stderr, timed_out, stdout_truncated, stderr_truncated })
303}
304
305fn drain<R: Read>(mut src: R, cap: usize) -> (Vec<u8>, bool) {
309 let mut kept = Vec::new();
310 let mut buf = [0u8; 16 * 1024];
311 let mut truncated = false;
312 loop {
313 match src.read(&mut buf) {
314 Ok(0) => break,
315 Ok(n) => {
316 if kept.len() < cap {
317 let room = cap - kept.len();
318 let take = room.min(n);
319 kept.extend_from_slice(&buf[..take]);
320 if take < n {
321 truncated = true;
322 }
323 } else {
324 truncated = true;
325 }
326 }
327 Err(ref e) if e.kind() == io::ErrorKind::Interrupted => continue,
328 Err(_) => break,
329 }
330 }
331 (kept, truncated)
332}
333
334fn wait_bounded(child: &mut Child, timeout: Option<Duration>) -> io::Result<(ExitStatus, bool)> {
340 let Some(limit) = timeout else {
341 return Ok((child.wait()?, false));
342 };
343 let deadline = Instant::now() + limit;
344 let mut nap = Duration::from_millis(1);
345 loop {
346 if let Some(status) = child.try_wait()? {
347 return Ok((status, false));
348 }
349 if Instant::now() >= deadline {
350 let _ = child.kill();
351 let status = child.wait()?;
353 return Ok((status, true));
354 }
355 std::thread::sleep(nap);
356 nap = (nap * 2).min(Duration::from_millis(50));
357 }
358}
359
360#[cfg(test)]
361mod tests {
362 use super::*;
363
364 fn sh(script: &str) -> Command {
365 let mut c = Command::new("/bin/sh");
366 c.arg("-c").arg(script);
367 c
368 }
369
370 #[test]
371 #[cfg_attr(windows, ignore = "uses /bin/sh")]
372 fn captures_stdout_and_exit_status() {
373 let out = run(sh("printf hello"), None, &[], &SpawnPolicy::default()).unwrap();
374 assert!(out.status.success());
375 assert_eq!(out.stdout, b"hello");
376 assert!(!out.timed_out);
377 assert!(!out.stdout_truncated);
378 }
379
380 #[test]
381 #[cfg_attr(windows, ignore = "uses /bin/sh")]
382 fn stdin_reaches_the_child() {
383 let out = run(sh("cat"), Some(b"payload"), &[], &SpawnPolicy::default()).unwrap();
384 assert_eq!(out.stdout, b"payload");
385 }
386
387 #[test]
388 #[cfg_attr(windows, ignore = "uses /bin/sh")]
389 fn timeout_kills_a_hung_child() {
390 let policy = SpawnPolicy::default().timeout(Some(Duration::from_millis(150)));
391 let out = run(sh("sleep 30"), None, &[], &policy).unwrap();
392 assert!(out.timed_out, "expected the child to be killed");
393 assert!(!out.status.success());
394 assert!(out.failure("tool").unwrap().contains("timed out"));
395 }
396
397 #[test]
398 #[cfg_attr(windows, ignore = "uses /bin/sh")]
399 fn output_cap_truncates_without_hanging() {
400 let policy = SpawnPolicy { max_output_bytes: 1024, ..SpawnPolicy::default() };
402 let out = run(sh("head -c 200000 /dev/zero"), None, &[], &policy).unwrap();
403 assert_eq!(out.stdout.len(), 1024);
404 assert!(out.stdout_truncated);
405 assert!(!out.timed_out, "draining past the cap must not stall the child");
406 }
407
408 #[test]
409 #[cfg_attr(windows, ignore = "uses /bin/sh")]
410 fn large_stdin_does_not_deadlock() {
411 let big = vec![b'x'; 4 * 1024 * 1024];
414 let policy = SpawnPolicy::default().timeout(Some(Duration::from_secs(20)));
415 let out = run(sh("cat"), Some(&big), &[], &policy).unwrap();
416 assert!(!out.timed_out, "write-then-read deadlock");
417 assert_eq!(out.stdout.len(), big.len());
418 }
419
420 #[test]
421 #[cfg_attr(windows, ignore = "uses /bin/sh")]
422 fn denied_vars_do_not_reach_the_child() {
423 std::env::set_var("AREEV_TEST_SECRET", "hunter2");
424 let policy = SpawnPolicy::default().deny_vars(["AREEV_TEST_SECRET".to_string()]);
425 let out = run(sh("printf %s \"${AREEV_TEST_SECRET:-absent}\""), None, &[], &policy).unwrap();
426 std::env::remove_var("AREEV_TEST_SECRET");
427 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent");
428 }
429
430 #[test]
431 #[cfg_attr(windows, ignore = "uses /bin/sh")]
432 fn inherited_vars_still_reach_the_child() {
433 std::env::set_var("AREEV_TEST_KEEP", "kept");
436 std::env::set_var("AREEV_TEST_DROP", "dropped");
437 let policy = SpawnPolicy::default().deny_vars(["AREEV_TEST_DROP".to_string()]);
438 let out = run(sh("printf %s \"${AREEV_TEST_KEEP:-absent}\""), None, &[], &policy).unwrap();
439 std::env::remove_var("AREEV_TEST_KEEP");
440 std::env::remove_var("AREEV_TEST_DROP");
441 assert_eq!(String::from_utf8_lossy(&out.stdout), "kept");
442 }
443
444 #[test]
445 #[cfg_attr(windows, ignore = "uses /bin/sh")]
446 fn clear_except_drops_everything_unlisted_but_keeps_extras() {
447 std::env::set_var("AREEV_TEST_AMBIENT", "ambient");
448 let policy = SpawnPolicy {
449 env: EnvPolicy::ClearExcept { allow: EnvPolicy::minimal_allow() },
450 ..SpawnPolicy::default()
451 };
452 let out = run(
453 sh("printf %s \"${AREEV_TEST_AMBIENT:-absent}/${AREEV_EXTRA:-none}\""),
454 None,
455 &[("AREEV_EXTRA", "set")],
456 &policy,
457 )
458 .unwrap();
459 std::env::remove_var("AREEV_TEST_AMBIENT");
460 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent/set");
461 }
462
463 #[test]
464 #[cfg_attr(windows, ignore = "uses /bin/sh")]
465 fn nonzero_exit_reports_stderr() {
466 let out = run(sh("echo boom >&2; exit 3"), None, &[], &SpawnPolicy::default()).unwrap();
467 let msg = out.failure("embed command").unwrap();
468 assert!(msg.contains("embed command"), "{msg}");
469 assert!(msg.contains("boom"), "{msg}");
470 }
471
472 #[test]
473 #[cfg_attr(windows, ignore = "uses /bin/sh")]
474 fn registered_secrets_are_scrubbed_without_the_seam_asking() {
475 std::env::set_var("AREEV_TEST_REGISTERED", "hunter2");
478 deny_env_var("AREEV_TEST_REGISTERED");
479 let out = run(
480 sh("printf %s \"${AREEV_TEST_REGISTERED:-absent}\""),
481 None,
482 &[],
483 &SpawnPolicy::default(),
484 )
485 .unwrap();
486 std::env::remove_var("AREEV_TEST_REGISTERED");
487 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent");
488 }
489
490 #[test]
491 fn deny_env_var_ignores_blanks() {
492 deny_env_var(" ");
493 deny_env_var("");
494 assert!(!secret_env_vars().iter().any(|v| v.trim().is_empty()));
495 }
496
497 #[test]
498 fn spawn_failure_is_an_error_not_a_panic() {
499 let cmd = Command::new("areev-no-such-binary-eaf1");
500 assert!(run(cmd, None, &[], &SpawnPolicy::default()).is_err());
501 }
502}