1use crate::tool::metadata::ToolExposure;
2use crate::tool::{ToolDefinition, ToolKind};
3use serde::{Deserialize, Serialize};
4use std::collections::HashMap;
5
6#[derive(Debug, Clone, Serialize, Deserialize)]
11pub struct ToolRegistry {
12 pub tools: HashMap<String, RegisteredTool>,
14}
15
16#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct RegisteredTool {
19 pub definition: ToolDefinition,
21 pub exposure: ToolExposure,
23 #[serde(default)]
25 pub phases: Vec<String>,
26}
27
28pub const MCP_TOOL_DEFER_THRESHOLD: usize = 100;
32
33impl ToolRegistry {
34 pub fn new() -> Self {
36 Self {
37 tools: HashMap::new(),
38 }
39 }
40
41 pub fn register(&mut self, definition: ToolDefinition) {
43 let exposure = definition.metadata.exposure;
44 self.tools.insert(
45 definition.name.clone(),
46 RegisteredTool {
47 definition,
48 exposure,
49 phases: Vec::new(), },
51 );
52 }
53
54 pub fn register_with(
56 &mut self,
57 definition: ToolDefinition,
58 exposure: ToolExposure,
59 phases: Vec<String>,
60 ) {
61 self.tools.insert(
62 definition.name.clone(),
63 RegisteredTool {
64 definition,
65 exposure,
66 phases,
67 },
68 );
69 }
70
71 pub fn visible_definitions(&self) -> Vec<ToolDefinition> {
77 let mut defs: Vec<ToolDefinition> = self
78 .tools
79 .values()
80 .filter(|t| matches!(t.exposure, ToolExposure::Direct | ToolExposure::ModelOnly))
81 .map(|t| t.definition.clone())
82 .collect();
83 defs.sort_by(|a, b| a.name.cmp(&b.name));
86 defs
87 }
88
89 pub fn visible_tool_names(&self) -> Vec<String> {
91 self.visible_definitions()
92 .into_iter()
93 .map(|d| d.name)
94 .collect()
95 }
96
97 pub fn for_phase(&self, phase: &str) -> Vec<ToolDefinition> {
102 self.tools
103 .values()
104 .filter(|t| t.exposure == ToolExposure::Direct)
105 .filter(|t| t.phases.is_empty() || t.phases.iter().any(|p| p == phase))
106 .map(|t| t.definition.clone())
107 .collect()
108 }
109
110 pub fn names(&self) -> Vec<String> {
112 let mut names: Vec<String> = self.tools.keys().cloned().collect();
113 names.sort();
114 names
115 }
116
117 pub fn unregister_prefix(&mut self, prefix: &str) {
119 self.tools.retain(|name, _| !name.starts_with(prefix));
120 }
121
122 pub fn retain_tools<F>(&mut self, mut pred: F)
124 where
125 F: FnMut(&str) -> bool,
126 {
127 self.tools.retain(|name, _| pred(name));
128 }
129
130 pub fn clear(&mut self) {
132 self.tools.clear();
133 }
134
135 pub fn get(&self, name: &str) -> Option<&RegisteredTool> {
137 self.tools.get(name)
138 }
139
140 pub fn deferred_definitions(&self) -> Vec<ToolDefinition> {
142 self.tools
143 .values()
144 .filter(|t| t.exposure == ToolExposure::Deferred)
145 .map(|t| t.definition.clone())
146 .collect()
147 }
148
149 pub fn search(&self, query: &str, max_results: usize) -> Vec<ToolDefinition> {
152 let query = query.to_lowercase();
153 let query_terms: Vec<&str> = query.split_whitespace().collect();
154 if query_terms.is_empty() || max_results == 0 {
155 return Vec::new();
156 }
157
158 let mut scored: Vec<(i32, &ToolDefinition)> = self
159 .tools
160 .values()
161 .filter(|t| {
162 matches!(
164 t.exposure,
165 ToolExposure::Direct | ToolExposure::Deferred | ToolExposure::ModelOnly
166 )
167 })
168 .map(|t| {
169 let def = &t.definition;
170 let score = compute_search_score(def, &query_terms);
171 (score, def)
172 })
173 .filter(|(score, _)| *score > 0)
174 .collect();
175
176 scored.sort_by(|(left_score, left), (right_score, right)| {
177 right_score
178 .cmp(left_score)
179 .then_with(|| left.name.cmp(&right.name))
180 });
181 scored
182 .into_iter()
183 .take(max_results)
184 .map(|(_, def)| def.clone())
185 .collect()
186 }
187}
188
189impl Default for ToolRegistry {
190 fn default() -> Self {
191 Self::new()
192 }
193}
194
195fn compute_search_score(def: &ToolDefinition, terms: &[&str]) -> i32 {
197 let mut score = 0i32;
198
199 for term in terms {
200 if def.name.to_lowercase() == *term {
202 score += 100;
203 continue;
204 }
205 if def.name.to_lowercase().contains(term) {
207 score += 50;
208 }
209 if def.description.to_lowercase().contains(term) {
211 score += 20;
212 }
213 if def.metadata.namespace.to_lowercase().contains(term) {
214 score += 18;
215 }
216 for tag in &def.metadata.tags {
218 if tag.to_lowercase().contains(term) {
219 score += 15;
220 }
221 }
222 for cap in &def.metadata.capabilities {
224 if cap.to_lowercase().contains(term) {
225 score += 10;
226 }
227 }
228 for example in &def.metadata.examples {
229 if example.to_string().to_lowercase().contains(term) {
230 score += 6;
231 }
232 }
233 let kind_str = match def.kind {
235 ToolKind::Read => "read",
236 ToolKind::Write => "write",
237 ToolKind::Command => "command",
238 ToolKind::Custom => "custom",
239 };
240 if kind_str.contains(term) {
241 score += 5;
242 }
243 }
244
245 score
246}
247
248#[derive(Debug, Clone)]
250pub struct ToolSet {
251 pub phase: String,
253 pub definitions: Vec<ToolDefinition>,
255}
256
257impl ToolSet {
258 pub fn for_phase(registry: &ToolRegistry, phase: &str) -> Self {
260 Self {
261 phase: phase.to_string(),
262 definitions: registry.for_phase(phase),
263 }
264 }
265
266 pub fn names(&self) -> Vec<String> {
268 self.definitions.iter().map(|d| d.name.clone()).collect()
269 }
270}
271
272pub mod phases {
274 pub const PLANNING: &str = "planning";
276 pub const READING: &str = "reading";
278 pub const EDITING: &str = "editing";
280 pub const VERIFYING: &str = "verifying";
282 pub const REVIEWING: &str = "reviewing";
284 pub const RECOVERY: &str = "recovery";
286}
287
288#[cfg(test)]
289mod tests {
290 use super::*;
291 use crate::tool::ToolMetadata;
292 use serde_json::json;
293
294 fn make_def(name: &str, kind: ToolKind, tags: &[&str], caps: &[&str]) -> ToolDefinition {
295 ToolDefinition {
296 name: name.to_string(),
297 description: format!("Tool that does {}", name),
298 kind,
299 input_schema: json!({"type": "object"}),
300 metadata: ToolMetadata {
301 tags: tags.iter().map(|s| s.to_string()).collect(),
302 capabilities: caps.iter().map(|s| s.to_string()).collect(),
303 ..ToolMetadata::default()
304 },
305 }
306 }
307
308 #[test]
309 fn registry_empty_by_default() {
310 let reg = ToolRegistry::new();
311 assert!(reg.visible_definitions().is_empty());
312 }
313
314 #[test]
315 fn visible_definitions_are_sorted_by_name() {
316 let mut reg = ToolRegistry::new();
317 for name in ["write_file", "bash", "read_file", "apply_patch"] {
319 reg.register(make_def(name, ToolKind::Read, &[], &[]));
320 }
321 let names: Vec<String> = reg
322 .visible_definitions()
323 .into_iter()
324 .map(|d| d.name)
325 .collect();
326 assert_eq!(
327 names,
328 vec!["apply_patch", "bash", "read_file", "write_file"]
329 );
330 }
331
332 #[test]
333 fn registry_register_and_retrieve() {
334 let mut reg = ToolRegistry::new();
335 let def = make_def(
336 "read_file",
337 ToolKind::Read,
338 &["file", "read"],
339 &["repo.read"],
340 );
341 reg.register(def.clone());
342 assert_eq!(reg.visible_definitions().len(), 1);
343 assert!(reg.get("read_file").is_some());
344 }
345
346 #[test]
347 fn registry_deferred_tools_not_visible() {
348 let mut reg = ToolRegistry::new();
349 reg.register(make_def("read_file", ToolKind::Read, &["tool"], &[]));
351 let mut def = make_def("secret_tool", ToolKind::Custom, &["power"], &[]);
352 def.metadata.exposure = ToolExposure::Deferred;
353 reg.register(def);
354 let visible: Vec<String> = reg
355 .visible_definitions()
356 .into_iter()
357 .map(|d| d.name)
358 .collect();
359 assert_eq!(visible, vec!["read_file".to_string()]);
360 assert_eq!(reg.deferred_definitions().len(), 1);
361 let found = reg.search("secret", 5);
363 assert_eq!(found.len(), 1);
364 assert_eq!(found[0].name, "secret_tool");
365 }
366
367 #[test]
368 fn registry_never_auto_promotes_deferred_tools() {
369 let mut reg = ToolRegistry::new();
370 for name in ["bash", "edit", "read_file"] {
371 reg.register(make_def(name, ToolKind::Read, &["core"], &[]));
372 }
373 for name in ["code", "package_manager", "browser"] {
374 let mut def = make_def(name, ToolKind::Custom, &["power"], &[]);
375 def.metadata.exposure = ToolExposure::Deferred;
376 reg.register(def);
377 }
378 let visible = reg.visible_tool_names();
379 assert_eq!(visible.len(), 3);
380 assert!(!visible.iter().any(|n| n == "code"));
381 assert!(!visible.iter().any(|n| n == "package_manager"));
382 assert!(!visible.iter().any(|n| n == "browser"));
383 assert_eq!(MCP_TOOL_DEFER_THRESHOLD, 100);
385 }
386
387 #[test]
388 fn registry_hidden_tools_not_searchable() {
389 let mut reg = ToolRegistry::new();
390 let mut def = make_def("internal_tool", ToolKind::Custom, &["internal"], &[]);
391 def.metadata.exposure = ToolExposure::Hidden;
392 reg.register(def);
393 assert!(reg.search("internal", 10).is_empty());
394 }
395
396 #[test]
397 fn registry_search_ranking() {
398 let mut reg = ToolRegistry::new();
399 reg.register(make_def(
400 "read_file",
401 ToolKind::Read,
402 &["file", "read"],
403 &["repo.read"],
404 ));
405 reg.register(make_def(
406 "write_file",
407 ToolKind::Write,
408 &["file", "write"],
409 &["repo.write"],
410 ));
411 reg.register(make_def(
412 "bash",
413 ToolKind::Command,
414 &["shell"],
415 &["shell.exec"],
416 ));
417
418 let results = reg.search("file", 10);
419 assert!(results.len() >= 2);
420 let names: Vec<&str> = results.iter().map(|d| d.name.as_str()).collect();
422 assert!(names.contains(&"read_file"));
424 assert!(names.contains(&"write_file"));
425 }
426
427 #[test]
428 fn registry_phase_filtering() {
429 let mut reg = ToolRegistry::new();
430 let read_def = make_def("read_file", ToolKind::Read, &[], &[]);
431 let write_def = make_def("write_file", ToolKind::Write, &[], &[]);
432 reg.register_with(read_def, ToolExposure::Direct, vec!["reading".to_string()]);
433 reg.register_with(write_def, ToolExposure::Direct, vec!["editing".to_string()]);
434
435 let reading_set = reg.for_phase("reading");
436 assert_eq!(reading_set.len(), 1);
437 assert_eq!(reading_set[0].name, "read_file");
438
439 let editing_set = reg.for_phase("editing");
440 assert_eq!(editing_set.len(), 1);
441 assert_eq!(editing_set[0].name, "write_file");
442 }
443
444 #[test]
445 fn registry_phase_empty_means_all_phases() {
446 let mut reg = ToolRegistry::new();
447 let def = make_def("bash", ToolKind::Command, &[], &[]);
448 reg.register_with(def, ToolExposure::Direct, vec![]); assert_eq!(reg.for_phase("planning").len(), 1);
451 assert_eq!(reg.for_phase("reading").len(), 1);
452 assert_eq!(reg.for_phase("editing").len(), 1);
453 }
454
455 #[test]
456 fn toolset_for_phase_creates_correct_set() {
457 let mut reg = ToolRegistry::new();
458 reg.register(make_def("read", ToolKind::Read, &[], &[]));
459 reg.register(make_def("write", ToolKind::Write, &[], &[]));
460
461 let ts = ToolSet::for_phase(®, "planning");
462 assert_eq!(ts.phase, "planning");
463 assert_eq!(ts.definitions.len(), 2);
464 }
465
466 #[test]
467 fn search_respects_max_results() {
468 let mut reg = ToolRegistry::new();
469 for i in 0..10 {
470 reg.register(make_def(
471 &format!("tool_{}", i),
472 ToolKind::Custom,
473 &["test"],
474 &[],
475 ));
476 }
477 let results = reg.search("test", 3);
478 assert_eq!(results.len(), 3);
479 }
480
481 #[test]
482 fn search_excludes_zero_score_tools() {
483 let mut reg = ToolRegistry::new();
484 reg.register(make_def(
485 "read_file",
486 ToolKind::Read,
487 &["file"],
488 &["repo.read"],
489 ));
490
491 assert!(reg.search("nonexistent-capability", 10).is_empty());
492 assert!(reg.search("", 10).is_empty());
493 }
494}