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 {
169 EnvPolicy::InheritExcept { deny } => deny.into_iter().collect(),
170 EnvPolicy::ClearExcept { .. } => BTreeSet::new(),
171 };
172 denied.extend(vars);
173 self.env = EnvPolicy::InheritExcept { deny: denied.into_iter().collect() };
174 self
175 }
176
177 pub fn timeout(mut self, timeout: Option<Duration>) -> Self {
178 self.timeout = timeout;
179 self
180 }
181
182 pub fn stderr(mut self, mode: StderrMode) -> Self {
183 self.stderr = mode;
184 self
185 }
186}
187
188#[derive(Debug)]
190pub struct SpawnOutput {
191 pub status: ExitStatus,
192 pub stdout: Vec<u8>,
193 pub stderr: Vec<u8>,
194 pub timed_out: bool,
197 pub stdout_truncated: bool,
198 pub stderr_truncated: bool,
199}
200
201impl SpawnOutput {
202 pub fn stderr_text(&self) -> String {
204 String::from_utf8_lossy(&self.stderr).trim().to_string()
205 }
206
207 pub fn failure(&self, what: &str) -> Option<String> {
210 if self.timed_out {
211 return Some(format!("{what} timed out and was killed"));
212 }
213 if !self.status.success() {
214 let err = self.stderr_text();
215 return Some(if err.is_empty() {
216 format!("{what} exited with {}", self.status)
217 } else {
218 format!("{what} exited with {}: {err}", self.status)
219 });
220 }
221 None
222 }
223}
224
225pub fn run(
235 mut cmd: Command,
236 stdin: Option<&[u8]>,
237 extra_env: &[(&str, &str)],
238 policy: &SpawnPolicy,
239) -> io::Result<SpawnOutput> {
240 match &policy.env {
241 EnvPolicy::InheritExcept { deny } => {
242 for var in deny {
243 cmd.env_remove(var);
244 }
245 }
246 EnvPolicy::ClearExcept { allow } => {
247 cmd.env_clear();
248 for var in allow {
249 if let Ok(val) = std::env::var(var) {
250 cmd.env(var, val);
251 }
252 }
253 }
254 }
255 for (k, v) in extra_env {
256 cmd.env(k, v);
257 }
258 if let Some(dir) = &policy.current_dir {
259 cmd.current_dir(dir);
260 }
261
262 cmd.stdin(if stdin.is_some() { Stdio::piped() } else { Stdio::null() })
263 .stdout(Stdio::piped())
264 .stderr(match policy.stderr {
265 StderrMode::Pipe => Stdio::piped(),
266 StderrMode::Inherit => Stdio::inherit(),
267 });
268
269 let mut child = cmd.spawn()?;
270
271 let stdin_thread = child.stdin.take().map(|mut pipe| {
275 let payload = stdin.unwrap_or_default().to_vec();
276 std::thread::spawn(move || {
277 let _ = pipe.write_all(&payload);
278 })
280 });
281
282 let cap = policy.max_output_bytes;
283 let out_thread = child.stdout.take().map(|pipe| std::thread::spawn(move || drain(pipe, cap)));
284 let err_thread = child.stderr.take().map(|pipe| std::thread::spawn(move || drain(pipe, cap)));
285
286 let (status, timed_out) = wait_bounded(&mut child, policy.timeout)?;
287
288 if let Some(t) = stdin_thread {
289 let _ = t.join();
290 }
291 let (stdout, stdout_truncated) = out_thread.and_then(|t| t.join().ok()).unwrap_or((Vec::new(), false));
292 let (stderr, stderr_truncated) = err_thread.and_then(|t| t.join().ok()).unwrap_or((Vec::new(), false));
293
294 Ok(SpawnOutput { status, stdout, stderr, timed_out, stdout_truncated, stderr_truncated })
295}
296
297fn drain<R: Read>(mut src: R, cap: usize) -> (Vec<u8>, bool) {
301 let mut kept = Vec::new();
302 let mut buf = [0u8; 16 * 1024];
303 let mut truncated = false;
304 loop {
305 match src.read(&mut buf) {
306 Ok(0) => break,
307 Ok(n) => {
308 if kept.len() < cap {
309 let room = cap - kept.len();
310 let take = room.min(n);
311 kept.extend_from_slice(&buf[..take]);
312 if take < n {
313 truncated = true;
314 }
315 } else {
316 truncated = true;
317 }
318 }
319 Err(ref e) if e.kind() == io::ErrorKind::Interrupted => continue,
320 Err(_) => break,
321 }
322 }
323 (kept, truncated)
324}
325
326fn wait_bounded(child: &mut Child, timeout: Option<Duration>) -> io::Result<(ExitStatus, bool)> {
332 let Some(limit) = timeout else {
333 return Ok((child.wait()?, false));
334 };
335 let deadline = Instant::now() + limit;
336 let mut nap = Duration::from_millis(1);
337 loop {
338 if let Some(status) = child.try_wait()? {
339 return Ok((status, false));
340 }
341 if Instant::now() >= deadline {
342 let _ = child.kill();
343 let status = child.wait()?;
345 return Ok((status, true));
346 }
347 std::thread::sleep(nap);
348 nap = (nap * 2).min(Duration::from_millis(50));
349 }
350}
351
352#[cfg(test)]
353mod tests {
354 use super::*;
355
356 fn sh(script: &str) -> Command {
357 let mut c = Command::new("/bin/sh");
358 c.arg("-c").arg(script);
359 c
360 }
361
362 #[test]
363 #[cfg_attr(windows, ignore = "uses /bin/sh")]
364 fn captures_stdout_and_exit_status() {
365 let out = run(sh("printf hello"), None, &[], &SpawnPolicy::default()).unwrap();
366 assert!(out.status.success());
367 assert_eq!(out.stdout, b"hello");
368 assert!(!out.timed_out);
369 assert!(!out.stdout_truncated);
370 }
371
372 #[test]
373 #[cfg_attr(windows, ignore = "uses /bin/sh")]
374 fn stdin_reaches_the_child() {
375 let out = run(sh("cat"), Some(b"payload"), &[], &SpawnPolicy::default()).unwrap();
376 assert_eq!(out.stdout, b"payload");
377 }
378
379 #[test]
380 #[cfg_attr(windows, ignore = "uses /bin/sh")]
381 fn timeout_kills_a_hung_child() {
382 let policy = SpawnPolicy::default().timeout(Some(Duration::from_millis(150)));
383 let out = run(sh("sleep 30"), None, &[], &policy).unwrap();
384 assert!(out.timed_out, "expected the child to be killed");
385 assert!(!out.status.success());
386 assert!(out.failure("tool").unwrap().contains("timed out"));
387 }
388
389 #[test]
390 #[cfg_attr(windows, ignore = "uses /bin/sh")]
391 fn output_cap_truncates_without_hanging() {
392 let policy = SpawnPolicy { max_output_bytes: 1024, ..SpawnPolicy::default() };
394 let out = run(sh("head -c 200000 /dev/zero"), None, &[], &policy).unwrap();
395 assert_eq!(out.stdout.len(), 1024);
396 assert!(out.stdout_truncated);
397 assert!(!out.timed_out, "draining past the cap must not stall the child");
398 }
399
400 #[test]
401 #[cfg_attr(windows, ignore = "uses /bin/sh")]
402 fn large_stdin_does_not_deadlock() {
403 let big = vec![b'x'; 4 * 1024 * 1024];
406 let policy = SpawnPolicy::default().timeout(Some(Duration::from_secs(20)));
407 let out = run(sh("cat"), Some(&big), &[], &policy).unwrap();
408 assert!(!out.timed_out, "write-then-read deadlock");
409 assert_eq!(out.stdout.len(), big.len());
410 }
411
412 #[test]
413 #[cfg_attr(windows, ignore = "uses /bin/sh")]
414 fn denied_vars_do_not_reach_the_child() {
415 std::env::set_var("AREEV_TEST_SECRET", "hunter2");
416 let policy = SpawnPolicy::default().deny_vars(["AREEV_TEST_SECRET".to_string()]);
417 let out = run(sh("printf %s \"${AREEV_TEST_SECRET:-absent}\""), None, &[], &policy).unwrap();
418 std::env::remove_var("AREEV_TEST_SECRET");
419 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent");
420 }
421
422 #[test]
423 #[cfg_attr(windows, ignore = "uses /bin/sh")]
424 fn inherited_vars_still_reach_the_child() {
425 std::env::set_var("AREEV_TEST_KEEP", "kept");
428 std::env::set_var("AREEV_TEST_DROP", "dropped");
429 let policy = SpawnPolicy::default().deny_vars(["AREEV_TEST_DROP".to_string()]);
430 let out = run(sh("printf %s \"${AREEV_TEST_KEEP:-absent}\""), None, &[], &policy).unwrap();
431 std::env::remove_var("AREEV_TEST_KEEP");
432 std::env::remove_var("AREEV_TEST_DROP");
433 assert_eq!(String::from_utf8_lossy(&out.stdout), "kept");
434 }
435
436 #[test]
437 #[cfg_attr(windows, ignore = "uses /bin/sh")]
438 fn clear_except_drops_everything_unlisted_but_keeps_extras() {
439 std::env::set_var("AREEV_TEST_AMBIENT", "ambient");
440 let policy = SpawnPolicy {
441 env: EnvPolicy::ClearExcept { allow: EnvPolicy::minimal_allow() },
442 ..SpawnPolicy::default()
443 };
444 let out = run(
445 sh("printf %s \"${AREEV_TEST_AMBIENT:-absent}/${AREEV_EXTRA:-none}\""),
446 None,
447 &[("AREEV_EXTRA", "set")],
448 &policy,
449 )
450 .unwrap();
451 std::env::remove_var("AREEV_TEST_AMBIENT");
452 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent/set");
453 }
454
455 #[test]
456 #[cfg_attr(windows, ignore = "uses /bin/sh")]
457 fn nonzero_exit_reports_stderr() {
458 let out = run(sh("echo boom >&2; exit 3"), None, &[], &SpawnPolicy::default()).unwrap();
459 let msg = out.failure("embed command").unwrap();
460 assert!(msg.contains("embed command"), "{msg}");
461 assert!(msg.contains("boom"), "{msg}");
462 }
463
464 #[test]
465 #[cfg_attr(windows, ignore = "uses /bin/sh")]
466 fn registered_secrets_are_scrubbed_without_the_seam_asking() {
467 std::env::set_var("AREEV_TEST_REGISTERED", "hunter2");
470 deny_env_var("AREEV_TEST_REGISTERED");
471 let out = run(
472 sh("printf %s \"${AREEV_TEST_REGISTERED:-absent}\""),
473 None,
474 &[],
475 &SpawnPolicy::default(),
476 )
477 .unwrap();
478 std::env::remove_var("AREEV_TEST_REGISTERED");
479 assert_eq!(String::from_utf8_lossy(&out.stdout), "absent");
480 }
481
482 #[test]
483 fn deny_env_var_ignores_blanks() {
484 deny_env_var(" ");
485 deny_env_var("");
486 assert!(!secret_env_vars().iter().any(|v| v.trim().is_empty()));
487 }
488
489 #[test]
490 fn spawn_failure_is_an_error_not_a_panic() {
491 let cmd = Command::new("areev-no-such-binary-eaf1");
492 assert!(run(cmd, None, &[], &SpawnPolicy::default()).is_err());
493 }
494}