1use crate::handle::{OutputChunk, ProcessHandle, SpawnedProcess, Stream};
2use crate::{CommandOptions, HeadTailBuffer, lock_or_recover, subprocess};
3use std::collections::HashMap;
4use std::error::Error;
5use std::fmt;
6use std::sync::atomic::{AtomicU64, Ordering};
7use std::sync::{Arc, Mutex as StdMutex};
8use std::time::Duration;
9use tokio::sync::{Mutex, Notify, broadcast};
10use tokio::task::JoinHandle;
11use tokio::time::{Instant, sleep_until};
12
13const EXIT_DRAIN_GRACE: Duration = Duration::from_millis(50);
14const RECENT_PROTECTION_COUNT: usize = 8;
15
16#[derive(Clone, Debug, Eq, PartialEq)]
17pub struct PoolConfig {
18 pub max_processes: usize,
19 pub max_output_bytes: usize,
20 pub default_yield_ms: u64,
21 pub max_yield_ms: u64,
22 pub background_timeout_ms: u64,
23}
24
25impl Default for PoolConfig {
26 fn default() -> Self {
27 Self {
28 max_processes: 64,
29 max_output_bytes: 1024 * 1024,
30 default_yield_ms: 250,
31 max_yield_ms: 30_000,
32 background_timeout_ms: 300_000,
33 }
34 }
35}
36
37#[derive(Clone, Debug)]
38pub struct ExecRequest {
39 pub command: CommandOptions,
40 pub yield_time_ms: u64,
41 pub max_output_bytes: Option<usize>,
42}
43
44impl ExecRequest {
45 pub fn new(command: CommandOptions) -> Self {
46 Self { command, yield_time_ms: 0, max_output_bytes: None }
47 }
48
49 pub fn with_yield_time_ms(mut self, yield_time_ms: u64) -> Self {
50 self.yield_time_ms = yield_time_ms;
51 self
52 }
53
54 pub fn with_max_output_bytes(mut self, max_output_bytes: usize) -> Self {
55 self.max_output_bytes = Some(max_output_bytes);
56 self
57 }
58}
59
60#[derive(Clone, Debug)]
61pub struct PollRequest<'a> {
62 pub process_id: &'a str,
63 pub yield_time_ms: u64,
64 pub max_output_bytes: Option<usize>,
65}
66
67impl<'a> PollRequest<'a> {
68 pub fn new(process_id: &'a str) -> Self {
69 Self { process_id, yield_time_ms: 0, max_output_bytes: None }
70 }
71
72 pub fn with_yield_time_ms(mut self, yield_time_ms: u64) -> Self {
73 self.yield_time_ms = yield_time_ms;
74 self
75 }
76
77 pub fn with_max_output_bytes(mut self, max_output_bytes: usize) -> Self {
78 self.max_output_bytes = Some(max_output_bytes);
79 self
80 }
81}
82
83#[derive(Clone, Debug)]
84pub struct StdinRequest<'a> {
85 pub process_id: &'a str,
86 pub input: &'a [u8],
87 pub yield_time_ms: u64,
88 pub max_output_bytes: Option<usize>,
89}
90
91impl<'a> StdinRequest<'a> {
92 pub fn new(process_id: &'a str, input: &'a [u8]) -> Self {
93 Self { process_id, input, yield_time_ms: 0, max_output_bytes: None }
94 }
95
96 pub fn with_yield_time_ms(mut self, yield_time_ms: u64) -> Self {
97 self.yield_time_ms = yield_time_ms;
98 self
99 }
100
101 pub fn with_max_output_bytes(mut self, max_output_bytes: usize) -> Self {
102 self.max_output_bytes = Some(max_output_bytes);
103 self
104 }
105}
106
107#[derive(Clone, Debug, Eq, PartialEq)]
108pub struct ExecResponse {
109 pub output: Vec<u8>,
110 pub stderr: Vec<u8>,
111 pub process_id: Option<String>,
112 pub exit_code: Option<i32>,
113 pub wall_time: Duration,
114}
115
116#[derive(Debug, Clone, Eq, PartialEq)]
117pub enum ExecError {
118 SpawnFailed(String),
119 UnknownProcess { process_id: String },
120 StdinClosed,
121 PoolFull,
122}
123
124impl fmt::Display for ExecError {
125 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
126 match self {
127 Self::SpawnFailed(message) => write!(f, "failed to spawn process: {message}"),
128 Self::UnknownProcess { process_id } => write!(f, "unknown process: {process_id}"),
129 Self::StdinClosed => write!(f, "stdin is closed for this process"),
130 Self::PoolFull => write!(f, "process pool is full"),
131 }
132 }
133}
134
135impl Error for ExecError {}
136
137#[derive(Clone)]
138pub struct ProcessPool {
139 inner: Arc<PoolInner>,
140}
141
142struct PoolInner {
143 config: PoolConfig,
144 entries: Mutex<HashMap<String, Arc<ProcessEntry>>>,
145 spawn_gate: Mutex<()>,
146 next_process_id: AtomicU64,
147}
148
149struct ProcessEntry {
150 handle: ProcessHandle,
151 output: StdMutex<HeadTailBuffer>,
152 stderr: StdMutex<HeadTailBuffer>,
153 notify: Arc<Notify>,
154 last_used: StdMutex<Instant>,
155 interaction: Mutex<()>,
156 buffer_task: StdMutex<Option<JoinHandle<()>>>,
157}
158
159impl ProcessPool {
160 pub fn new(config: PoolConfig) -> Self {
161 let default_yield_ms = config.default_yield_ms.max(1);
162 let max_yield_ms = config.max_yield_ms.max(default_yield_ms);
163 let normalized = PoolConfig {
164 max_processes: config.max_processes.max(1),
165 max_output_bytes: config.max_output_bytes,
166 default_yield_ms,
167 max_yield_ms,
168 background_timeout_ms: config.background_timeout_ms.max(1),
169 };
170
171 Self {
172 inner: Arc::new(PoolInner {
173 config: normalized,
174 entries: Mutex::new(HashMap::new()),
175 spawn_gate: Mutex::new(()),
176 next_process_id: AtomicU64::new(1000),
177 }),
178 }
179 }
180
181 pub async fn exec(&self, request: ExecRequest) -> Result<ExecResponse, ExecError> {
182 let _spawn_guard = self.inner.spawn_gate.lock().await;
183 self.ensure_capacity().await?;
184
185 let process_id = self.next_process_id();
186 let entry = ProcessEntry::new(
187 subprocess::spawn(request.command)
188 .await
189 .map_err(|err| ExecError::SpawnFailed(err.to_string()))?,
190 self.inner.config.max_output_bytes,
191 );
192
193 self.inner.entries.lock().await.insert(process_id.clone(), Arc::clone(&entry));
194 drop(_spawn_guard);
195
196 self.interact(&entry, &process_id, None, request.yield_time_ms, request.max_output_bytes)
197 .await
198 }
199
200 pub async fn poll_output(&self, request: PollRequest<'_>) -> Result<ExecResponse, ExecError> {
201 self.prune_expired_entries().await;
202 let entry = self.entry(request.process_id).await?;
203 self.interact(
204 &entry,
205 request.process_id,
206 None,
207 request.yield_time_ms,
208 request.max_output_bytes,
209 )
210 .await
211 }
212
213 pub async fn write_stdin(&self, request: StdinRequest<'_>) -> Result<ExecResponse, ExecError> {
214 self.prune_expired_entries().await;
215 let entry = self.entry(request.process_id).await?;
216 self.interact(
217 &entry,
218 request.process_id,
219 Some(request.input),
220 request.yield_time_ms,
221 request.max_output_bytes,
222 )
223 .await
224 }
225
226 pub async fn kill(&self, process_id: &str) -> Result<(), ExecError> {
227 let removed = self.inner.entries.lock().await.remove(process_id);
228 let Some(entry) = removed else {
229 return Err(ExecError::UnknownProcess { process_id: process_id.to_string() });
230 };
231
232 shutdown_entry(entry, true);
233 Ok(())
234 }
235
236 pub async fn terminate_all(&self) {
237 let removed =
238 self.inner.entries.lock().await.drain().map(|(_, entry)| entry).collect::<Vec<_>>();
239 for entry in removed {
240 shutdown_entry(entry, true);
241 }
242 }
243
244 async fn interact(
245 &self,
246 entry: &Arc<ProcessEntry>,
247 process_id: &str,
248 input: Option<&[u8]>,
249 yield_time_ms: u64,
250 max_output_bytes: Option<usize>,
251 ) -> Result<ExecResponse, ExecError> {
252 let _interaction_guard = entry.interaction.lock().await;
253 entry.touch();
254
255 let stdin_closed = match input {
256 Some(input) => entry.handle.write_stdin(input).await.is_err(),
257 None => false,
258 };
259
260 if stdin_closed {
261 if entry.handle.has_exited() {
262 self.remove_if_same(process_id, entry, false).await;
263 }
264 return Err(ExecError::StdinClosed);
265 }
266
267 let response =
268 self.collect_locked(entry, process_id, yield_time_ms, max_output_bytes).await;
269
270 if response.process_id.is_none() {
271 self.remove_if_same(process_id, entry, false).await;
272 }
273
274 Ok(response)
275 }
276
277 async fn collect_locked(
278 &self,
279 entry: &Arc<ProcessEntry>,
280 process_id: &str,
281 yield_time_ms: u64,
282 max_output_bytes: Option<usize>,
283 ) -> ExecResponse {
284 let started = Instant::now();
285 let deadline =
286 Instant::now() + Duration::from_millis(self.normalize_yield_ms(yield_time_ms));
287 let output_limit = self.normalize_output_bytes(max_output_bytes);
288 let mut output = HeadTailBuffer::new(output_limit);
289 let mut stderr = HeadTailBuffer::new(output_limit);
290 let mut exit_grace_deadline = None;
291
292 loop {
293 entry.drain_into(&mut output, &mut stderr);
294
295 let now = Instant::now();
296 if let Some(grace_deadline) = exit_grace_deadline {
297 if now >= grace_deadline {
298 break;
299 }
300 } else if entry.handle.has_exited() {
301 exit_grace_deadline = Some(deadline.min(now + EXIT_DRAIN_GRACE));
302 } else if now >= deadline {
303 break;
304 }
305
306 let wait_until = exit_grace_deadline.unwrap_or(deadline);
307 if Instant::now() >= wait_until {
308 continue;
309 }
310
311 let notified = entry.notify.notified();
312 tokio::pin!(notified);
313
314 tokio::select! {
315 _ = &mut notified => {}
316 _ = entry.handle.wait_for_exit(), if exit_grace_deadline.is_none() => {
317 exit_grace_deadline = Some(deadline.min(Instant::now() + EXIT_DRAIN_GRACE));
318 }
319 _ = sleep_until(wait_until) => {
320 break;
321 }
322 }
323 }
324
325 entry.touch();
326 let exit_code = entry.handle.exit_code();
327 let process_id =
328 if entry.handle.has_exited() { None } else { Some(process_id.to_string()) };
329
330 ExecResponse {
331 output: output.to_bytes(),
332 stderr: stderr.to_bytes(),
333 process_id,
334 exit_code,
335 wall_time: started.elapsed(),
336 }
337 }
338
339 async fn entry(&self, process_id: &str) -> Result<Arc<ProcessEntry>, ExecError> {
340 self.inner
341 .entries
342 .lock()
343 .await
344 .get(process_id)
345 .cloned()
346 .ok_or_else(|| ExecError::UnknownProcess { process_id: process_id.to_string() })
347 }
348
349 async fn ensure_capacity(&self) -> Result<(), ExecError> {
350 let mut removed = self.take_expired_entries().await;
351 {
352 let mut entries = self.inner.entries.lock().await;
353 while entries.len() >= self.inner.config.max_processes {
354 let Some(process_id) = eviction_candidate(&entries) else {
355 break;
356 };
357
358 if let Some(entry) = entries.remove(&process_id) {
359 removed.push(entry);
360 }
361 }
362
363 if entries.len() >= self.inner.config.max_processes {
364 return Err(ExecError::PoolFull);
365 }
366 }
367
368 shutdown_entries(removed, true);
369
370 Ok(())
371 }
372
373 async fn prune_expired_entries(&self) {
374 shutdown_entries(self.take_expired_entries().await, true);
375 }
376
377 async fn remove_if_same(
378 &self,
379 process_id: &str,
380 expected: &Arc<ProcessEntry>,
381 terminate: bool,
382 ) {
383 let removed = {
384 let mut entries = self.inner.entries.lock().await;
385 match entries.get(process_id) {
386 Some(entry) if Arc::ptr_eq(entry, expected) => entries.remove(process_id),
387 _ => None,
388 }
389 };
390
391 if let Some(entry) = removed {
392 shutdown_entry(entry, terminate);
393 }
394 }
395
396 fn next_process_id(&self) -> String {
397 self.inner.next_process_id.fetch_add(1, Ordering::Relaxed).to_string()
398 }
399
400 fn normalize_output_bytes(&self, max_output_bytes: Option<usize>) -> usize {
401 max_output_bytes
402 .unwrap_or(self.inner.config.max_output_bytes)
403 .min(self.inner.config.max_output_bytes)
404 }
405
406 fn normalize_yield_ms(&self, yield_time_ms: u64) -> u64 {
407 let yield_time_ms =
408 if yield_time_ms == 0 { self.inner.config.default_yield_ms } else { yield_time_ms };
409
410 yield_time_ms.clamp(self.inner.config.default_yield_ms, self.inner.config.max_yield_ms)
411 }
412
413 async fn take_expired_entries(&self) -> Vec<Arc<ProcessEntry>> {
414 let timeout = Duration::from_millis(self.inner.config.background_timeout_ms);
415 let now = Instant::now();
416 let mut entries = self.inner.entries.lock().await;
417 let expired_ids = entries
418 .iter()
419 .filter(|(_, entry)| now.duration_since(entry.last_used()) > timeout)
420 .map(|(process_id, _)| process_id.clone())
421 .collect::<Vec<_>>();
422
423 expired_ids.into_iter().filter_map(|process_id| entries.remove(&process_id)).collect()
424 }
425}
426
427impl ProcessEntry {
428 fn new(spawned: SpawnedProcess, max_output_bytes: usize) -> Arc<Self> {
429 let (handle, rx) = spawned;
430 let entry = Arc::new(Self {
431 handle,
432 output: StdMutex::new(HeadTailBuffer::new(max_output_bytes)),
433 stderr: StdMutex::new(HeadTailBuffer::new(max_output_bytes)),
434 notify: Arc::new(Notify::new()),
435 last_used: StdMutex::new(Instant::now()),
436 interaction: Mutex::new(()),
437 buffer_task: StdMutex::new(None),
438 });
439
440 let task_entry = Arc::clone(&entry);
441 let task = tokio::spawn(async move {
442 task_entry.buffer_output(rx).await;
443 });
444 *lock_or_recover(&entry.buffer_task) = Some(task);
445
446 entry
447 }
448
449 async fn buffer_output(self: Arc<Self>, mut rx: broadcast::Receiver<OutputChunk>) {
450 loop {
451 match rx.recv().await {
452 Ok(chunk) => {
453 self.push_chunk(chunk);
454 self.notify.notify_waiters();
455 }
456 Err(broadcast::error::RecvError::Lagged(_)) => continue,
457 Err(broadcast::error::RecvError::Closed) => return,
458 }
459 }
460 }
461
462 fn push_chunk(&self, chunk: OutputChunk) {
463 let OutputChunk { stream, data } = chunk;
464
465 if stream == Stream::Stderr {
466 lock_or_recover(&self.stderr).push_chunk(data.clone());
467 }
468
469 lock_or_recover(&self.output).push_chunk(data);
470 }
471
472 fn drain_into(&self, output: &mut HeadTailBuffer, stderr: &mut HeadTailBuffer) {
473 lock_or_recover(&self.output).drain_into(output);
474 lock_or_recover(&self.stderr).drain_into(stderr);
475 }
476
477 fn touch(&self) {
478 *lock_or_recover(&self.last_used) = Instant::now();
479 }
480
481 fn last_used(&self) -> Instant {
482 *lock_or_recover(&self.last_used)
483 }
484
485 fn abort_buffer_task(&self) {
486 if let Some(task) = lock_or_recover(&self.buffer_task).take() {
487 task.abort();
488 }
489 }
490}
491
492fn eviction_candidate(entries: &HashMap<String, Arc<ProcessEntry>>) -> Option<String> {
493 let mut candidates = entries
494 .iter()
495 .map(|(process_id, entry)| {
496 (process_id.clone(), entry.last_used(), entry.handle.has_exited())
497 })
498 .collect::<Vec<_>>();
499
500 if let Some((process_id, _, _)) = candidates
501 .iter()
502 .filter(|(_, _, exited)| *exited)
503 .min_by_key(|(_, last_used, _)| *last_used)
504 {
505 return Some(process_id.clone());
506 }
507
508 candidates.sort_by_key(|(_, last_used, _)| *last_used);
509 if candidates.is_empty() {
510 return None;
511 }
512
513 let protected = RECENT_PROTECTION_COUNT.min(candidates.len().saturating_sub(1));
514 let unprotected_end = candidates.len().saturating_sub(protected);
515
516 candidates.into_iter().take(unprotected_end).next().map(|(process_id, _, _)| process_id)
517}
518
519fn shutdown_entry(entry: Arc<ProcessEntry>, terminate: bool) {
520 if terminate && !entry.handle.has_exited() {
521 entry.handle.terminate();
522 }
523
524 entry.abort_buffer_task();
525}
526
527fn shutdown_entries(entries: Vec<Arc<ProcessEntry>>, terminate: bool) {
528 for entry in entries {
529 shutdown_entry(entry, terminate);
530 }
531}