mermaid_cli/providers/tool/
mod.rs1pub mod computer_use;
19pub mod exec;
20pub mod filesystem;
21pub mod mcp;
22pub mod memory;
23pub mod path_safety;
24pub mod policy_gate;
25pub mod subagent;
26pub mod web;
27pub mod web_client;
28
29use async_trait::async_trait;
30use std::collections::HashMap;
31use std::sync::Arc;
32
33use crate::domain::{ToolDefinition, ToolOutcome};
34
35use super::ctx::ExecContext;
36
37#[async_trait]
41pub trait ToolExecutor: Send + Sync {
42 fn name(&self) -> &'static str;
45
46 fn schema(&self) -> ToolDefinition;
52
53 fn is_internal(&self) -> bool {
59 false
60 }
61
62 async fn execute(&self, args: serde_json::Value, ctx: ExecContext) -> ToolOutcome;
66}
67
68pub struct ToolRegistry {
72 entries: HashMap<&'static str, Arc<dyn ToolExecutor>>,
73}
74
75impl ToolRegistry {
76 pub fn new() -> Self {
77 Self {
78 entries: HashMap::new(),
79 }
80 }
81
82 pub fn register(&mut self, tool: Arc<dyn ToolExecutor>) {
83 self.entries.insert(tool.name(), tool);
84 }
85
86 pub fn get(&self, name: &str) -> Option<Arc<dyn ToolExecutor>> {
87 self.entries.get(name).cloned()
88 }
89
90 pub fn len(&self) -> usize {
91 self.entries.len()
92 }
93
94 pub fn is_empty(&self) -> bool {
95 self.entries.is_empty()
96 }
97
98 pub fn names(&self) -> impl Iterator<Item = &'static str> + '_ {
99 self.entries.keys().copied()
100 }
101
102 pub fn describe_all(&self) -> Vec<ToolDefinition> {
108 self.entries
109 .values()
110 .filter(|t| !t.is_internal())
111 .map(|t| t.schema())
112 .collect()
113 }
114}
115
116impl Default for ToolRegistry {
117 fn default() -> Self {
118 let mut r = Self::new();
119 r.register(Arc::new(filesystem::ReadFileTool));
120 r.register(Arc::new(filesystem::WriteFileTool));
121 r.register(Arc::new(filesystem::EditFileTool));
122 r.register(Arc::new(filesystem::DeleteFileTool));
123 r.register(Arc::new(filesystem::CreateDirectoryTool));
124 r.register(Arc::new(exec::ExecuteCommandTool));
125 r.register(Arc::new(memory::MemoryTool));
126 r.register(Arc::new(mcp::McpToolProxy));
130 r
131 }
132}
133
134#[derive(Debug, Clone, Copy, PartialEq, Eq)]
140pub enum TuiMode {
141 Interactive,
142 Headless,
143}
144
145impl ToolRegistry {
146 fn register_computer_use_tools(&mut self, backend: computer_use::Backend) {
152 let driver = Arc::new(computer_use::ComputerUseDriver::new(backend));
153 self.register(Arc::new(computer_use::ScreenshotTool::new(driver.clone())));
154 if backend.supports_input_injection() {
155 self.register(Arc::new(computer_use::ClickTool::new(driver.clone())));
156 self.register(Arc::new(computer_use::TypeTextTool::new(driver.clone())));
157 self.register(Arc::new(computer_use::PressKeyTool::new(driver.clone())));
158 self.register(Arc::new(computer_use::ScrollTool::new(driver.clone())));
159 self.register(Arc::new(computer_use::MouseMoveTool::new(driver.clone())));
160 }
161 if backend.supports_window_listing() {
162 self.register(Arc::new(computer_use::ListWindowsTool::new(driver.clone())));
163 }
164 }
165
166 pub fn build(
184 _config: &crate::app::Config,
185 mode: TuiMode,
186 providers: Arc<crate::providers::ProviderFactory>,
187 ) -> Arc<Self> {
188 let mut r = Self::new();
189 r.register(Arc::new(filesystem::ReadFileTool));
190 r.register(Arc::new(filesystem::WriteFileTool));
191 r.register(Arc::new(filesystem::EditFileTool));
192 r.register(Arc::new(filesystem::DeleteFileTool));
193 r.register(Arc::new(filesystem::CreateDirectoryTool));
194 r.register(Arc::new(exec::ExecuteCommandTool));
195 r.register(Arc::new(memory::MemoryTool));
196 r.register(Arc::new(mcp::McpToolProxy));
197
198 if let Some(key) = crate::utils::resolve_api_key("OLLAMA_API_KEY", None) {
199 r.register(Arc::new(web::WebSearchTool::new(key.clone())));
200 r.register(Arc::new(web::WebFetchTool::new(key)));
201 }
202
203 if mode == TuiMode::Interactive {
208 let backend = computer_use::probe();
209 if backend.is_usable() {
210 r.register_computer_use_tools(backend);
211 }
212 }
213
214 let spawner = Arc::new(subagent::SubagentSpawner::new(providers));
219 r.register(Arc::new(subagent::SubagentTool::new(spawner)));
220
221 Arc::new(r)
222 }
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228
229 #[test]
230 fn default_registry_has_builtin_tools() {
231 let r = ToolRegistry::default();
232 for name in &[
233 "read_file",
234 "write_file",
235 "edit_file",
236 "delete_file",
237 "create_directory",
238 "execute_command",
239 "memory",
240 ] {
241 assert!(r.get(name).is_some(), "missing: {}", name);
242 }
243 assert!(r.get("not_a_tool").is_none());
244 assert!(r.len() >= 6);
245 }
246
247 #[test]
248 fn computer_use_registration_is_selective_per_backend() {
249 use computer_use::Backend;
250 let reg = |b: Backend| {
251 let mut r = ToolRegistry::new();
252 r.register_computer_use_tools(b);
253 r
254 };
255
256 let mac = reg(Backend::MacOS);
259 assert!(mac.get("screenshot").is_some());
260 for t in [
261 "click",
262 "type_text",
263 "press_key",
264 "scroll",
265 "mouse_move",
266 "list_windows",
267 ] {
268 assert!(mac.get(t).is_none(), "macOS must not advertise {t}");
269 }
270
271 let way = reg(Backend::Wayland);
273 assert!(way.get("click").is_some());
274 assert!(way.get("list_windows").is_none());
275
276 let x11 = reg(Backend::X11);
278 for t in [
279 "screenshot",
280 "click",
281 "type_text",
282 "press_key",
283 "scroll",
284 "mouse_move",
285 "list_windows",
286 ] {
287 assert!(x11.get(t).is_some(), "X11 missing {t}");
288 }
289 }
290
291 #[test]
292 fn describe_all_returns_one_per_user_facing_tool() {
293 let r = ToolRegistry::default();
294 let schemas = r.describe_all();
295 let visible = r
298 .names()
299 .filter(|n| r.get(n).map(|t| !t.is_internal()).unwrap_or(false))
300 .count();
301 assert_eq!(schemas.len(), visible);
302 for schema in &schemas {
303 assert!(
304 r.get(&schema.name).is_some(),
305 "schema for unknown tool: {}",
306 schema.name
307 );
308 }
309 }
310
311 #[test]
312 fn mcp_proxy_is_registered_but_internal() {
313 let r = ToolRegistry::default();
314 let proxy = r.get("mcp_proxy").expect("mcp_proxy registered");
315 assert!(proxy.is_internal());
316 assert!(!r.describe_all().iter().any(|s| s.name == "mcp_proxy"));
317 }
318
319 #[test]
320 fn schema_name_matches_executor_name() {
321 let r = ToolRegistry::default();
322 for name in r.names() {
323 let tool = r.get(name).unwrap();
324 assert_eq!(tool.name(), tool.schema().name.as_str());
325 }
326 }
327
328 static ENV_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
333
334 #[test]
335 fn build_registers_web_tools_when_key_present() {
336 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
337 let prior = std::env::var("OLLAMA_API_KEY").ok();
338 unsafe {
339 std::env::set_var("OLLAMA_API_KEY", "test-key-build");
340 }
341 let cfg = crate::app::Config::default();
342 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
343 let r = ToolRegistry::build(&cfg, TuiMode::Interactive, providers);
344 assert!(r.get("web_search").is_some(), "web_search registered");
345 assert!(r.get("web_fetch").is_some(), "web_fetch registered");
346 unsafe {
347 match prior {
348 Some(v) => std::env::set_var("OLLAMA_API_KEY", v),
349 None => std::env::remove_var("OLLAMA_API_KEY"),
350 }
351 }
352 }
353
354 #[test]
355 fn build_skips_web_tools_without_key() {
356 let _guard = ENV_LOCK.lock().unwrap_or_else(|e| e.into_inner());
357 let prior = std::env::var("OLLAMA_API_KEY").ok();
358 unsafe {
359 std::env::remove_var("OLLAMA_API_KEY");
360 }
361 let cfg = crate::app::Config::default();
362 let providers = Arc::new(crate::providers::ProviderFactory::new(cfg.clone()));
363 let r = ToolRegistry::build(&cfg, TuiMode::Headless, providers);
364 assert!(r.get("web_search").is_none(), "web_search skipped");
365 assert!(r.get("web_fetch").is_none(), "web_fetch skipped");
366 assert!(r.get("read_file").is_some());
367 assert!(r.get("execute_command").is_some());
368 unsafe {
369 if let Some(v) = prior {
370 std::env::set_var("OLLAMA_API_KEY", v);
371 }
372 }
373 }
374}