use std::collections::HashMap; use std::sync::{Arc, RwLock}; use crate::domain::tools::{Tool, ToolFunction}; use super::traits::Tool as ToolTrait; pub struct ToolRegistry { tools: RwLock>>, } impl ToolRegistry { pub fn new() -> Self { Self { tools: RwLock::new(HashMap::new()), } } pub fn register(&self, tool: T) { self.tools .write() .expect("ToolRegistry lock poisoned") .insert(tool.name().to_string(), Arc::new(tool)); } pub fn get(&self, name: &str) -> Option> { self.tools .read() .expect("ToolRegistry lock poisoned") .get(name) .cloned() } /// Get all registered tools. /// Used for concurrent tool execution when we need to look up tools by name. pub fn get_all(&self) -> Vec> { self.tools .read() .expect("ToolRegistry lock poisoned") .values() .cloned() .collect() } pub fn get_definitions(&self) -> Vec { self.tools .read() .expect("ToolRegistry lock poisoned") .values() .map(|tool| Tool { tool_type: "function".to_string(), function: ToolFunction { name: tool.name().to_string(), description: tool.description().to_string(), parameters: tool.parameters_schema(), }, }) .collect() } pub fn has_tools(&self) -> bool { !self.tools .read() .expect("ToolRegistry lock poisoned") .is_empty() } pub fn tool_names(&self) -> Vec { self.tools .read() .expect("ToolRegistry lock poisoned") .keys() .cloned() .collect() } /// 创建一个排除指定工具的新 registry 副本 pub fn without(&self, exclude: &[&str]) -> Self { let exclude_set: std::collections::HashSet<&str> = exclude.iter().copied().collect(); let tools = self.tools.read().expect("ToolRegistry lock poisoned"); let filtered: HashMap> = tools .iter() .filter(|(name, _)| !exclude_set.contains(name.as_str())) .map(|(k, v)| (k.clone(), v.clone())) .collect(); let new_registry = ToolRegistry::new(); *new_registry.tools.write().expect("ToolRegistry lock poisoned") = filtered; new_registry } /// 创建一个仅包含指定工具的新 registry 副本(白名单)。 /// include 中不存在于当前 registry 的名称会被静默跳过(取交集语义)。 pub fn only(&self, include: &[&str]) -> Self { let include_set: std::collections::HashSet<&str> = include.iter().copied().collect(); let tools = self.tools.read().expect("ToolRegistry lock poisoned"); let filtered: HashMap> = tools .iter() .filter(|(name, _)| include_set.contains(name.as_str())) .map(|(k, v)| (k.clone(), v.clone())) .collect(); let new_registry = ToolRegistry::new(); *new_registry.tools.write().expect("ToolRegistry lock poisoned") = filtered; new_registry } } impl Default for ToolRegistry { fn default() -> Self { Self::new() } } #[cfg(test)] mod tests { use super::*; use crate::tools::traits::ToolResult; use async_trait::async_trait; /// 仅用于测试的占位工具,按构造名注册 struct FakeTool { tool_name: String, } #[async_trait] impl ToolTrait for FakeTool { fn name(&self) -> &str { &self.tool_name } fn description(&self) -> &str { "fake" } fn parameters_schema(&self) -> serde_json::Value { serde_json::json!({}) } async fn execute(&self, _args: serde_json::Value) -> anyhow::Result { Ok(ToolResult { success: true, output: String::new(), error: None, }) } } fn registry_with(names: &[&str]) -> ToolRegistry { let reg = ToolRegistry::new(); for n in names { reg.register(FakeTool { tool_name: n.to_string(), }); } reg } fn sorted_names(reg: &ToolRegistry) -> Vec { let mut v = reg.tool_names(); v.sort(); v } #[test] fn only_keeps_listed_tools() { let reg = registry_with(&["read", "edit", "write", "bash"]); let filtered = reg.only(&["read", "bash"]); assert_eq!(sorted_names(&filtered), vec!["bash", "read"]); } #[test] fn only_silently_skips_missing_names() { let reg = registry_with(&["read", "edit"]); let filtered = reg.only(&["read", "nonexistent", "glob"]); assert_eq!(sorted_names(&filtered), vec!["read"]); } #[test] fn only_with_empty_include_returns_empty() { let reg = registry_with(&["read", "edit"]); let filtered = reg.only(&[]); assert!(filtered.tool_names().is_empty()); } #[test] fn only_does_not_mutate_source() { let reg = registry_with(&["read", "edit", "write"]); let _ = reg.only(&["read"]); // 源 registry 不受影响 let mut v = reg.tool_names(); v.sort(); assert_eq!(v, vec!["edit", "read", "write"]); } }