1use crate::McpToolInfo;
20use anyhow::Result;
21use serde_json::Value;
22use std::cmp::Ordering;
23use std::sync::Arc;
24use tracing::{debug, info};
25
26#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
28pub enum DetailLevel {
29 NameOnly,
31 NameAndDescription,
33 Full,
35}
36
37impl DetailLevel {
38 pub fn as_str(&self) -> &'static str {
39 match self {
40 Self::NameOnly => "name-only",
41 Self::NameAndDescription => "name-and-description",
42 Self::Full => "full",
43 }
44 }
45}
46
47#[derive(Debug, Clone, serde::Serialize)]
49pub struct ToolDiscoveryResult {
50 pub name: String,
51 pub provider: String,
52 description: String,
53 relevance_score: f32,
54 input_schema: Option<Value>,
56 output_schema: Option<Value>,
58}
59
60impl ToolDiscoveryResult {
61 pub fn to_json(&self, detail_level: DetailLevel) -> Value {
63 match detail_level {
64 DetailLevel::NameOnly => serde_json::json!({
65 "name": self.name,
66 "provider": self.provider,
67 }),
68 DetailLevel::NameAndDescription => serde_json::json!({
69 "name": self.name,
70 "provider": self.provider,
71 "description": self.description,
72 }),
73 DetailLevel::Full => {
74 let mut item = serde_json::json!({
75 "name": self.name,
76 "provider": self.provider,
77 "description": self.description,
78 "input_schema": self.input_schema,
79 });
80 if let Some(schema) = self.output_schema.as_ref()
81 && let Some(object) = item.as_object_mut()
82 {
83 drop(object.insert("output_schema".to_string(), schema.clone()));
84 }
85 item
86 }
87 }
88 }
89}
90
91pub struct ToolDiscovery {
93 mcp_client: Arc<dyn crate::McpToolExecutor>,
94}
95
96fn group_results_by_provider_preserving_order(
97 tools: impl IntoIterator<Item = ToolDiscoveryResult>,
98) -> Vec<(String, Vec<ToolDiscoveryResult>)> {
99 let mut grouped: Vec<(String, Vec<ToolDiscoveryResult>)> = Vec::new();
100
101 for tool in tools {
102 let provider = tool.provider.clone();
103 if let Some((_, provider_tools)) =
104 grouped.iter_mut().find(|(existing_provider, _)| *existing_provider == provider)
105 {
106 provider_tools.push(tool);
107 } else {
108 grouped.push((provider, vec![tool]));
109 }
110 }
111
112 grouped
113}
114
115impl ToolDiscovery {
116 pub fn new(mcp_client: Arc<dyn crate::McpToolExecutor>) -> Self {
118 Self { mcp_client }
119 }
120
121 pub async fn search_tools(&self, keyword: &str, detail_level: DetailLevel) -> Result<Vec<ToolDiscoveryResult>> {
129 let tools = self.mcp_client.list_mcp_tools().await?;
130
131 debug!(keyword = keyword, count = tools.len(), "Searching tools for keyword");
132
133 let mut scored: Vec<(&McpToolInfo, f32)> = Vec::with_capacity(tools.len() / 4);
137 for tool in &tools {
138 let relevance_score = self.calculate_relevance(tool, keyword);
139
140 if relevance_score > 0.0 {
142 scored.push((tool, relevance_score));
143 }
144 }
145
146 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal));
150
151 let total_results = scored.len();
153 if total_results > 5 {
154 info!(
155 keyword = keyword,
156 matched = total_results,
157 displayed = 5,
158 overflow = total_results - 5,
159 detail_level = detail_level.as_str(),
160 "Tool search completed with overflow"
161 );
162 scored.truncate(5);
163 } else {
164 info!(
165 keyword = keyword,
166 matched = total_results,
167 detail_level = detail_level.as_str(),
168 "Tool search completed"
169 );
170 }
171
172 let mut results = Vec::with_capacity(scored.len());
174 for (tool, relevance_score) in scored {
175 let (input_schema, output_schema) = match detail_level {
177 DetailLevel::Full => (Some(tool.input_schema.clone()), tool.output_schema.clone()),
178 _ => (None, None),
179 };
180
181 results.push(ToolDiscoveryResult {
182 name: tool.name.clone(),
183 provider: tool.provider.clone(),
184 description: tool.description.clone(),
185 relevance_score,
186 input_schema,
187 output_schema,
188 });
189 }
190
191 Ok(results)
192 }
193
194 pub async fn get_tool_detail(&self, tool_name: &str) -> Result<Option<ToolDiscoveryResult>> {
196 let tools = self.mcp_client.list_mcp_tools().await?;
197
198 for tool in tools {
199 if tool.name.eq_ignore_ascii_case(tool_name) {
200 return Ok(Some(ToolDiscoveryResult {
201 name: tool.name.clone(),
202 provider: tool.provider.clone(),
203 description: tool.description.clone(),
204 relevance_score: 1.0,
205 input_schema: Some(tool.input_schema),
206 output_schema: tool.output_schema,
207 }));
208 }
209 }
210
211 Ok(None)
212 }
213
214 async fn list_tools_by_provider(&self) -> Result<Vec<(String, Vec<ToolDiscoveryResult>)>> {
216 let tools = self.mcp_client.list_mcp_tools().await?;
217
218 Ok(group_results_by_provider_preserving_order(tools.into_iter().map(|tool| ToolDiscoveryResult {
219 name: tool.name,
220 provider: tool.provider,
221 description: tool.description,
222 relevance_score: 1.0,
223 input_schema: None,
224 output_schema: None,
225 })))
226 }
227
228 fn calculate_relevance(&self, tool: &McpToolInfo, keyword: &str) -> f32 {
232 let keyword_lower = keyword.to_lowercase();
233
234 if tool.name.eq_ignore_ascii_case(keyword) {
236 return 1.0;
237 }
238
239 if tool.name.to_lowercase().contains(&keyword_lower) {
241 return 0.8;
242 }
243
244 if tool.description.to_lowercase().contains(&keyword_lower) {
246 return 0.6;
247 }
248
249 let name_fuzzy = self.fuzzy_score(&tool.name.to_lowercase(), &keyword_lower);
251 if name_fuzzy > 0.3 {
252 return 0.5 * name_fuzzy;
253 }
254
255 let desc_fuzzy = self.fuzzy_score(&tool.description.to_lowercase(), &keyword_lower);
257 if desc_fuzzy > 0.2 {
258 return 0.3 * desc_fuzzy;
259 }
260
261 0.0
262 }
263
264 #[expect(
270 clippy::cast_possible_truncation,
271 reason = "Sørensen-Dice is normalized to the [0, 1] range, so the f32 score remains bounded."
272 )]
273 fn fuzzy_score(&self, haystack: &str, needle: &str) -> f32 {
274 if needle.is_empty() {
275 return 1.0;
276 }
277 if haystack.is_empty() {
278 return 0.0;
279 }
280 strsim::sorensen_dice(haystack, needle) as f32
281 }
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287 use serde_json::json;
288
289 fn mock_tool(provider: &str, name: &str, description: &str) -> McpToolInfo {
290 McpToolInfo {
291 name: name.to_string(),
292 description: description.to_string(),
293 provider: provider.to_string(),
294 input_schema: json!({}),
295 output_schema: None,
296 }
297 }
298
299 #[test]
300 fn fuzzy_score_exact_match() {
301 let discovery = ToolDiscovery::new(Arc::new(MockMcpClient::default()));
302 assert!((discovery.fuzzy_score("read_file", "read_file") - 1.0).abs() < f32::EPSILON);
303 }
304
305 #[test]
306 fn fuzzy_score_partial_match() {
307 let discovery = ToolDiscovery::new(Arc::new(MockMcpClient::default()));
309 let score = discovery.fuzzy_score("read_file", "read");
310 assert!(score > 0.5 && score <= 1.0, "expected >0.5, got {score}");
311 }
312
313 #[test]
314 fn fuzzy_score_no_match() {
315 let discovery = ToolDiscovery::new(Arc::new(MockMcpClient::default()));
316 assert!(discovery.fuzzy_score("read_file", "xyz").abs() < f32::EPSILON);
317 }
318
319 #[test]
320 fn full_detail_json_includes_output_schema_only_when_advertised() {
321 let with_schema = ToolDiscoveryResult {
322 name: "ask".to_string(),
323 provider: "deepwiki".to_string(),
324 description: "Ask.".to_string(),
325 relevance_score: 1.0,
326 input_schema: Some(json!({"type": "object"})),
327 output_schema: Some(json!({"type": "object"})),
328 };
329 assert_eq!(with_schema.to_json(DetailLevel::Full)["output_schema"], json!({"type": "object"}));
330
331 let without_schema = ToolDiscoveryResult { output_schema: None, ..with_schema.clone() };
332 let full = without_schema.to_json(DetailLevel::Full);
333 assert!(full.get("output_schema").is_none(), "absent schema must stay absent");
334 assert!(
335 with_schema
336 .to_json(DetailLevel::NameAndDescription)
337 .get("output_schema")
338 .is_none(),
339 "compact levels must not carry schemas"
340 );
341 }
342
343 #[tokio::test]
344 async fn list_tools_by_provider_preserves_first_seen_provider_and_tool_order() {
345 let discovery = ToolDiscovery::new(Arc::new(MockMcpClient {
346 tools: vec![
347 mock_tool("gmail", "send_email", "Send an email."),
348 mock_tool("calendar", "create_event", "Create a calendar event."),
349 mock_tool("gmail", "read_email", "Read an email."),
350 mock_tool("docs", "search", "Search docs."),
351 mock_tool("calendar", "list_events", "List calendar events."),
352 ],
353 }));
354
355 let grouped = discovery.list_tools_by_provider().await.expect("grouped tools");
356
357 let providers = grouped.iter().map(|(provider, _)| provider.as_str()).collect::<Vec<_>>();
358 assert_eq!(providers, vec!["gmail", "calendar", "docs"]);
359
360 let tool_names = grouped
361 .into_iter()
362 .map(|(_, tools)| tools.into_iter().map(|tool| tool.name).collect::<Vec<_>>())
363 .collect::<Vec<_>>();
364 assert_eq!(
365 tool_names,
366 vec![
367 vec!["send_email".to_string(), "read_email".to_string()],
368 vec!["create_event".to_string(), "list_events".to_string()],
369 vec!["search".to_string()],
370 ]
371 );
372 }
373
374 #[derive(Default)]
376 struct MockMcpClient {
377 tools: Vec<McpToolInfo>,
378 }
379
380 #[tokio::test]
381 async fn search_tools_keeps_highest_scores_despite_late_position_and_ties() {
382 let discovery = ToolDiscovery::new(Arc::new(MockMcpClient {
389 tools: vec![
390 mock_tool("prov", "calendar", "Show the calendar."),
391 mock_tool("prov", "docs", "Search the docs."),
392 mock_tool("prov", "forward_mail", "Forward a message."),
393 mock_tool("prov", "send_mail", "Send a message."),
394 mock_tool("prov", "mail", "Mail things."),
395 mock_tool("prov", "read_mail", "Read a message."),
396 mock_tool("prov", "delete_mail", "Delete a message."),
397 mock_tool("prov", "archive", "Archive old mail threads."),
398 ],
399 }));
400
401 let results = discovery.search_tools("mail", DetailLevel::Full).await.expect("search tools");
403
404 let names = results.iter().map(|result| result.name.as_str()).collect::<Vec<_>>();
406 assert_eq!(names, vec!["mail", "forward_mail", "send_mail", "read_mail", "delete_mail"]);
407 let scores = results.iter().map(|result| result.relevance_score).collect::<Vec<_>>();
408 assert_eq!(scores, vec![1.0, 0.8, 0.8, 0.8, 0.8]);
409 assert!(results.iter().all(|result| result.input_schema.is_some()));
410
411 let compact = discovery
413 .search_tools("mail", DetailLevel::NameAndDescription)
414 .await
415 .expect("compact search");
416 let compact_names = compact.iter().map(|result| result.name.as_str()).collect::<Vec<_>>();
417 assert_eq!(compact_names, names);
418 assert!(
419 compact
420 .iter()
421 .all(|result| result.input_schema.is_none() && result.output_schema.is_none())
422 );
423 }
424
425 #[async_trait::async_trait]
426 impl crate::McpToolExecutor for MockMcpClient {
427 async fn execute_mcp_tool(&self, _tool_name: &str, _args: &Value) -> Result<Value> {
428 Ok(Value::Null)
429 }
430
431 async fn list_mcp_tools(&self) -> Result<Vec<McpToolInfo>> {
432 Ok(self.tools.clone())
433 }
434
435 async fn has_mcp_tool(&self, _tool_name: &str) -> Result<bool> {
436 Ok(false)
437 }
438
439 fn get_status(&self) -> crate::McpClientStatus {
440 crate::McpClientStatus {
441 enabled: true,
442 provider_count: 0,
443 active_connections: 0,
444 configured_providers: vec![],
445 }
446 }
447 }
448}