Skip to main content

vtcode_config/mcp/
transport.rs

1//! MCP provider transport configuration and deserialization.
2//!
3//! Extracted from the `mcp` module so the transport wire-decoding lives behind
4//! a strict interface: [`McpProviderConfig`] is the only public entry point and
5//! [`McpProviderConfigWire`] is private to this module. Callers never see the
6//! flat wire shape — they construct or deserialize [`McpProviderConfig`] and
7//! pattern-match on [`McpTransportConfig`].
8//!
9//! The manual [`Deserialize`] impl avoids the `#[serde(flatten)]` +
10//! `#[serde(untagged)]` map-buffering overhead: every transport field is
11//! decoded in a single pass and the transport enum is constructed afterwards,
12//! preserving the untagged Stdio-then-Http precedence. See the regression tests
13//! at the bottom of this file for the exact dispatch contract.
14//!
15//! # Validation semantics (intentional guardrail)
16//!
17//! Because every recognized field is decoded against its declared type in that
18//! single pass, a malformed *known* field is rejected even when it belongs to
19//! the transport variant that was not selected. The previous `#[serde(untagged)]`
20//! path trial-deserialized Stdio first and silently ignored irrelevant HTTP
21//! fields (and vice-versa), so a stray wrong-typed `endpoint` on a valid Stdio
22//! provider used to parse. Under the flat wire it now errors. This stricter
23//! behavior is an intentional interface guardrail — surfacing malformed known
24//! configuration early — and is not a behavior-preserving no-op relative to the
25//! derived path. Forward compatibility for *unknown* fields is unaffected: the
26//! wire ignores any field it does not declare.
27
28use crate::env_helpers::default_enabled;
29use hashbrown::HashMap;
30use serde::{Deserialize, Deserializer, Serialize};
31use std::collections::BTreeMap;
32use vtcode_auth::McpOAuthConfig;
33
34/// Transport configuration for MCP providers
35#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
36#[allow(
37    clippy::large_enum_variant,
38    reason = "Intentional compatibility, platform, or test-only suppression."
39)]
40#[derive(Debug, Clone, Deserialize, Serialize)]
41#[serde(untagged)]
42pub enum McpTransportConfig {
43    /// Standard I/O transport (stdio)
44    Stdio(McpStdioServerConfig),
45    /// HTTP transport
46    Http(McpHttpServerConfig),
47}
48
49/// Configuration for stdio-based MCP servers
50#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
51#[derive(Debug, Clone, Deserialize, Serialize, Default)]
52pub struct McpStdioServerConfig {
53    /// Command to execute
54    pub command: String,
55
56    /// Command arguments
57    pub args: Vec<String>,
58
59    /// Working directory for the command
60    #[serde(default)]
61    pub working_directory: Option<String>,
62}
63
64/// Pinned MCP protocol versions (single source of truth for version strings).
65///
66/// The typed rmcp counterparts live in `vtcode-mcp` (`stable_protocol_version`,
67/// `legacy_fallback_protocol_version`); keep both sides aligned when the spec
68/// publishes a new stable revision.
69pub const MCP_STABLE_PROTOCOL_VERSION: &str = "2025-11-25";
70pub const MCP_LEGACY_PROTOCOL_VERSION: &str = "2024-11-05";
71
72/// Handshake strategy for HTTP-based MCP servers.
73///
74/// `Legacy` performs the `initialize` / `notifications/initialized` handshake
75/// directly (rmcp `ClientLifecycleMode::Initialize`). `Auto` first probes
76/// `server/discover` and falls back to the legacy handshake when the peer
77/// reports a legacy server or stays silent past the discover timeout (rmcp
78/// `ClientLifecycleMode::Auto`). Legacy is the default: it avoids the
79/// discover-timeout penalty against legacy-only servers.
80#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
81#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, Deserialize, Serialize)]
82#[serde(rename_all = "snake_case")]
83pub enum McpHttpHandshakeMode {
84    #[default]
85    Legacy,
86    Auto,
87}
88
89/// Configuration for HTTP-based MCP servers
90///
91/// Note: HTTP transport is partially implemented. Basic connectivity testing is supported,
92/// but full streamable HTTP MCP server support requires additional implementation
93/// using Server-Sent Events (SSE) or WebSocket connections.
94#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
95#[derive(Debug, Clone, Deserialize, Serialize)]
96pub struct McpHttpServerConfig {
97    /// Server endpoint URL
98    pub endpoint: String,
99
100    /// API key environment variable name
101    #[serde(default)]
102    pub api_key_env: Option<String>,
103
104    /// Optional OAuth configuration for providers that issue bearer tokens dynamically.
105    #[serde(default)]
106    pub oauth: Option<McpOAuthConfig>,
107
108    /// Protocol version
109    #[serde(default = "default_mcp_protocol_version")]
110    pub protocol_version: String,
111
112    /// Handshake strategy (`legacy` direct initialize, or `auto` discover
113    /// with legacy fallback). Defaults to `legacy`.
114    #[serde(default)]
115    pub handshake: McpHttpHandshakeMode,
116
117    /// Headers to include in requests
118    #[serde(default, alias = "headers")]
119    #[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
120    pub http_headers: HashMap<String, String>,
121
122    /// Headers whose values are sourced from environment variables
123    /// (`{ header-name = "ENV_VAR" }`). Empty values are ignored.
124    #[serde(default)]
125    #[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
126    pub env_http_headers: HashMap<String, String>,
127}
128
129impl Default for McpHttpServerConfig {
130    fn default() -> Self {
131        Self {
132            endpoint: String::new(),
133            api_key_env: None,
134            oauth: None,
135            protocol_version: default_mcp_protocol_version(),
136            handshake: McpHttpHandshakeMode::default(),
137            http_headers: HashMap::new(),
138            env_http_headers: HashMap::new(),
139        }
140    }
141}
142
143/// Configuration for a single MCP provider
144#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
145#[derive(Debug, Clone, Serialize)]
146pub struct McpProviderConfig {
147    /// Provider name (used for identification)
148    pub name: String,
149
150    /// Transport configuration
151    #[serde(flatten)]
152    pub transport: McpTransportConfig,
153
154    /// Provider-specific environment variables
155    #[serde(default)]
156    #[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
157    pub env: HashMap<String, String>,
158
159    /// Whether this provider is enabled
160    #[serde(default = "default_provider_enabled")]
161    pub enabled: bool,
162
163    /// Maximum number of concurrent requests to this provider
164    #[serde(default = "default_provider_max_concurrent")]
165    pub max_concurrent_requests: usize,
166
167    /// Startup timeout in milliseconds for this provider
168    #[serde(default)]
169    pub startup_timeout_ms: Option<u64>,
170}
171
172impl<'de> Deserialize<'de> for McpProviderConfig {
173    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
174    where
175        D: Deserializer<'de>,
176    {
177        // Deserialize every provider + transport field in a single pass.
178        // `#[serde(flatten)]` + `#[serde(untagged)]` on `transport` would force
179        // Serde to buffer the whole provider object into a `Map<String, Value>`
180        // and then trial-deserialize each transport variant. A provider has only
181        // two mutually-exclusive transport forms — stdio (`command` + `args`) or
182        // HTTP (`endpoint`) — so the fields are decoded directly and the
183        // transport enum is constructed afterwards, matching the untagged
184        // Stdio-then-Http precedence.
185        //
186        // Note: because every declared field is type-checked in this pass, a
187        // malformed known field is rejected even on the non-selected variant
188        // (stricter than the old untagged trial-deserialize, which ignored
189        // irrelevant fields). See the module-level "Validation semantics" note.
190        let wire = McpProviderConfigWire::deserialize(deserializer)?;
191
192        let transport = if let (Some(command), Some(args)) = (wire.command, wire.args) {
193            McpTransportConfig::Stdio(McpStdioServerConfig {
194                command,
195                args,
196                working_directory: wire.working_directory,
197            })
198        } else if let Some(endpoint) = wire.endpoint {
199            McpTransportConfig::Http(McpHttpServerConfig {
200                endpoint,
201                api_key_env: wire.api_key_env,
202                oauth: wire.oauth,
203                protocol_version: wire.protocol_version,
204                handshake: wire.handshake,
205                http_headers: wire.http_headers,
206                env_http_headers: wire.env_http_headers,
207            })
208        } else {
209            return Err(serde::de::Error::custom(
210                "MCP provider must specify either a stdio `command` (with `args`) or an HTTP `endpoint`",
211            ));
212        };
213
214        Ok(McpProviderConfig {
215            name: wire.name,
216            transport,
217            env: wire.env,
218            enabled: wire.enabled,
219            max_concurrent_requests: wire.max_concurrent_requests,
220            startup_timeout_ms: wire.startup_timeout_ms,
221        })
222    }
223}
224
225/// Flat wire shape for [`McpProviderConfig`] (see [`McpProviderConfig::deserialize`]).
226///
227/// Every transport field is declared directly so no intermediate map is
228/// buffered. Fields belonging to the unused transport variant simply stay at
229/// their `Option`/default values and are ignored when constructing the
230/// [`McpTransportConfig`].
231#[derive(Deserialize)]
232struct McpProviderConfigWire {
233    name: String,
234    // stdio transport
235    #[serde(default)]
236    command: Option<String>,
237    #[serde(default)]
238    args: Option<Vec<String>>,
239    #[serde(default)]
240    working_directory: Option<String>,
241    // http transport
242    #[serde(default)]
243    endpoint: Option<String>,
244    #[serde(default)]
245    api_key_env: Option<String>,
246    #[serde(default)]
247    oauth: Option<McpOAuthConfig>,
248    #[serde(default = "default_mcp_protocol_version")]
249    protocol_version: String,
250    #[serde(default)]
251    handshake: McpHttpHandshakeMode,
252    #[serde(default, alias = "headers")]
253    http_headers: HashMap<String, String>,
254    #[serde(default)]
255    env_http_headers: HashMap<String, String>,
256    // provider-level
257    #[serde(default)]
258    env: HashMap<String, String>,
259    #[serde(default = "default_provider_enabled")]
260    enabled: bool,
261    #[serde(default = "default_provider_max_concurrent")]
262    max_concurrent_requests: usize,
263    #[serde(default)]
264    startup_timeout_ms: Option<u64>,
265}
266
267impl Default for McpProviderConfig {
268    fn default() -> Self {
269        Self {
270            name: String::new(),
271            transport: McpTransportConfig::Stdio(McpStdioServerConfig::default()),
272            env: HashMap::new(),
273            enabled: default_provider_enabled(),
274            max_concurrent_requests: default_provider_max_concurrent(),
275            startup_timeout_ms: None,
276        }
277    }
278}
279
280fn default_provider_enabled() -> bool {
281    default_enabled()
282}
283
284fn default_provider_max_concurrent() -> usize {
285    3
286}
287
288fn default_mcp_protocol_version() -> String {
289    // Default to the current stable spec revision, not the draft. Draft
290    // revisions are in-progress and not ready for consumption
291    // (https://modelcontextprotocol.io/docs/draft/learn/versioning.md);
292    // explicit opt-in to a draft remains possible via config and is still
293    // clamped to the last stable version on the legacy handshake path.
294    MCP_STABLE_PROTOCOL_VERSION.into()
295}
296
297#[cfg(test)]
298mod tests {
299    use super::*;
300
301    #[test]
302    fn test_mcp_provider_config_http_transport_from_toml() {
303        // HTTP provider, discriminated by `endpoint`. Locks in the flat-wire
304        // deserialize path (no `#[serde(flatten)]` + `#[serde(untagged)]`
305        // buffering) for the HTTP transport branch.
306        let toml_str = r#"
307name = "deepwiki"
308enabled = true
309endpoint = "https://mcp.deepwiki.com/mcp"
310protocol_version = "2024-11-05"
311max_concurrent_requests = 3
312
313[http_headers]
314Authorization = "Bearer token"
315"#;
316        let provider: McpProviderConfig = toml::from_str(toml_str).expect("http provider must parse");
317
318        assert_eq!(provider.name, "deepwiki");
319        assert!(provider.enabled);
320        assert_eq!(provider.max_concurrent_requests, 3);
321        match provider.transport {
322            McpTransportConfig::Http(http) => {
323                assert_eq!(http.endpoint, "https://mcp.deepwiki.com/mcp");
324                assert_eq!(http.protocol_version, "2024-11-05");
325                assert_eq!(http.http_headers.get("Authorization"), Some(&"Bearer token".to_string()));
326            }
327            McpTransportConfig::Stdio(_) => panic!("expected HTTP transport"),
328        }
329    }
330
331    #[test]
332    fn test_mcp_provider_config_stdio_transport_from_toml() {
333        let toml_str = r#"
334name = "time"
335command = "uvx"
336args = ["mcp-server-time"]
337working_directory = "/tmp"
338"#;
339        let provider: McpProviderConfig = toml::from_str(toml_str).expect("stdio provider must parse");
340        match provider.transport {
341            McpTransportConfig::Stdio(stdio) => {
342                assert_eq!(stdio.command, "uvx");
343                assert_eq!(stdio.args, vec!["mcp-server-time"]);
344                assert_eq!(stdio.working_directory.as_deref(), Some("/tmp"));
345            }
346            McpTransportConfig::Http(_) => panic!("expected stdio transport"),
347        }
348    }
349
350    #[test]
351    fn test_mcp_provider_config_stdio_wins_when_both_transports_present() {
352        // Replicates the untagged Stdio-then-Http precedence: `command` + `args`
353        // select stdio even when `endpoint` is also set.
354        let toml_str = r#"
355name = "mixed"
356command = "uvx"
357args = ["mcp-server-time"]
358endpoint = "https://example.com/mcp"
359"#;
360        let provider: McpProviderConfig = toml::from_str(toml_str).expect("mixed provider must parse");
361        assert!(matches!(provider.transport, McpTransportConfig::Stdio(_)));
362    }
363
364    #[test]
365    fn test_mcp_provider_config_http_fallback_when_command_lacks_args() {
366        // `command` without `args` cannot form a stdio transport, so the
367        // presence of `endpoint` selects HTTP (matches untagged fallback).
368        let toml_str = r#"
369name = "fallback"
370command = "uvx"
371endpoint = "https://example.com/mcp"
372"#;
373        let provider: McpProviderConfig = toml::from_str(toml_str).expect("fallback provider must parse");
374        assert!(matches!(provider.transport, McpTransportConfig::Http(_)));
375    }
376
377    #[test]
378    fn test_mcp_provider_config_rejects_missing_transport() {
379        let toml_str = "name = \"bare\"";
380        let result: Result<McpProviderConfig, _> = toml::from_str(toml_str);
381        assert!(result.is_err(), "provider without command/endpoint must error");
382    }
383
384    #[test]
385    fn test_mcp_provider_config_rejects_malformed_known_field_on_non_selected_transport() {
386        // Intentional strict guardrail: a wrong-typed *known* field is rejected
387        // even when it belongs to the transport variant that was not selected.
388        // The old `#[serde(untagged)]` path would select Stdio (valid
389        // `command` + `args`) and silently ignore the malformed `endpoint`; the
390        // flat wire type-checks every declared field in one pass, so this now
391        // errors. See the module-level "Validation semantics" note.
392        let toml_str = r#"
393name = "strict"
394command = "uvx"
395args = ["mcp-server-time"]
396endpoint = 42
397"#;
398        let result: Result<McpProviderConfig, _> = toml::from_str(toml_str);
399        assert!(result.is_err(), "malformed known field on non-selected transport must error under the flat wire");
400    }
401
402    #[test]
403    fn test_mcp_http_handshake_defaults_to_legacy() {
404        let config = McpHttpServerConfig::default();
405        assert_eq!(config.handshake, McpHttpHandshakeMode::Legacy);
406
407        let toml_str = r#"
408name = "plain"
409endpoint = "https://example.com/mcp"
410"#;
411        let provider: McpProviderConfig = toml::from_str(toml_str).expect("http provider must parse");
412        match provider.transport {
413            McpTransportConfig::Http(http) => assert_eq!(http.handshake, McpHttpHandshakeMode::Legacy),
414            McpTransportConfig::Stdio(_) => panic!("expected HTTP transport"),
415        }
416    }
417
418    #[test]
419    fn test_mcp_http_handshake_parses_auto() {
420        let toml_str = r#"
421name = "modern"
422endpoint = "https://example.com/mcp"
423handshake = "auto"
424"#;
425        let provider: McpProviderConfig = toml::from_str(toml_str).expect("http provider must parse");
426        match provider.transport {
427            McpTransportConfig::Http(http) => assert_eq!(http.handshake, McpHttpHandshakeMode::Auto),
428            McpTransportConfig::Stdio(_) => panic!("expected HTTP transport"),
429        }
430    }
431}