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 config: McpProviderConfig,
41 elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
43 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 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 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 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 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 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 pub(super) async fn is_healthy(&self) -> bool {
546 let client = self.client.load_full();
547 client.is_healthy().await
548 }
549
550 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 {
564 let old = self.client.load_full();
565 drop(old.shutdown().await);
566 }
567
568 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 {
576 let new_client = new_provider.client.load_full();
577 self.client.store(new_client);
578 }
579
580 self.invalidate_caches();
582
583 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
610fn 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 #[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 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}