Skip to main content

prism_mcp_rs/core/
tool_discovery.rs

1//! Tool Discovery and Management System
2//!
3//! Module provides complete tool discovery, filtering, and management capabilities
4//! based on the improved metadata system. It allows smart tool selection,
5//! categorization, performance monitoring, and lifecycle management.
6
7use 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
16/// Tool discovery and management system
17pub struct ToolRegistry {
18    /// Registered tools indexed by name
19    tools: HashMap<String, Tool>,
20    /// Tool execution statistics
21    global_stats: GlobalToolStats,
22}
23
24/// Global statistics across all tools
25#[derive(Debug, Clone)]
26pub struct GlobalToolStats {
27    /// Total number of registered tools
28    pub total_tools: usize,
29    /// Number of deprecated tools
30    pub deprecated_tools: usize,
31    /// Number of disabled tools
32    pub disabled_tools: usize,
33    /// Total executions across all tools
34    pub total_executions: u64,
35    /// Total successful executions
36    pub total_successes: u64,
37    /// Overall success rate
38    pub overall_success_rate: f64,
39    /// Most frequently used tool
40    pub most_used_tool: Option<String>,
41    /// Most reliable tool (highest success rate)
42    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/// Tool discovery result with ranking information
61#[derive(Debug, Clone)]
62pub struct DiscoveryResult {
63    /// Tool name
64    pub name: String,
65    /// Match score (0.0 to 1.0, higher is better)
66    pub match_score: f64,
67    /// Reason for recommendation
68    pub recommendation_reason: String,
69    /// Tool metadata snapshot
70    pub metadata: ImprovedToolMetadata,
71    /// Whether tool is deprecated
72    pub is_deprecated: bool,
73    /// Whether tool is enabled
74    pub is_enabled: bool,
75}
76
77/// Tool discovery criteria
78#[derive(Debug, Clone, Default)]
79pub struct DiscoveryCriteria {
80    /// Category filter
81    pub category_filter: Option<CategoryFilter>,
82    /// Required behavior hints
83    pub required_hints: ToolBehaviorHints,
84    /// Preferred behavior hints (for ranking)
85    pub preferred_hints: ToolBehaviorHints,
86    /// Exclude deprecated tools
87    pub exclude_deprecated: bool,
88    /// Exclude disabled tools
89    pub exclude_disabled: bool,
90    /// Minimum success rate (0.0 to 1.0)
91    pub min_success_rate: Option<f64>,
92    /// Maximum average execution time
93    pub max_execution_time: Option<Duration>,
94    /// Text search in name/description
95    pub text_search: Option<String>,
96    /// Minimum number of executions (for reliability filtering)
97    pub min_executions: Option<u64>,
98}
99
100impl Default for ToolRegistry {
101    fn default() -> Self {
102        Self::new()
103    }
104}
105
106impl ToolRegistry {
107    /// Create a new tool registry
108    pub fn new() -> Self {
109        Self {
110            tools: HashMap::new(),
111            global_stats: GlobalToolStats::default(),
112        }
113    }
114
115    /// Register a tool in the registry
116    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    /// Unregister a tool from the registry
131    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    /// Get a tool by name
142    pub fn get_tool(&self, name: &str) -> Option<&Tool> {
143        self.tools.get(name)
144    }
145
146    /// Get a mutable reference to a tool by name
147    pub fn get_tool_mut(&mut self, name: &str) -> Option<&mut Tool> {
148        self.tools.get_mut(name)
149    }
150
151    /// List all tool names
152    pub fn list_tool_names(&self) -> Vec<String> {
153        self.tools.keys().cloned().collect()
154    }
155
156    /// Discover tools based on criteria
157    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        // Sort by match score (descending)
167        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    /// Get tools by category
177    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    /// Get deprecated tools
186    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    /// Get disabled tools
195    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    /// Get performance report for all tools
204    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    /// Get global statistics
214    pub fn get_global_stats(&self) -> &GlobalToolStats {
215        &self.global_stats
216    }
217
218    /// Recommend best tool for a specific use case
219    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        // Add text search based on use case
227        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    /// Clean up deprecated tools based on policy
234    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                    // Check severity-based removal
249                    if matches!(deprecation.severity, DeprecationSeverity::Critical) {
250                        return true;
251                    }
252
253                    // Check time-based removal
254                    if let Some(removal_date) = deprecation.removal_date {
255                        if current_time >= removal_date {
256                            return true;
257                        }
258                    }
259
260                    // Check age-based removal
261                    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    /// Update global statistics
287    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            // Track most used tool
313            if metrics.execution_count > max_executions {
314                max_executions = metrics.execution_count;
315                most_used = Some(name.clone());
316            }
317
318            // Track most reliable tool (with minimum executions)
319            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    /// Evaluate how well a tool matches the discovery criteria
336    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        // Filter out tools that don't meet basic criteria
346        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        // Filter by minimum success rate
357        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        // Filter by maximum execution time
364        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        // Filter by minimum executions
371        if let Some(min_execs) = criteria.min_executions {
372            if metrics.execution_count < min_execs {
373                return None;
374            }
375        }
376
377        // Category matching
378        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        // Text search matching
388        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                // If text search is specified but doesn't match, exclude this tool
403                return None;
404            }
405        }
406
407        // Behavior hints matching - check required hints first
408        let hints = tool.behavior_hints();
409
410        // Filter out tools that don't meet required hints
411        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        // Add score bonuses for meeting required hints
433        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        // Preferred hints bonus
448        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        // Performance-based scoring
459        if metrics.execution_count > 0 {
460            // Success rate bonus
461            let success_bonus = (metrics.success_rate / 100.0) * 0.2;
462            score += success_bonus;
463
464            // Usage frequency bonus (logarithmic scale)
465            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        // Deprecation penalty
477        if tool.is_deprecated() {
478            score *= 0.5;
479            reasons.push("deprecated (reduced score)".to_string());
480        }
481
482        // Disabled penalty
483        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/// Policy for cleaning up deprecated tools
500#[derive(Debug, Clone)]
501pub struct DeprecationCleanupPolicy {
502    /// Maximum number of days to keep deprecated tools
503    pub max_deprecated_days: u32,
504    /// Remove tools marked as critical immediately
505    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        // Register tool
561        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        // Try to register duplicate - should fail
566        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        // Unregister tool
574        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        // Add tools with different categories
584        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        // Test category filtering
604        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        // Add tools with different characteristics
620        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        // Test discovery with read-only requirement
651        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        // Test discovery excluding deprecated
663        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); // Should exclude deprecated tool
670        assert!(!results.iter().any(|r| r.name == "old_tool"));
671
672        // Test text search
673        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        // Recommend tool for file processing
735        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        // Add tools with different deprecation states
750        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        // Clean up with default policy (should remove critical tools)
783        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}