1use 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#[derive(Clone)]
18pub struct BloomFilter {
19 bits: Vec<bool>,
21 num_hashes: usize,
23 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 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 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 pub fn clear(&mut self) {
61 self.bits.fill(false);
62 }
63
64 #[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 #[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 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#[derive(Debug, Clone, Hash, PartialEq, Eq)]
106struct ToolDiscoveryCacheKey {
107 provider_name: String,
108 keyword: String,
109 detail_level: DetailLevel,
110}
111
112#[derive(Clone)]
114struct CachedToolDiscoveryEntry {
115 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#[derive(Debug, Clone)]
129pub struct ToolDiscoveryResult {
130 tool: McpToolInfo,
131 relevance_score: f64,
132 detail_level: DetailLevel,
133}
134
135pub(crate) struct ToolDiscoveryCache {
137 inner: Arc<RwLock<DiscoveryCacheInner>>,
138 config: CacheConfig,
140}
141
142#[derive(Clone)]
143struct CacheConfig {
144 max_age: Duration,
146 provider_refresh_interval: Duration,
148 expected_tool_count: usize,
150 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), provider_refresh_interval: Duration::from_secs(60), expected_tool_count: 1000,
160 false_positive_rate: 0.01, };
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 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 fn get_cached_discovery(
190 &self,
191 provider_name: &str,
192 keyword: &str,
193 detail_level: DetailLevel,
194 ) -> Option<Arc<Vec<ToolDiscoveryResult>>> {
195 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 if cached.timestamp.elapsed() < self.config.max_age {
213 return Some(Arc::clone(&cached.results));
214 } else {
215 drop(inner.detailed_cache.pop(&key));
217 }
218 }
219
220 None
221 }
222
223 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 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 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 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; }
285
286 inner.all_tools_cache.get(provider_name).cloned()
287 }
288
289 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 inner.bloom_filter.clear(); 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 pub fn cache_tool_result(&self, _cache_key: String, _result: serde_json::Value) {
318 }
322
323 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 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#[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
369pub 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 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 if !self.cache.might_have_tool(keyword) && !keyword.is_empty() {
391 return Arc::new(Vec::new());
392 }
393
394 if let Some(cached) = self.cache.get_cached_discovery(provider_name, keyword, detail_level) {
396 return cached;
397 }
398
399 let results = Arc::new(self.perform_search(&all_tools, keyword, detail_level));
401
402 self.cache
404 .cache_discovery_shared(provider_name, keyword, detail_level, Arc::clone(&results));
405
406 results
407 }
408
409 pub fn get_all_tools_cached(&self, provider_name: &str, all_tools: Vec<McpToolInfo>) -> Vec<McpToolInfo> {
411 if let Some(cached) = self.cache.get_all_tools(provider_name, true) {
413 return cached;
414 }
415
416 self.cache.cache_all_tools(provider_name, all_tools.clone());
418
419 all_tools
420 }
421
422 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 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 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 if name_lower == keyword {
460 score += 1.0;
461 }
462 else if name_lower.starts_with(keyword) {
464 score += 0.8;
465 }
466 else if name_lower.contains(keyword) {
468 score += 0.6;
469 }
470
471 if description_lower.contains(keyword) {
473 score += 0.3;
474 }
475
476 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 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 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 assert!(cache.get_cached_discovery(provider_name, keyword, detail_level).is_none());
546
547 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 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}