Skip to main content

vtcode_mcp/
client.rs

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
31/// Stable provider order keeps the aggregated catalog (and LLM prompt-cache prefix) deterministic.
32fn 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    /// Create a new MCP client from the configuration.
47    pub fn new(config: McpClientConfig) -> Self {
48        Self::with_sandbox_context(config, None)
49    }
50
51    /// Create an MCP client with an optional inherited stdio sandbox context.
52    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    /// Register a handler used to satisfy elicitation requests from providers.
70    pub fn set_elicitation_handler(&mut self, handler: Arc<dyn McpElicitationHandler>) {
71        self.elicitation_handler = Some(handler);
72    }
73
74    /// Establish connections to all configured providers and complete the
75    /// MCP handshake.
76    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    /// Validate tool arguments based on security configuration
115    fn validate_tool_arguments(&self, _tool_name: &str, args: &Value) -> Result<()> {
116        // Check argument size
117        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        // Check for path traversal in file-related arguments
129        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    /// Execute a tool call after validating arguments.
140    ///
141    /// Public-facing version that takes ownership of `args` for compatibility
142    /// with existing callers. Delegates to the reference-taking implementation
143    /// to avoid unnecessary cloning when the caller already has a reference.
144    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    // Internal reference-taking implementation to avoid cloning when not necessary.
149    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    /// Refresh the internal allow list at runtime.
166    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    /// Current allow list snapshot.
182    pub fn current_allowlist(&self) -> McpAllowListConfig {
183        self.state.read().allowlist.clone()
184    }
185
186    /// Return the provider name serving the given tool if previously cached.
187    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    /// Return the provider responsible for the given resource URI if known.
192    pub fn provider_for_resource(&self, uri: &str) -> Option<String> {
193        self.state.read().resource_provider_index.get(uri).cloned()
194    }
195
196    /// Return the provider that exposes the given prompt if known.
197    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    /// Execute a tool call on the appropriate provider.
202    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    /// List all tools from all active providers.
207    pub async fn list_tools(&self) -> Result<Vec<McpToolInfo>> {
208        self.collect_tools(false).await
209    }
210
211    /// List all resources exposed by connected MCP providers.
212    pub async fn list_resources(&self) -> Result<Vec<McpResourceInfo>> {
213        self.collect_resources(false).await
214    }
215
216    /// Force refresh and list resources from providers.
217    pub async fn refresh_resources(&self) -> Result<Vec<McpResourceInfo>> {
218        self.collect_resources(true).await
219    }
220
221    /// List all prompts advertised by connected MCP providers.
222    pub async fn list_prompts(&self) -> Result<Vec<McpPromptInfo>> {
223        self.collect_prompts(false).await
224    }
225
226    /// Force refresh and list prompts from providers.
227    pub async fn refresh_prompts(&self) -> Result<Vec<McpPromptInfo>> {
228        self.collect_prompts(true).await
229    }
230
231    /// Read a single resource from its originating provider.
232    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    /// Retrieve a rendered prompt from its originating provider.
242    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    /// Shutdown all active provider connections.
263    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    /// Current status snapshot for UI/debugging purposes.
289    pub fn get_status(&self) -> McpClientStatus {
290        let state = self.state.read();
291        let providers = &state.providers;
292        // Use iterator to collect keys directly without intermediate push
293        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    /// Return configured MCP servers and their current connection state.
304    ///
305    /// Connected servers additionally report the protocol version settled
306    /// during the handshake (`negotiated_protocol_version`), which may differ
307    /// from the configured version when the server negotiates down.
308    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    /// Return whether model-callable lifecycle tools are enabled by config.
343    pub fn allow_model_lifecycle_control(&self) -> bool {
344        self.config.lifecycle.allow_model_control
345    }
346
347    /// Connect one configured MCP server by name.
348    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    /// Disconnect one active MCP server by name.
398    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    /// Sync MCP tool descriptions to files for dynamic context discovery
422    ///
423    /// This implements Cursor-style dynamic context discovery:
424    /// - Tool descriptions are written to `.vtcode/mcp/tools/{provider}/{tool}.md`
425    /// - Status is written to `.vtcode/mcp/status.json`
426    /// - Agents can discover tools via grep/read_file without loading all schemas
427    ///
428    /// Returns the paths to written files (index path, tool count)
429    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        // Create directories
435        ensure_dir_exists(&tools_dir)
436            .await
437            .with_context(|| format!("Failed to create MCP tools directory: {}", tools_dir.display()))?;
438
439        // Group tools by provider
440        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        // Write tool files per provider
446        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        // Write index file
462        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        // Write status file
469        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    /// Generate INDEX.md content for MCP tools
487    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            // Summary table
500            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            // Per-provider sections
516            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    /// Generate status.json content
539    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        // Collect provider references in one pass
552        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        // Collect provider references in one pass
600        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        // Collect provider references in one pass
642        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            // Elicitation is now a first-class capability in rmcp
993            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        // Convert result to JSON to access fields flexibly
1017        let result_json = serde_json::to_value(&result)?;
1018        let result_obj = result_json.as_object();
1019
1020        // Check for error - handle both rmcp's is_error field and meta message
1021        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            // Try to find text content in the content array
1036            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        // Add meta if present
1053        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        // Add content if present
1063        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        // Values are the original indices of Alpha, alpha, alpha2, beta, zeta.
1244        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}