1use crate::core::error::{McpError, McpResult};
8use crate::core::tool::Tool;
9use crate::core::tool_metadata::{
10 CategoryFilter, DeprecationSeverity, ImprovedToolMetadata, ToolBehaviorHints,
11};
12use chrono::Utc;
13use std::collections::HashMap;
14use std::time::Duration;
15
16pub struct ToolRegistry {
18 tools: HashMap<String, Tool>,
20 global_stats: GlobalToolStats,
22}
23
24#[derive(Debug, Clone)]
26pub struct GlobalToolStats {
27 pub total_tools: usize,
29 pub deprecated_tools: usize,
31 pub disabled_tools: usize,
33 pub total_executions: u64,
35 pub total_successes: u64,
37 pub overall_success_rate: f64,
39 pub most_used_tool: Option<String>,
41 pub most_reliable_tool: Option<String>,
43}
44
45impl Default for GlobalToolStats {
46 fn default() -> Self {
47 Self {
48 total_tools: 0,
49 deprecated_tools: 0,
50 disabled_tools: 0,
51 total_executions: 0,
52 total_successes: 0,
53 overall_success_rate: 0.0,
54 most_used_tool: None,
55 most_reliable_tool: None,
56 }
57 }
58}
59
60#[derive(Debug, Clone)]
62pub struct DiscoveryResult {
63 pub name: String,
65 pub match_score: f64,
67 pub recommendation_reason: String,
69 pub metadata: ImprovedToolMetadata,
71 pub is_deprecated: bool,
73 pub is_enabled: bool,
75}
76
77#[derive(Debug, Clone, Default)]
79pub struct DiscoveryCriteria {
80 pub category_filter: Option<CategoryFilter>,
82 pub required_hints: ToolBehaviorHints,
84 pub preferred_hints: ToolBehaviorHints,
86 pub exclude_deprecated: bool,
88 pub exclude_disabled: bool,
90 pub min_success_rate: Option<f64>,
92 pub max_execution_time: Option<Duration>,
94 pub text_search: Option<String>,
96 pub min_executions: Option<u64>,
98}
99
100impl Default for ToolRegistry {
101 fn default() -> Self {
102 Self::new()
103 }
104}
105
106impl ToolRegistry {
107 pub fn new() -> Self {
109 Self {
110 tools: HashMap::new(),
111 global_stats: GlobalToolStats::default(),
112 }
113 }
114
115 pub fn register_tool(&mut self, tool: Tool) -> McpResult<()> {
117 let name = tool.info.name.clone();
118
119 if self.tools.contains_key(&name) {
120 return Err(McpError::validation(format!(
121 "Tool '{name}' is already registered"
122 )));
123 }
124
125 self.tools.insert(name, tool);
126 self.update_global_stats();
127 Ok(())
128 }
129
130 pub fn unregister_tool(&mut self, name: &str) -> McpResult<Tool> {
132 let tool = self
133 .tools
134 .remove(name)
135 .ok_or_else(|| McpError::validation(format!("Tool '{name}' not found")))?;
136
137 self.update_global_stats();
138 Ok(tool)
139 }
140
141 pub fn get_tool(&self, name: &str) -> Option<&Tool> {
143 self.tools.get(name)
144 }
145
146 pub fn get_tool_mut(&mut self, name: &str) -> Option<&mut Tool> {
148 self.tools.get_mut(name)
149 }
150
151 pub fn list_tool_names(&self) -> Vec<String> {
153 self.tools.keys().cloned().collect()
154 }
155
156 pub fn discover_tools(&self, criteria: &DiscoveryCriteria) -> Vec<DiscoveryResult> {
158 let mut results = Vec::new();
159
160 for (name, tool) in &self.tools {
161 if let Some(result) = self.evaluate_tool_match(name, tool, criteria) {
162 results.push(result);
163 }
164 }
165
166 results.sort_by(|a, b| {
168 b.match_score
169 .partial_cmp(&a.match_score)
170 .unwrap_or(std::cmp::Ordering::Equal)
171 });
172
173 results
174 }
175
176 pub fn get_tools_by_category(&self, filter: &CategoryFilter) -> Vec<String> {
178 self.tools
179 .iter()
180 .filter(|(_, tool)| tool.matches_category_filter(filter))
181 .map(|(name, _)| name.clone())
182 .collect()
183 }
184
185 pub fn get_deprecated_tools(&self) -> Vec<String> {
187 self.tools
188 .iter()
189 .filter(|(_, tool)| tool.is_deprecated())
190 .map(|(name, _)| name.clone())
191 .collect()
192 }
193
194 pub fn get_disabled_tools(&self) -> Vec<String> {
196 self.tools
197 .iter()
198 .filter(|(_, tool)| !tool.is_enabled())
199 .map(|(name, _)| name.clone())
200 .collect()
201 }
202
203 pub fn get_performance_report(
205 &self,
206 ) -> HashMap<String, crate::core::tool_metadata::ToolPerformanceMetrics> {
207 self.tools
208 .iter()
209 .map(|(name, tool)| (name.clone(), tool.performance_metrics()))
210 .collect()
211 }
212
213 pub fn get_global_stats(&self) -> &GlobalToolStats {
215 &self.global_stats
216 }
217
218 pub fn recommend_tool(
220 &self,
221 use_case: &str,
222 criteria: &DiscoveryCriteria,
223 ) -> Option<DiscoveryResult> {
224 let mut improved_criteria = criteria.clone();
225
226 improved_criteria.text_search = Some(use_case.to_string());
228
229 let results = self.discover_tools(&improved_criteria);
230 results.into_iter().next()
231 }
232
233 pub fn cleanup_deprecated_tools(&mut self, policy: &DeprecationCleanupPolicy) -> Vec<String> {
235 let mut removed_tools = Vec::new();
236
237 let current_time = Utc::now();
238
239 let tools_to_remove: Vec<String> = self
240 .tools
241 .iter()
242 .filter(|(_, tool)| {
243 if let Some(ref deprecation) = tool.improved_metadata.deprecation {
244 if !deprecation.deprecated {
245 return false;
246 }
247
248 if matches!(deprecation.severity, DeprecationSeverity::Critical) {
250 return true;
251 }
252
253 if let Some(removal_date) = deprecation.removal_date {
255 if current_time >= removal_date {
256 return true;
257 }
258 }
259
260 if let Some(deprecated_date) = deprecation.deprecated_date {
262 let age = current_time.signed_duration_since(deprecated_date);
263 if age.num_days() > policy.max_deprecated_days as i64 {
264 return true;
265 }
266 }
267 }
268 false
269 })
270 .map(|(name, _)| name.clone())
271 .collect();
272
273 for name in tools_to_remove {
274 if self.tools.remove(&name).is_some() {
275 removed_tools.push(name);
276 }
277 }
278
279 if !removed_tools.is_empty() {
280 self.update_global_stats();
281 }
282
283 removed_tools
284 }
285
286 fn update_global_stats(&mut self) {
288 let mut stats = GlobalToolStats {
289 total_tools: self.tools.len(),
290 ..Default::default()
291 };
292
293 let mut max_executions = 0u64;
294 let mut max_success_rate = 0.0f64;
295 let mut most_used = None;
296 let mut most_reliable = None;
297
298 for (name, tool) in &self.tools {
299 let metrics = tool.performance_metrics();
300
301 if tool.is_deprecated() {
302 stats.deprecated_tools += 1;
303 }
304
305 if !tool.is_enabled() {
306 stats.disabled_tools += 1;
307 }
308
309 stats.total_executions += metrics.execution_count;
310 stats.total_successes += metrics.success_count;
311
312 if metrics.execution_count > max_executions {
314 max_executions = metrics.execution_count;
315 most_used = Some(name.clone());
316 }
317
318 if metrics.execution_count >= 5 && metrics.success_rate > max_success_rate {
320 max_success_rate = metrics.success_rate;
321 most_reliable = Some(name.clone());
322 }
323 }
324
325 if stats.total_executions > 0 {
326 stats.overall_success_rate =
327 (stats.total_successes as f64 / stats.total_executions as f64) * 100.0;
328 }
329
330 stats.most_used_tool = most_used;
331 stats.most_reliable_tool = most_reliable;
332 self.global_stats = stats;
333 }
334
335 fn evaluate_tool_match(
337 &self,
338 name: &str,
339 tool: &Tool,
340 criteria: &DiscoveryCriteria,
341 ) -> Option<DiscoveryResult> {
342 let mut score = 0.0f64;
343 let mut reasons = Vec::new();
344
345 if criteria.exclude_deprecated && tool.is_deprecated() {
347 return None;
348 }
349
350 if criteria.exclude_disabled && !tool.is_enabled() {
351 return None;
352 }
353
354 let metrics = tool.performance_metrics();
355
356 if let Some(min_rate) = criteria.min_success_rate {
358 if metrics.execution_count > 0 && metrics.success_rate < min_rate * 100.0 {
359 return None;
360 }
361 }
362
363 if let Some(max_time) = criteria.max_execution_time {
365 if metrics.execution_count > 0 && metrics.average_execution_time > max_time {
366 return None;
367 }
368 }
369
370 if let Some(min_execs) = criteria.min_executions {
372 if metrics.execution_count < min_execs {
373 return None;
374 }
375 }
376
377 if let Some(ref filter) = criteria.category_filter {
379 if tool.matches_category_filter(filter) {
380 score += 0.3;
381 reasons.push("matches category criteria".to_string());
382 } else {
383 return None;
384 }
385 }
386
387 if let Some(ref search_text) = criteria.text_search {
389 let search_lower = search_text.to_lowercase();
390 let name_match = name.to_lowercase().contains(&search_lower);
391 let desc_match = tool
392 .info
393 .description
394 .as_ref()
395 .map(|d| d.to_lowercase().contains(&search_lower))
396 .unwrap_or(false);
397
398 if name_match || desc_match {
399 score += if name_match { 0.4 } else { 0.2 };
400 reasons.push("matches text search".to_string());
401 } else {
402 return None;
404 }
405 }
406
407 let hints = tool.behavior_hints();
409
410 if criteria.required_hints.read_only.unwrap_or(false) && !hints.read_only.unwrap_or(false) {
412 return None;
413 }
414 if criteria.required_hints.idempotent.unwrap_or(false) && !hints.idempotent.unwrap_or(false)
415 {
416 return None;
417 }
418 if criteria.required_hints.cacheable.unwrap_or(false) && !hints.cacheable.unwrap_or(false) {
419 return None;
420 }
421 if criteria.required_hints.destructive.unwrap_or(false)
422 && !hints.destructive.unwrap_or(false)
423 {
424 return None;
425 }
426 if criteria.required_hints.requires_auth.unwrap_or(false)
427 && !hints.requires_auth.unwrap_or(false)
428 {
429 return None;
430 }
431
432 if criteria.required_hints.read_only.unwrap_or(false) && hints.read_only.unwrap_or(false) {
434 score += 0.2;
435 reasons.push("read-only as required".to_string());
436 }
437 if criteria.required_hints.idempotent.unwrap_or(false) && hints.idempotent.unwrap_or(false)
438 {
439 score += 0.2;
440 reasons.push("idempotent as required".to_string());
441 }
442 if criteria.required_hints.cacheable.unwrap_or(false) && hints.cacheable.unwrap_or(false) {
443 score += 0.15;
444 reasons.push("cacheable as required".to_string());
445 }
446
447 if criteria.preferred_hints.read_only.unwrap_or(false) && hints.read_only.unwrap_or(false) {
449 score += 0.1;
450 reasons.push("preferred: read-only".to_string());
451 }
452 if criteria.preferred_hints.idempotent.unwrap_or(false) && hints.idempotent.unwrap_or(false)
453 {
454 score += 0.1;
455 reasons.push("preferred: idempotent".to_string());
456 }
457
458 if metrics.execution_count > 0 {
460 let success_bonus = (metrics.success_rate / 100.0) * 0.2;
462 score += success_bonus;
463
464 let usage_bonus = (metrics.execution_count as f64).ln() * 0.05;
466 score += usage_bonus.min(0.15);
467
468 if metrics.success_rate > 95.0 {
469 reasons.push("high reliability".to_string());
470 }
471 if metrics.execution_count > 100 {
472 reasons.push("well-tested".to_string());
473 }
474 }
475
476 if tool.is_deprecated() {
478 score *= 0.5;
479 reasons.push("deprecated (reduced score)".to_string());
480 }
481
482 if !tool.is_enabled() {
484 score *= 0.1;
485 reasons.push("disabled (reduced score)".to_string());
486 }
487
488 Some(DiscoveryResult {
489 name: name.to_string(),
490 match_score: score.min(1.0),
491 recommendation_reason: reasons.join(", "),
492 metadata: tool.improved_metadata.clone(),
493 is_deprecated: tool.is_deprecated(),
494 is_enabled: tool.is_enabled(),
495 })
496 }
497}
498
499#[derive(Debug, Clone)]
501pub struct DeprecationCleanupPolicy {
502 pub max_deprecated_days: u32,
504 pub remove_critical_immediately: bool,
506}
507
508impl Default for DeprecationCleanupPolicy {
509 fn default() -> Self {
510 Self {
511 max_deprecated_days: 90,
512 remove_critical_immediately: true,
513 }
514 }
515}
516
517#[cfg(test)]
518mod tests {
519 use super::*;
520 use crate::core::tool::{ToolBuilder, ToolHandler};
521 use crate::core::tool_metadata::*;
522 use async_trait::async_trait;
523 use serde_json::Value;
524 use std::collections::HashMap;
525
526 struct MockHandler {
527 result: String,
528 }
529
530 #[async_trait]
531 impl ToolHandler for MockHandler {
532 async fn call(
533 &self,
534 _args: HashMap<String, Value>,
535 ) -> McpResult<crate::protocol::types::ToolResult> {
536 Ok(crate::protocol::types::ToolResult {
537 content: vec![crate::protocol::types::ContentBlock::Text {
538 text: self.result.clone(),
539 annotations: None,
540 meta: None,
541 }],
542 is_error: None,
543 structured_content: None,
544 meta: None,
545 })
546 }
547 }
548
549 #[test]
550 fn test_tool_registry_basic_operations() {
551 let mut registry = ToolRegistry::new();
552
553 let tool = ToolBuilder::new("test_tool")
554 .description("A test tool")
555 .build(MockHandler {
556 result: "test".to_string(),
557 })
558 .unwrap();
559
560 registry.register_tool(tool).unwrap();
562 assert_eq!(registry.list_tool_names().len(), 1);
563 assert!(registry.get_tool("test_tool").is_some());
564
565 let duplicate_tool = ToolBuilder::new("test_tool")
567 .build(MockHandler {
568 result: "duplicate".to_string(),
569 })
570 .unwrap();
571 assert!(registry.register_tool(duplicate_tool).is_err());
572
573 let removed = registry.unregister_tool("test_tool").unwrap();
575 assert_eq!(removed.info.name, "test_tool");
576 assert_eq!(registry.list_tool_names().len(), 0);
577 }
578
579 #[test]
580 fn test_tool_discovery_by_category() {
581 let mut registry = ToolRegistry::new();
582
583 let file_tool = ToolBuilder::new("file_reader")
585 .category_simple("file".to_string(), Some("read".to_string()))
586 .tag("filesystem".to_string())
587 .build(MockHandler {
588 result: "file".to_string(),
589 })
590 .unwrap();
591
592 let network_tool = ToolBuilder::new("http_client")
593 .category_simple("network".to_string(), Some("http".to_string()))
594 .tag("client".to_string())
595 .build(MockHandler {
596 result: "network".to_string(),
597 })
598 .unwrap();
599
600 registry.register_tool(file_tool).unwrap();
601 registry.register_tool(network_tool).unwrap();
602
603 let file_filter = CategoryFilter::new().with_primary("file".to_string());
605 let file_tools = registry.get_tools_by_category(&file_filter);
606 assert_eq!(file_tools.len(), 1);
607 assert!(file_tools.contains(&"file_reader".to_string()));
608
609 let network_filter = CategoryFilter::new().with_primary("network".to_string());
610 let network_tools = registry.get_tools_by_category(&network_filter);
611 assert_eq!(network_tools.len(), 1);
612 assert!(network_tools.contains(&"http_client".to_string()));
613 }
614
615 #[test]
616 fn test_tool_discovery_criteria() {
617 let mut registry = ToolRegistry::new();
618
619 let read_only_tool = ToolBuilder::new("reader")
621 .description("Reads data")
622 .read_only()
623 .idempotent()
624 .cacheable()
625 .build(MockHandler {
626 result: "read".to_string(),
627 })
628 .unwrap();
629
630 let destructive_tool = ToolBuilder::new("deleter")
631 .description("Deletes data")
632 .destructive()
633 .build(MockHandler {
634 result: "delete".to_string(),
635 })
636 .unwrap();
637
638 let deprecated_tool = ToolBuilder::new("old_tool")
639 .description("Old tool")
640 .deprecated_simple("Use new_tool instead")
641 .build(MockHandler {
642 result: "old".to_string(),
643 })
644 .unwrap();
645
646 registry.register_tool(read_only_tool).unwrap();
647 registry.register_tool(destructive_tool).unwrap();
648 registry.register_tool(deprecated_tool).unwrap();
649
650 let criteria = DiscoveryCriteria {
652 required_hints: ToolBehaviorHints::new().read_only(),
653 exclude_deprecated: false,
654 exclude_disabled: false,
655 ..Default::default()
656 };
657
658 let results = registry.discover_tools(&criteria);
659 assert_eq!(results.len(), 1);
660 assert_eq!(results[0].name, "reader");
661
662 let criteria = DiscoveryCriteria {
664 exclude_deprecated: true,
665 ..Default::default()
666 };
667
668 let results = registry.discover_tools(&criteria);
669 assert_eq!(results.len(), 2); assert!(!results.iter().any(|r| r.name == "old_tool"));
671
672 let criteria = DiscoveryCriteria {
674 text_search: Some("delete".to_string()),
675 exclude_deprecated: false,
676 ..Default::default()
677 };
678
679 let results = registry.discover_tools(&criteria);
680 assert_eq!(results.len(), 1);
681 assert_eq!(results[0].name, "deleter");
682 }
683
684 #[test]
685 fn test_global_statistics() {
686 let mut registry = ToolRegistry::new();
687
688 let tool1 = ToolBuilder::new("tool1")
689 .build(MockHandler {
690 result: "1".to_string(),
691 })
692 .unwrap();
693
694 let tool2 = ToolBuilder::new("tool2")
695 .deprecated_simple("Old tool")
696 .build(MockHandler {
697 result: "2".to_string(),
698 })
699 .unwrap();
700
701 registry.register_tool(tool1).unwrap();
702 registry.register_tool(tool2).unwrap();
703
704 let stats = registry.get_global_stats();
705 assert_eq!(stats.total_tools, 2);
706 assert_eq!(stats.deprecated_tools, 1);
707 assert_eq!(stats.disabled_tools, 0);
708 }
709
710 #[test]
711 fn test_tool_recommendation() {
712 let mut registry = ToolRegistry::new();
713
714 let file_tool = ToolBuilder::new("file_processor")
715 .description("Processes files efficiently")
716 .category_simple("file".to_string(), Some("process".to_string()))
717 .read_only()
718 .build(MockHandler {
719 result: "processed".to_string(),
720 })
721 .unwrap();
722
723 let network_tool = ToolBuilder::new("network_handler")
724 .description("Handles network requests")
725 .category_simple("network".to_string(), None)
726 .build(MockHandler {
727 result: "handled".to_string(),
728 })
729 .unwrap();
730
731 registry.register_tool(file_tool).unwrap();
732 registry.register_tool(network_tool).unwrap();
733
734 let criteria = DiscoveryCriteria::default();
736 let recommendation = registry.recommend_tool("file", &criteria);
737
738 assert!(recommendation.is_some());
739 let result = recommendation.unwrap();
740 assert_eq!(result.name, "file_processor");
741 assert!(result.match_score > 0.0);
742 assert!(result.recommendation_reason.contains("matches text search"));
743 }
744
745 #[test]
746 fn test_deprecation_cleanup() {
747 let mut registry = ToolRegistry::new();
748
749 let normal_tool = ToolBuilder::new("normal")
751 .build(MockHandler {
752 result: "normal".to_string(),
753 })
754 .unwrap();
755
756 let deprecated_tool = ToolBuilder::new("deprecated")
757 .deprecated(
758 ToolDeprecation::new("Old version".to_string())
759 .with_severity(DeprecationSeverity::Low),
760 )
761 .build(MockHandler {
762 result: "deprecated".to_string(),
763 })
764 .unwrap();
765
766 let critical_tool = ToolBuilder::new("critical")
767 .deprecated(
768 ToolDeprecation::new("Security issue".to_string())
769 .with_severity(DeprecationSeverity::Critical),
770 )
771 .build(MockHandler {
772 result: "critical".to_string(),
773 })
774 .unwrap();
775
776 registry.register_tool(normal_tool).unwrap();
777 registry.register_tool(deprecated_tool).unwrap();
778 registry.register_tool(critical_tool).unwrap();
779
780 assert_eq!(registry.list_tool_names().len(), 3);
781
782 let policy = DeprecationCleanupPolicy::default();
784 let removed = registry.cleanup_deprecated_tools(&policy);
785
786 assert_eq!(removed.len(), 1);
787 assert!(removed.contains(&"critical".to_string()));
788 assert_eq!(registry.list_tool_names().len(), 2);
789 }
790}