Skip to main content

vtcode_mcp/
tool_discovery_cache.rs

1//! Tool discovery caching system for MCP to avoid redundant tool searches
2//!
3//! This module provides multi-level caching for MCP tool discovery with
4//! bloom filters for fast negative lookups and LRU cache for positive results.
5
6use lru::LruCache;
7use rustc_hash::FxHashMap;
8use std::num::NonZeroUsize;
9use std::sync::{Arc, RwLock};
10use std::time::{Duration, Instant};
11use tracing::error;
12
13use super::McpToolInfo;
14use super::tool_discovery::DetailLevel;
15
16/// Bloom filter for fast negative lookups (tool doesn't exist)
17#[derive(Clone)]
18pub struct BloomFilter {
19    /// Bit array for the filter
20    bits: Vec<bool>,
21    /// Number of hash functions
22    num_hashes: usize,
23    /// Size of the bit array
24    size: usize,
25}
26
27impl BloomFilter {
28    fn new(expected_items: usize, false_positive_rate: f64) -> Self {
29        let expected_items = expected_items.max(1);
30        let size = Self::optimal_size(expected_items, false_positive_rate).max(1);
31        let num_hashes = Self::optimal_num_hashes(size, expected_items).max(1);
32
33        Self { bits: vec![false; size], num_hashes, size }
34    }
35
36    /// Add an item to the bloom filter
37    fn insert(&mut self, item: &str) {
38        for i in 0..self.num_hashes {
39            let hash = self.hash(item, i);
40            let index = hash % self.size;
41            if let Some(bit) = self.bits.get_mut(index) {
42                *bit = true;
43            }
44        }
45    }
46
47    /// Check if an item might be in the set
48    fn contains(&self, item: &str) -> bool {
49        for i in 0..self.num_hashes {
50            let hash = self.hash(item, i);
51            let index = hash % self.size;
52            if !self.bits.get(index).copied().unwrap_or(false) {
53                return false;
54            }
55        }
56        true
57    }
58
59    /// Clear the bloom filter
60    pub fn clear(&mut self) {
61        self.bits.fill(false);
62    }
63
64    /// Calculate optimal size for bloom filter
65    #[allow(
66        clippy::cast_sign_loss,
67        reason = "Intentional compatibility, platform, or test-only suppression."
68    )]
69    #[expect(
70        clippy::cast_possible_truncation,
71        reason = "The calculated Bloom filter size is bounded by the target usize range before conversion."
72    )]
73    fn optimal_size(expected_items: usize, false_positive_rate: f64) -> usize {
74        let size = -(expected_items as f64 * false_positive_rate.ln() / (2.0_f64.ln().powi(2)));
75        size.max(0.0).ceil() as usize
76    }
77
78    /// Calculate optimal number of hash functions
79    #[allow(
80        clippy::cast_sign_loss,
81        reason = "Intentional compatibility, platform, or test-only suppression."
82    )]
83    #[expect(
84        clippy::cast_possible_truncation,
85        reason = "The calculated Bloom filter hash count is bounded by the target usize range before conversion."
86    )]
87    fn optimal_num_hashes(size: usize, expected_items: usize) -> usize {
88        let num_hashes = (size as f64 / expected_items as f64) * 2.0_f64.ln();
89        num_hashes.max(0.0).ceil() as usize
90    }
91
92    /// Simple hash function for bloom filter
93    fn hash(&self, item: &str, seed: usize) -> usize {
94        use std::collections::hash_map::DefaultHasher;
95        use std::hash::{Hash, Hasher};
96
97        let mut hasher = DefaultHasher::new();
98        item.hash(&mut hasher);
99        seed.hash(&mut hasher);
100        usize::try_from(hasher.finish()).unwrap_or(usize::MAX)
101    }
102}
103
104/// Cache key for tool discovery results
105#[derive(Debug, Clone, Hash, PartialEq, Eq)]
106struct ToolDiscoveryCacheKey {
107    provider_name: String,
108    keyword: String,
109    detail_level: DetailLevel,
110}
111
112/// Cached tool discovery result (internal cache entry)
113#[derive(Clone)]
114struct CachedToolDiscoveryEntry {
115    // OPTIMIZATION: Use Arc to avoid cloning large vectors on cache hits
116    results: Arc<Vec<ToolDiscoveryResult>>,
117    timestamp: Instant,
118}
119
120struct DiscoveryCacheInner {
121    bloom_filter: BloomFilter,
122    detailed_cache: LruCache<ToolDiscoveryCacheKey, CachedToolDiscoveryEntry>,
123    all_tools_cache: FxHashMap<String, Vec<McpToolInfo>>,
124    last_refresh: FxHashMap<String, Instant>,
125}
126
127/// Cached tool discovery result (matches actual API)
128#[derive(Debug, Clone)]
129pub struct ToolDiscoveryResult {
130    tool: McpToolInfo,
131    relevance_score: f64,
132    detail_level: DetailLevel,
133}
134
135/// Multi-level caching system for tool discovery
136pub(crate) struct ToolDiscoveryCache {
137    inner: Arc<RwLock<DiscoveryCacheInner>>,
138    /// Cache configuration
139    config: CacheConfig,
140}
141
142#[derive(Clone)]
143struct CacheConfig {
144    /// Maximum age for cached entries
145    max_age: Duration,
146    /// Maximum age for provider tool lists
147    provider_refresh_interval: Duration,
148    /// Expected number of tools for bloom filter sizing
149    expected_tool_count: usize,
150    /// Acceptable false positive rate for bloom filter
151    false_positive_rate: f64,
152}
153
154impl ToolDiscoveryCache {
155    pub(crate) fn new(capacity: usize) -> Self {
156        let config = CacheConfig {
157            max_age: Duration::from_secs(300),                  // 5 minutes
158            provider_refresh_interval: Duration::from_secs(60), // 1 minute
159            expected_tool_count: 1000,
160            false_positive_rate: 0.01, // 1% false positive rate
161        };
162
163        let bloom_filter = BloomFilter::new(config.expected_tool_count, config.false_positive_rate);
164        let cache_size = NonZeroUsize::new(capacity).or(NonZeroUsize::new(100));
165
166        Self {
167            inner: Arc::new(RwLock::new(DiscoveryCacheInner {
168                bloom_filter,
169                detailed_cache: LruCache::new(cache_size.unwrap_or(NonZeroUsize::MIN)),
170                all_tools_cache: FxHashMap::default(),
171                last_refresh: FxHashMap::default(),
172            })),
173            config,
174        }
175    }
176
177    /// Check if a tool might exist (fast negative lookup)
178    pub fn might_have_tool(&self, tool_name: &str) -> bool {
179        match self.inner.read() {
180            Ok(inner) => inner.bloom_filter.contains(tool_name),
181            Err(_) => {
182                tracing::warn!("Bloom filter lock poisoned, assuming tool might exist");
183                true
184            }
185        }
186    }
187
188    /// Get cached tool discovery results
189    fn get_cached_discovery(
190        &self,
191        provider_name: &str,
192        keyword: &str,
193        detail_level: DetailLevel,
194    ) -> Option<Arc<Vec<ToolDiscoveryResult>>> {
195        // OPTIMIZATION: Use to_owned() for explicit String allocation
196        let key = ToolDiscoveryCacheKey {
197            provider_name: provider_name.to_owned(),
198            keyword: keyword.to_owned(),
199            detail_level,
200        };
201
202        let mut inner = match self.inner.write() {
203            Ok(inner) => inner,
204            Err(e) => {
205                tracing::error!("Detailed cache lock poisoned: {}", e);
206                return None;
207            }
208        };
209
210        if let Some(cached) = inner.detailed_cache.get(&key) {
211            // Check if the cached entry is still fresh
212            if cached.timestamp.elapsed() < self.config.max_age {
213                return Some(Arc::clone(&cached.results));
214            } else {
215                // Entry is stale, remove it
216                drop(inner.detailed_cache.pop(&key));
217            }
218        }
219
220        None
221    }
222
223    /// Cache tool discovery results
224    fn cache_discovery(
225        &self,
226        provider_name: &str,
227        keyword: &str,
228        detail_level: DetailLevel,
229        results: Vec<ToolDiscoveryResult>,
230    ) {
231        self.cache_discovery_shared(provider_name, keyword, detail_level, Arc::new(results));
232    }
233
234    fn cache_discovery_shared(
235        &self,
236        provider_name: &str,
237        keyword: &str,
238        detail_level: DetailLevel,
239        results: Arc<Vec<ToolDiscoveryResult>>,
240    ) {
241        // OPTIMIZATION: Use to_owned() for explicit String allocation
242        let key = ToolDiscoveryCacheKey {
243            provider_name: provider_name.to_owned(),
244            keyword: keyword.to_owned(),
245            detail_level,
246        };
247
248        let cached = CachedToolDiscoveryEntry {
249            // OPTIMIZATION: Wrap in Arc once, share across cache hits
250            results: Arc::clone(&results),
251            timestamp: Instant::now(),
252        };
253
254        let Ok(mut inner) = self.inner.write() else {
255            tracing::error!("Failed to acquire discovery cache lock for writing");
256            return;
257        };
258
259        drop(inner.detailed_cache.put(key, cached));
260
261        for result in results.iter() {
262            inner.bloom_filter.insert(&result.tool.name);
263        }
264    }
265
266    /// Get all cached tools for a provider (with refresh checking)
267    pub fn get_all_tools(&self, provider_name: &str, refresh_if_stale: bool) -> Option<Vec<McpToolInfo>> {
268        let inner = match self.inner.read() {
269            Ok(inner) => inner,
270            Err(e) => {
271                error!("Discovery cache lock poisoned: {}", e);
272                return None;
273            }
274        };
275
276        let should_refresh = if let Some(last) = inner.last_refresh.get(provider_name) {
277            last.elapsed() > self.config.provider_refresh_interval
278        } else {
279            true
280        };
281
282        if should_refresh && refresh_if_stale {
283            return None; // Signal that refresh is needed
284        }
285
286        inner.all_tools_cache.get(provider_name).cloned()
287    }
288
289    /// Cache all tools for a provider
290    pub fn cache_all_tools(&self, provider_name: &str, tools: Vec<McpToolInfo>) {
291        let mut inner = match self.inner.write() {
292            Ok(inner) => inner,
293            Err(e) => {
294                tracing::error!("Discovery cache lock poisoned: {}", e);
295                return;
296            }
297        };
298
299        drop(inner.all_tools_cache.insert(provider_name.to_owned(), tools.clone()));
300        let _previous = inner.last_refresh.insert(provider_name.to_owned(), Instant::now());
301
302        // Update bloom filter with all tool names
303        inner.bloom_filter.clear(); // Clear and rebuild for accuracy
304
305        let all_tool_names: Vec<String> = inner
306            .all_tools_cache
307            .values()
308            .flat_map(|provider_tools| provider_tools.iter().map(|tool| tool.name.clone()))
309            .collect();
310
311        for tool_name in all_tool_names {
312            inner.bloom_filter.insert(&tool_name);
313        }
314    }
315
316    /// Cache a single tool result (for read-only tools)
317    pub fn cache_tool_result(&self, _cache_key: String, _result: serde_json::Value) {
318        // This would be implemented for caching individual tool execution results
319        // For now, we'll just store it in a simple cache
320        // In a full implementation, this would use a separate cache with different TTL
321    }
322
323    /// Clear all caches
324    pub fn clear(&self) {
325        if let Ok(mut inner) = self.inner.write() {
326            inner.bloom_filter.clear();
327            inner.detailed_cache.clear();
328            inner.all_tools_cache.clear();
329            inner.last_refresh.clear();
330        }
331    }
332
333    /// Get cache statistics
334    pub(crate) fn stats(&self) -> ToolCacheStats {
335        let (detailed_entries, detailed_capacity, all_tools_entries, bf_size, bf_hashes) = self
336            .inner
337            .read()
338            .map(|inner| {
339                (
340                    inner.detailed_cache.len(),
341                    inner.detailed_cache.cap().get(),
342                    inner.all_tools_cache.len(),
343                    inner.bloom_filter.size,
344                    inner.bloom_filter.num_hashes,
345                )
346            })
347            .unwrap_or((0, 0, 0, 0, 0));
348
349        ToolCacheStats {
350            detailed_cache_entries: detailed_entries,
351            detailed_cache_capacity: detailed_capacity,
352            all_tools_cache_entries: all_tools_entries,
353            bloom_filter_size: bf_size,
354            bloom_filter_hashes: bf_hashes,
355        }
356    }
357}
358
359/// Cache statistics for monitoring
360#[derive(Debug, Clone)]
361pub struct ToolCacheStats {
362    detailed_cache_entries: usize,
363    detailed_cache_capacity: usize,
364    all_tools_cache_entries: usize,
365    bloom_filter_size: usize,
366    bloom_filter_hashes: usize,
367}
368
369/// Enhanced tool discovery with caching
370pub struct CachedToolDiscovery {
371    cache: Arc<ToolDiscoveryCache>,
372}
373
374impl CachedToolDiscovery {
375    pub fn new(cache_capacity: usize) -> Self {
376        Self {
377            cache: Arc::new(ToolDiscoveryCache::new(cache_capacity)),
378        }
379    }
380
381    /// Search for tools with multi-level caching
382    pub fn search_tools(
383        &self,
384        provider_name: &str,
385        keyword: &str,
386        detail_level: DetailLevel,
387        all_tools: Vec<McpToolInfo>,
388    ) -> Arc<Vec<ToolDiscoveryResult>> {
389        // Check bloom filter first (fast negative lookup)
390        if !self.cache.might_have_tool(keyword) && !keyword.is_empty() {
391            return Arc::new(Vec::new());
392        }
393
394        // Check detailed cache
395        if let Some(cached) = self.cache.get_cached_discovery(provider_name, keyword, detail_level) {
396            return cached;
397        }
398
399        // Perform the search
400        let results = Arc::new(self.perform_search(&all_tools, keyword, detail_level));
401
402        // Cache the results
403        self.cache
404            .cache_discovery_shared(provider_name, keyword, detail_level, Arc::clone(&results));
405
406        results
407    }
408
409    /// Get all tools for a provider with caching
410    pub fn get_all_tools_cached(&self, provider_name: &str, all_tools: Vec<McpToolInfo>) -> Vec<McpToolInfo> {
411        // Check cache first
412        if let Some(cached) = self.cache.get_all_tools(provider_name, true) {
413            return cached;
414        }
415
416        // Cache the results
417        self.cache.cache_all_tools(provider_name, all_tools.clone());
418
419        all_tools
420    }
421
422    /// Perform the actual search on tool list
423    fn perform_search(
424        &self,
425        tools: &[McpToolInfo],
426        keyword: &str,
427        detail_level: DetailLevel,
428    ) -> Vec<ToolDiscoveryResult> {
429        let keyword_lower = keyword.to_lowercase();
430        let mut results = Vec::new();
431
432        for tool in tools {
433            let relevance_score = self.calculate_relevance(tool, &keyword_lower);
434
435            if relevance_score > 0.0 {
436                let result = ToolDiscoveryResult { tool: tool.clone(), relevance_score, detail_level };
437                results.push(result);
438            }
439        }
440
441        // Sort by relevance score
442        results.sort_by(|a, b| {
443            b.relevance_score
444                .partial_cmp(&a.relevance_score)
445                .unwrap_or(std::cmp::Ordering::Equal)
446        });
447
448        results
449    }
450
451    /// Calculate relevance score for a tool
452    fn calculate_relevance(&self, tool: &McpToolInfo, keyword: &str) -> f64 {
453        let name_lower = tool.name.to_lowercase();
454        let description_lower = tool.description.to_lowercase();
455
456        let mut score: f64 = 0.0;
457
458        // Name exact match
459        if name_lower == keyword {
460            score += 1.0;
461        }
462        // Name starts with keyword
463        else if name_lower.starts_with(keyword) {
464            score += 0.8;
465        }
466        // Name contains keyword
467        else if name_lower.contains(keyword) {
468            score += 0.6;
469        }
470
471        // Description contains keyword
472        if description_lower.contains(keyword) {
473            score += 0.3;
474        }
475
476        // Input schema contains keyword
477        let schema_str = serde_json::to_string(&tool.input_schema).unwrap_or_default().to_lowercase();
478        if schema_str.contains(keyword) {
479            score += 0.2;
480        }
481
482        // Fuzzy fallback using Sørensen-Dice for partial keyword matches
483        if score == 0.0 {
484            let sd_name = strsim::sorensen_dice(&name_lower, keyword);
485            let sd_desc = strsim::sorensen_dice(&description_lower, keyword);
486            let max_sd = sd_name.max(sd_desc);
487            if max_sd > 0.3 {
488                score = max_sd * 0.5;
489            }
490        }
491
492        score.min(1.0)
493    }
494
495    /// Get cache statistics
496    pub fn stats(&self) -> ToolCacheStats {
497        self.cache.stats()
498    }
499}
500
501#[cfg(test)]
502mod tests {
503    use super::*;
504
505    #[test]
506    fn test_bloom_filter() {
507        let mut filter = BloomFilter::new(100, 0.01);
508
509        filter.insert("tool1");
510        filter.insert("tool2");
511        filter.insert("tool3");
512
513        assert!(filter.contains("tool1"));
514        assert!(filter.contains("tool2"));
515        assert!(filter.contains("tool3"));
516        assert!(!filter.contains("tool4"));
517    }
518
519    #[test]
520    fn test_cache_key_equality() {
521        let key1 = ToolDiscoveryCacheKey {
522            provider_name: "test".to_string(),
523            keyword: "search".to_string(),
524            detail_level: DetailLevel::Full,
525        };
526
527        let key2 = ToolDiscoveryCacheKey {
528            provider_name: "test".to_string(),
529            keyword: "search".to_string(),
530            detail_level: DetailLevel::Full,
531        };
532
533        assert_eq!(key1, key2);
534    }
535
536    #[test]
537    fn test_tool_discovery_cache() {
538        let cache = ToolDiscoveryCache::new(10);
539
540        let provider_name = "test_provider";
541        let keyword = "search";
542        let detail_level = DetailLevel::Full;
543
544        // Cache miss
545        assert!(cache.get_cached_discovery(provider_name, keyword, detail_level).is_none());
546
547        // Cache some results
548        let results = vec![ToolDiscoveryResult {
549            tool: McpToolInfo {
550                name: "search_files".to_string(),
551                description: "Search for files".to_string(),
552                provider: "test".to_string(),
553                input_schema: serde_json::json!({}),
554                output_schema: None,
555            },
556            relevance_score: 0.9,
557            detail_level,
558        }];
559
560        cache.cache_discovery(provider_name, keyword, detail_level, results.clone());
561
562        // Cache hit
563        let cached = cache.get_cached_discovery(provider_name, keyword, detail_level);
564        assert!(cached.is_some());
565        assert_eq!(cached.unwrap().len(), 1);
566    }
567}