1use crate::handle::{
2 ChildTerminator, OUTPUT_CHANNEL_CAPACITY, OutputChunk, ProcessHandle, STDIN_CHANNEL_CAPACITY,
3 SpawnedProcess, Stream,
4};
5use crate::process_group;
6use crate::{HeadTailBuffer, OutputReceiver, lock_or_recover};
7use anyhow::{Context, Result, anyhow};
8use std::collections::HashMap;
9#[cfg(unix)]
10use std::os::unix::process::ExitStatusExt;
11use std::path::PathBuf;
12use std::process::{ExitStatus, Stdio};
13use std::sync::{Arc, Mutex as StdMutex};
14use std::time::Duration;
15use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt};
16use tokio::process::{Child, ChildStdin, Command};
17use tokio::sync::{Notify, broadcast, mpsc};
18use tokio::time::Instant;
19
20const EXIT_OUTPUT_GRACE: Duration = Duration::from_millis(50);
21
22#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
23pub enum StdinMode {
24 #[default]
25 Piped,
26 Null,
27}
28
29#[derive(Clone, Debug, Default, Eq, PartialEq)]
30pub struct CommandOptions {
31 program: String,
32 args: Vec<String>,
33 cwd: PathBuf,
34 env: HashMap<String, String>,
35 stdin: StdinMode,
36}
37
38impl CommandOptions {
39 pub fn new(program: impl Into<String>, cwd: impl Into<PathBuf>) -> Self {
40 Self {
41 program: program.into(),
42 args: Vec::new(),
43 cwd: cwd.into(),
44 env: HashMap::new(),
45 stdin: StdinMode::Piped,
46 }
47 }
48
49 pub fn arg(mut self, arg: impl Into<String>) -> Self {
50 self.args.push(arg.into());
51 self
52 }
53
54 pub fn args<I, S>(mut self, args: I) -> Self
55 where
56 I: IntoIterator<Item = S>,
57 S: Into<String>,
58 {
59 self.args.extend(args.into_iter().map(Into::into));
60 self
61 }
62
63 pub fn env(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
64 self.env.insert(key.into(), value.into());
65 self
66 }
67
68 pub fn envs<I, K, V>(mut self, envs: I) -> Self
69 where
70 I: IntoIterator<Item = (K, V)>,
71 K: Into<String>,
72 V: Into<String>,
73 {
74 self.env.extend(envs.into_iter().map(|(key, value)| (key.into(), value.into())));
75 self
76 }
77
78 pub fn stdin(mut self, stdin: StdinMode) -> Self {
79 self.stdin = stdin;
80 self
81 }
82
83 pub fn no_stdin(self) -> Self {
84 self.stdin(StdinMode::Null)
85 }
86}
87
88#[derive(Clone, Debug, Eq, PartialEq)]
89pub struct RunOutput {
90 pub stdout: Vec<u8>,
91 pub stderr: Vec<u8>,
92 pub stdout_omitted_bytes: usize,
93 pub stderr_omitted_bytes: usize,
94 pub stdout_head_bytes: usize,
95 pub stderr_head_bytes: usize,
96 pub exit_code: Option<i32>,
97 pub timed_out: bool,
98 pub wall_time: Duration,
99}
100
101pub async fn run_with_timeout(
102 options: CommandOptions,
103 timeout: Duration,
104 max_output_bytes: usize,
105) -> Result<RunOutput> {
106 let started = Instant::now();
107 let (handle, rx) = spawn(options).await?;
108 let mut receiver = OutputReceiver::from(rx);
109 let mut stdout = HeadTailBuffer::new(max_output_bytes);
110 let mut stderr = HeadTailBuffer::new(max_output_bytes);
111 let deadline = Instant::now() + timeout;
112 let timeout_sleep = tokio::time::sleep_until(deadline);
113 tokio::pin!(timeout_sleep);
114 let mut timed_out = false;
115 let mut exit_seen = false;
116
117 loop {
118 drain_receiver(&mut receiver, &mut stdout, &mut stderr);
119
120 if exit_seen {
121 break;
122 }
123
124 if handle.has_exited() {
125 exit_seen = true;
126 tokio::time::sleep(EXIT_OUTPUT_GRACE).await;
127 continue;
128 }
129
130 tokio::select! {
131 chunk = receiver.recv() => {
132 if let Some(chunk) = chunk {
133 push_output_chunk(chunk, &mut stdout, &mut stderr);
134 }
135 }
136 _ = handle.wait_for_exit() => {
137 exit_seen = true;
138 tokio::time::sleep(EXIT_OUTPUT_GRACE).await;
139 }
140 _ = &mut timeout_sleep => {
141 timed_out = true;
142 handle.terminate();
143 tokio::time::sleep(EXIT_OUTPUT_GRACE).await;
144 break;
145 }
146 }
147 }
148
149 drain_receiver(&mut receiver, &mut stdout, &mut stderr);
150
151 Ok(RunOutput {
152 stdout: stdout.to_bytes(),
153 stderr: stderr.to_bytes(),
154 stdout_omitted_bytes: stdout.omitted_bytes(),
155 stderr_omitted_bytes: stderr.omitted_bytes(),
156 stdout_head_bytes: stdout.head_bytes(),
157 stderr_head_bytes: stderr.head_bytes(),
158 exit_code: if timed_out { None } else { handle.exit_code() },
159 timed_out,
160 wall_time: started.elapsed(),
161 })
162}
163
164pub async fn spawn(options: CommandOptions) -> Result<SpawnedProcess> {
165 let with_stdin = matches!(options.stdin, StdinMode::Piped);
166 let mut command = Command::new(&options.program);
167 command
168 .args(&options.args)
169 .current_dir(&options.cwd)
170 .envs(&options.env)
171 .stdout(Stdio::piped())
172 .stderr(Stdio::piped())
173 .kill_on_drop(false);
174
175 if with_stdin {
176 command.stdin(Stdio::piped());
177 } else {
178 command.stdin(Stdio::null());
179 }
180
181 #[cfg(unix)]
182 {
183 let parent_pid = unsafe { libc::getpid() };
184 unsafe {
185 command.pre_exec(move || {
186 process_group::detach_from_tty()?;
187 process_group::set_parent_death_signal(parent_pid)?;
188 Ok(())
189 });
190 }
191 }
192
193 let mut child = command
194 .spawn()
195 .with_context(|| format!("failed to spawn process `{}`", options.program))?;
196 let pid = child.id().ok_or_else(|| anyhow!("spawned process is missing a pid"))?;
197
198 let writer = if with_stdin { child.stdin.take().map(spawn_stdin_writer) } else { None };
199 let stdout = child.stdout.take().ok_or_else(|| anyhow!("failed to capture stdout"))?;
200 let stderr = child.stderr.take().ok_or_else(|| anyhow!("failed to capture stderr"))?;
201
202 let (output_tx, output_rx) = broadcast::channel(OUTPUT_CHANNEL_CAPACITY);
203 let exit_code = Arc::new(StdMutex::new(None));
204 let exit_notify = Arc::new(Notify::new());
205
206 spawn_reader(stdout, Stream::Stdout, output_tx.clone());
207 spawn_reader(stderr, Stream::Stderr, output_tx.clone());
208 spawn_exit_watcher(child, Arc::clone(&exit_code), Arc::clone(&exit_notify));
209
210 let handle = ProcessHandle::from_parts(
211 output_tx,
212 writer,
213 exit_code,
214 exit_notify,
215 Box::new(PidTerminator { pid }),
216 );
217
218 Ok((handle, output_rx))
219}
220
221pub fn terminate_child_process_group(child: &mut Child) {
223 if let Some(pid) = child.id() {
224 let _ = process_group::kill_by_pid(pid);
225 }
226
227 let _ = child.start_kill();
228}
229
230fn spawn_stdin_writer(mut stdin: ChildStdin) -> mpsc::Sender<Vec<u8>> {
231 let (tx, mut rx) = mpsc::channel::<Vec<u8>>(STDIN_CHANNEL_CAPACITY);
232 tokio::spawn(async move {
233 while let Some(data) = rx.recv().await {
234 if stdin.write_all(&data).await.is_err() {
235 return;
236 }
237 if stdin.flush().await.is_err() {
238 return;
239 }
240 }
241 });
242 tx
243}
244
245struct PidTerminator {
246 pid: u32,
247}
248
249impl ChildTerminator for PidTerminator {
250 fn terminate(&mut self) {
251 let _ = process_group::kill_by_pid(self.pid);
252 }
253}
254
255fn spawn_reader<R>(mut reader: R, stream: Stream, output_tx: broadcast::Sender<OutputChunk>)
256where
257 R: AsyncRead + Unpin + Send + 'static,
258{
259 tokio::spawn(async move {
260 let mut buffer = [0u8; 8192];
261 loop {
262 match reader.read(&mut buffer).await {
263 Ok(0) => return,
264 Ok(n) => {
265 let chunk = OutputChunk { stream, data: buffer[..n].to_vec() };
266 let _ = output_tx.send(chunk);
267 }
268 Err(_) => return,
269 }
270 }
271 });
272}
273
274fn spawn_exit_watcher(
275 mut child: Child,
276 exit_code: Arc<StdMutex<Option<i32>>>,
277 exit_notify: Arc<Notify>,
278) {
279 tokio::spawn(async move {
280 let code = match child.wait().await {
281 Ok(status) => Some(normalize_exit_code(status)),
282 Err(_) => Some(-1),
283 };
284
285 *lock_or_recover(&exit_code) = code;
286 exit_notify.notify_waiters();
287 });
288}
289
290fn normalize_exit_code(status: ExitStatus) -> i32 {
291 if let Some(code) = status.code() {
292 return code;
293 }
294
295 #[cfg(unix)]
296 if let Some(signal) = status.signal() {
297 return 128 + signal;
298 }
299
300 -1
301}
302
303fn drain_receiver(
304 receiver: &mut OutputReceiver,
305 stdout: &mut HeadTailBuffer,
306 stderr: &mut HeadTailBuffer,
307) {
308 receiver.drain_with(|chunk| push_output_chunk(chunk, stdout, stderr));
309}
310
311fn push_output_chunk(chunk: OutputChunk, stdout: &mut HeadTailBuffer, stderr: &mut HeadTailBuffer) {
312 let OutputChunk { stream, data } = chunk;
313 match stream {
314 Stream::Stdout => stdout.push_chunk(data),
315 Stream::Stderr => stderr.push_chunk(data),
316 }
317}