mermaid_cli/providers/tool/
mod.rs1pub mod apply_patch;
19pub mod ask_user_question;
20pub mod computer_use;
21pub mod enter_plan_mode;
22pub mod exec;
23pub mod exit_plan_mode;
24pub mod filesystem;
25pub mod mcp;
26pub mod memory;
27pub mod path_lock;
28pub mod path_safety;
29pub mod policy_gate;
30pub mod subagent;
31pub mod tasks;
32pub mod web;
33pub mod web_client;
34pub mod workspace;
35
36use async_trait::async_trait;
37use std::collections::HashMap;
38use std::sync::Arc;
39
40use mermaid_domain::{ToolDefinition, ToolOutcome};
41
42use super::ctx::ExecContext;
43
44#[async_trait]
48pub trait ToolExecutor: Send + Sync {
49 fn name(&self) -> &'static str;
52
53 fn schema(&self) -> ToolDefinition;
59
60 fn is_internal(&self) -> bool {
66 false
67 }
68
69 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome;
73}
74
75pub struct ToolRegistry {
79 entries: HashMap<&'static str, Arc<dyn ToolExecutor>>,
80 unavailable: HashMap<&'static str, String>,
88 web_capabilities: Option<Arc<web::WebCapabilities>>,
93 subagent_spawner: Option<Arc<subagent::SubagentSpawner>>,
98}
99
100impl ToolRegistry {
101 #[must_use]
102 pub fn new() -> Self {
103 Self {
104 entries: HashMap::new(),
105 unavailable: HashMap::new(),
106 web_capabilities: None,
107 subagent_spawner: None,
108 }
109 }
110
111 #[must_use]
112 pub fn web_capabilities(&self) -> Option<&web::WebCapabilities> {
113 self.web_capabilities.as_deref()
114 }
115
116 #[must_use]
117 pub fn subagent_spawner(&self) -> Option<&Arc<subagent::SubagentSpawner>> {
118 self.subagent_spawner.as_ref()
119 }
120
121 pub fn register(&mut self, tool: Arc<dyn ToolExecutor>) {
122 self.entries.insert(tool.name(), tool);
123 }
124
125 pub fn note_unavailable(&mut self, tool: &'static str, reason: impl Into<String>) {
130 self.unavailable.insert(tool, reason.into());
131 }
132
133 #[must_use]
134 pub fn unavailable_reason(&self, name: &str) -> Option<&str> {
135 self.unavailable.get(name).map(String::as_str)
136 }
137
138 #[must_use]
144 pub fn unknown_tool_outcome(&self, tool_key: &str, called_name: &str) -> ToolOutcome {
145 self.unavailable.get(tool_key).map_or_else(
146 || ToolOutcome::error(format!("unknown tool: {called_name}"), 0.0),
147 |reason| ToolOutcome::error(format!("{called_name} is not available: {reason}"), 0.0),
148 )
149 }
150
151 #[must_use]
152 pub fn get(&self, name: &str) -> Option<Arc<dyn ToolExecutor>> {
153 self.entries.get(name).cloned()
154 }
155
156 #[must_use]
157 pub fn len(&self) -> usize {
158 self.entries.len()
159 }
160
161 #[must_use]
162 pub fn is_empty(&self) -> bool {
163 self.entries.is_empty()
164 }
165
166 pub fn names(&self) -> impl Iterator<Item = &'static str> + '_ {
167 self.entries.keys().copied()
168 }
169
170 #[must_use]
176 pub fn describe_all(&self) -> Vec<ToolDefinition> {
177 self.entries
178 .values()
179 .filter(|t| !t.is_internal())
180 .map(|t| t.schema())
181 .collect()
182 }
183}
184
185impl Default for ToolRegistry {
186 fn default() -> Self {
187 let mut r = Self::new();
188 r.register(Arc::new(filesystem::ReadFileTool));
189 r.register(Arc::new(filesystem::WriteFileTool));
190 r.register(Arc::new(apply_patch::ApplyPatchTool));
191 r.register(Arc::new(filesystem::DeleteFileTool));
192 r.register(Arc::new(filesystem::CreateDirectoryTool));
193 r.register(Arc::new(exec::ExecuteCommandTool));
194 r.register(Arc::new(memory::MemoryTool));
195 r.register(Arc::new(ask_user_question::AskUserQuestionTool));
196 r.register(Arc::new(enter_plan_mode::EnterPlanModeTool));
199 r.register(Arc::new(exit_plan_mode::ExitPlanModeTool));
200 r.register(Arc::new(tasks::TaskCreateTool));
201 r.register(Arc::new(tasks::TaskUpdateTool));
202 r.register(Arc::new(tasks::TaskListTool));
203 r.register(Arc::new(mcp::McpToolProxy));
207 r
208 }
209}
210
211#[derive(Debug, Clone, Copy, PartialEq, Eq)]
217pub enum TuiMode {
218 Interactive,
219 Headless,
220}
221
222impl ToolRegistry {
223 fn register_computer_use_tools(&mut self, backend: computer_use::Backend) {
229 let driver = Arc::new(computer_use::ComputerUseDriver::new(backend));
230 self.register(Arc::new(computer_use::ScreenshotTool::new(driver.clone())));
231 if backend.supports_input_injection() {
232 self.register(Arc::new(computer_use::ClickTool::new(driver.clone())));
233 self.register(Arc::new(computer_use::TypeTextTool::new(driver.clone())));
234 self.register(Arc::new(computer_use::PressKeyTool::new(driver.clone())));
235 self.register(Arc::new(computer_use::ScrollTool::new(driver.clone())));
236 self.register(Arc::new(computer_use::MouseMoveTool::new(driver.clone())));
237 }
238 if backend.supports_window_listing() {
239 self.register(Arc::new(computer_use::ListWindowsTool::new(driver.clone())));
240 }
241 }
242
243 pub fn build(
259 config: &mermaid_domain::Config,
260 mode: TuiMode,
261 providers: Arc<crate::providers::ProviderFactory>,
262 ) -> Arc<Self> {
263 let mut r = Self::new();
264 let web_capabilities = Arc::new(web::WebCapabilities::resolve(&config.web));
265 r.register(Arc::new(filesystem::ReadFileTool));
266 r.register(Arc::new(filesystem::WriteFileTool));
267 r.register(Arc::new(apply_patch::ApplyPatchTool));
268 r.register(Arc::new(filesystem::DeleteFileTool));
269 r.register(Arc::new(filesystem::CreateDirectoryTool));
270 r.register(Arc::new(exec::ExecuteCommandTool));
271 r.register(Arc::new(memory::MemoryTool));
272 r.register(Arc::new(ask_user_question::AskUserQuestionTool));
273 r.register(Arc::new(enter_plan_mode::EnterPlanModeTool));
274 r.register(Arc::new(exit_plan_mode::ExitPlanModeTool));
275 r.register(Arc::new(tasks::TaskCreateTool));
276 r.register(Arc::new(tasks::TaskUpdateTool));
277 r.register(Arc::new(tasks::TaskListTool));
278 r.register(Arc::new(mcp::McpToolProxy));
279
280 if config.safety.network == mermaid_domain::NetworkPolicy::Allow {
286 match web_capabilities.fetch_tool() {
287 Some(tool) => r.register(Arc::new(tool)),
288 None => r.note_unavailable(
289 "web_fetch",
290 web_capabilities.fetch.absence_reason("web_fetch"),
291 ),
292 }
293 match web_capabilities.search_tool() {
294 Some(tool) => r.register(Arc::new(tool)),
295 None => r.note_unavailable(
296 "web_search",
297 web_capabilities.search.absence_reason("web_search"),
298 ),
299 }
300 } else {
301 for tool in ["web_fetch", "web_search"] {
302 r.note_unavailable(
303 tool,
304 format!(
305 "{tool} is disabled: network access is off \
306 (safety.network = \"deny\" / --no-network)"
307 ),
308 );
309 }
310 }
311
312 if mode == TuiMode::Interactive {
317 let backend = computer_use::probe();
318 if backend.is_usable() {
319 r.register_computer_use_tools(backend);
320 }
321 }
322
323 let spawner = Arc::new(subagent::SubagentSpawner::new(
328 providers,
329 Arc::clone(&web_capabilities),
330 ));
331 r.register(Arc::new(subagent::SubagentTool::new(spawner.clone())));
332 r.subagent_spawner = Some(spawner);
333 r.web_capabilities = Some(web_capabilities);
334
335 Arc::new(r)
336 }
337}
338
339#[cfg(test)]
340mod tests {
341 use super::*;
342
343 #[test]
344 fn default_registry_has_builtin_tools() {
345 let r = ToolRegistry::default();
346 for name in &[
347 "read_file",
348 "write_file",
349 "apply_patch",
350 "delete_file",
351 "create_directory",
352 "execute_command",
353 "memory",
354 ] {
355 assert!(r.get(name).is_some(), "missing: {name}");
356 }
357 assert!(r.get("not_a_tool").is_none());
358 assert!(r.len() >= 6);
359 }
360
361 #[test]
362 fn computer_use_registration_is_selective_per_backend() {
363 use computer_use::Backend;
364 let reg = |b: Backend| {
365 let mut r = ToolRegistry::new();
366 r.register_computer_use_tools(b);
367 r
368 };
369
370 let mac = reg(Backend::MacOS);
373 assert!(mac.get("screenshot").is_some());
374 for t in [
375 "click",
376 "type_text",
377 "press_key",
378 "scroll",
379 "mouse_move",
380 "list_windows",
381 ] {
382 assert!(mac.get(t).is_none(), "macOS must not advertise {t}");
383 }
384
385 let way = reg(Backend::Wayland);
387 assert!(way.get("click").is_some());
388 assert!(way.get("list_windows").is_none());
389
390 let x11 = reg(Backend::X11);
392 for t in [
393 "screenshot",
394 "click",
395 "type_text",
396 "press_key",
397 "scroll",
398 "mouse_move",
399 "list_windows",
400 ] {
401 assert!(x11.get(t).is_some(), "X11 missing {t}");
402 }
403 }
404
405 #[test]
406 fn describe_all_returns_one_per_user_facing_tool() {
407 let r = ToolRegistry::default();
408 let schemas = r.describe_all();
409 let visible = r
412 .names()
413 .filter(|n| r.get(n).map(|t| !t.is_internal()).unwrap_or(false))
414 .count();
415 assert_eq!(schemas.len(), visible);
416 for schema in &schemas {
417 assert!(
418 r.get(&schema.name).is_some(),
419 "schema for unknown tool: {}",
420 schema.name
421 );
422 }
423 }
424
425 #[test]
426 fn mcp_proxy_is_registered_but_internal() {
427 let r = ToolRegistry::default();
428 let proxy = r.get("mcp_proxy").expect("mcp_proxy registered");
429 assert!(proxy.is_internal());
430 assert!(!r.describe_all().iter().any(|s| s.name == "mcp_proxy"));
431 }
432
433 #[test]
434 fn schema_name_matches_executor_name() {
435 let r = ToolRegistry::default();
436 for name in r.names() {
437 let tool = r.get(name).unwrap();
438 assert_eq!(tool.name(), tool.schema().name.as_str());
439 }
440 }
441
442 static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
447
448 #[test]
449 fn build_registers_zero_config_web_tools_without_key() {
450 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
454 let prior = std::env::var("OLLAMA_API_KEY").ok();
455 unsafe {
456 std::env::remove_var("OLLAMA_API_KEY");
457 }
458 let cfg = mermaid_domain::Config::default();
459 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
460 let r = ToolRegistry::build(&cfg, TuiMode::Headless, providers);
461 assert!(
462 r.get("web_fetch").is_some(),
463 "native web_fetch registers without a key"
464 );
465 assert_eq!(
466 r.get("web_search").is_some(),
467 crate::searxng::managed_backend_viability().is_ok(),
468 "auto web_search registers only when managed SearXNG is viable"
469 );
470 assert!(r.get("read_file").is_some());
471 assert!(r.get("execute_command").is_some());
472 let web = r
473 .web_capabilities()
474 .expect("config-aware registries retain the resolved web status");
475 assert_eq!(web.fetch.backend, "native");
476 assert_eq!(web.search.backend, "managed_searxng");
477 unsafe {
478 if let Some(v) = prior {
479 std::env::set_var("OLLAMA_API_KEY", v);
480 }
481 }
482 }
483
484 #[test]
485 fn build_registers_ollama_web_search_with_key() {
486 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
489 let prior = std::env::var("OLLAMA_API_KEY").ok();
490 unsafe {
491 std::env::set_var("OLLAMA_API_KEY", "test-key-build");
492 }
493 let mut cfg = mermaid_domain::Config::default();
494 cfg.web.search_backend = mermaid_domain::SearchBackend::Ollama;
495 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
496 let r = ToolRegistry::build(&cfg, TuiMode::Interactive, providers);
497 assert!(r.get("web_search").is_some(), "web_search registered");
498 assert!(r.get("web_fetch").is_some(), "web_fetch registered");
499 unsafe {
500 match prior {
501 Some(v) => std::env::set_var("OLLAMA_API_KEY", v),
502 None => std::env::remove_var("OLLAMA_API_KEY"),
503 }
504 }
505 }
506
507 #[test]
508 fn auto_search_never_selects_cloud_just_because_a_key_exists() {
509 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
510 let prior = std::env::var("OLLAMA_API_KEY").ok();
511 unsafe {
512 std::env::set_var("OLLAMA_API_KEY", "test-key-must-not-route");
513 }
514 let cfg = mermaid_domain::Config::default();
515 let capabilities = web::WebCapabilities::resolve(&cfg.web);
516 assert_eq!(capabilities.search.backend, "managed_searxng");
517 assert_eq!(
518 capabilities.search.available,
519 crate::searxng::managed_backend_viability().is_ok()
520 );
521 unsafe {
522 match prior {
523 Some(value) => std::env::set_var("OLLAMA_API_KEY", value),
524 None => std::env::remove_var("OLLAMA_API_KEY"),
525 }
526 }
527 }
528
529 #[test]
530 fn auto_search_fallback_engages_only_when_opted_in_with_a_key() {
531 let _guard = ENV_LOCK
536 .lock()
537 .unwrap_or_else(std::sync::PoisonError::into_inner);
538 let prior = std::env::var("OLLAMA_API_KEY").ok();
539 unsafe {
540 std::env::set_var("OLLAMA_API_KEY", "test-key-fallback");
541 }
542 let mut cfg = mermaid_domain::Config::default();
543 cfg.web.allow_ollama_search_fallback = true;
544 let capabilities = web::WebCapabilities::resolve(&cfg.web);
545 if crate::searxng::managed_backend_viability().is_ok() {
546 assert_eq!(capabilities.search.backend, "managed_searxng");
547 assert!(capabilities.search.available);
548 } else {
549 assert_eq!(capabilities.search.backend, "ollama_cloud");
550 assert!(capabilities.search.available);
551 assert_eq!(capabilities.search.egress, web::Egress::OffMachine);
552 assert!(
553 capabilities.search_tool().is_some(),
554 "the fallback must produce a registrable tool"
555 );
556 }
557 unsafe {
558 match prior {
559 Some(v) => std::env::set_var("OLLAMA_API_KEY", v),
560 None => std::env::remove_var("OLLAMA_API_KEY"),
561 }
562 }
563 }
564
565 #[test]
566 fn auto_search_fallback_without_a_key_reports_the_whole_chain() {
567 let _guard = ENV_LOCK
568 .lock()
569 .unwrap_or_else(std::sync::PoisonError::into_inner);
570 let prior = std::env::var("OLLAMA_API_KEY").ok();
571 unsafe {
572 std::env::remove_var("OLLAMA_API_KEY");
573 }
574 let mut cfg = mermaid_domain::Config::default();
575 cfg.web.allow_ollama_search_fallback = true;
576 let capabilities = web::WebCapabilities::resolve(&cfg.web);
577 if crate::searxng::managed_backend_viability().is_err() {
578 assert!(!capabilities.search.available);
579 assert_eq!(capabilities.search.backend, "ollama_cloud");
580 let reason = capabilities.search.reason.as_deref().unwrap_or_default();
581 assert!(reason.contains("managed bundle"), "{reason}");
582 assert!(reason.contains("OLLAMA_API_KEY"), "{reason}");
583 }
584 unsafe {
585 if let Some(v) = prior {
586 std::env::set_var("OLLAMA_API_KEY", v);
587 }
588 }
589 }
590
591 #[test]
592 fn build_registers_searxng_web_search_without_key() {
593 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
596 let prior = std::env::var("OLLAMA_API_KEY").ok();
597 unsafe {
598 std::env::remove_var("OLLAMA_API_KEY");
599 }
600 let mut cfg = mermaid_domain::Config::default();
601 cfg.web.search_backend = mermaid_domain::SearchBackend::Searxng;
602 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
603 let r = ToolRegistry::build(&cfg, TuiMode::Headless, providers);
604 assert!(
605 r.get("web_search").is_some(),
606 "searxng web_search registers without a key"
607 );
608 assert!(
609 r.get("web_fetch").is_some(),
610 "native web_fetch still present"
611 );
612 unsafe {
613 if let Some(v) = prior {
614 std::env::set_var("OLLAMA_API_KEY", v);
615 }
616 }
617 }
618
619 #[test]
620 fn network_deny_omits_all_web_capabilities() {
621 let mut cfg = mermaid_domain::Config::default();
622 cfg.safety.network = mermaid_domain::NetworkPolicy::Deny;
623 cfg.web.search_backend = mermaid_domain::SearchBackend::Searxng;
624 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
625 let registry = ToolRegistry::build(&cfg, TuiMode::Headless, providers);
626 assert!(registry.get("web_fetch").is_none());
627 assert!(registry.get("web_search").is_none());
628 assert!(registry.get("read_file").is_some());
629 for tool in ["web_fetch", "web_search"] {
632 let outcome = registry.unknown_tool_outcome(tool, tool);
633 let msg = outcome.error_message().unwrap_or_default();
634 assert!(msg.contains("safety.network"), "{tool}: {msg}");
635 }
636 assert!(registry.unavailable_reason("read_file").is_none());
639 let outcome = registry.unknown_tool_outcome("frobnicate", "frobnicate");
640 assert_eq!(
641 outcome.error_message().unwrap_or_default(),
642 "unknown tool: frobnicate"
643 );
644 }
645
646 #[test]
647 fn unavailable_search_backend_reason_reaches_the_model() {
648 let _guard = ENV_LOCK
654 .lock()
655 .unwrap_or_else(std::sync::PoisonError::into_inner);
656 let prior = std::env::var("OLLAMA_API_KEY").ok();
657 unsafe {
658 std::env::remove_var("OLLAMA_API_KEY");
659 }
660 let cfg = mermaid_domain::Config::default();
661 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
662 let registry = ToolRegistry::build(&cfg, TuiMode::Headless, providers);
663 match crate::searxng::managed_backend_viability() {
664 Ok(_) => {
665 assert!(registry.get("web_search").is_some());
666 assert!(registry.unavailable_reason("web_search").is_none());
667 },
668 Err(viability_reason) => {
669 assert!(registry.get("web_search").is_none());
670 let reason = registry
671 .unavailable_reason("web_search")
672 .expect("absence reason recorded");
673 assert!(
674 reason.contains(&viability_reason),
675 "must carry the real cause: {reason}"
676 );
677 assert!(
678 reason.contains("search_backend"),
679 "must carry the remediation: {reason}"
680 );
681 },
682 }
683 unsafe {
684 if let Some(v) = prior {
685 std::env::set_var("OLLAMA_API_KEY", v);
686 }
687 }
688 }
689}