1use crate::env_helpers::default_enabled;
29use hashbrown::HashMap;
30use serde::{Deserialize, Deserializer, Serialize};
31use std::collections::BTreeMap;
32use vtcode_auth::McpOAuthConfig;
33
34#[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 Stdio(McpStdioServerConfig),
45 Http(McpHttpServerConfig),
47}
48
49#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
51#[derive(Debug, Clone, Deserialize, Serialize, Default)]
52pub struct McpStdioServerConfig {
53 pub command: String,
55
56 pub args: Vec<String>,
58
59 #[serde(default)]
61 pub working_directory: Option<String>,
62}
63
64pub const MCP_STABLE_PROTOCOL_VERSION: &str = "2025-11-25";
70pub const MCP_LEGACY_PROTOCOL_VERSION: &str = "2024-11-05";
71
72#[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#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
95#[derive(Debug, Clone, Deserialize, Serialize)]
96pub struct McpHttpServerConfig {
97 pub endpoint: String,
99
100 #[serde(default)]
102 pub api_key_env: Option<String>,
103
104 #[serde(default)]
106 pub oauth: Option<McpOAuthConfig>,
107
108 #[serde(default = "default_mcp_protocol_version")]
110 pub protocol_version: String,
111
112 #[serde(default)]
115 pub handshake: McpHttpHandshakeMode,
116
117 #[serde(default, alias = "headers")]
119 #[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
120 pub http_headers: HashMap<String, String>,
121
122 #[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#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
145#[derive(Debug, Clone, Serialize)]
146pub struct McpProviderConfig {
147 pub name: String,
149
150 #[serde(flatten)]
152 pub transport: McpTransportConfig,
153
154 #[serde(default)]
156 #[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
157 pub env: HashMap<String, String>,
158
159 #[serde(default = "default_provider_enabled")]
161 pub enabled: bool,
162
163 #[serde(default = "default_provider_max_concurrent")]
165 pub max_concurrent_requests: usize,
166
167 #[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 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#[derive(Deserialize)]
232struct McpProviderConfigWire {
233 name: String,
234 #[serde(default)]
236 command: Option<String>,
237 #[serde(default)]
238 args: Option<Vec<String>>,
239 #[serde(default)]
240 working_directory: Option<String>,
241 #[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 #[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 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 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 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 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 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}