1use anyhow::{Context, Result, anyhow, bail};
2use async_trait::async_trait;
3use chrono::Utc;
4use parking_lot::RwLock;
5use rmcp::model::{CallToolResult, ClientCapabilities, InitializeRequestParams, RootsCapabilities};
6use rustc_hash::FxHashMap;
7use serde_json::{Map, Value, json};
8use std::collections::BTreeMap;
9use std::path::{Path, PathBuf};
10use std::sync::Arc;
11use std::time::Duration;
12use tracing::{info, warn};
13use vtcode_commons::fs::{ensure_dir_exists, write_file_with_context};
14use vtcode_config::mcp::{McpAllowListConfig, McpClientConfig, McpProviderConfig, McpTransportConfig};
15
16use super::McpSandboxContext;
17use super::{
18 McpClientStatus, McpElicitationHandler, McpPromptDetail, McpPromptInfo, McpProvider, McpResourceData,
19 McpResourceInfo, McpToolExecutor, McpToolInfo, format_tool_markdown, sanitize_filename,
20};
21use crate::connection_pool::McpConnectionPool;
22
23struct McpClientState {
24 providers: FxHashMap<String, Arc<McpProvider>>,
25 allowlist: McpAllowListConfig,
26 tool_provider_index: FxHashMap<String, String>,
27 resource_provider_index: FxHashMap<String, String>,
28 prompt_provider_index: FxHashMap<String, String>,
29}
30
31fn providers_sorted_by_name<V: Clone>(providers: &FxHashMap<String, V>) -> Vec<V> {
33 let mut entries: Vec<(&String, &V)> = providers.iter().collect();
34 entries.sort_unstable_by(|a, b| a.0.cmp(b.0));
35 entries.into_iter().map(|(_, provider)| provider.clone()).collect()
36}
37
38pub struct McpClient {
39 config: McpClientConfig,
40 state: RwLock<McpClientState>,
41 elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
42 sandbox_context: Option<McpSandboxContext>,
43}
44
45impl McpClient {
46 pub fn new(config: McpClientConfig) -> Self {
48 Self::with_sandbox_context(config, None)
49 }
50
51 pub fn with_sandbox_context(config: McpClientConfig, sandbox_context: Option<McpSandboxContext>) -> Self {
53 let allowlist = config.allowlist.clone();
54
55 Self {
56 config,
57 state: RwLock::new(McpClientState {
58 providers: FxHashMap::default(),
59 allowlist,
60 tool_provider_index: FxHashMap::default(),
61 resource_provider_index: FxHashMap::default(),
62 prompt_provider_index: FxHashMap::default(),
63 }),
64 elicitation_handler: None,
65 sandbox_context,
66 }
67 }
68
69 pub fn set_elicitation_handler(&mut self, handler: Arc<dyn McpElicitationHandler>) {
71 self.elicitation_handler = Some(handler);
72 }
73
74 pub async fn initialize(&mut self) -> Result<()> {
77 if !self.config.enabled {
78 info!("MCP client is disabled in configuration");
79 return Ok(());
80 }
81
82 info!("Initializing MCP client with {} configured providers", self.config.providers.len());
83
84 let max_connections = self.config.max_concurrent_connections.max(1);
85 let pool = McpConnectionPool::new(max_connections, self.config.request_timeout_seconds);
86 let tool_timeout = Some(Duration::from_secs(self.config.request_timeout_seconds));
87 let allowlist_snapshot = self.state.read().allowlist.clone();
88
89 let provider_configs: Vec<McpProviderConfig> =
90 self.config.providers.iter().filter(|c| c.enabled).cloned().collect();
91
92 let results = pool
93 .initialize_providers_parallel(
94 provider_configs,
95 self.elicitation_handler.clone(),
96 self.sandbox_context.clone(),
97 tool_timeout,
98 &allowlist_snapshot,
99 )
100 .await
101 .map_err(|e| anyhow::anyhow!("MCP connection pool initialization failed: {e}"))?;
102
103 let mut initialized = FxHashMap::default();
104 for (name, provider) in results {
105 drop(initialized.insert(name, provider));
106 }
107
108 self.state.write().providers = initialized;
109 info!("MCP client initialization complete. Active providers: {}", self.state.read().providers.len());
110
111 Ok(())
112 }
113
114 fn validate_tool_arguments(&self, _tool_name: &str, args: &Value) -> Result<()> {
116 if self.config.security.validation.max_argument_size > 0 {
118 let args_size = serde_json::to_string(args).map_or(0, |s| u32::try_from(s.len()).unwrap_or(u32::MAX));
119
120 if args_size > self.config.security.validation.max_argument_size {
121 return Err(anyhow::anyhow!(
122 "Tool arguments exceed maximum size of {} bytes",
123 self.config.security.validation.max_argument_size
124 ));
125 }
126 }
127
128 if self.config.security.validation.path_traversal_protection
130 && let Some(path) = args.get("path").and_then(|v| v.as_str())
131 && (path.contains("../") || path.starts_with("../") || path.contains("..\\") || path.starts_with("..\\"))
132 {
133 return Err(anyhow::anyhow!("Path traversal detected in arguments"));
134 }
135
136 Ok(())
137 }
138
139 async fn execute_tool_with_validation(&self, tool_name: &str, args: Value) -> Result<Value> {
145 self.execute_tool_with_validation_ref(tool_name, &args).await
146 }
147
148 async fn execute_tool_with_validation_ref(&self, tool_name: &str, args: &Value) -> Result<Value> {
150 if !self.config.enabled {
151 return Err(anyhow!("MCP support is disabled in the current configuration"));
152 }
153
154 self.validate_tool_arguments(tool_name, args)?;
155
156 let provider = self.resolve_provider_for_tool(tool_name).await?;
157 let allowlist_snapshot = self.state.read().allowlist.clone();
158 let result = provider
159 .call_tool(tool_name, args, self.tool_timeout(), &allowlist_snapshot)
160 .await?;
161
162 Self::format_tool_result(&provider.name, tool_name, result)
163 }
164
165 pub fn update_allowlist(&self, allowlist: McpAllowListConfig) {
167 let providers: Vec<Arc<McpProvider>> = {
168 let mut state = self.state.write();
169 state.allowlist = allowlist;
170 state.tool_provider_index.clear();
171 state.resource_provider_index.clear();
172 state.prompt_provider_index.clear();
173 state.providers.values().cloned().collect()
174 };
175
176 for provider in providers {
177 provider.invalidate_caches();
178 }
179 }
180
181 pub fn current_allowlist(&self) -> McpAllowListConfig {
183 self.state.read().allowlist.clone()
184 }
185
186 pub fn provider_for_tool(&self, tool_name: &str) -> Option<String> {
188 self.state.read().tool_provider_index.get(tool_name).cloned()
189 }
190
191 pub fn provider_for_resource(&self, uri: &str) -> Option<String> {
193 self.state.read().resource_provider_index.get(uri).cloned()
194 }
195
196 pub fn provider_for_prompt(&self, prompt_name: &str) -> Option<String> {
198 self.state.read().prompt_provider_index.get(prompt_name).cloned()
199 }
200
201 pub async fn execute_tool(&self, tool_name: &str, args: Value) -> Result<Value> {
203 self.execute_tool_with_validation(tool_name, args).await
204 }
205
206 pub async fn list_tools(&self) -> Result<Vec<McpToolInfo>> {
208 self.collect_tools(false).await
209 }
210
211 pub async fn list_resources(&self) -> Result<Vec<McpResourceInfo>> {
213 self.collect_resources(false).await
214 }
215
216 pub async fn refresh_resources(&self) -> Result<Vec<McpResourceInfo>> {
218 self.collect_resources(true).await
219 }
220
221 pub async fn list_prompts(&self) -> Result<Vec<McpPromptInfo>> {
223 self.collect_prompts(false).await
224 }
225
226 pub async fn refresh_prompts(&self) -> Result<Vec<McpPromptInfo>> {
228 self.collect_prompts(true).await
229 }
230
231 pub async fn read_resource(&self, uri: &str) -> Result<McpResourceData> {
233 let provider = self.resolve_provider_for_resource(uri).await?;
234 let provider_name = provider.name.clone();
235 let allowlist_snapshot = self.state.read().allowlist.clone();
236 let data = provider.read_resource(uri, self.request_timeout(), &allowlist_snapshot).await?;
237 drop(self.state.write().resource_provider_index.insert(uri.into(), provider_name));
238 Ok(data)
239 }
240
241 pub async fn get_prompt(
243 &self,
244 prompt_name: &str,
245 arguments: Option<hashbrown::HashMap<String, String>>,
246 ) -> Result<McpPromptDetail> {
247 let provider = self.resolve_provider_for_prompt(prompt_name).await?;
248 let provider_name = provider.name.clone();
249 let allowlist_snapshot = self.state.read().allowlist.clone();
250 let prompt = provider
251 .get_prompt(prompt_name, arguments.unwrap_or_default(), self.request_timeout(), &allowlist_snapshot)
252 .await?;
253 drop(
254 self.state
255 .write()
256 .prompt_provider_index
257 .insert(prompt_name.into(), provider_name),
258 );
259 Ok(prompt)
260 }
261
262 pub async fn shutdown(&self) -> Result<()> {
264 let providers: Vec<Arc<McpProvider>> = {
265 let mut state = self.state.write();
266 let values: Vec<_> = state.providers.values().cloned().collect();
267 state.providers.clear();
268 state.tool_provider_index.clear();
269 state.resource_provider_index.clear();
270 state.prompt_provider_index.clear();
271 values
272 };
273
274 if providers.is_empty() {
275 info!("No active MCP connections to shutdown");
276 return Ok(());
277 }
278
279 info!("Shutting down {} MCP providers", providers.len());
280 for provider in providers {
281 if let Err(err) = provider.shutdown().await {
282 warn!("Provider '{}' shutdown returned error: {err}", provider.name);
283 }
284 }
285 Ok(())
286 }
287
288 pub fn get_status(&self) -> McpClientStatus {
290 let state = self.state.read();
291 let providers = &state.providers;
292 let mut configured_providers: Vec<String> = providers.keys().cloned().collect();
294 configured_providers.sort_unstable();
295 McpClientStatus {
296 enabled: self.config.enabled,
297 provider_count: providers.len(),
298 active_connections: providers.len(),
299 configured_providers,
300 }
301 }
302
303 pub async fn list_servers(&self) -> Vec<Value> {
309 let live: Vec<(McpProviderConfig, Option<Arc<McpProvider>>)> = {
310 let state = self.state.read();
311 self.config
312 .providers
313 .iter()
314 .map(|provider_config| (provider_config.clone(), state.providers.get(&provider_config.name).cloned()))
315 .collect()
316 };
317 let mut servers = Vec::with_capacity(live.len());
318 for (provider_config, provider) in live {
319 let negotiated = match &provider {
320 Some(provider) => provider.negotiated_protocol_version().await,
321 None => None,
322 };
323 let connected = provider.is_some();
324 let (transport, target) = match &provider_config.transport {
325 McpTransportConfig::Stdio(stdio) => ("stdio", Value::String(stdio.command.clone())),
326 McpTransportConfig::Http(http) => ("http", Value::String(http.endpoint.clone())),
327 };
328
329 servers.push(json!({
330 "name": provider_config.name,
331 "enabled": provider_config.enabled,
332 "connected": connected,
333 "connection_state": if connected { "connected" } else { "disconnected" },
334 "transport": transport,
335 "target": target,
336 "negotiated_protocol_version": negotiated.map(Value::String).unwrap_or(Value::Null),
337 }));
338 }
339 servers
340 }
341
342 pub fn allow_model_lifecycle_control(&self) -> bool {
344 self.config.lifecycle.allow_model_control
345 }
346
347 pub async fn connect_server(&self, server_name: &str) -> Result<()> {
349 if !self.config.enabled {
350 bail!("MCP support is disabled in the current configuration");
351 }
352
353 if self.state.read().providers.contains_key(server_name) {
354 return Ok(());
355 }
356
357 let provider_config = self
358 .config
359 .providers
360 .iter()
361 .find(|provider| provider.name == server_name)
362 .cloned()
363 .ok_or_else(|| anyhow!("MCP server '{server_name}' is not configured"))?;
364
365 if !provider_config.enabled {
366 bail!("MCP server '{server_name}' is configured but disabled");
367 }
368
369 if let Some(reason) = self.requirement_mismatch_reason(&provider_config) {
370 bail!("Cannot connect MCP server '{}': {}", provider_config.name, reason);
371 }
372
373 if matches!(provider_config.transport, McpTransportConfig::Http(_)) && !self.config.experimental_use_rmcp_client
374 {
375 bail!(
376 "Cannot connect MCP HTTP server '{}' while experimental_use_rmcp_client is disabled",
377 provider_config.name
378 );
379 }
380
381 let allowlist_snapshot = self.state.read().allowlist.clone();
382 let tool_timeout = self.tool_timeout();
383 let provider = self
384 .connect_and_initialize_provider(&provider_config, &allowlist_snapshot, tool_timeout)
385 .await?;
386
387 if let Err(err) = provider.cached_tools_or_refresh_shared(&allowlist_snapshot, tool_timeout).await {
388 warn!("Connected MCP server '{}' but failed to refresh tools: {err}", server_name);
389 } else if let Some(cache) = provider.cached_tools_shared().await {
390 self.record_tool_provider(&provider.name, &cache);
391 }
392
393 drop(self.state.write().providers.insert(provider.name.clone(), Arc::new(provider)));
394 Ok(())
395 }
396
397 pub async fn disconnect_server(&self, server_name: &str) -> Result<()> {
399 let provider = {
400 let mut state = self.state.write();
401 let provider = state
402 .providers
403 .remove(server_name)
404 .ok_or_else(|| anyhow!("MCP server '{server_name}' is not connected"))?;
405 state
406 .tool_provider_index
407 .retain(|_, provider_name| provider_name != server_name);
408 state
409 .resource_provider_index
410 .retain(|_, provider_name| provider_name != server_name);
411 state
412 .prompt_provider_index
413 .retain(|_, provider_name| provider_name != server_name);
414 provider
415 };
416
417 provider.shutdown().await?;
418 Ok(())
419 }
420
421 pub async fn sync_tools_to_files(&self, workspace_root: &Path) -> Result<(PathBuf, usize)> {
430 let tools = self.list_tools().await?;
431 let mcp_dir = workspace_root.join(".vtcode").join("mcp");
432 let tools_dir = mcp_dir.join("tools");
433
434 ensure_dir_exists(&tools_dir)
436 .await
437 .with_context(|| format!("Failed to create MCP tools directory: {}", tools_dir.display()))?;
438
439 let mut by_provider: BTreeMap<String, Vec<&McpToolInfo>> = BTreeMap::new();
441 for tool in &tools {
442 by_provider.entry(tool.provider.clone()).or_default().push(tool);
443 }
444
445 for (provider, provider_tools) in &by_provider {
447 let provider_dir = tools_dir.join(sanitize_filename(provider));
448 ensure_dir_exists(&provider_dir)
449 .await
450 .with_context(|| format!("Failed to create provider directory: {}", provider_dir.display()))?;
451
452 for tool in provider_tools {
453 let tool_content = format_tool_markdown(tool);
454 let tool_path = provider_dir.join(format!("{}.md", sanitize_filename(&tool.name)));
455 write_file_with_context(&tool_path, &tool_content, "MCP tool file")
456 .await
457 .with_context(|| format!("Failed to write tool file: {}", tool_path.display()))?;
458 }
459 }
460
461 let index_content = self.generate_tools_index(&tools, &by_provider);
463 let index_path = tools_dir.join("INDEX.md");
464 write_file_with_context(&index_path, &index_content, "MCP tools index")
465 .await
466 .with_context(|| format!("Failed to write MCP tools index: {}", index_path.display()))?;
467
468 let status = self.generate_status_json();
470 let status_path = mcp_dir.join("status.json");
471 let status_json = serde_json::to_string_pretty(&status)?;
472 write_file_with_context(&status_path, &status_json, "MCP status")
473 .await
474 .with_context(|| format!("Failed to write MCP status: {}", status_path.display()))?;
475
476 info!(
477 tools = tools.len(),
478 providers = by_provider.len(),
479 index = %index_path.display(),
480 "Synced MCP tool descriptions to files"
481 );
482
483 Ok((index_path, tools.len()))
484 }
485
486 fn generate_tools_index(&self, tools: &[McpToolInfo], by_provider: &BTreeMap<String, Vec<&McpToolInfo>>) -> String {
488 let mut content = String::new();
489 content.push_str("# MCP Tools Index\n\n");
490 content.push_str("This file lists all available MCP tools for dynamic discovery.\n");
491 content.push_str("Use `read_file` on individual tool files for full schema details.\n\n");
492
493 if tools.is_empty() {
494 content.push_str("*No MCP tools available.*\n\n");
495 content.push_str("Configure MCP servers in `vtcode.toml` or `.mcp.json`.\n");
496 } else {
497 content.push_str(&format!("**Total Tools**: {}\n\n", tools.len()));
498
499 content.push_str("## Quick Reference\n\n");
501 content.push_str("| Provider | Tool | Description |\n");
502 content.push_str("|----------|------|-------------|\n");
503
504 for tool in tools {
505 let desc = tool.description.lines().next().unwrap_or(&tool.description);
506 let desc_truncated = vtcode_commons::formatting::truncate_byte_budget(desc, 57, "...");
507 content.push_str(&format!(
508 "| {} | `{}` | {} |\n",
509 tool.provider,
510 tool.name,
511 desc_truncated.replace('|', "\\|")
512 ));
513 }
514
515 content.push_str("\n## Tools by Provider\n\n");
517 for (provider, provider_tools) in by_provider {
518 content.push_str(&format!("### {provider}\n\n"));
519 for tool in provider_tools {
520 content.push_str(&format!(
521 "- **{}**: {}\n - Path: `.vtcode/mcp/tools/{}/{}.md`\n",
522 tool.name,
523 tool.description.lines().next().unwrap_or(&tool.description),
524 sanitize_filename(provider),
525 sanitize_filename(&tool.name)
526 ));
527 }
528 content.push('\n');
529 }
530 }
531
532 content.push_str("\n---\n");
533 content.push_str("*Generated automatically. Do not edit manually.*\n");
534
535 content
536 }
537
538 fn generate_status_json(&self) -> Value {
540 let status = self.get_status();
541 json!({
542 "enabled": status.enabled,
543 "provider_count": status.provider_count,
544 "active_connections": status.active_connections,
545 "configured_providers": status.configured_providers,
546 "last_updated": Utc::now().to_rfc3339(),
547 })
548 }
549
550 async fn collect_tools(&self, force_refresh: bool) -> Result<Vec<McpToolInfo>> {
551 let (providers, allowlist) = {
553 let state = self.state.read();
554 (providers_sorted_by_name(&state.providers), state.allowlist.clone())
555 };
556
557 if providers.is_empty() {
558 return Ok(Vec::new());
559 }
560
561 let timeout = self.tool_timeout();
562 let mut all_tools = Vec::with_capacity(128);
563 let mut index_updates: FxHashMap<String, String> = FxHashMap::with_capacity_and_hasher(128, Default::default());
564
565 for provider in providers {
566 let provider_name = provider.name.clone();
567 let tools = if force_refresh {
568 provider.refresh_tools(&allowlist, timeout).await
569 } else {
570 provider.list_tools(&allowlist, timeout).await
571 };
572
573 match tools {
574 Ok(tools) => {
575 for tool in &tools {
576 let _ = index_updates.entry(tool.name.clone()).or_insert_with(|| provider_name.clone());
577 }
578 all_tools.extend(tools);
579 }
580 Err(err) => {
581 warn!("Failed to list tools for provider '{}': {err}", provider_name);
582 }
583 }
584 }
585
586 if !index_updates.is_empty() || force_refresh {
587 let mut state = self.state.write();
588 if index_updates.is_empty() {
589 state.tool_provider_index.clear();
590 } else {
591 state.tool_provider_index = index_updates;
592 }
593 }
594
595 Ok(all_tools)
596 }
597
598 async fn collect_resources(&self, force_refresh: bool) -> Result<Vec<McpResourceInfo>> {
599 let (providers, allowlist) = {
601 let state = self.state.read();
602 (providers_sorted_by_name(&state.providers), state.allowlist.clone())
603 };
604
605 if providers.is_empty() {
606 self.state.write().resource_provider_index.clear();
607 return Ok(Vec::new());
608 }
609
610 let timeout = self.request_timeout();
611 let mut all_resources = Vec::with_capacity(64);
612
613 for provider in providers {
614 let resources = if force_refresh {
615 provider.refresh_resources(&allowlist, timeout).await
616 } else {
617 provider.list_resources(&allowlist, timeout).await
618 };
619
620 match resources {
621 Ok(resources) => {
622 all_resources.extend(resources);
623 }
624 Err(err) => {
625 warn!("Failed to list resources for provider '{}': {err}", provider.name);
626 }
627 }
628 }
629
630 let mut state = self.state.write();
631 let index = &mut state.resource_provider_index;
632 index.clear();
633 for resource in &all_resources {
634 let _ = index.entry(resource.uri.clone()).or_insert_with(|| resource.provider.clone());
635 }
636
637 Ok(all_resources)
638 }
639
640 async fn collect_prompts(&self, force_refresh: bool) -> Result<Vec<McpPromptInfo>> {
641 let (providers, allowlist) = {
643 let state = self.state.read();
644 (providers_sorted_by_name(&state.providers), state.allowlist.clone())
645 };
646
647 if providers.is_empty() {
648 self.state.write().prompt_provider_index.clear();
649 return Ok(Vec::new());
650 }
651
652 let timeout = self.request_timeout();
653 let mut all_prompts = Vec::with_capacity(32);
654
655 for provider in providers {
656 let prompts = if force_refresh {
657 provider.refresh_prompts(&allowlist, timeout).await
658 } else {
659 provider.list_prompts(&allowlist, timeout).await
660 };
661
662 match prompts {
663 Ok(prompts) => {
664 all_prompts.extend(prompts);
665 }
666 Err(err) => {
667 warn!("Failed to list prompts for provider '{}': {err}", provider.name);
668 }
669 }
670 }
671
672 let mut state = self.state.write();
673 let index = &mut state.prompt_provider_index;
674 index.clear();
675 for prompt in &all_prompts {
676 let _ = index.entry(prompt.name.clone()).or_insert_with(|| prompt.provider.clone());
677 }
678
679 Ok(all_prompts)
680 }
681
682 async fn resolve_provider_for_tool(&self, tool_name: &str) -> Result<Arc<McpProvider>> {
683 if !self.config.enabled {
684 return Err(anyhow!("MCP support is disabled in the current configuration"));
685 }
686
687 if let Some(provider) = self.provider_for_tool(tool_name)
688 && let Some(found) = self.state.read().providers.get(&provider)
689 {
690 return Ok(found.clone());
691 }
692
693 let (allowlist, providers) = {
694 let state = self.state.read();
695 (state.allowlist.clone(), providers_sorted_by_name(&state.providers))
696 };
697 let timeout = self.tool_timeout();
698
699 if providers.is_empty() {
700 if self.config.providers.is_empty() {
701 return Err(anyhow!(
702 "No MCP providers are configured. Use `vtcode mcp add` or update vtcode.toml to register one."
703 ));
704 }
705
706 return Err(anyhow!(
707 "No MCP providers are currently connected. Ensure MCP initialization completed successfully."
708 ));
709 }
710
711 for provider in providers {
712 match provider.has_tool(tool_name, &allowlist, timeout).await {
713 Ok(true) => {
714 drop(
715 self.state
716 .write()
717 .tool_provider_index
718 .insert(tool_name.into(), provider.name.clone()),
719 );
720 return Ok(provider);
721 }
722 Ok(false) => continue,
723 Err(err) => {
724 warn!("Error checking tool '{}' on provider '{}': {err}", tool_name, provider.name);
725 }
726 }
727 }
728
729 match self.collect_tools(true).await {
730 Ok(_) => {
731 if let Some(provider) = self.provider_for_tool(tool_name)
732 && let Some(found) = self.state.read().providers.get(&provider)
733 {
734 return Ok(found.clone());
735 }
736 }
737 Err(err) => {
738 warn!("Failed to refresh MCP tool caches while resolving '{}': {err}", tool_name);
739 }
740 }
741
742 Err(anyhow!(
743 "Tool '{tool_name}' not found on any MCP provider.\n\n\
744 To use this tool:\n\
745 1. Install the MCP server: `uv tool install mcp-server-{tool_name}`\n\
746 2. Add to vtcode.toml:\n \
747 [[mcp.providers]]\n \
748 name = \"{tool_name}\"\n \
749 command = \"uvx\"\n \
750 args = [\"mcp-server-{tool_name}\"]\n\
751 3. Restart VT Code\n\n\
752 Or use the built-in alternative if available (e.g., web_fetch instead of mcp_fetch)"
753 ))
754 }
755
756 async fn resolve_provider_for_resource(&self, uri: &str) -> Result<Arc<McpProvider>> {
757 if let Some(provider) = self.provider_for_resource(uri)
758 && let Some(found) = self.state.read().providers.get(&provider)
759 {
760 return Ok(found.clone());
761 }
762
763 let (allowlist, providers) = {
764 let state = self.state.read();
765 (state.allowlist.clone(), providers_sorted_by_name(&state.providers))
766 };
767 let timeout = self.request_timeout();
768
769 for provider in providers {
770 match provider.has_resource(uri, &allowlist, timeout).await {
771 Ok(true) => {
772 drop(
773 self.state
774 .write()
775 .resource_provider_index
776 .insert(uri.into(), provider.name.clone()),
777 );
778 return Ok(provider);
779 }
780 Ok(false) => continue,
781 Err(err) => {
782 warn!("Error checking resource '{}' on provider '{}': {err}", uri, provider.name);
783 }
784 }
785 }
786
787 Err(anyhow!("Resource '{uri}' not found on any MCP provider"))
788 }
789
790 async fn resolve_provider_for_prompt(&self, prompt_name: &str) -> Result<Arc<McpProvider>> {
791 if let Some(provider) = self.provider_for_prompt(prompt_name)
792 && let Some(found) = self.state.read().providers.get(&provider)
793 {
794 return Ok(found.clone());
795 }
796
797 let (allowlist, providers) = {
798 let state = self.state.read();
799 (state.allowlist.clone(), providers_sorted_by_name(&state.providers))
800 };
801 let timeout = self.request_timeout();
802
803 for provider in providers {
804 match provider.has_prompt(prompt_name, &allowlist, timeout).await {
805 Ok(true) => {
806 drop(
807 self.state
808 .write()
809 .prompt_provider_index
810 .insert(prompt_name.into(), provider.name.clone()),
811 );
812 return Ok(provider);
813 }
814 Ok(false) => continue,
815 Err(err) => {
816 warn!("Error checking prompt '{}' on provider '{}': {err}", prompt_name, provider.name);
817 }
818 }
819 }
820
821 Err(anyhow!("Prompt '{prompt_name}' not found on any MCP provider"))
822 }
823
824 fn record_tool_provider(&self, provider: &str, tools: &[McpToolInfo]) {
825 let mut state = self.state.write();
826 let index = &mut state.tool_provider_index;
827 for tool in tools {
828 let _ = index
829 .entry(tool.name.clone())
830 .and_modify(|owner| {
831 if provider < owner.as_str() {
832 *owner = provider.to_string();
833 }
834 })
835 .or_insert_with(|| provider.to_string());
836 }
837 }
838
839 async fn connect_and_initialize_provider(
840 &self,
841 provider_config: &McpProviderConfig,
842 allowlist_snapshot: &McpAllowListConfig,
843 tool_timeout: Option<Duration>,
844 ) -> Result<McpProvider> {
845 let total_attempts = self.provider_retry_attempts();
846 let mut last_error: Option<anyhow::Error> = None;
847
848 for attempt_idx in 0..total_attempts {
849 let attempt_number = attempt_idx + 1;
850 match self
851 .connect_and_initialize_provider_once(provider_config, allowlist_snapshot, tool_timeout)
852 .await
853 {
854 Ok(provider) => return Ok(provider),
855 Err(err) => {
856 if attempt_number == total_attempts {
857 return Err(err);
858 }
859
860 let retries_remaining = total_attempts - attempt_number;
861 warn!(
862 provider = provider_config.name.as_str(),
863 attempt = attempt_number,
864 retries_remaining,
865 error = %err,
866 "MCP provider initialization failed; retrying"
867 );
868 last_error = Some(err);
869 tokio::time::sleep(Self::provider_retry_delay(attempt_idx)).await;
870 }
871 }
872 }
873
874 Err(last_error.unwrap_or_else(|| anyhow!("Failed to initialize MCP provider '{}'", provider_config.name)))
875 }
876
877 async fn connect_and_initialize_provider_once(
878 &self,
879 provider_config: &McpProviderConfig,
880 allowlist_snapshot: &McpAllowListConfig,
881 tool_timeout: Option<Duration>,
882 ) -> Result<McpProvider> {
883 let provider = McpProvider::connect(
884 provider_config.clone(),
885 self.elicitation_handler.clone(),
886 self.sandbox_context.clone(),
887 )
888 .await
889 .with_context(|| format!("Failed to connect to MCP provider '{}'", provider_config.name))?;
890 let provider_startup_timeout = self.resolve_startup_timeout(provider_config);
891 provider
892 .initialize(
893 self.build_initialize_params(&provider),
894 provider_startup_timeout,
895 tool_timeout,
896 allowlist_snapshot,
897 )
898 .await
899 .with_context(|| format!("Failed to initialize MCP provider '{}'", provider_config.name))?;
900 Ok(provider)
901 }
902
903 fn startup_timeout(&self) -> Option<Duration> {
904 match self.config.startup_timeout_seconds {
905 Some(0) => None,
906 Some(value) => Some(Duration::from_secs(value)),
907 None => self.request_timeout(),
908 }
909 }
910
911 fn requirement_mismatch_reason(&self, provider_config: &McpProviderConfig) -> Option<String> {
912 let requirements = &self.config.requirements;
913 if !requirements.enforce {
914 return None;
915 }
916
917 match &provider_config.transport {
918 McpTransportConfig::Stdio(stdio) => {
919 if requirements
920 .allowed_stdio_commands
921 .iter()
922 .any(|allowed| allowed == &stdio.command)
923 {
924 None
925 } else {
926 Some(format!("stdio command '{}' is not allowlisted", stdio.command))
927 }
928 }
929 McpTransportConfig::Http(http) => {
930 if requirements
931 .allowed_http_endpoints
932 .iter()
933 .any(|allowed| allowed == &http.endpoint)
934 {
935 None
936 } else {
937 Some(format!("HTTP endpoint '{}' is not allowlisted", http.endpoint))
938 }
939 }
940 }
941 }
942
943 fn resolve_startup_timeout(&self, provider_config: &McpProviderConfig) -> Option<Duration> {
944 if let Some(timeout_ms) = provider_config.startup_timeout_ms {
945 if timeout_ms == 0 {
946 None
947 } else {
948 Some(Duration::from_millis(timeout_ms))
949 }
950 } else {
951 self.startup_timeout()
952 }
953 }
954
955 fn provider_retry_attempts(&self) -> usize {
956 self.config.retry_attempts.try_into().unwrap_or(usize::MAX).saturating_add(1)
957 }
958
959 fn provider_retry_delay(attempt_idx: usize) -> Duration {
960 let base_ms = 250u64;
961 let max_ms = 5000u64;
962 let exp = base_ms.saturating_mul(2u64.saturating_pow(u32::try_from(attempt_idx).unwrap_or(u32::MAX)));
963 let delay = exp.min(max_ms);
964 Duration::from_millis(delay)
965 }
966
967 fn tool_timeout(&self) -> Option<Duration> {
968 match self.config.tool_timeout_seconds {
969 Some(0) => None,
970 Some(value) => Some(Duration::from_secs(value)),
971 None => self.request_timeout(),
972 }
973 }
974
975 fn request_timeout(&self) -> Option<Duration> {
976 if self.config.request_timeout_seconds == 0 {
977 None
978 } else {
979 Some(Duration::from_secs(self.config.request_timeout_seconds))
980 }
981 }
982
983 fn build_initialize_params(&self, _provider: &McpProvider) -> InitializeRequestParams {
984 let mut capabilities = ClientCapabilities::default();
985 {
986 let mut roots_cap = RootsCapabilities::default();
987 roots_cap.list_changed = Some(true);
988 capabilities.roots = Some(roots_cap);
989 }
990
991 if self.elicitation_handler.is_some() {
992 capabilities.elicitation = Some(
994 rmcp::model::ElicitationCapability::new()
995 .with_form(rmcp::model::FormElicitationCapability::new().with_schema_validation(true)),
996 );
997 }
998
999 InitializeRequestParams::new(capabilities, super::utils::build_client_implementation())
1000 .with_protocol_version(super::rmcp_client::latest_protocol_version())
1001 }
1002
1003 pub(super) fn normalize_arguments(args: &Value) -> Map<String, Value> {
1004 match args {
1005 Value::Null => Map::new(),
1006 Value::Object(map) => map.clone(),
1007 other => {
1008 let mut map = Map::new();
1009 drop(map.insert("value".to_owned(), other.clone()));
1010 map
1011 }
1012 }
1013 }
1014
1015 fn format_tool_result(provider_name: &str, tool_name: &str, result: CallToolResult) -> Result<Value> {
1016 let result_json = serde_json::to_value(&result)?;
1018 let result_obj = result_json.as_object();
1019
1020 let is_error = result_obj
1022 .and_then(|o| o.get("isError"))
1023 .or_else(|| result_obj.and_then(|o| o.get("is_error")))
1024 .and_then(Value::as_bool)
1025 .unwrap_or(false);
1026
1027 if is_error {
1028 let mut message = result_obj
1029 .and_then(|o| o.get("_meta"))
1030 .or_else(|| result_obj.and_then(|o| o.get("meta")))
1031 .and_then(|m| m.get("message"))
1032 .and_then(Value::as_str)
1033 .map(str::to_owned);
1034
1035 if message.is_none()
1037 && let Some(content) = result_obj.and_then(|o| o.get("content")).and_then(Value::as_array)
1038 {
1039 message = content
1040 .iter()
1041 .find_map(|block| block.get("text").and_then(Value::as_str).map(str::to_owned));
1042 }
1043
1044 let message = message.unwrap_or_else(|| "Unknown MCP tool error".to_owned());
1045 return Err(anyhow!("MCP tool '{tool_name}' on provider '{provider_name}' reported an error: {message}"));
1046 }
1047
1048 let mut payload = Map::new();
1049 drop(payload.insert("provider".into(), Value::String(provider_name.to_string())));
1050 drop(payload.insert("tool".into(), Value::String(tool_name.to_string())));
1051
1052 if let Some(meta) = result_obj
1054 .and_then(|o| o.get("_meta"))
1055 .or_else(|| result_obj.and_then(|o| o.get("meta")))
1056 .and_then(Value::as_object)
1057 && !meta.is_empty()
1058 {
1059 drop(payload.insert("meta".into(), Value::Object(meta.clone())));
1060 }
1061
1062 if let Some(content) = result_obj.and_then(|o| o.get("content"))
1064 && !content.is_null()
1065 && !content.as_array().map(|a| a.is_empty()).unwrap_or(true)
1066 {
1067 drop(payload.insert("content".into(), content.clone()));
1068 }
1069
1070 Ok(Value::Object(payload))
1071 }
1072}
1073
1074#[async_trait]
1075impl McpToolExecutor for McpClient {
1076 async fn execute_mcp_tool(&self, tool_name: &str, args: &Value) -> Result<Value> {
1077 self.execute_tool_with_validation_ref(tool_name, args).await
1078 }
1079
1080 async fn list_mcp_tools(&self) -> Result<Vec<McpToolInfo>> {
1081 self.collect_tools(false).await
1082 }
1083
1084 async fn has_mcp_tool(&self, tool_name: &str) -> Result<bool> {
1085 if !self.config.enabled {
1086 return Ok(false);
1087 }
1088
1089 if self.provider_for_tool(tool_name).is_some() {
1090 return Ok(true);
1091 }
1092
1093 if self.state.read().providers.is_empty() {
1094 if self.config.providers.is_empty() {
1095 return Ok(false);
1096 }
1097
1098 bail!("No MCP providers are currently connected. Ensure MCP initialization completed successfully.");
1099 }
1100
1101 let tools = self.collect_tools(false).await?;
1102 Ok(tools.iter().any(|tool| tool.name == tool_name))
1103 }
1104
1105 fn get_status(&self) -> McpClientStatus {
1106 self.get_status()
1107 }
1108}
1109
1110#[cfg(test)]
1111mod tests {
1112 use super::McpClient;
1113 use rustc_hash::FxHashMap;
1114 use vtcode_config::mcp::{
1115 McpClientConfig, McpHttpServerConfig, McpProviderConfig, McpRequirementsConfig, McpStdioServerConfig,
1116 McpTransportConfig,
1117 };
1118
1119 fn base_config() -> McpClientConfig {
1120 McpClientConfig {
1121 enabled: true,
1122 requirements: McpRequirementsConfig {
1123 enforce: true,
1124 allowed_stdio_commands: vec!["uvx".to_string()],
1125 allowed_http_endpoints: vec!["https://allowed.example/mcp".to_string()],
1126 },
1127 ..McpClientConfig::default()
1128 }
1129 }
1130
1131 #[test]
1132 fn requirements_allow_matching_stdio_command() {
1133 let client = McpClient::new(base_config());
1134 let provider = McpProviderConfig {
1135 name: "time".to_string(),
1136 transport: McpTransportConfig::Stdio(McpStdioServerConfig {
1137 command: "uvx".to_string(),
1138 args: vec![],
1139 working_directory: None,
1140 }),
1141 ..McpProviderConfig::default()
1142 };
1143
1144 assert!(client.requirement_mismatch_reason(&provider).is_none());
1145 }
1146
1147 #[test]
1148 fn requirements_block_unmatched_stdio_command() {
1149 let client = McpClient::new(base_config());
1150 let provider = McpProviderConfig {
1151 name: "time".to_string(),
1152 transport: McpTransportConfig::Stdio(McpStdioServerConfig {
1153 command: "npx".to_string(),
1154 args: vec![],
1155 working_directory: None,
1156 }),
1157 ..McpProviderConfig::default()
1158 };
1159
1160 assert!(
1161 client
1162 .requirement_mismatch_reason(&provider)
1163 .is_some_and(|reason| reason.contains("not allowlisted"))
1164 );
1165 }
1166
1167 #[test]
1168 fn requirements_block_unmatched_http_endpoint() {
1169 let client = McpClient::new(base_config());
1170 let provider = McpProviderConfig {
1171 name: "remote".to_string(),
1172 transport: McpTransportConfig::Http(McpHttpServerConfig {
1173 endpoint: "https://blocked.example/mcp".to_string(),
1174 ..McpHttpServerConfig::default()
1175 }),
1176 ..McpProviderConfig::default()
1177 };
1178
1179 assert!(
1180 client
1181 .requirement_mismatch_reason(&provider)
1182 .is_some_and(|reason| reason.contains("not allowlisted"))
1183 );
1184 }
1185
1186 #[tokio::test]
1187 async fn list_servers_includes_configured_provider_metadata() {
1188 let mut config = base_config();
1189 config.providers = vec![McpProviderConfig {
1190 name: "calendar".to_string(),
1191 transport: McpTransportConfig::Http(McpHttpServerConfig {
1192 endpoint: "https://calendar.example/mcp".to_string(),
1193 ..McpHttpServerConfig::default()
1194 }),
1195 ..McpProviderConfig::default()
1196 }];
1197
1198 let client = McpClient::new(config);
1199 let servers = client.list_servers().await;
1200 assert_eq!(servers.len(), 1);
1201 assert_eq!(servers[0]["name"], "calendar");
1202 assert_eq!(servers[0]["connected"], false);
1203 assert_eq!(servers[0]["connection_state"], "disconnected");
1204 assert_eq!(servers[0]["transport"], "http");
1205 assert_eq!(servers[0]["target"], "https://calendar.example/mcp");
1206 assert!(
1207 servers[0]["negotiated_protocol_version"].is_null(),
1208 "disconnected servers have no negotiated version"
1209 );
1210 }
1211
1212 #[tokio::test]
1213 async fn connect_server_rejects_unknown_server_name() {
1214 let client = McpClient::new(base_config());
1215 let err = Box::pin(client.connect_server("missing"))
1216 .await
1217 .expect_err("missing server should error");
1218 assert!(err.to_string().contains("not configured"));
1219 }
1220
1221 #[tokio::test]
1222 async fn disconnect_server_rejects_unknown_server_name() {
1223 let client = McpClient::new(base_config());
1224 let err = client
1225 .disconnect_server("missing")
1226 .await
1227 .expect_err("missing server should error");
1228 assert!(err.to_string().contains("not connected"));
1229 }
1230
1231 #[test]
1232 fn providers_sorted_by_name_ignores_insertion_order() {
1233 let names = ["zeta", "Alpha", "alpha", "beta", "alpha2"];
1234 let forward: FxHashMap<String, usize> =
1235 names.iter().enumerate().map(|(idx, name)| ((*name).to_owned(), idx)).collect();
1236 let reverse: FxHashMap<String, usize> = names
1237 .iter()
1238 .enumerate()
1239 .rev()
1240 .map(|(idx, name)| ((*name).to_owned(), idx))
1241 .collect();
1242
1243 let expected = vec![1, 2, 4, 3, 0];
1245 assert_eq!(super::providers_sorted_by_name(&forward), expected);
1246 assert_eq!(super::providers_sorted_by_name(&reverse), expected);
1247 }
1248
1249 #[test]
1250 fn record_tool_provider_keeps_first_provider_by_name_regardless_of_connect_order() {
1251 use super::McpToolInfo;
1252 use serde_json::json;
1253
1254 let tool = |provider: &str, name: &str| McpToolInfo {
1255 name: name.to_owned(),
1256 description: String::new(),
1257 provider: provider.to_owned(),
1258 input_schema: json!({}),
1259 output_schema: None,
1260 };
1261
1262 for order in [["alpha", "beta"], ["beta", "alpha"]] {
1263 let client = McpClient::new(base_config());
1264 for provider in order {
1265 client.record_tool_provider(
1266 provider,
1267 &[tool(provider, "shared"), tool(provider, &format!("{provider}_only"))],
1268 );
1269 }
1270 assert_eq!(client.provider_for_tool("shared").as_deref(), Some("alpha"), "order {order:?}");
1271 assert_eq!(client.provider_for_tool("beta_only").as_deref(), Some("beta"));
1272 }
1273 }
1274}