1use crate::tool::Tool;
8use pe_core::error::PeError;
9use pe_core::llm::ToolSchema;
10use std::collections::HashMap;
11use std::sync::Arc;
12
13pub 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 pub fn new() -> Self {
44 Self {
45 tools: HashMap::new(),
46 }
47 }
48
49 pub fn register(&mut self, tool: impl Tool + 'static) -> Result<&mut Self, PeError> {
51 self.try_register(tool)
52 }
53
54 pub fn register_arc(&mut self, tool: Arc<dyn Tool>) -> Result<&mut Self, PeError> {
56 self.try_register_arc(tool)
57 }
58
59 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 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 pub fn get(&self, name: &str) -> Option<Arc<dyn Tool>> {
81 self.tools.get(name).cloned()
82 }
83
84 pub fn schemas(&self) -> Vec<ToolSchema> {
86 self.tools.values().map(|t| t.schema()).collect()
87 }
88
89 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 pub fn names(&self) -> Vec<&str> {
102 self.tools.keys().map(String::as_str).collect()
103 }
104
105 pub fn len(&self) -> usize {
107 self.tools.len()
108 }
109
110 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}