Skip to main content

ante_exec/
pool.rs

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}