mermaid_cli/providers/tool/
mod.rs1pub mod apply_patch;
20pub mod ask_user_question;
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 Default for ToolRegistry {
105 fn default() -> Self {
106 Self::new()
107 }
108}
109
110impl ToolRegistry {
111 #[must_use]
112 pub fn new() -> Self {
113 Self {
114 entries: HashMap::new(),
115 unavailable: HashMap::new(),
116 web_capabilities: None,
117 subagent_spawner: None,
118 }
119 }
120
121 #[must_use]
122 pub fn web_capabilities(&self) -> Option<&web::WebCapabilities> {
123 self.web_capabilities.as_deref()
124 }
125
126 #[must_use]
127 pub fn subagent_spawner(&self) -> Option<&Arc<subagent::SubagentSpawner>> {
128 self.subagent_spawner.as_ref()
129 }
130
131 pub fn register(&mut self, tool: Arc<dyn ToolExecutor>) {
132 self.entries.insert(tool.name(), tool);
133 }
134
135 pub fn note_unavailable(&mut self, tool: &'static str, reason: impl Into<String>) {
140 self.unavailable.insert(tool, reason.into());
141 }
142
143 #[must_use]
144 pub fn unavailable_reason(&self, name: &str) -> Option<&str> {
145 self.unavailable.get(name).map(String::as_str)
146 }
147
148 #[must_use]
154 pub fn unknown_tool_outcome(&self, tool_key: &str, called_name: &str) -> ToolOutcome {
155 self.unavailable.get(tool_key).map_or_else(
156 || ToolOutcome::error(format!("unknown tool: {called_name}"), 0.0),
157 |reason| ToolOutcome::error(format!("{called_name} is not available: {reason}"), 0.0),
158 )
159 }
160
161 #[must_use]
162 pub fn get(&self, name: &str) -> Option<Arc<dyn ToolExecutor>> {
163 self.entries.get(name).cloned()
164 }
165
166 #[must_use]
167 pub fn len(&self) -> usize {
168 self.entries.len()
169 }
170
171 #[must_use]
172 pub fn is_empty(&self) -> bool {
173 self.entries.is_empty()
174 }
175
176 pub fn names(&self) -> impl Iterator<Item = &'static str> + '_ {
177 self.entries.keys().copied()
178 }
179
180 #[must_use]
186 pub fn describe_all(&self) -> Vec<ToolDefinition> {
187 self.entries
188 .values()
189 .filter(|t| !t.is_internal())
190 .map(|t| t.schema())
191 .collect()
192 }
193}
194
195impl ToolRegistry {
196 pub fn build(
209 config: &mermaid_domain::Config,
210 providers: Arc<crate::providers::ProviderFactory>,
211 ) -> Arc<Self> {
212 let mut r = Self::new();
213 let web_capabilities = Arc::new(web::WebCapabilities::resolve(&config.web));
214 r.register(Arc::new(filesystem::ReadFileTool));
215 r.register(Arc::new(filesystem::WriteFileTool));
216 r.register(Arc::new(filesystem::EditFileTool));
217 r.register(Arc::new(apply_patch::ApplyPatchTool));
218 r.register(Arc::new(filesystem::DeleteFileTool));
219 r.register(Arc::new(filesystem::CreateDirectoryTool));
220 r.register(Arc::new(exec::ExecuteCommandTool));
221 r.register(Arc::new(memory::MemoryTool));
222 r.register(Arc::new(ask_user_question::AskUserQuestionTool));
223 r.register(Arc::new(enter_plan_mode::EnterPlanModeTool));
224 r.register(Arc::new(exit_plan_mode::ExitPlanModeTool));
225 r.register(Arc::new(tasks::TaskCreateTool));
226 r.register(Arc::new(tasks::TaskUpdateTool));
227 r.register(Arc::new(tasks::TaskListTool));
228 r.register(Arc::new(mcp::McpToolProxy));
229
230 if config.safety.network == mermaid_domain::NetworkPolicy::Allow {
236 match web_capabilities.fetch_tool() {
237 Some(tool) => r.register(Arc::new(tool)),
238 None => r.note_unavailable(
239 "web_fetch",
240 web_capabilities.fetch.absence_reason("web_fetch"),
241 ),
242 }
243 match web_capabilities.search_tool() {
244 Some(tool) => r.register(Arc::new(tool)),
245 None => r.note_unavailable(
246 "web_search",
247 web_capabilities.search.absence_reason("web_search"),
248 ),
249 }
250 } else {
251 for tool in ["web_fetch", "web_search"] {
252 r.note_unavailable(
253 tool,
254 format!(
255 "{tool} is disabled: network access is off \
256 (safety.network = \"deny\" / --no-network)"
257 ),
258 );
259 }
260 }
261
262 let spawner = Arc::new(subagent::SubagentSpawner::new(
267 providers,
268 Arc::clone(&web_capabilities),
269 ));
270 r.register(Arc::new(subagent::SubagentTool::new(spawner.clone())));
271 r.subagent_spawner = Some(spawner);
272 r.web_capabilities = Some(web_capabilities);
273
274 Arc::new(r)
275 }
276}
277
278#[cfg(test)]
279mod tests {
280 use super::*;
281
282 fn headless_registry() -> Arc<ToolRegistry> {
287 let cfg = mermaid_domain::Config::default();
288 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
289 ToolRegistry::build(&cfg, providers)
290 }
291
292 #[test]
293 fn default_registry_has_builtin_tools() {
294 let r = headless_registry();
295 for name in &[
296 "read_file",
297 "write_file",
298 "edit_file",
299 "apply_patch",
300 "delete_file",
301 "create_directory",
302 "execute_command",
303 "memory",
304 ] {
305 assert!(r.get(name).is_some(), "missing: {name}");
306 }
307 assert!(r.get("not_a_tool").is_none());
308 assert!(r.len() >= 6);
309 }
310
311 #[test]
312 fn describe_all_returns_one_per_user_facing_tool() {
313 let r = headless_registry();
314 let schemas = r.describe_all();
315 let visible = r
318 .names()
319 .filter(|n| r.get(n).map(|t| !t.is_internal()).unwrap_or(false))
320 .count();
321 assert_eq!(schemas.len(), visible);
322 for schema in &schemas {
323 assert!(
324 r.get(&schema.name).is_some(),
325 "schema for unknown tool: {}",
326 schema.name
327 );
328 }
329 }
330
331 #[test]
332 fn mcp_proxy_is_registered_but_internal() {
333 let r = headless_registry();
334 let proxy = r.get("mcp_proxy").expect("mcp_proxy registered");
335 assert!(proxy.is_internal());
336 assert!(!r.describe_all().iter().any(|s| s.name == "mcp_proxy"));
337 }
338
339 #[test]
340 fn schema_name_matches_executor_name() {
341 let r = headless_registry();
342 for name in r.names() {
343 let tool = r.get(name).unwrap();
344 assert_eq!(tool.name(), tool.schema().name.as_str());
345 }
346 }
347
348 static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
353
354 #[test]
355 fn build_registers_zero_config_web_tools_without_key() {
356 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
360 let prior = std::env::var("OLLAMA_API_KEY").ok();
361 unsafe {
362 std::env::remove_var("OLLAMA_API_KEY");
363 }
364 let cfg = mermaid_domain::Config::default();
365 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
366 let r = ToolRegistry::build(&cfg, providers);
367 assert!(
368 r.get("web_fetch").is_some(),
369 "native web_fetch registers without a key"
370 );
371 assert_eq!(
372 r.get("web_search").is_some(),
373 crate::searxng::managed_backend_viability().is_ok(),
374 "auto web_search registers only when managed SearXNG is viable"
375 );
376 assert!(r.get("read_file").is_some());
377 assert!(r.get("execute_command").is_some());
378 let web = r
379 .web_capabilities()
380 .expect("config-aware registries retain the resolved web status");
381 assert_eq!(web.fetch.backend, "native");
382 assert_eq!(web.search.backend, "managed_searxng");
383 unsafe {
384 if let Some(v) = prior {
385 std::env::set_var("OLLAMA_API_KEY", v);
386 }
387 }
388 }
389
390 #[test]
391 fn build_registers_ollama_web_search_with_key() {
392 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
395 let prior = std::env::var("OLLAMA_API_KEY").ok();
396 unsafe {
397 std::env::set_var("OLLAMA_API_KEY", "test-key-build");
398 }
399 let mut cfg = mermaid_domain::Config::default();
400 cfg.web.search_backend = mermaid_domain::SearchBackend::Ollama;
401 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
402 let r = ToolRegistry::build(&cfg, providers);
403 assert!(r.get("web_search").is_some(), "web_search registered");
404 assert!(r.get("web_fetch").is_some(), "web_fetch registered");
405 unsafe {
406 match prior {
407 Some(v) => std::env::set_var("OLLAMA_API_KEY", v),
408 None => std::env::remove_var("OLLAMA_API_KEY"),
409 }
410 }
411 }
412
413 #[test]
414 fn auto_search_never_selects_cloud_just_because_a_key_exists() {
415 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
416 let prior = std::env::var("OLLAMA_API_KEY").ok();
417 unsafe {
418 std::env::set_var("OLLAMA_API_KEY", "test-key-must-not-route");
419 }
420 let cfg = mermaid_domain::Config::default();
421 let capabilities = web::WebCapabilities::resolve(&cfg.web);
422 assert_eq!(capabilities.search.backend, "managed_searxng");
423 assert_eq!(
424 capabilities.search.available,
425 crate::searxng::managed_backend_viability().is_ok()
426 );
427 unsafe {
428 match prior {
429 Some(value) => std::env::set_var("OLLAMA_API_KEY", value),
430 None => std::env::remove_var("OLLAMA_API_KEY"),
431 }
432 }
433 }
434
435 #[test]
436 fn auto_search_fallback_engages_only_when_opted_in_with_a_key() {
437 let _guard = ENV_LOCK
442 .lock()
443 .unwrap_or_else(std::sync::PoisonError::into_inner);
444 let prior = std::env::var("OLLAMA_API_KEY").ok();
445 unsafe {
446 std::env::set_var("OLLAMA_API_KEY", "test-key-fallback");
447 }
448 let mut cfg = mermaid_domain::Config::default();
449 cfg.web.allow_ollama_search_fallback = true;
450 let capabilities = web::WebCapabilities::resolve(&cfg.web);
451 if crate::searxng::managed_backend_viability().is_ok() {
452 assert_eq!(capabilities.search.backend, "managed_searxng");
453 assert!(capabilities.search.available);
454 } else {
455 assert_eq!(capabilities.search.backend, "ollama_cloud");
456 assert!(capabilities.search.available);
457 assert_eq!(capabilities.search.egress, web::Egress::OffMachine);
458 assert!(
459 capabilities.search_tool().is_some(),
460 "the fallback must produce a registrable tool"
461 );
462 }
463 unsafe {
464 match prior {
465 Some(v) => std::env::set_var("OLLAMA_API_KEY", v),
466 None => std::env::remove_var("OLLAMA_API_KEY"),
467 }
468 }
469 }
470
471 #[test]
472 fn auto_search_fallback_without_a_key_reports_the_whole_chain() {
473 let _guard = ENV_LOCK
474 .lock()
475 .unwrap_or_else(std::sync::PoisonError::into_inner);
476 let prior = std::env::var("OLLAMA_API_KEY").ok();
477 unsafe {
478 std::env::remove_var("OLLAMA_API_KEY");
479 }
480 let mut cfg = mermaid_domain::Config::default();
481 cfg.web.allow_ollama_search_fallback = true;
482 let capabilities = web::WebCapabilities::resolve(&cfg.web);
483 if crate::searxng::managed_backend_viability().is_err() {
484 assert!(!capabilities.search.available);
485 assert_eq!(capabilities.search.backend, "ollama_cloud");
486 let reason = capabilities.search.reason.as_deref().unwrap_or_default();
487 assert!(reason.contains("managed bundle"), "{reason}");
488 assert!(reason.contains("OLLAMA_API_KEY"), "{reason}");
489 }
490 unsafe {
491 if let Some(v) = prior {
492 std::env::set_var("OLLAMA_API_KEY", v);
493 }
494 }
495 }
496
497 #[test]
498 fn build_registers_searxng_web_search_without_key() {
499 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
502 let prior = std::env::var("OLLAMA_API_KEY").ok();
503 unsafe {
504 std::env::remove_var("OLLAMA_API_KEY");
505 }
506 let mut cfg = mermaid_domain::Config::default();
507 cfg.web.search_backend = mermaid_domain::SearchBackend::Searxng;
508 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
509 let r = ToolRegistry::build(&cfg, providers);
510 assert!(
511 r.get("web_search").is_some(),
512 "searxng web_search registers without a key"
513 );
514 assert!(
515 r.get("web_fetch").is_some(),
516 "native web_fetch still present"
517 );
518 unsafe {
519 if let Some(v) = prior {
520 std::env::set_var("OLLAMA_API_KEY", v);
521 }
522 }
523 }
524
525 #[test]
526 fn network_deny_omits_all_web_capabilities() {
527 let mut cfg = mermaid_domain::Config::default();
528 cfg.safety.network = mermaid_domain::NetworkPolicy::Deny;
529 cfg.web.search_backend = mermaid_domain::SearchBackend::Searxng;
530 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
531 let registry = ToolRegistry::build(&cfg, providers);
532 assert!(registry.get("web_fetch").is_none());
533 assert!(registry.get("web_search").is_none());
534 assert!(registry.get("read_file").is_some());
535 for tool in ["web_fetch", "web_search"] {
538 let outcome = registry.unknown_tool_outcome(tool, tool);
539 let msg = outcome.error_message().unwrap_or_default();
540 assert!(msg.contains("safety.network"), "{tool}: {msg}");
541 }
542 assert!(registry.unavailable_reason("read_file").is_none());
545 let outcome = registry.unknown_tool_outcome("frobnicate", "frobnicate");
546 assert_eq!(
547 outcome.error_message().unwrap_or_default(),
548 "unknown tool: frobnicate"
549 );
550 }
551
552 #[test]
553 fn unavailable_search_backend_reason_reaches_the_model() {
554 let _guard = ENV_LOCK
560 .lock()
561 .unwrap_or_else(std::sync::PoisonError::into_inner);
562 let prior = std::env::var("OLLAMA_API_KEY").ok();
563 unsafe {
564 std::env::remove_var("OLLAMA_API_KEY");
565 }
566 let cfg = mermaid_domain::Config::default();
567 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
568 let registry = ToolRegistry::build(&cfg, providers);
569 match crate::searxng::managed_backend_viability() {
570 Ok(_) => {
571 assert!(registry.get("web_search").is_some());
572 assert!(registry.unavailable_reason("web_search").is_none());
573 },
574 Err(viability_reason) => {
575 assert!(registry.get("web_search").is_none());
576 let reason = registry
577 .unavailable_reason("web_search")
578 .expect("absence reason recorded");
579 assert!(
580 reason.contains(&viability_reason),
581 "must carry the real cause: {reason}"
582 );
583 assert!(
584 reason.contains("search_backend"),
585 "must carry the remediation: {reason}"
586 );
587 },
588 }
589 unsafe {
590 if let Some(v) = prior {
591 std::env::set_var("OLLAMA_API_KEY", v);
592 }
593 }
594 }
595}