Skip to main content

pe_tools/
registry.rs

1//! Tool registry — stores and retrieves tools by name.
2//!
3//! The [`ToolRegistry`] is the central collection of tools available to an agent.
4//! It provides name-based lookup for [`ToolNode`](super::ToolNode) execution
5//! and schema listing for LLM prompt construction.
6
7use crate::tool::Tool;
8use pe_core::error::PeError;
9use pe_core::llm::ToolSchema;
10use std::collections::HashMap;
11use std::sync::Arc;
12
13/// Registry that stores tools by name and provides lookup + schema listing.
14///
15/// # Example
16///
17/// ```ignore
18/// let mut registry = ToolRegistry::new();
19/// registry.register(my_search_tool)?;
20/// registry.register(my_calculator_tool)?;
21///
22/// // Get schemas for LLM prompt
23/// let schemas = registry.schemas();
24///
25/// // Look up tool for execution
26/// let tool = registry.get("search").unwrap();
27/// let result = tool.execute(input).await?;
28/// ```
29pub struct ToolRegistry {
30    tools: HashMap<String, Arc<dyn Tool>>,
31}
32
33impl Clone for ToolRegistry {
34    fn clone(&self) -> Self {
35        Self {
36            tools: self.tools.clone(),
37        }
38    }
39}
40
41impl ToolRegistry {
42    /// Create an empty registry.
43    pub fn new() -> Self {
44        Self {
45            tools: HashMap::new(),
46        }
47    }
48
49    /// Register a tool with typed duplicate-name handling.
50    pub fn register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
51        self.try_register(tool)
52    }
53
54    /// Register a pre-wrapped `Arc<dyn Tool>` with typed duplicate-name handling.
55    pub fn register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
56        self.try_register_arc(tool)
57    }
58
59    /// Fallible registration path for dynamic/runtime-owned insertion.
60    pub fn try_register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
61        let name = tool.name().to_string();
62        if self.tools.contains_key(&name) {
63            return Err(PeError::ToolAlreadyRegistered { tool: name });
64        }
65        self.tools.insert(name, Arc::new(tool));
66        Ok(self)
67    }
68
69    /// Fallible registration path for pre-wrapped `Arc<dyn Tool>`.
70    pub fn try_register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
71        let name = tool.name().to_string();
72        if self.tools.contains_key(&name) {
73            return Err(PeError::ToolAlreadyRegistered { tool: name });
74        }
75        self.tools.insert(name, tool);
76        Ok(self)
77    }
78
79    /// Get a tool by name.
80    pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
81        self.tools.get(name).cloned()
82    }
83
84    /// Produce [`ToolSchema`] list for the LLM — all registered tools.
85    pub fn schemas(&self) -> Vec<ToolSchema> {
86        self.tools.values().map(|t| t.schema()).collect()
87    }
88
89    /// Create a subset registry containing only tools with the given names.
90    /// Used by `ToolPolicy::filter()` to enforce per-agent tool allowlists.
91    /// Tools not found in this registry are silently skipped.
92    pub fn filter(&self, names: &[&str]) -> ToolRegistry {
93        let tools = names
94            .iter()
95            .filter_map(|n| self.tools.get(*n).map(|t| (n.to_string(), t.clone())))
96            .collect();
97        ToolRegistry { tools }
98    }
99
100    /// All registered tool names.
101    pub fn names(&self) -> Vec<&str> {
102        self.tools.keys().map(String::as_str).collect()
103    }
104
105    /// Number of registered tools.
106    pub fn len(&self) -> usize {
107        self.tools.len()
108    }
109
110    /// Whether the registry is empty.
111    pub fn is_empty(&self) -> bool {
112        self.tools.is_empty()
113    }
114}
115
116impl Default for ToolRegistry {
117    fn default() -> Self {
118        Self::new()
119    }
120}
121
122impl std::fmt::Debug for ToolRegistry {
123    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
124        f.debug_struct("ToolRegistry")
125            .field("tools", &self.tools.keys().collect::<Vec<_>>())
126            .finish()
127    }
128}
129
130#[cfg(test)]
131mod tests {
132    use super::*;
133    use crate::tool::FunctionTool;
134
135    fn make_tool(name: &str) -> FunctionTool {
136        FunctionTool::new(
137            name,
138            format!("Tool {name}"),
139            serde_json::json!({"type": "object"}),
140            |_| Box::pin(async { Ok(serde_json::json!("ok")) }),
141        )
142    }
143
144    #[test]
145    fn register_and_get_round_trip() {
146        let mut reg = ToolRegistry::new();
147        reg.register(make_tool("search")).unwrap();
148
149        let tool = reg.get("search");
150        assert!(tool.is_some());
151        assert_eq!(tool.unwrap().name(), "search");
152    }
153
154    #[test]
155    fn get_nonexistent_returns_none() {
156        let reg = ToolRegistry::new();
157        assert!(reg.get("missing").is_none());
158    }
159
160    #[test]
161    fn duplicate_try_registration_returns_typed_error() {
162        let mut reg = ToolRegistry::new();
163        reg.try_register(make_tool("dup")).unwrap();
164        let err = reg.try_register(make_tool("dup")).unwrap_err();
165
166        assert!(matches!(err, PeError::ToolAlreadyRegistered { .. }));
167    }
168
169    #[test]
170    fn schemas_returns_all_tool_schemas() {
171        let mut reg = ToolRegistry::new();
172        reg.register(make_tool("alpha")).unwrap();
173        reg.register(make_tool("beta")).unwrap();
174
175        let schemas = reg.schemas();
176        assert_eq!(schemas.len(), 2);
177
178        let names: Vec<&str> = schemas.iter().map(|s| s.name.as_str()).collect();
179        assert!(names.contains(&"alpha"));
180        assert!(names.contains(&"beta"));
181    }
182
183    #[test]
184    fn filter_returns_subset() {
185        let mut reg = ToolRegistry::new();
186        reg.register(make_tool("a")).unwrap();
187        reg.register(make_tool("b")).unwrap();
188        reg.register(make_tool("c")).unwrap();
189
190        let filtered = reg.filter(&["a", "c"]);
191        assert_eq!(filtered.len(), 2);
192        assert!(filtered.get("a").is_some());
193        assert!(filtered.get("b").is_none());
194        assert!(filtered.get("c").is_some());
195    }
196
197    #[test]
198    fn filter_skips_missing_names() {
199        let mut reg = ToolRegistry::new();
200        reg.register(make_tool("x")).unwrap();
201
202        let filtered = reg.filter(&["x", "y", "z"]);
203        assert_eq!(filtered.len(), 1);
204        assert!(filtered.get("x").is_some());
205    }
206
207    #[test]
208    fn names_returns_all_names() {
209        let mut reg = ToolRegistry::new();
210        reg.register(make_tool("one")).unwrap();
211        reg.register(make_tool("two")).unwrap();
212
213        let mut names = reg.names();
214        names.sort();
215        assert_eq!(names, vec!["one", "two"]);
216    }
217
218    #[test]
219    fn empty_registry() {
220        let reg = ToolRegistry::new();
221        assert!(reg.is_empty());
222        assert_eq!(reg.len(), 0);
223        assert!(reg.schemas().is_empty());
224        assert!(reg.names().is_empty());
225    }
226}