mermaid_cli/providers/tool/
subagent.rs1use std::sync::Arc;
26use std::time::{Duration, Instant};
27
28use async_trait::async_trait;
29use serde_json::Value;
30use tokio::sync::{Semaphore, mpsc};
31use tokio::time::timeout;
32use tokio_util::sync::CancellationToken;
33
34use crate::domain::{
35 Msg, State, ToolDefinition, ToolMetadata, ToolOutcome, ToolRunMetadata, TurnState, update,
36};
37use crate::effect::{EffectRunner, MSG_CHANNEL_CAPACITY};
38use crate::models::MessageRole;
39use crate::providers::ProviderFactory;
40use crate::providers::ctx::{ExecContext, ProgressEvent, SubagentPhase};
41
42use super::ToolExecutor;
43use super::ToolRegistry;
44
45pub const MAX_DEPTH: usize = 3;
48
49pub const MAX_INFLIGHT: usize = 10;
54
55pub const DEFAULT_TIMEOUT_SECS: u64 = 20 * 60;
58
59tokio::task_local! {
60 static SUBAGENT_DEPTH: usize;
65}
66
67pub struct SubagentSpawner {
69 providers: Arc<ProviderFactory>,
70 inflight: Arc<Semaphore>,
71}
72
73impl SubagentSpawner {
74 pub fn new(providers: Arc<ProviderFactory>) -> Self {
75 Self {
76 providers,
77 inflight: Arc::new(Semaphore::new(MAX_INFLIGHT)),
78 }
79 }
80}
81
82pub struct SubagentTool {
84 spawner: Arc<SubagentSpawner>,
85}
86
87impl SubagentTool {
88 pub fn new(spawner: Arc<SubagentSpawner>) -> Self {
89 Self { spawner }
90 }
91}
92
93#[async_trait]
94impl ToolExecutor for SubagentTool {
95 fn name(&self) -> &'static str {
96 "agent"
97 }
98
99 fn schema(&self) -> ToolDefinition {
100 ToolDefinition {
101 name: "agent".to_string(),
102 description: format!(
103 "Spawn a child agent with its own context and tool access to work on an \
104 independent sub-task. Useful for parallel fan-out (emit multiple `agent` \
105 calls in the same turn to run them concurrently) or for scoping a noisy \
106 sub-task (the child's tool output doesn't clutter the parent's turn). \
107 Depth-capped at {max_depth}; breadth-capped at {max_breadth} concurrent. \
108 Subagents don't get GUI (screenshot/click/…) access because coordinate \
109 metadata can't be shared cleanly.",
110 max_depth = MAX_DEPTH,
111 max_breadth = MAX_INFLIGHT,
112 ),
113 input_schema: serde_json::json!({
114 "type": "object",
115 "properties": {
116 "prompt": {
117 "type": "string",
118 "description": "The task for the subagent. Self-contained; the subagent has no access to the parent's conversation."
119 },
120 "description": {
121 "type": "string",
122 "description": "Short label shown in the parent's status line (e.g. 'list domain files')."
123 }
124 },
125 "required": ["prompt"]
126 }),
127 }
128 }
129
130 async fn execute(&self, args: Value, ctx: ExecContext) -> ToolOutcome {
131 let started = Instant::now();
132
133 let current_depth = SUBAGENT_DEPTH.try_with(|d| *d).unwrap_or(0);
135 if current_depth >= MAX_DEPTH {
136 return ToolOutcome::error(format!("subagent depth limit {} reached", MAX_DEPTH), 0.0);
137 }
138
139 let prompt = match args.get("prompt").and_then(|v| v.as_str()) {
141 Some(s) if !s.trim().is_empty() => s.to_string(),
142 _ => {
143 return ToolOutcome::error("agent requires non-empty `prompt`", 0.0);
144 },
145 };
146 let description = args
147 .get("description")
148 .and_then(|v| v.as_str())
149 .unwrap_or("subagent")
150 .to_string();
151
152 if let Some(blocked) = super::policy_gate::gate_external(
156 &ctx,
157 "agent",
158 crate::runtime::ToolCategory::Subagent,
159 format!("subagent: {}", description),
160 &args,
161 )
162 .await
163 {
164 return blocked;
165 }
166
167 let permit = tokio::select! {
171 biased;
172 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
173 p = self.spawner.inflight.clone().acquire_owned() => match p {
174 Ok(permit) => permit,
175 Err(_) => return ToolOutcome::error(
176 "subagent semaphore closed",
177 started.elapsed().as_secs_f64(),
178 ),
179 },
180 };
181
182 let config = (*ctx.config).clone();
192 let cwd = ctx.workdir.clone();
193 let model_id = if ctx.model_id.is_empty() {
194 default_model_id(&config)
195 } else {
196 ctx.model_id.clone()
197 };
198 let child_model_id = model_id.clone();
199 let child_state = State::new(config.clone(), cwd.clone(), model_id);
200
201 let child_tools = build_child_registry(self.spawner.providers.clone());
202
203 let child_token = ctx.token.child_token();
207 let (child_tx, child_rx) = mpsc::channel(MSG_CHANNEL_CAPACITY);
208 let child_runner =
209 EffectRunner::new_child(child_tx, cwd, self.spawner.providers.clone(), child_tools);
210
211 let drive = drive_child(
214 child_state,
215 child_runner,
216 child_rx,
217 ctx.progress.clone(),
218 prompt,
219 description.clone(),
220 child_token,
221 );
222 let depth_scoped = SUBAGENT_DEPTH.scope(current_depth + 1, drive);
223
224 let result = timeout(Duration::from_secs(DEFAULT_TIMEOUT_SECS), depth_scoped).await;
225 drop(permit);
226
227 let elapsed = started.elapsed().as_secs_f64();
228 match result {
229 Ok(Ok(summary)) => ToolOutcome::success(summary, "subagent completed", elapsed)
230 .with_metadata(subagent_metadata(child_model_id)),
231 Ok(Err(DriveError::Cancelled)) => ToolOutcome::cancelled(),
232 Ok(Err(DriveError::Errored(e))) => {
233 ToolOutcome::error(format!("subagent ({}): {}", description, e), elapsed)
234 .with_metadata(subagent_metadata(child_model_id))
235 },
236 Err(_) => ToolOutcome::error(
237 format!(
238 "subagent ({}) exceeded {}s timeout",
239 description, DEFAULT_TIMEOUT_SECS
240 ),
241 elapsed,
242 )
243 .with_metadata(subagent_metadata(child_model_id)),
244 }
245 }
246}
247
248fn subagent_metadata(model_id: String) -> ToolRunMetadata {
249 ToolRunMetadata {
250 detail: ToolMetadata::Subagent { model_id },
251 ..ToolRunMetadata::default()
252 }
253}
254
255enum DriveError {
256 Cancelled,
257 Errored(String),
258}
259
260async fn drive_child(
264 mut state: State,
265 mut runner: EffectRunner,
266 mut msg_rx: mpsc::Receiver<Msg>,
267 parent_progress: mpsc::Sender<ProgressEvent>,
268 prompt: String,
269 description: String,
270 token: CancellationToken,
271) -> Result<String, DriveError> {
272 let _ = parent_progress
274 .send(ProgressEvent::SubagentText(format!(
275 "▶ {} — {}",
276 description,
277 prompt.chars().take(80).collect::<String>()
278 )))
279 .await;
280
281 runner.dispatch(crate::domain::Cmd::RefreshInstructions);
283
284 let seed = Msg::SubmitPrompt {
286 text: prompt,
287 attachment_ids: vec![],
288 };
289 let (new_state, cmds) = update(state, seed);
290 state = new_state;
291 for cmd in cmds {
292 runner.dispatch(cmd);
293 }
294
295 loop {
297 if token.is_cancelled() {
298 runner.shutdown().await;
299 return Err(DriveError::Cancelled);
300 }
301 if matches!(state.turn, TurnState::Idle) && state.ui.queued_messages.is_empty() {
302 break;
303 }
304
305 let msg = tokio::select! {
306 biased;
307 _ = token.cancelled() => {
308 runner.shutdown().await;
309 return Err(DriveError::Cancelled);
310 },
311 recv = msg_rx.recv() => match recv {
312 Some(m) => m,
313 None => {
314 break;
316 },
317 },
318 };
319
320 forward_child_event(&msg, &parent_progress, &state).await;
324
325 let (new_state, cmds) = update(state, msg);
326 state = new_state;
327 for cmd in cmds {
328 runner.dispatch(cmd);
329 }
330 if state.should_exit {
331 break;
332 }
333 }
334
335 runner.shutdown().await;
336
337 let summary = state
339 .session
340 .messages()
341 .iter()
342 .rev()
343 .find(|m| m.role == MessageRole::Assistant)
344 .map(|m| m.content.clone())
345 .unwrap_or_default();
346 if summary.trim().is_empty() {
347 return Err(DriveError::Errored(
348 "subagent produced no assistant output".to_string(),
349 ));
350 }
351 Ok(summary)
352}
353
354async fn forward_child_event(msg: &Msg, progress: &mpsc::Sender<ProgressEvent>, state: &State) {
359 match msg {
360 Msg::ToolStarted {
361 turn: _, call_id, ..
362 } => {
363 let tool_name = lookup_tool_name(state, *call_id).unwrap_or_else(|| "tool".to_string());
364 let _ = progress
365 .send(ProgressEvent::SubagentToolCall {
366 child_call_id: *call_id,
367 tool_name,
368 phase: SubagentPhase::Started,
369 })
370 .await;
371 },
372 Msg::ToolFinished {
373 turn: _,
374 call_id,
375 outcome,
376 } => {
377 let tool_name = lookup_tool_name(state, *call_id).unwrap_or_else(|| "tool".to_string());
378 let phase = if outcome.is_success() {
379 SubagentPhase::Finished
380 } else {
381 SubagentPhase::Errored
382 };
383 let _ = progress
384 .send(ProgressEvent::SubagentToolCall {
385 child_call_id: *call_id,
386 tool_name,
387 phase,
388 })
389 .await;
390 },
391 Msg::StreamText { chunk, .. } => {
392 if !chunk.trim().is_empty() {
395 let snippet: String = chunk.chars().take(120).collect();
396 let _ = progress.send(ProgressEvent::SubagentText(snippet)).await;
397 }
398 },
399 _ => {},
400 }
401}
402
403fn lookup_tool_name(state: &State, call_id: crate::domain::ToolCallId) -> Option<String> {
406 match &state.turn {
407 TurnState::ExecutingTools { calls, .. } => calls
408 .iter()
409 .find(|c| c.call_id == call_id)
410 .map(|c| c.source.function.name.clone()),
411 _ => None,
412 }
413}
414
415fn build_child_registry(providers: Arc<ProviderFactory>) -> Arc<ToolRegistry> {
429 use super::{
430 computer_use, exec, filesystem, mcp,
431 web::{WebFetchTool, WebSearchTool},
432 };
433 let mut r = ToolRegistry::new();
434 r.register(Arc::new(filesystem::ReadFileTool));
435 r.register(Arc::new(filesystem::WriteFileTool));
436 r.register(Arc::new(filesystem::EditFileTool));
437 r.register(Arc::new(filesystem::DeleteFileTool));
438 r.register(Arc::new(filesystem::CreateDirectoryTool));
439 r.register(Arc::new(exec::ExecuteCommandTool));
440 r.register(Arc::new(mcp::McpToolProxy));
441 if let Some(key) = crate::utils::resolve_api_key("OLLAMA_API_KEY", None) {
442 r.register(Arc::new(WebSearchTool::new(key.clone())));
443 r.register(Arc::new(WebFetchTool::new(key)));
444 }
445 let _ = computer_use::probe;
449 let _ = providers;
450 Arc::new(r)
451}
452
453fn default_model_id(config: &crate::app::Config) -> String {
458 if !config.default_model.provider.is_empty() && !config.default_model.name.is_empty() {
459 format!(
460 "{}/{}",
461 config.default_model.provider, config.default_model.name
462 )
463 } else {
464 config.default_model.name.clone()
465 }
466}
467
468#[cfg(test)]
469mod tests {
470 use super::*;
471 use crate::domain::{ToolCallId, TurnId};
472 use crate::providers::ctx::test_exec_context;
473 use std::path::PathBuf;
474
475 #[tokio::test]
476 async fn depth_cap_rejects_when_at_max() {
477 let spawner = Arc::new(SubagentSpawner::new(Arc::new(ProviderFactory::new(
478 crate::app::Config::default(),
479 ))));
480 let tool = SubagentTool::new(spawner);
481 let (ctx, _rx) = test_exec_context(TurnId(1), ToolCallId(1), PathBuf::from("/tmp"));
482
483 let outcome = SUBAGENT_DEPTH
484 .scope(
485 MAX_DEPTH,
486 tool.execute(serde_json::json!({"prompt": "hi"}), ctx),
487 )
488 .await;
489 let error = outcome.error_message().expect("expected error");
490 assert!(
491 error.contains("depth limit"),
492 "expected depth-limit error, got: {}",
493 error
494 );
495 }
496
497 #[tokio::test]
498 async fn empty_prompt_is_rejected() {
499 let spawner = Arc::new(SubagentSpawner::new(Arc::new(ProviderFactory::new(
500 crate::app::Config::default(),
501 ))));
502 let tool = SubagentTool::new(spawner);
503 let (ctx, _rx) = test_exec_context(TurnId(1), ToolCallId(1), PathBuf::from("/tmp"));
504 let outcome = tool.execute(serde_json::json!({"prompt": " "}), ctx).await;
505 assert_eq!(outcome.status, crate::domain::ToolStatus::Error);
506 }
507
508 #[test]
512 fn default_model_id_reads_config_provider_and_name() {
513 let mut cfg = crate::app::Config::default();
514 cfg.default_model.provider = "ollama".to_string();
515 cfg.default_model.name = "qwen3-coder:30b".to_string();
516 assert_eq!(default_model_id(&cfg), "ollama/qwen3-coder:30b");
517 }
518
519 #[test]
520 fn default_model_id_returns_bare_name_when_provider_empty() {
521 let mut cfg = crate::app::Config::default();
522 cfg.default_model.name = "just-a-name".to_string();
523 assert_eq!(default_model_id(&cfg), "just-a-name");
526 }
527
528 #[test]
529 fn build_child_registry_excludes_gui_and_self() {
530 let providers = Arc::new(ProviderFactory::new(crate::app::Config::default()));
531 let r = build_child_registry(providers);
532 assert!(r.get("screenshot").is_none());
534 assert!(r.get("click").is_none());
535 assert!(r.get("type_text").is_none());
536 assert!(r.get("press_key").is_none());
537 assert!(r.get("scroll").is_none());
538 assert!(r.get("mouse_move").is_none());
539 assert!(r.get("list_windows").is_none());
540 assert!(r.get("agent").is_none());
542 assert!(r.get("read_file").is_some());
544 assert!(r.get("execute_command").is_some());
545 }
546}