mermaid_cli/providers/tool/
subagent.rs1use std::sync::Arc;
25use std::time::{Duration, Instant};
26
27use async_trait::async_trait;
28use serde_json::Value;
29use tokio::sync::{Semaphore, mpsc};
30use tokio_util::sync::CancellationToken;
31
32use crate::domain::{
33 Msg, State, ToolDefinition, ToolMetadata, ToolOutcome, ToolRunMetadata, TurnState, update,
34};
35use crate::effect::{EffectRunner, MSG_CHANNEL_CAPACITY};
36use crate::models::MessageRole;
37use crate::providers::ProviderFactory;
38use crate::providers::ctx::{ExecContext, ProgressEvent, SubagentPhase};
39
40use super::ToolExecutor;
41use super::ToolRegistry;
42
43pub const MAX_INFLIGHT: usize = 10;
48
49pub const DEFAULT_TIMEOUT_SECS: u64 = 20 * 60;
52
53pub struct SubagentSpawner {
55 providers: Arc<ProviderFactory>,
56 inflight: Arc<Semaphore>,
57}
58
59impl SubagentSpawner {
60 pub fn new(providers: Arc<ProviderFactory>) -> Self {
61 Self {
62 providers,
63 inflight: Arc::new(Semaphore::new(MAX_INFLIGHT)),
64 }
65 }
66}
67
68pub struct SubagentTool {
70 spawner: Arc<SubagentSpawner>,
71}
72
73impl SubagentTool {
74 pub fn new(spawner: Arc<SubagentSpawner>) -> Self {
75 Self { spawner }
76 }
77}
78
79#[async_trait]
80impl ToolExecutor for SubagentTool {
81 fn name(&self) -> &'static str {
82 "agent"
83 }
84
85 fn schema(&self) -> ToolDefinition {
86 ToolDefinition {
87 name: "agent".to_string(),
88 description: format!(
89 "Spawn a child agent with its own context and tool access to work on an \
90 independent sub-task. Useful for parallel fan-out (emit multiple `agent` \
91 calls in the same turn to run them concurrently) or for scoping a noisy \
92 sub-task (the child's tool output doesn't clutter the parent's turn). \
93 Breadth-capped at {max_breadth} concurrent; subagents can't themselves \
94 spawn subagents. Subagents don't get GUI (screenshot/click/…) access \
95 because coordinate metadata can't be shared cleanly.",
96 max_breadth = MAX_INFLIGHT,
97 ),
98 input_schema: serde_json::json!({
99 "type": "object",
100 "properties": {
101 "prompt": {
102 "type": "string",
103 "description": "The task for the subagent. Self-contained; the subagent has no access to the parent's conversation."
104 },
105 "description": {
106 "type": "string",
107 "description": "Short label shown in the parent's status line (e.g. 'list domain files')."
108 }
109 },
110 "required": ["prompt"]
111 }),
112 }
113 }
114
115 async fn execute(&self, args: Value, ctx: ExecContext) -> ToolOutcome {
116 let started = Instant::now();
117
118 let prompt = match args.get("prompt").and_then(|v| v.as_str()) {
120 Some(s) if !s.trim().is_empty() => s.to_string(),
121 _ => {
122 return ToolOutcome::error("agent requires non-empty `prompt`", 0.0);
123 },
124 };
125 let description = args
126 .get("description")
127 .and_then(|v| v.as_str())
128 .unwrap_or("subagent")
129 .to_string();
130
131 if let Some(blocked) = super::policy_gate::gate_external(
135 &ctx,
136 "agent",
137 crate::runtime::ToolCategory::Subagent,
138 format!("subagent: {}", description),
139 &args,
140 )
141 .await
142 {
143 return blocked;
144 }
145
146 let permit = tokio::select! {
150 biased;
151 _ = ctx.token.cancelled() => return ToolOutcome::cancelled(),
152 p = self.spawner.inflight.clone().acquire_owned() => match p {
153 Ok(permit) => permit,
154 Err(_) => return ToolOutcome::error(
155 "subagent semaphore closed",
156 started.elapsed().as_secs_f64(),
157 ),
158 },
159 };
160
161 let config = (*ctx.config).clone();
171 let cwd = ctx.workdir.clone();
172 let model_id = if ctx.model_id.is_empty() {
173 default_model_id(&config)
174 } else {
175 ctx.model_id.clone()
176 };
177 let child_model_id = model_id.clone();
178 let mut child_state =
185 State::new(config.clone(), cwd.clone(), model_id, chrono::Local::now());
186 child_state.session.safety_mode = ctx.safety_mode;
187 let (instructions, memory) =
194 crate::app::instructions::load_project_context(&cwd, &config.memory);
195 child_state.instructions = instructions;
196 child_state.memory = memory;
197
198 let child_tools = build_child_registry(self.spawner.providers.clone());
199
200 let child_token = ctx.token.child_token();
204 let (child_tx, child_rx) = mpsc::channel(MSG_CHANNEL_CAPACITY);
205 let child_runner =
206 EffectRunner::new_child(child_tx, cwd, self.spawner.providers.clone(), child_tools);
207
208 let result = drive_child(
212 child_state,
213 child_runner,
214 child_rx,
215 ctx.progress.clone(),
216 prompt,
217 description.clone(),
218 child_token,
219 )
220 .await;
221 drop(permit);
222
223 let elapsed = started.elapsed().as_secs_f64();
224 match result {
225 Ok(summary) => ToolOutcome::success(summary, "subagent completed", elapsed)
226 .with_metadata(subagent_metadata(child_model_id)),
227 Err(DriveError::Cancelled) => ToolOutcome::cancelled(),
228 Err(DriveError::TimedOut) => ToolOutcome::error(
229 format!(
230 "subagent ({}) exceeded {}s timeout",
231 description, DEFAULT_TIMEOUT_SECS
232 ),
233 elapsed,
234 )
235 .with_metadata(subagent_metadata(child_model_id)),
236 Err(DriveError::Errored(e)) => {
237 ToolOutcome::error(format!("subagent ({}): {}", description, e), elapsed)
238 .with_metadata(subagent_metadata(child_model_id))
239 },
240 }
241 }
242}
243
244fn subagent_metadata(model_id: String) -> ToolRunMetadata {
245 ToolRunMetadata {
246 detail: ToolMetadata::Subagent { model_id },
247 ..ToolRunMetadata::default()
248 }
249}
250
251enum DriveError {
252 Cancelled,
253 TimedOut,
254 Errored(String),
255}
256
257async fn drive_child(
261 mut state: State,
262 mut runner: EffectRunner,
263 mut msg_rx: mpsc::Receiver<Msg>,
264 parent_progress: mpsc::Sender<ProgressEvent>,
265 prompt: String,
266 description: String,
267 token: CancellationToken,
268) -> Result<String, DriveError> {
269 let _ = parent_progress
271 .send(ProgressEvent::SubagentText(format!(
272 "{} — {}",
273 description,
274 prompt.chars().take(80).collect::<String>()
275 )))
276 .await;
277
278 let seed = Msg::SubmitPrompt {
285 text: prompt,
286 attachment_ids: vec![],
287 };
288 let (new_state, cmds) = update(state, seed);
289 state = new_state;
290 for cmd in cmds {
291 runner.dispatch(cmd);
292 }
293
294 let deadline = tokio::time::sleep(Duration::from_secs(DEFAULT_TIMEOUT_SECS));
300 tokio::pin!(deadline);
301
302 let mut outcome: Result<(), DriveError> = Ok(());
303 loop {
304 if token.is_cancelled() {
305 outcome = Err(DriveError::Cancelled);
306 break;
307 }
308 if matches!(state.turn, TurnState::Idle) && state.ui.queued_messages.is_empty() {
309 break;
310 }
311
312 let msg = tokio::select! {
313 biased;
314 _ = token.cancelled() => {
315 outcome = Err(DriveError::Cancelled);
316 break;
317 },
318 _ = &mut deadline => {
319 outcome = Err(DriveError::TimedOut);
320 break;
321 },
322 recv = msg_rx.recv() => match recv {
323 Some(m) => m,
324 None => break, },
326 };
327
328 forward_child_event(&msg, &parent_progress, &state).await;
332
333 let (new_state, cmds) = update(state, msg);
334 state = new_state;
335 for cmd in cmds {
336 runner.dispatch(cmd);
337 }
338 if state.should_exit {
339 break;
340 }
341 }
342
343 runner.shutdown().await;
346
347 outcome?;
348
349 let summary = state
351 .session
352 .messages()
353 .iter()
354 .rev()
355 .find(|m| m.role == MessageRole::Assistant)
356 .map(|m| m.content.clone())
357 .unwrap_or_default();
358 if summary.trim().is_empty() {
359 return Err(DriveError::Errored(
360 "subagent produced no assistant output".to_string(),
361 ));
362 }
363 Ok(summary)
364}
365
366async fn forward_child_event(msg: &Msg, progress: &mpsc::Sender<ProgressEvent>, state: &State) {
371 match msg {
372 Msg::ToolStarted {
373 turn: _, call_id, ..
374 } => {
375 let tool_name = lookup_tool_name(state, *call_id).unwrap_or_else(|| "tool".to_string());
376 let _ = progress
377 .send(ProgressEvent::SubagentToolCall {
378 child_call_id: *call_id,
379 tool_name,
380 phase: SubagentPhase::Started,
381 })
382 .await;
383 },
384 Msg::ToolFinished {
385 turn: _,
386 call_id,
387 outcome,
388 } => {
389 let tool_name = lookup_tool_name(state, *call_id).unwrap_or_else(|| "tool".to_string());
390 let phase = if outcome.is_success() {
391 SubagentPhase::Finished
392 } else {
393 SubagentPhase::Errored
394 };
395 let _ = progress
396 .send(ProgressEvent::SubagentToolCall {
397 child_call_id: *call_id,
398 tool_name,
399 phase,
400 })
401 .await;
402 },
403 Msg::StreamText { chunk, .. } => {
404 if !chunk.trim().is_empty() {
407 let snippet: String = chunk.chars().take(120).collect();
408 let _ = progress.send(ProgressEvent::SubagentText(snippet)).await;
409 }
410 },
411 _ => {},
412 }
413}
414
415fn lookup_tool_name(state: &State, call_id: crate::domain::ToolCallId) -> Option<String> {
418 match &state.turn {
419 TurnState::ExecutingTools { calls, .. } => calls
420 .iter()
421 .find(|c| c.call_id == call_id)
422 .map(|c| c.source.function.name.clone()),
423 _ => None,
424 }
425}
426
427fn build_child_registry(providers: Arc<ProviderFactory>) -> Arc<ToolRegistry> {
441 use super::{
442 computer_use, exec, filesystem, mcp,
443 web::{WebFetchTool, WebSearchTool},
444 };
445 let mut r = ToolRegistry::new();
446 r.register(Arc::new(filesystem::ReadFileTool));
447 r.register(Arc::new(filesystem::WriteFileTool));
448 r.register(Arc::new(filesystem::EditFileTool));
449 r.register(Arc::new(filesystem::DeleteFileTool));
450 r.register(Arc::new(filesystem::CreateDirectoryTool));
451 r.register(Arc::new(exec::ExecuteCommandTool));
452 r.register(Arc::new(mcp::McpToolProxy));
453 if let Some(key) = crate::utils::resolve_api_key("OLLAMA_API_KEY", None) {
454 r.register(Arc::new(WebSearchTool::new(key.clone())));
455 r.register(Arc::new(WebFetchTool::new(key)));
456 }
457 let _ = computer_use::probe;
462 let _ = providers;
463 Arc::new(r)
464}
465
466fn default_model_id(config: &crate::app::Config) -> String {
471 if !config.default_model.provider.is_empty() && !config.default_model.name.is_empty() {
472 format!(
473 "{}/{}",
474 config.default_model.provider, config.default_model.name
475 )
476 } else {
477 config.default_model.name.clone()
478 }
479}
480
481#[cfg(test)]
482mod tests {
483 use super::*;
484 use crate::domain::{ToolCallId, TurnId};
485 use crate::providers::ctx::test_exec_context;
486 use std::path::PathBuf;
487
488 #[tokio::test]
489 async fn empty_prompt_is_rejected() {
490 let spawner = Arc::new(SubagentSpawner::new(Arc::new(ProviderFactory::new(
491 crate::app::Config::default(),
492 ))));
493 let tool = SubagentTool::new(spawner);
494 let (ctx, _rx) = test_exec_context(TurnId(1), ToolCallId(1), PathBuf::from("/tmp"));
495 let outcome = tool.execute(serde_json::json!({"prompt": " "}), ctx).await;
496 assert_eq!(outcome.status, crate::domain::ToolStatus::Error);
497 }
498
499 #[test]
500 fn child_state_inherits_live_safety_mode_over_config_default() {
501 use crate::runtime::SafetyMode;
505 let mut config = crate::app::Config::default();
506 config.safety.mode = SafetyMode::FullAccess; let mut child_state = State::new(
508 config,
509 PathBuf::from("/tmp"),
510 "ollama/test".to_string(),
511 chrono::Local::now(),
512 );
513 assert_eq!(child_state.session.safety_mode, SafetyMode::FullAccess);
515 child_state.session.safety_mode = SafetyMode::Ask;
517 assert_eq!(child_state.session.safety_mode, SafetyMode::Ask);
518 }
519
520 #[test]
524 fn default_model_id_reads_config_provider_and_name() {
525 let mut cfg = crate::app::Config::default();
526 cfg.default_model.provider = "ollama".to_string();
527 cfg.default_model.name = "qwen3-coder:30b".to_string();
528 assert_eq!(default_model_id(&cfg), "ollama/qwen3-coder:30b");
529 }
530
531 #[test]
532 fn default_model_id_returns_bare_name_when_provider_empty() {
533 let mut cfg = crate::app::Config::default();
534 cfg.default_model.name = "just-a-name".to_string();
535 assert_eq!(default_model_id(&cfg), "just-a-name");
538 }
539
540 #[test]
541 fn build_child_registry_excludes_gui_and_self() {
542 let providers = Arc::new(ProviderFactory::new(crate::app::Config::default()));
543 let r = build_child_registry(providers);
544 assert!(r.get("screenshot").is_none());
546 assert!(r.get("click").is_none());
547 assert!(r.get("type_text").is_none());
548 assert!(r.get("press_key").is_none());
549 assert!(r.get("scroll").is_none());
550 assert!(r.get("mouse_move").is_none());
551 assert!(r.get("list_windows").is_none());
552 assert!(r.get("agent").is_none());
554 assert!(r.get("read_file").is_some());
556 assert!(r.get("execute_command").is_some());
557 }
558}