Skip to main content

gate4agent_shell_capabilities/
lib.rs

1//! Native, bounded execution authority for pure capability-probe plans.
2
3use gate4agent_catalog::{
4    parse_capability_models_for, resolve_capability_probe_for, AgentSpec,
5    CAPABILITY_PROBE_OUTPUT_MAX_BYTES,
6};
7use gate4agent_types::{CapabilityModelSummary, CapabilityProbeFailure, CapabilityProbeRequest};
8use std::collections::HashMap;
9use std::future::pending;
10use std::process::Stdio;
11use std::sync::atomic::{AtomicUsize, Ordering};
12use std::sync::Arc;
13use std::time::Duration;
14use tokio::io::{AsyncRead, AsyncReadExt};
15use tokio::process::Command;
16use tokio::task::JoinHandle;
17use tokio::time::{sleep_until, timeout, Instant};
18
19#[derive(Clone, Copy, Debug, Eq, PartialEq)]
20pub struct NativeCapabilityProbeConfig {
21    pub timeout: Duration,
22    pub output_max_bytes: usize,
23}
24
25impl Default for NativeCapabilityProbeConfig {
26    fn default() -> Self {
27        Self {
28            timeout: Duration::from_secs(60),
29            output_max_bytes: CAPABILITY_PROBE_OUTPUT_MAX_BYTES,
30        }
31    }
32}
33
34#[derive(Clone, Debug, Eq, Hash, PartialEq)]
35struct ProbeCacheKey {
36    agent_id: String,
37    program: String,
38    args: Vec<String>,
39}
40
41pub struct NativeCapabilityProbeAuthority {
42    config: NativeCapabilityProbeConfig,
43    cache: HashMap<ProbeCacheKey, Result<Vec<CapabilityModelSummary>, CapabilityProbeFailure>>,
44}
45
46impl Default for NativeCapabilityProbeAuthority {
47    fn default() -> Self {
48        Self::new(NativeCapabilityProbeConfig::default())
49    }
50}
51
52impl NativeCapabilityProbeAuthority {
53    pub fn new(mut config: NativeCapabilityProbeConfig) -> Self {
54        config.output_max_bytes = config
55            .output_max_bytes
56            .clamp(1, CAPABILITY_PROBE_OUTPUT_MAX_BYTES);
57        Self {
58            config,
59            cache: HashMap::new(),
60        }
61    }
62
63    /// Runs at most once for an agent launch contract in this native host.
64    /// Both success and failure are cached to preserve Orca's once-per-host
65    /// fallback semantics.
66    pub async fn probe(
67        &mut self,
68        spec: &AgentSpec,
69        working_directory: &str,
70    ) -> Result<Vec<CapabilityModelSummary>, CapabilityProbeFailure> {
71        CapabilityProbeRequest {
72            working_directory: working_directory.to_owned(),
73        }
74        .validate()
75        .map_err(|_| CapabilityProbeFailure::AuthorityRejected)?;
76        let plan = resolve_capability_probe_for(spec)
77            .map_err(|_| CapabilityProbeFailure::AuthorityRejected)?;
78        let key = ProbeCacheKey {
79            agent_id: spec.id.to_string(),
80            program: plan.program.clone(),
81            args: plan.args.clone(),
82        };
83        if let Some(cached) = self.cache.get(&key) {
84            return cached.clone();
85        }
86        let result = execute_probe(
87            spec,
88            &plan.program,
89            &plan.args,
90            working_directory,
91            self.config,
92        )
93        .await;
94        self.cache.insert(key, result.clone());
95        result
96    }
97}
98
99async fn execute_probe(
100    spec: &AgentSpec,
101    program: &str,
102    args: &[String],
103    working_directory: &str,
104    config: NativeCapabilityProbeConfig,
105) -> Result<Vec<CapabilityModelSummary>, CapabilityProbeFailure> {
106    let mut command = native_command(program, args);
107    command
108        .current_dir(working_directory)
109        .stdin(Stdio::null())
110        .stdout(Stdio::piped())
111        .stderr(Stdio::piped())
112        .kill_on_drop(true);
113    configure_process_group(&mut command);
114    let mut child = command
115        .spawn()
116        .map_err(|_| CapabilityProbeFailure::SpawnUnavailable)?;
117    let process_id = child.id().ok_or(CapabilityProbeFailure::SpawnUnavailable)?;
118    let stdout = child
119        .stdout
120        .take()
121        .ok_or(CapabilityProbeFailure::SpawnUnavailable)?;
122    let stderr = child
123        .stderr
124        .take()
125        .ok_or(CapabilityProbeFailure::SpawnUnavailable)?;
126    let total = Arc::new(AtomicUsize::new(0));
127    let mut stdout_task = Some(tokio::spawn(read_bounded(
128        stdout,
129        Arc::clone(&total),
130        config.output_max_bytes,
131    )));
132    let mut stderr_task = Some(tokio::spawn(read_bounded(
133        stderr,
134        total,
135        config.output_max_bytes,
136    )));
137    let mut wait_task = Some(tokio::spawn(async move { child.wait().await }));
138    let deadline = Instant::now() + config.timeout;
139    let mut stdout_bytes = None;
140    let mut stderr_bytes = None;
141    let mut status = None;
142
143    while stdout_bytes.is_none() || stderr_bytes.is_none() || status.is_none() {
144        tokio::select! {
145            _ = sleep_until(deadline) => {
146                terminate_process_tree(process_id).await;
147                abort_tasks(&mut wait_task, &mut stdout_task, &mut stderr_task);
148                return Err(CapabilityProbeFailure::TimedOut);
149            }
150            result = join_optional(&mut stdout_task), if stdout_task.is_some() => {
151                stdout_task = None;
152                stdout_bytes = Some(resolve_reader(result, process_id, &mut wait_task, &mut stderr_task).await?);
153            }
154            result = join_optional(&mut stderr_task), if stderr_task.is_some() => {
155                stderr_task = None;
156                stderr_bytes = Some(resolve_reader(result, process_id, &mut wait_task, &mut stdout_task).await?);
157            }
158            result = join_optional(&mut wait_task), if wait_task.is_some() => {
159                wait_task = None;
160                status = Some(result
161                    .map_err(|_| CapabilityProbeFailure::ExecutorUnavailable)?
162                    .map_err(|_| CapabilityProbeFailure::SpawnUnavailable)?);
163            }
164        }
165    }
166
167    let status = status.expect("wait task completed");
168    if !status.success() {
169        return Err(CapabilityProbeFailure::NonZeroExit {
170            exit_code: status.code(),
171        });
172    }
173    let stdout = String::from_utf8_lossy(stdout_bytes.as_deref().unwrap_or_default());
174    parse_capability_models_for(spec, &stdout)
175        .map_err(|_| CapabilityProbeFailure::AuthorityRejected)
176}
177
178async fn join_optional<T>(task: &mut Option<JoinHandle<T>>) -> Result<T, tokio::task::JoinError> {
179    match task.as_mut() {
180        Some(task) => task.await,
181        None => pending().await,
182    }
183}
184
185async fn resolve_reader(
186    result: Result<Result<Vec<u8>, CapabilityProbeFailure>, tokio::task::JoinError>,
187    process_id: u32,
188    wait_task: &mut Option<JoinHandle<std::io::Result<std::process::ExitStatus>>>,
189    peer_task: &mut Option<JoinHandle<Result<Vec<u8>, CapabilityProbeFailure>>>,
190) -> Result<Vec<u8>, CapabilityProbeFailure> {
191    match result {
192        Ok(Ok(bytes)) => Ok(bytes),
193        Ok(Err(failure)) => {
194            terminate_process_tree(process_id).await;
195            if let Some(task) = wait_task.take() {
196                task.abort();
197            }
198            if let Some(task) = peer_task.take() {
199                task.abort();
200            }
201            Err(failure)
202        }
203        Err(_) => {
204            terminate_process_tree(process_id).await;
205            if let Some(task) = wait_task.take() {
206                task.abort();
207            }
208            if let Some(task) = peer_task.take() {
209                task.abort();
210            }
211            Err(CapabilityProbeFailure::ExecutorUnavailable)
212        }
213    }
214}
215
216async fn read_bounded(
217    mut reader: impl AsyncRead + Unpin,
218    total: Arc<AtomicUsize>,
219    limit: usize,
220) -> Result<Vec<u8>, CapabilityProbeFailure> {
221    let mut output = Vec::new();
222    let mut chunk = [0_u8; 8 * 1024];
223    loop {
224        let read = reader
225            .read(&mut chunk)
226            .await
227            .map_err(|_| CapabilityProbeFailure::ExecutorUnavailable)?;
228        if read == 0 {
229            return Ok(output);
230        }
231        let previous = total.fetch_add(read, Ordering::AcqRel);
232        if previous > limit || read > limit.saturating_sub(previous) {
233            return Err(CapabilityProbeFailure::OutputLimitExceeded);
234        }
235        output.extend_from_slice(&chunk[..read]);
236    }
237}
238
239fn abort_tasks(
240    wait_task: &mut Option<JoinHandle<std::io::Result<std::process::ExitStatus>>>,
241    stdout_task: &mut Option<JoinHandle<Result<Vec<u8>, CapabilityProbeFailure>>>,
242    stderr_task: &mut Option<JoinHandle<Result<Vec<u8>, CapabilityProbeFailure>>>,
243) {
244    for task in [
245        wait_task.take().map(|task| task.abort_handle()),
246        stdout_task.take().map(|task| task.abort_handle()),
247        stderr_task.take().map(|task| task.abort_handle()),
248    ]
249    .into_iter()
250    .flatten()
251    {
252        task.abort();
253    }
254}
255
256fn native_command(program: &str, args: &[String]) -> Command {
257    #[cfg(windows)]
258    if program.ends_with(".cmd") || program.ends_with(".bat") {
259        let mut command = Command::new("cmd.exe");
260        command.args(["/D", "/S", "/C", program]).args(args);
261        return command;
262    }
263    let mut command = Command::new(program);
264    command.args(args);
265    command
266}
267
268fn configure_process_group(command: &mut Command) {
269    #[cfg(unix)]
270    {
271        use std::os::unix::process::CommandExt;
272        command.as_std_mut().process_group(0);
273    }
274    #[cfg(windows)]
275    {
276        use std::os::windows::process::CommandExt;
277        const CREATE_NO_WINDOW: u32 = 0x0800_0000;
278        const CREATE_NEW_PROCESS_GROUP: u32 = 0x0000_0200;
279        command
280            .as_std_mut()
281            .creation_flags(CREATE_NO_WINDOW | CREATE_NEW_PROCESS_GROUP);
282    }
283}
284
285async fn terminate_process_tree(process_id: u32) {
286    #[cfg(windows)]
287    let mut command = {
288        let mut command = Command::new("taskkill.exe");
289        command.args(["/PID", &process_id.to_string(), "/T", "/F"]);
290        command
291    };
292    #[cfg(unix)]
293    let mut command = {
294        let mut command = Command::new("kill");
295        command.args(["-KILL", "--", &format!("-{process_id}")]);
296        command
297    };
298    command
299        .stdin(Stdio::null())
300        .stdout(Stdio::null())
301        .stderr(Stdio::null());
302    let _ = timeout(Duration::from_secs(2), command.status()).await;
303}
304
305#[cfg(test)]
306mod tests {
307    use super::*;
308    use gate4agent_catalog::builtin_registry;
309    use std::fs;
310    use std::time::{SystemTime, UNIX_EPOCH};
311
312    /// A minimal stand-in spec for `execute_probe` fixtures below. `id`,
313    /// `launch`, and `capabilities.adapters.capability_probe` are
314    /// overwritten by the caller; the rest just needs to be a valid spec.
315    /// No provider in the current fleet declares a capability-probe
316    /// adapter -- the feature existed to serve `cursor`'s `--list-models`,
317    /// and `cursor` is not part of the fleet -- so these fixtures exercise
318    /// `execute_probe` directly (spawn, timeout, output bounding) rather
319    /// than through the public `probe()`, which fails closed for every
320    /// real spec before ever reaching it.
321    fn probe_fixture_spec() -> AgentSpec {
322        builtin_registry().get_by_id("codex").unwrap().clone()
323    }
324
325    fn fixture_dir(name: &str) -> std::path::PathBuf {
326        let nonce = SystemTime::now()
327            .duration_since(UNIX_EPOCH)
328            .unwrap()
329            .as_nanos();
330        let path = std::env::temp_dir().join(format!(
331            "gate4agent-capability-{name}-{}-{nonce}",
332            std::process::id()
333        ));
334        fs::create_dir(&path).unwrap();
335        path
336    }
337
338    #[tokio::test]
339    async fn combined_output_limit_fails_typed() {
340        let directory = fixture_dir("limit");
341        let script_path = directory.join(if cfg!(windows) {
342            "probe.ps1"
343        } else {
344            "probe.sh"
345        });
346        #[cfg(windows)]
347        let script = "[Console]::Write('auto - A label that exceeds the configured limit')";
348        #[cfg(not(windows))]
349        let script = "printf 'auto - A label that exceeds the configured limit'";
350        fs::write(&script_path, script).unwrap();
351        #[cfg(windows)]
352        let (program, args) = (
353            "powershell.exe".to_owned(),
354            vec![
355                "-NoProfile".to_owned(),
356                "-NonInteractive".to_owned(),
357                "-File".to_owned(),
358                script_path.display().to_string(),
359            ],
360        );
361        #[cfg(not(windows))]
362        let (program, args) = ("sh".to_owned(), vec![script_path.display().to_string()]);
363        let spec = probe_fixture_spec();
364        let result = execute_probe(
365            &spec,
366            &program,
367            &args,
368            directory.to_str().unwrap(),
369            NativeCapabilityProbeConfig {
370                timeout: Duration::from_secs(5),
371                output_max_bytes: 16,
372            },
373        )
374        .await;
375        assert_eq!(result, Err(CapabilityProbeFailure::OutputLimitExceeded));
376        fs::remove_dir_all(directory).unwrap();
377    }
378
379    #[tokio::test]
380    async fn timeout_fails_typed_and_terminates_the_probe_process() {
381        let directory = fixture_dir("timeout");
382        let script_path = directory.join(if cfg!(windows) {
383            "probe.cmd"
384        } else {
385            "probe.sh"
386        });
387        #[cfg(windows)]
388        let script = "@echo off\r\ntimeout /T 10 /NOBREAK >NUL\r\n";
389        #[cfg(not(windows))]
390        let script = "sleep 10";
391        fs::write(&script_path, script).unwrap();
392        #[cfg(windows)]
393        let (program, args) = (script_path.display().to_string(), Vec::new());
394        #[cfg(not(windows))]
395        let (program, args) = ("sh".to_owned(), vec![script_path.display().to_string()]);
396        let spec = probe_fixture_spec();
397        let result = execute_probe(
398            &spec,
399            &program,
400            &args,
401            directory.to_str().unwrap(),
402            NativeCapabilityProbeConfig {
403                timeout: Duration::from_millis(50),
404                output_max_bytes: 1_024,
405            },
406        )
407        .await;
408        assert_eq!(result, Err(CapabilityProbeFailure::TimedOut));
409        fs::remove_dir_all(directory).unwrap();
410    }
411}