gate4agent_shell_capabilities/
lib.rs1use 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 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 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}