Skip to main content

vtcode_mcp/
provider.rs

1use anyhow::{Context, Result, anyhow};
2use arc_swap::ArcSwap;
3use hashbrown::HashMap;
4use rmcp::model::{
5    CallToolRequestParams, CallToolResult, GetPromptRequestParams, InitializeRequestParams, Prompt,
6    ReadResourceRequestParams, Resource, ServerPeerInfo, Tool,
7};
8use rmcp::service::ClientLifecycleMode;
9use serde_json::{Map, Value};
10use std::ffi::OsString;
11use std::path::PathBuf;
12use std::sync::Arc;
13use std::time::Duration;
14use tokio::sync::{Mutex, Semaphore};
15use tracing::{Instrument, Span, warn};
16use url::Url;
17
18use super::rmcp_client::{auto_lifecycle_mode, latest_protocol_version};
19use super::{LATEST_PROTOCOL_VERSION, SUPPORTED_PROTOCOL_VERSIONS};
20
21use vtcode_config::auth::McpOAuthService;
22use vtcode_config::mcp::{McpAllowListConfig, McpHttpHandshakeMode, McpProviderConfig, McpTransportConfig};
23use vtcode_utility_tool_specs::parse_mcp_tool;
24
25use super::{McpClient, McpSandboxContext, RmcpClient};
26use super::{
27    McpElicitationHandler, McpPromptDetail, McpPromptInfo, McpResourceData, McpResourceInfo, McpToolInfo,
28    TIMEZONE_ARGUMENT, build_headers, ensure_timezone_argument, schema_requires_field,
29};
30
31pub struct McpProvider {
32    pub(super) name: String,
33    #[expect(
34        dead_code,
35        reason = "Intentional compatibility, platform, test, or API-shape suppression."
36    )]
37    protocol_version: String,
38    client: ArcSwap<RmcpClient>,
39    /// Stored config so we can reconnect after disconnection.
40    config: McpProviderConfig,
41    /// Stored elicitation handler for reconnection.
42    elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
43    /// Sandbox inherited by stdio launches and reconnects.
44    sandbox_context: Option<McpSandboxContext>,
45    pub(crate) semaphore: Arc<Semaphore>,
46    caches: Mutex<ProviderCaches>,
47    initialize_result: Mutex<Option<ServerPeerInfo>>,
48}
49
50#[derive(Default)]
51struct ProviderCaches {
52    tools: Option<Arc<Vec<McpToolInfo>>>,
53    resources: Option<Arc<Vec<McpResourceInfo>>>,
54    prompts: Option<Arc<Vec<McpPromptInfo>>>,
55}
56
57impl McpProvider {
58    pub(super) async fn connect(
59        config: McpProviderConfig,
60        elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
61        sandbox_context: Option<McpSandboxContext>,
62    ) -> Result<Self> {
63        if config.name.trim().is_empty() {
64            return Err(anyhow!("MCP provider name cannot be empty"));
65        }
66
67        let max_requests = std::cmp::max(1, config.max_concurrent_requests);
68
69        let (client, protocol_version) = match &config.transport {
70            McpTransportConfig::Stdio(stdio) => {
71                let program = OsString::from(&stdio.command);
72                let args: Vec<OsString> = stdio.args.iter().map(OsString::from).collect();
73                let working_dir = stdio.working_directory.as_ref().map(PathBuf::from);
74                let env: HashMap<OsString, OsString> = config
75                    .env
76                    .iter()
77                    .map(|(key, value)| (OsString::from(key), OsString::from(value)))
78                    .collect();
79                let client = RmcpClient::new_stdio_client(
80                    config.name.clone(),
81                    program,
82                    args,
83                    working_dir,
84                    Some(env),
85                    elicitation_handler.clone(),
86                    sandbox_context.clone(),
87                )
88                .await?;
89                (client, LATEST_PROTOCOL_VERSION.to_string())
90            }
91            McpTransportConfig::Http(http) => {
92                if !SUPPORTED_PROTOCOL_VERSIONS
93                    .iter()
94                    .any(|supported| supported == &http.protocol_version)
95                {
96                    return Err(anyhow!(
97                        "MCP HTTP provider '{}' requested unsupported protocol version '{}'",
98                        config.name,
99                        http.protocol_version
100                    ));
101                }
102
103                let bearer_token = if let Some(oauth) = http.oauth.as_ref() {
104                    McpOAuthService::new()
105                        .resolve_access_token(&config.name, oauth)
106                        .await?
107                        .ok_or_else(|| {
108                            anyhow!(
109                                "MCP HTTP provider '{}' requires OAuth login. Run `vtcode mcp login {}`.",
110                                config.name,
111                                config.name
112                            )
113                        })
114                        .map(Some)?
115                } else {
116                    match http.api_key_env.as_ref() {
117                        Some(var) => Some(
118                            std::env::var(var)
119                                .with_context(|| format!("Missing MCP API key environment variable: {var}"))?,
120                        ),
121                        None => None,
122                    }
123                };
124
125                let headers = build_headers(&http.http_headers, &http.env_http_headers);
126                let client = RmcpClient::new_streamable_http_client(
127                    config.name.clone(),
128                    &http.endpoint,
129                    bearer_token,
130                    headers,
131                    elicitation_handler.clone(),
132                )
133                .await?;
134                (client, http.protocol_version.clone())
135            }
136        };
137
138        Ok(Self {
139            name: config.name.clone(),
140            protocol_version,
141            client: ArcSwap::from_pointee(client),
142            config,
143            elicitation_handler,
144            sandbox_context,
145            semaphore: Arc::new(Semaphore::new(max_requests)),
146            caches: Mutex::new(ProviderCaches::default()),
147            initialize_result: Mutex::new(None),
148        })
149    }
150
151    pub(super) fn invalidate_caches(&self) {
152        if let Ok(mut caches) = self.caches.try_lock() {
153            caches.tools = None;
154            caches.resources = None;
155            caches.prompts = None;
156        }
157    }
158
159    pub(super) async fn initialize(
160        &self,
161        params: InitializeRequestParams,
162        startup_timeout: Option<Duration>,
163        _tool_timeout: Option<Duration>,
164        _allowlist: &McpAllowListConfig,
165    ) -> Result<()> {
166        let client = self.client.load_full();
167        let result = client.initialize(params, startup_timeout, self.handshake_lifecycle()).await?;
168
169        let protocol_version_str = result.protocol_version.to_string();
170        if !SUPPORTED_PROTOCOL_VERSIONS
171            .iter()
172            .any(|supported| *supported == protocol_version_str)
173        {
174            return Err(anyhow!(
175                "MCP server for '{}' negotiated unsupported protocol version '{}'",
176                self.name,
177                protocol_version_str
178            ));
179        }
180
181        *self.initialize_result.lock().await = Some(result);
182        Ok(())
183    }
184
185    /// Lifecycle mode for the rmcp handshake, derived from the transport.
186    ///
187    /// stdio always probes `server/discover` first (`Auto`) since a local
188    /// child process answers immediately either way. HTTP defaults to the
189    /// direct legacy handshake so legacy-only servers (e.g. DeepWiki) are not
190    /// penalized by the discover-timeout fallback; opt in to `Auto` per
191    /// provider via `handshake = "auto"` for modern servers.
192    fn handshake_lifecycle(&self) -> ClientLifecycleMode {
193        match &self.config.transport {
194            McpTransportConfig::Stdio(_) => auto_lifecycle_mode(),
195            McpTransportConfig::Http(http) => match http.handshake {
196                McpHttpHandshakeMode::Legacy => ClientLifecycleMode::Initialize,
197                McpHttpHandshakeMode::Auto => auto_lifecycle_mode(),
198            },
199        }
200    }
201
202    /// Protocol version negotiated during the last successful handshake, if any.
203    pub(super) async fn negotiated_protocol_version(&self) -> Option<String> {
204        self.initialize_result
205            .lock()
206            .await
207            .as_ref()
208            .map(|info| info.protocol_version.to_string())
209    }
210
211    pub(super) async fn list_tools(
212        &self,
213        allowlist: &McpAllowListConfig,
214        timeout: Option<Duration>,
215    ) -> Result<Vec<McpToolInfo>> {
216        Ok(self.list_tools_shared(allowlist, timeout).await?.as_ref().clone())
217    }
218
219    async fn list_tools_shared(
220        &self,
221        allowlist: &McpAllowListConfig,
222        timeout: Option<Duration>,
223    ) -> Result<Arc<Vec<McpToolInfo>>> {
224        let mut caches = self.caches.lock().await;
225        if self.client.load_full().take_tool_list_changed() {
226            caches.tools = None;
227        }
228
229        if let Some(cache) = &caches.tools {
230            return Ok(Arc::clone(cache));
231        }
232        drop(caches);
233
234        self.refresh_tools_shared(allowlist, timeout).await
235    }
236
237    pub(super) async fn refresh_tools(
238        &self,
239        allowlist: &McpAllowListConfig,
240        timeout: Option<Duration>,
241    ) -> Result<Vec<McpToolInfo>> {
242        Ok(self.refresh_tools_shared(allowlist, timeout).await?.as_ref().clone())
243    }
244
245    async fn refresh_tools_shared(
246        &self,
247        allowlist: &McpAllowListConfig,
248        timeout: Option<Duration>,
249    ) -> Result<Arc<Vec<McpToolInfo>>> {
250        let client = self.client.load_full();
251        let tools = client.list_all_tools(timeout).await?;
252        let filtered = Arc::new(self.filter_tools(tools, allowlist));
253        self.caches.lock().await.tools = Some(Arc::clone(&filtered));
254        Ok(filtered)
255    }
256
257    pub(super) async fn has_tool(
258        &self,
259        tool_name: &str,
260        allowlist: &McpAllowListConfig,
261        timeout: Option<Duration>,
262    ) -> Result<bool> {
263        let tools = self.list_tools_shared(allowlist, timeout).await?;
264        Ok(tools.iter().any(|tool| tool.name == tool_name))
265    }
266
267    pub(super) async fn call_tool(
268        &self,
269        tool_name: &str,
270        args: &Value,
271        timeout: Option<Duration>,
272        allowlist: &McpAllowListConfig,
273    ) -> Result<CallToolResult> {
274        if !allowlist.is_tool_allowed(&self.name, tool_name) {
275            return Err(anyhow!("Tool '{}' is blocked by the MCP allow list for provider '{}'", tool_name, self.name));
276        }
277
278        let _permit = self
279            .semaphore
280            .clone()
281            .acquire_owned()
282            .await
283            .context("Failed to acquire MCP request slot")?;
284        let mut arguments = McpClient::normalize_arguments(args);
285        self.add_argument_defaults(tool_name, &mut arguments, allowlist, timeout)
286            .await
287            .with_context(|| {
288                format!("failed to prepare arguments for MCP tool '{}' on provider '{}'", tool_name, self.name)
289            })?;
290        let params = CallToolRequestParams::new(tool_name.to_string()).with_arguments(arguments);
291        let client = self.client.load_full();
292        async move { client.call_tool(params, timeout).await }
293            .instrument(mcp_tool_call_span(&self.name, tool_name, &self.config.transport))
294            .await
295    }
296
297    async fn add_argument_defaults(
298        &self,
299        tool_name: &str,
300        arguments: &mut Map<String, Value>,
301        allowlist: &McpAllowListConfig,
302        timeout: Option<Duration>,
303    ) -> Result<()> {
304        let requires_timezone = self
305            .tool_requires_field(tool_name, TIMEZONE_ARGUMENT, allowlist, timeout)
306            .await?;
307        ensure_timezone_argument(arguments, requires_timezone)?;
308        Ok(())
309    }
310
311    async fn tool_requires_field(
312        &self,
313        tool_name: &str,
314        field: &str,
315        allowlist: &McpAllowListConfig,
316        timeout: Option<Duration>,
317    ) -> Result<bool> {
318        if let Some(tools) = &self.caches.lock().await.tools
319            && let Some(tool) = tools.iter().find(|tool| tool.name == tool_name)
320        {
321            return Ok(schema_requires_field(&tool.input_schema, field));
322        }
323
324        match self.refresh_tools_shared(allowlist, timeout).await {
325            Ok(tools) => Ok(tools
326                .iter()
327                .find(|tool| tool.name == tool_name)
328                .map(|tool| schema_requires_field(&tool.input_schema, field))
329                .unwrap_or(false)),
330            Err(err) => {
331                warn!(
332                    "Failed to refresh tools while inspecting schema for '{}' on provider '{}': {err}",
333                    tool_name, self.name
334                );
335                Ok(false)
336            }
337        }
338    }
339
340    pub(super) async fn list_resources(
341        &self,
342        allowlist: &McpAllowListConfig,
343        timeout: Option<Duration>,
344    ) -> Result<Vec<McpResourceInfo>> {
345        Ok(self.list_resources_shared(allowlist, timeout).await?.as_ref().clone())
346    }
347
348    async fn list_resources_shared(
349        &self,
350        allowlist: &McpAllowListConfig,
351        timeout: Option<Duration>,
352    ) -> Result<Arc<Vec<McpResourceInfo>>> {
353        let mut caches = self.caches.lock().await;
354        if self.client.load_full().take_resource_list_changed() {
355            caches.resources = None;
356        }
357
358        if let Some(cache) = &caches.resources {
359            return Ok(Arc::clone(cache));
360        }
361        drop(caches);
362
363        self.refresh_resources_shared(allowlist, timeout).await
364    }
365
366    pub(super) async fn refresh_resources(
367        &self,
368        allowlist: &McpAllowListConfig,
369        timeout: Option<Duration>,
370    ) -> Result<Vec<McpResourceInfo>> {
371        Ok(self.refresh_resources_shared(allowlist, timeout).await?.as_ref().clone())
372    }
373
374    async fn refresh_resources_shared(
375        &self,
376        allowlist: &McpAllowListConfig,
377        timeout: Option<Duration>,
378    ) -> Result<Arc<Vec<McpResourceInfo>>> {
379        let client = self.client.load_full();
380        let resources = client.list_all_resources(timeout).await?;
381        let filtered = Arc::new(self.filter_resources(resources, allowlist));
382        self.caches.lock().await.resources = Some(Arc::clone(&filtered));
383        Ok(filtered)
384    }
385
386    pub(super) async fn has_resource(
387        &self,
388        uri: &str,
389        allowlist: &McpAllowListConfig,
390        timeout: Option<Duration>,
391    ) -> Result<bool> {
392        let resources = self.list_resources_shared(allowlist, timeout).await?;
393        Ok(resources.iter().any(|resource| resource.uri == uri))
394    }
395
396    pub(super) async fn read_resource(
397        &self,
398        uri: &str,
399        timeout: Option<Duration>,
400        allowlist: &McpAllowListConfig,
401    ) -> Result<McpResourceData> {
402        if !allowlist.is_resource_allowed(&self.name, uri) {
403            return Err(anyhow!("Resource '{}' is blocked by the MCP allow list for provider '{}'", uri, self.name));
404        }
405
406        let _permit = self
407            .semaphore
408            .clone()
409            .acquire_owned()
410            .await
411            .context("Failed to acquire MCP request slot")?;
412        let params = ReadResourceRequestParams::new(uri.to_string());
413        let client = self.client.load_full();
414        let result = client.read_resource(params, timeout).await?;
415        Ok(McpResourceData {
416            provider: self.name.clone(),
417            uri: uri.to_string(),
418            contents: result.contents,
419            meta: Map::new(),
420        })
421    }
422
423    pub(super) async fn list_prompts(
424        &self,
425        allowlist: &McpAllowListConfig,
426        timeout: Option<Duration>,
427    ) -> Result<Vec<McpPromptInfo>> {
428        Ok(self.list_prompts_shared(allowlist, timeout).await?.as_ref().clone())
429    }
430
431    async fn list_prompts_shared(
432        &self,
433        allowlist: &McpAllowListConfig,
434        timeout: Option<Duration>,
435    ) -> Result<Arc<Vec<McpPromptInfo>>> {
436        let mut caches = self.caches.lock().await;
437        if self.client.load_full().take_prompt_list_changed() {
438            caches.prompts = None;
439        }
440
441        if let Some(cache) = &caches.prompts {
442            return Ok(Arc::clone(cache));
443        }
444        drop(caches);
445
446        self.refresh_prompts_shared(allowlist, timeout).await
447    }
448
449    pub(super) async fn refresh_prompts(
450        &self,
451        allowlist: &McpAllowListConfig,
452        timeout: Option<Duration>,
453    ) -> Result<Vec<McpPromptInfo>> {
454        Ok(self.refresh_prompts_shared(allowlist, timeout).await?.as_ref().clone())
455    }
456
457    async fn refresh_prompts_shared(
458        &self,
459        allowlist: &McpAllowListConfig,
460        timeout: Option<Duration>,
461    ) -> Result<Arc<Vec<McpPromptInfo>>> {
462        let client = self.client.load_full();
463        let prompts = client.list_all_prompts(timeout).await?;
464        let filtered = Arc::new(self.filter_prompts(prompts, allowlist));
465        self.caches.lock().await.prompts = Some(Arc::clone(&filtered));
466        Ok(filtered)
467    }
468
469    pub(super) async fn has_prompt(
470        &self,
471        prompt_name: &str,
472        allowlist: &McpAllowListConfig,
473        timeout: Option<Duration>,
474    ) -> Result<bool> {
475        let prompts = self.list_prompts_shared(allowlist, timeout).await?;
476        Ok(prompts.iter().any(|prompt| prompt.name == prompt_name))
477    }
478
479    pub(super) async fn get_prompt(
480        &self,
481        prompt_name: &str,
482        arguments: HashMap<String, String>,
483        timeout: Option<Duration>,
484        allowlist: &McpAllowListConfig,
485    ) -> Result<McpPromptDetail> {
486        if !allowlist.is_prompt_allowed(&self.name, prompt_name) {
487            return Err(anyhow!(
488                "Prompt '{}' is blocked by the MCP allow list for provider '{}'",
489                prompt_name,
490                self.name
491            ));
492        }
493
494        let _permit = self
495            .semaphore
496            .clone()
497            .acquire_owned()
498            .await
499            .context("Failed to acquire MCP request slot")?;
500        // Convert HashMap<String, String> to JsonObject (BTreeMap<String, Value>)
501        let args_json: Map<String, Value> = arguments.into_iter().map(|(k, v)| (k, Value::String(v))).collect();
502
503        let params = GetPromptRequestParams::new(prompt_name.to_string()).with_arguments(args_json);
504        let client = self.client.load_full();
505        let result = client.get_prompt(params, timeout).await?;
506        Ok(McpPromptDetail {
507            provider: self.name.clone(),
508            name: prompt_name.to_string(),
509            description: result.description,
510            messages: result.messages,
511            meta: Map::new(),
512        })
513    }
514
515    /// Return the cached tool list as a shared handle.
516    ///
517    /// The cache already stores `Arc<Vec<McpToolInfo>>`; returning the `Arc`
518    /// avoids deep-cloning every tool description and JSON schema on each hit.
519    pub(super) async fn cached_tools_shared(&self) -> Option<Arc<Vec<McpToolInfo>>> {
520        self.caches.lock().await.tools.as_ref().map(Arc::clone)
521    }
522
523    /// Return cached tools or refresh them, handing back the shared cache entry.
524    ///
525    /// Callers that only need to read/record the list can keep the `Arc` and
526    /// skip the per-call deep clone that the owned-`Vec` return used to force.
527    pub(super) async fn cached_tools_or_refresh_shared(
528        &self,
529        allowlist: &McpAllowListConfig,
530        timeout: Option<Duration>,
531    ) -> Result<Arc<Vec<McpToolInfo>>> {
532        if let Some(tools) = self.cached_tools_shared().await {
533            return Ok(tools);
534        }
535
536        self.refresh_tools_shared(allowlist, timeout).await
537    }
538
539    pub(super) async fn shutdown(&self) -> Result<()> {
540        let client = self.client.load_full();
541        client.shutdown().await
542    }
543
544    /// Returns `true` when the underlying transport is still connected and responsive.
545    pub(super) async fn is_healthy(&self) -> bool {
546        let client = self.client.load_full();
547        client.is_healthy().await
548    }
549
550    /// Attempt to re-establish the MCP connection using the stored configuration.
551    ///
552    /// This replaces the inner [`RmcpClient`] with a freshly connected one, then
553    /// re-initialises the provider (tools/resources/prompts caches are invalidated).
554    pub(super) async fn reconnect(
555        &self,
556        startup_timeout: Option<Duration>,
557        tool_timeout: Option<Duration>,
558        allowlist: &McpAllowListConfig,
559    ) -> Result<()> {
560        tracing::info!(provider = self.name.as_str(), "Attempting MCP reconnection");
561
562        // Shut down the old client (best-effort).
563        {
564            let old = self.client.load_full();
565            drop(old.shutdown().await);
566        }
567
568        // Create a fresh client from the stored config.
569        let new_provider =
570            McpProvider::connect(self.config.clone(), self.elicitation_handler.clone(), self.sandbox_context.clone())
571                .await
572                .with_context(|| format!("MCP reconnect failed for provider '{}'", self.name))?;
573
574        // Swap inner client.
575        {
576            let new_client = new_provider.client.load_full();
577            self.client.store(new_client);
578        }
579
580        // Invalidate all caches before re-initialisation.
581        self.invalidate_caches();
582
583        // Re-initialise (handshake + tool refresh).
584        let init_params = InitializeRequestParams::new(
585            rmcp::model::ClientCapabilities::default(),
586            super::utils::build_client_implementation(),
587        )
588        .with_protocol_version(latest_protocol_version());
589        self.initialize(init_params, startup_timeout, tool_timeout, allowlist)
590            .await
591            .with_context(|| format!("MCP re-initialization failed for provider '{}'", self.name))?;
592
593        tracing::info!(provider = self.name.as_str(), "MCP reconnection successful");
594        Ok(())
595    }
596
597    fn filter_tools(&self, tools: Vec<Tool>, allowlist: &McpAllowListConfig) -> Vec<McpToolInfo> {
598        filter_tools_sorted(&self.name, tools, allowlist)
599    }
600
601    fn filter_resources(&self, resources: Vec<Resource>, allowlist: &McpAllowListConfig) -> Vec<McpResourceInfo> {
602        filter_resources_sorted(&self.name, resources, allowlist)
603    }
604
605    fn filter_prompts(&self, prompts: Vec<Prompt>, allowlist: &McpAllowListConfig) -> Vec<McpPromptInfo> {
606        filter_prompts_sorted(&self.name, prompts, allowlist)
607    }
608}
609
610// Servers may reorder lists between calls; a stable order keeps LLM prompt-cache prefixes valid.
611fn filter_tools_sorted(provider: &str, tools: Vec<Tool>, allowlist: &McpAllowListConfig) -> Vec<McpToolInfo> {
612    let mut filtered: Vec<McpToolInfo> = tools
613        .into_iter()
614        .filter(|tool| allowlist.is_tool_allowed(provider, &tool.name))
615        .map(|tool| {
616            let parsed = parse_mcp_tool(&tool);
617            McpToolInfo {
618                description: parsed.description,
619                input_schema: parsed.input_schema,
620                output_schema: parsed.output_schema,
621                provider: provider.to_owned(),
622                name: parsed.name,
623            }
624        })
625        .collect();
626    filtered.sort_by(|a, b| a.name.cmp(&b.name));
627    filtered
628}
629
630fn filter_resources_sorted(
631    provider: &str,
632    resources: Vec<Resource>,
633    allowlist: &McpAllowListConfig,
634) -> Vec<McpResourceInfo> {
635    let mut filtered: Vec<McpResourceInfo> = resources
636        .into_iter()
637        .filter(|resource| allowlist.is_resource_allowed(provider, &resource.uri))
638        .map(|resource| McpResourceInfo {
639            provider: provider.to_owned(),
640            uri: resource.uri.clone(),
641            name: resource.name.clone(),
642            description: resource.description.clone(),
643            mime_type: resource.mime_type.clone(),
644            size: resource.size.map(|s| i64::try_from(s).unwrap_or(i64::MAX)),
645        })
646        .collect();
647    filtered.sort_by(|a, b| a.uri.cmp(&b.uri));
648    filtered
649}
650
651fn filter_prompts_sorted(provider: &str, prompts: Vec<Prompt>, allowlist: &McpAllowListConfig) -> Vec<McpPromptInfo> {
652    let mut filtered: Vec<McpPromptInfo> = prompts
653        .into_iter()
654        .filter(|prompt| allowlist.is_prompt_allowed(provider, &prompt.name))
655        .map(|prompt| McpPromptInfo {
656            provider: provider.to_owned(),
657            name: prompt.name.clone(),
658            description: prompt.description.clone(),
659            arguments: prompt.arguments.clone().unwrap_or_default(),
660        })
661        .collect();
662    filtered.sort_by(|a, b| a.name.cmp(&b.name));
663    filtered
664}
665
666fn mcp_tool_call_span(provider_name: &str, tool_name: &str, transport: &McpTransportConfig) -> Span {
667    let (transport_label, server_address, server_port) = match transport {
668        McpTransportConfig::Stdio(_) => ("stdio", String::new(), 0_u16),
669        McpTransportConfig::Http(http) => {
670            let (server_address, server_port) = Url::parse(&http.endpoint)
671                .ok()
672                .and_then(|url| {
673                    url.host_str()
674                        .map(|host| (host.to_string(), url.port_or_known_default().unwrap_or_default()))
675                })
676                .unwrap_or_default();
677            ("streamable_http", server_address, server_port)
678        }
679    };
680
681    tracing::info_span!(
682        "mcp.tools.call",
683        provider = provider_name,
684        tool = tool_name,
685        rpc_system = "jsonrpc",
686        rpc_method = "tools/call",
687        transport = transport_label,
688        server_address = server_address.as_str(),
689        server_port = server_port,
690    )
691}
692
693#[cfg(test)]
694mod tests {
695    use super::mcp_tool_call_span;
696    use std::fs::{self, File, OpenOptions};
697    use std::io::{BufWriter, Write};
698    use std::sync::{Arc, Mutex};
699    use tempfile::tempdir;
700    use tracing_subscriber::{fmt::format::FmtSpan, prelude::*};
701    use vtcode_config::mcp::{McpHttpServerConfig, McpStdioServerConfig, McpTransportConfig};
702
703    /// Minimal clonable buffered writer for tests. Implements `Write` so it can
704    /// be passed to `tracing_subscriber::fmt::layer().with_writer(..)`.
705    #[derive(Clone)]
706    struct TestWriter(Arc<Mutex<BufWriter<File>>>);
707
708    impl TestWriter {
709        fn open(path: &std::path::Path) -> Self {
710            let file = OpenOptions::new()
711                .create(true)
712                .truncate(true)
713                .write(true)
714                .open(path)
715                .expect("open trace log");
716            Self(Arc::new(Mutex::new(BufWriter::new(file))))
717        }
718
719        fn flush(&self) {
720            if let Ok(mut w) = self.0.lock() {
721                drop(w.flush());
722            }
723        }
724    }
725
726    impl Write for TestWriter {
727        fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
728            self.0.lock().map_err(|e| std::io::Error::other(e.to_string()))?.write(buf)
729        }
730
731        fn flush(&mut self) -> std::io::Result<()> {
732            self.0.lock().map_err(|e| std::io::Error::other(e.to_string()))?.flush()
733        }
734    }
735
736    #[test]
737    fn mcp_tool_call_span_records_http_metadata() {
738        let tempdir = tempdir().expect("tempdir");
739        let log_file = tempdir.path().join("trace.log");
740        let writer = TestWriter::open(&log_file);
741        let writer_for_layer = writer.clone();
742        let subscriber = tracing_subscriber::registry().with(
743            tracing_subscriber::fmt::layer()
744                .with_writer(move || writer_for_layer.clone())
745                .with_span_events(FmtSpan::FULL)
746                .with_ansi(false),
747        );
748        let _guard = tracing::subscriber::set_default(subscriber);
749
750        {
751            let span = mcp_tool_call_span(
752                "calendar",
753                "get_events",
754                &McpTransportConfig::Http(McpHttpServerConfig {
755                    endpoint: "https://example.com:8443/mcp".to_string(),
756                    api_key_env: None,
757                    oauth: None,
758                    protocol_version: "2024-11-05".to_string(),
759                    handshake: McpHttpHandshakeMode::Legacy,
760                    http_headers: Default::default(),
761                    env_http_headers: Default::default(),
762                }),
763            );
764            let _entered = span.enter();
765        }
766
767        writer.flush();
768        let logs = fs::read_to_string(&log_file).expect("trace log");
769        assert!(logs.contains("mcp.tools.call"));
770        assert!(logs.contains("provider=\"calendar\""));
771        assert!(logs.contains("tool=\"get_events\""));
772        assert!(logs.contains("rpc_system=\"jsonrpc\""));
773        assert!(logs.contains("rpc_method=\"tools/call\""));
774        assert!(logs.contains("transport=\"streamable_http\""));
775        assert!(logs.contains("server_address=\"example.com\""));
776        assert!(logs.contains("server_port=8443"));
777    }
778
779    #[test]
780    fn mcp_tool_call_span_defaults_stdio_transport() {
781        let span = mcp_tool_call_span(
782            "filesystem",
783            "read_file",
784            &McpTransportConfig::Stdio(McpStdioServerConfig {
785                command: "rmcp-server".to_string(),
786                args: Vec::new(),
787                working_directory: None,
788            }),
789        );
790
791        assert_eq!(span.metadata().expect("metadata").name(), "mcp.tools.call");
792    }
793
794    use rmcp::model::{ProtocolVersion, ServerCapabilities, ServerPeerInfo};
795    use rmcp::service::ClientLifecycleMode;
796    use vtcode_config::mcp::{McpHttpHandshakeMode, McpProviderConfig};
797
798    async fn stdio_provider(name: &str) -> super::McpProvider {
799        let config = McpProviderConfig {
800            name: name.to_string(),
801            transport: McpTransportConfig::Stdio(McpStdioServerConfig {
802                command: "true".to_string(),
803                args: Vec::new(),
804                working_directory: None,
805            }),
806            ..McpProviderConfig::default()
807        };
808        super::McpProvider::connect(config, None, None)
809            .await
810            .expect("stdio provider connects")
811    }
812
813    async fn http_provider(name: &str, handshake: McpHttpHandshakeMode) -> super::McpProvider {
814        let config = McpProviderConfig {
815            name: name.to_string(),
816            transport: McpTransportConfig::Http(McpHttpServerConfig {
817                endpoint: "https://example.com/mcp".to_string(),
818                handshake,
819                ..McpHttpServerConfig::default()
820            }),
821            ..McpProviderConfig::default()
822        };
823        super::McpProvider::connect(config, None, None)
824            .await
825            .expect("http provider connects")
826    }
827
828    #[tokio::test]
829    async fn handshake_lifecycle_defaults_to_legacy_initialize_for_http() {
830        let provider = http_provider("legacy", McpHttpHandshakeMode::Legacy).await;
831        assert!(matches!(provider.handshake_lifecycle(), ClientLifecycleMode::Initialize));
832    }
833
834    #[tokio::test]
835    async fn handshake_lifecycle_uses_auto_for_stdio_and_opt_in_http() {
836        let stdio = stdio_provider("local").await;
837        match stdio.handshake_lifecycle() {
838            ClientLifecycleMode::Auto { preferred_versions, legacy_version } => {
839                let versions: Vec<String> = preferred_versions.iter().map(ToString::to_string).collect();
840                assert_eq!(versions, vec!["2025-11-25", "2025-06-18", "2025-03-26", "2024-11-05"]);
841                assert_eq!(legacy_version.map(|version| version.to_string()), Some("2024-11-05".to_string()));
842            }
843            other => panic!("expected Auto lifecycle, got {other:?}"),
844        }
845
846        let modern = http_provider("modern", McpHttpHandshakeMode::Auto).await;
847        assert!(matches!(modern.handshake_lifecycle(), ClientLifecycleMode::Auto { .. }));
848    }
849
850    #[tokio::test]
851    async fn negotiated_protocol_version_tracks_last_handshake() {
852        let provider = http_provider("probe", McpHttpHandshakeMode::Legacy).await;
853        assert_eq!(provider.negotiated_protocol_version().await, None);
854
855        *provider.initialize_result.lock().await =
856            Some(ServerPeerInfo::new(ProtocolVersion::V_2025_11_25, ServerCapabilities::default()));
857        assert_eq!(provider.negotiated_protocol_version().await.as_deref(), Some("2025-11-25"));
858    }
859
860    #[test]
861    fn filtered_catalogs_are_sorted_regardless_of_server_order() {
862        use super::{filter_prompts_sorted, filter_resources_sorted, filter_tools_sorted};
863        use rmcp::model::{Prompt, Resource, Tool};
864        use serde_json::Map;
865        use std::sync::Arc;
866        use vtcode_config::mcp::McpAllowListConfig;
867
868        let allowlist = McpAllowListConfig::default();
869        // Uppercase sorts before lowercase and a prefix sorts before its extension in byte order.
870        let tool_names = ["search_v2", "Zeta", "search", "alpha"];
871        let tools = tool_names
872            .iter()
873            .map(|name| Tool::new(*name, "desc", Arc::new(Map::new())))
874            .collect();
875        let sorted_tools: Vec<String> = filter_tools_sorted("p", tools, &allowlist)
876            .into_iter()
877            .map(|tool| tool.name)
878            .collect();
879        assert_eq!(sorted_tools, ["Zeta", "alpha", "search", "search_v2"]);
880
881        let resources = ["file:///b", "file:///a/z", "file:///a"]
882            .iter()
883            .map(|uri| Resource::new(*uri, "r"))
884            .collect();
885        let sorted_resources: Vec<String> = filter_resources_sorted("p", resources, &allowlist)
886            .into_iter()
887            .map(|resource| resource.uri)
888            .collect();
889        assert_eq!(sorted_resources, ["file:///a", "file:///a/z", "file:///b"]);
890
891        let prompts = ["review", "Explain", "review2"]
892            .iter()
893            .map(|name| Prompt::new(*name, None::<String>, None))
894            .collect();
895        let sorted_prompts: Vec<String> = filter_prompts_sorted("p", prompts, &allowlist)
896            .into_iter()
897            .map(|prompt| prompt.name)
898            .collect();
899        assert_eq!(sorted_prompts, ["Explain", "review", "review2"]);
900    }
901}