Skip to main content

mermaid_cli/providers/tool/
mod.rs

1//! Tool executors — one type per tool the model can call.
2//!
3//! The trait is small: `execute(args, ctx) -> ToolOutcome` for
4//! dispatch, plus `schema() -> ToolDefinition` for advertising the
5//! tool to the model. Everything else (cancellation, progress,
6//! identity, workdir) rides inside `ExecContext`.
7//!
8//! Adding a tool:
9//!   1. New file under `src/providers/tool/`.
10//!   2. Impl `ToolExecutor` for a unit struct — both `execute` and
11//!      `schema`.
12//!   3. Register it in `ToolRegistry::build()` — the ONE factory production
13//!      uses. There is no second registry to forget.
14//!
15//! Because `schema()` lives on the same trait as `execute()`, the
16//! name + JSON schema the model sees cannot drift from the handler
17//! that runs when the model calls it. Single source of truth.
18
19pub 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/// Implemented by every tool that the model can call. All tools are
45/// `Send + Sync` — they run across tokio `select!` branches inside
46/// the effect runner.
47#[async_trait]
48pub trait ToolExecutor: Send + Sync {
49    /// Canonical name the model uses to call this tool. Matches
50    /// `schema().name` exactly.
51    fn name(&self) -> &'static str;
52
53    /// JSON-schema description the model sees in the outgoing
54    /// request. Adapters translate this into provider-native shape
55    /// (Anthropic's `type: "custom"`, Gemini's `function_declarations`,
56    /// OpenAI's flat `tools`, Ollama's function calling). The same
57    /// `ToolDefinition` feeds all four.
58    fn schema(&self) -> ToolDefinition;
59
60    /// True for tools that exist for internal dispatch only and
61    /// should NOT be advertised to the model (e.g. the MCP proxy
62    /// router, which fronts every `mcp__server__tool` call — the
63    /// individual MCP tools are advertised separately from
64    /// `state.mcp.servers`). Default `false`.
65    fn is_internal(&self) -> bool {
66        false
67    }
68
69    /// Run the tool. The returned `ToolOutcome` is passed verbatim
70    /// into `Msg::ToolFinished` — there's no error-to-outcome
71    /// conversion happening outside this function.
72    async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome;
73}
74
75/// Registry of dispatchable tools. Single source of truth for what
76/// the model sees AND what handles a call when the model issues it.
77/// Built once at startup; read-only after that.
78pub struct ToolRegistry {
79    entries: HashMap<&'static str, Arc<dyn ToolExecutor>>,
80    /// Teaching errors for tools that were CONSIDERED at build time and
81    /// deliberately not registered (backend unavailable, network denied,
82    /// filtered out of a child registry). Dispatch returns the reason when
83    /// the model calls one, so the model learns why the tool is absent and
84    /// what to do about it — a bare "unknown tool" reads as a schema bug
85    /// and was observed driving models to fabricate results instead of
86    /// reporting the gap.
87    unavailable: HashMap<&'static str, String>,
88    /// Startup-resolved web routing/viability shared by the parent registry,
89    /// its provider-facing definitions, UI diagnostics, and every child
90    /// registry. Keeping the backend clients here prevents credentials or
91    /// environment changes from silently re-resolving a different route.
92    web_capabilities: Option<Arc<web::WebCapabilities>>,
93    /// Direct handle to the subagent spawner (also reachable through the
94    /// `agent` tool entry, but `dyn ToolExecutor` can't be downcast). The
95    /// effect layer uses it to service `Cmd::KillBackgroundAgent`. `None`
96    /// in registries built without a spawner (child registries, tests).
97    subagent_spawner: Option<Arc<subagent::SubagentSpawner>>,
98}
99
100/// An empty registry. Every session's registry comes from [`ToolRegistry::build`];
101/// this exists for callers that assemble one by hand (stubs, the subagent
102/// child registry) and for the `Default` convention, never as a second list
103/// of built-in tools.
104impl 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    /// Record why a tool that could exist in this registry deliberately does
136    /// not. The reason is model-facing: it must name the cause and the
137    /// remediation, because it is returned verbatim when the model calls the
138    /// absent tool.
139    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    /// The outcome for a call this registry cannot dispatch: the recorded
149    /// teaching error when the tool was deliberately omitted, else the plain
150    /// unknown-tool error. `called_name` is the name the model used — for
151    /// MCP calls it differs from the internal `mcp_proxy` routing key, and
152    /// the model should see the name it actually wrote.
153    #[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    /// Emit every user-facing tool's schema, for inclusion in an
181    /// outgoing `ChatRequest.tools`. Effect runner calls this before
182    /// dispatching `Cmd::CallModel` so the model always sees the
183    /// same list the runner can dispatch. Internal routers (the MCP
184    /// proxy) are filtered out.
185    #[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    /// Config-aware factory. Always registers filesystem + exec +
197    /// the MCP proxy + the subagent tool. Conditionally registers:
198    ///
199    ///   - Viable `web_fetch` and `web_search` capabilities resolved once by
200    ///     `web::WebCapabilities`. Global network denial omits both.
201    ///
202    /// `providers` is the shared `ProviderFactory` that the effect
203    /// runner also holds; the `SubagentSpawner` needs it so child
204    /// reducer loops hit the same provider cache.
205    ///
206    /// Returns `Arc<Self>` so the effect runner can share a handle
207    /// across turns without cloning the underlying `HashMap`.
208    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        // `safety.network = "deny"` is a global egress kill-switch, not only
231        // a shell sandbox flag. Omit web capabilities entirely so adapters and
232        // subagents cannot advertise or execute them — and record why, so a
233        // model that calls one anyway is taught the cause instead of shown a
234        // bare "unknown tool".
235        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        // Subagents: always register. Depth + breadth caps live on
263        // `SubagentSpawner`; the tool itself is harmless when nobody
264        // calls it. Headless runs do register the agent — a CI prompt
265        // may still delegate to subagents for batched work.
266        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    /// The production registry as a headless session sees it: `build()` with
283    /// a default config and a stub provider factory. Tests go through the same
284    /// factory production does, so a tool registered anywhere else does not
285    /// exist as far as they are concerned.
286    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        // mcp_proxy is registered but internal — filtered out of
316        // describe_all. So len() includes it but schemas don't.
317        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    /// Serialization guard for tests that mutate the `OLLAMA_API_KEY`
349    /// env var. Cargo's default test harness runs tests in parallel
350    /// threads inside one process; without this mutex two env-touching
351    /// tests would race and occasionally flip each other's expectations.
352    static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
353
354    #[test]
355    fn build_registers_zero_config_web_tools_without_key() {
356        // Both web tools register with no OLLAMA_API_KEY: web_fetch is native,
357        // and web_search defaults to `auto`, which falls back to a managed local
358        // SearXNG (the process starts lazily at call time, not here).
359        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        // Cloud routing is explicit: a key plus an explicit Ollama backend
393        // registers search without changing the native fetch default.
394        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        // The opt-in flips exactly one case: auto + no viable bundle + key.
438        // A viable sovereign default always wins over the fallback, and the
439        // not-opted-in side is pinned by
440        // `auto_search_never_selects_cloud_just_because_a_key_exists`.
441        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        // The SearXNG search backend registers regardless of OLLAMA_API_KEY —
500        // reachability is a call-time concern, not a registration one.
501        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        // Calling an omitted web tool is answered with the cause, not a bare
536        // "unknown tool" — that reply was observed driving models to guess.
537        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        // Matched pair: a registered tool carries no absence note, and a name
543        // never considered stays a plain unknown tool.
544        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        // The Windows field logs: `search_backend = "auto"` with no viable
555        // managed bundle registered nothing, and calling `web_search` got
556        // "unknown tool". The registry must instead carry the viability
557        // reason plus the remediation. On hosts where the managed bundle IS
558        // viable, the tool registers and no note exists — both sides pinned.
559        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}