use crate::env_helpers::default_enabled;
use hashbrown::HashMap;
use serde::{Deserialize, Deserializer, Serialize};
use std::collections::BTreeMap;
use vtcode_auth::McpOAuthConfig;
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[allow(clippy::large_enum_variant)]
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(untagged)]
pub enum McpTransportConfig {
Stdio(McpStdioServerConfig),
Http(McpHttpServerConfig),
}
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
pub struct McpStdioServerConfig {
pub command: String,
pub args: Vec<String>,
#[serde(default)]
pub working_directory: Option<String>,
}
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, Deserialize, Serialize)]
pub struct McpHttpServerConfig {
pub endpoint: String,
#[serde(default)]
pub api_key_env: Option<String>,
#[serde(default)]
pub oauth: Option<McpOAuthConfig>,
#[serde(default = "default_mcp_protocol_version")]
pub protocol_version: String,
#[serde(default, alias = "headers")]
#[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
pub http_headers: HashMap<String, String>,
#[serde(default)]
#[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
pub env_http_headers: HashMap<String, String>,
}
impl Default for McpHttpServerConfig {
fn default() -> Self {
Self {
endpoint: String::new(),
api_key_env: None,
oauth: None,
protocol_version: default_mcp_protocol_version(),
http_headers: HashMap::new(),
env_http_headers: HashMap::new(),
}
}
}
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
#[derive(Debug, Clone, Serialize)]
pub struct McpProviderConfig {
pub name: String,
#[serde(flatten)]
pub transport: McpTransportConfig,
#[serde(default)]
#[cfg_attr(feature = "schema", schemars(with = "BTreeMap<String, String>"))]
pub env: HashMap<String, String>,
#[serde(default = "default_provider_enabled")]
pub enabled: bool,
#[serde(default = "default_provider_max_concurrent")]
pub max_concurrent_requests: usize,
#[serde(default)]
pub startup_timeout_ms: Option<u64>,
}
impl<'de> Deserialize<'de> for McpProviderConfig {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let wire = McpProviderConfigWire::deserialize(deserializer)?;
let transport = if let (Some(command), Some(args)) = (wire.command, wire.args) {
McpTransportConfig::Stdio(McpStdioServerConfig {
command,
args,
working_directory: wire.working_directory,
})
} else if let Some(endpoint) = wire.endpoint {
McpTransportConfig::Http(McpHttpServerConfig {
endpoint,
api_key_env: wire.api_key_env,
oauth: wire.oauth,
protocol_version: wire.protocol_version,
http_headers: wire.http_headers,
env_http_headers: wire.env_http_headers,
})
} else {
return Err(serde::de::Error::custom(
"MCP provider must specify either a stdio `command` (with `args`) or an HTTP `endpoint`",
));
};
Ok(McpProviderConfig {
name: wire.name,
transport,
env: wire.env,
enabled: wire.enabled,
max_concurrent_requests: wire.max_concurrent_requests,
startup_timeout_ms: wire.startup_timeout_ms,
})
}
}
#[derive(Deserialize)]
struct McpProviderConfigWire {
name: String,
#[serde(default)]
command: Option<String>,
#[serde(default)]
args: Option<Vec<String>>,
#[serde(default)]
working_directory: Option<String>,
#[serde(default)]
endpoint: Option<String>,
#[serde(default)]
api_key_env: Option<String>,
#[serde(default)]
oauth: Option<McpOAuthConfig>,
#[serde(default = "default_mcp_protocol_version")]
protocol_version: String,
#[serde(default, alias = "headers")]
http_headers: HashMap<String, String>,
#[serde(default)]
env_http_headers: HashMap<String, String>,
#[serde(default)]
env: HashMap<String, String>,
#[serde(default = "default_provider_enabled")]
enabled: bool,
#[serde(default = "default_provider_max_concurrent")]
max_concurrent_requests: usize,
#[serde(default)]
startup_timeout_ms: Option<u64>,
}
impl Default for McpProviderConfig {
fn default() -> Self {
Self {
name: String::new(),
transport: McpTransportConfig::Stdio(McpStdioServerConfig::default()),
env: HashMap::new(),
enabled: default_provider_enabled(),
max_concurrent_requests: default_provider_max_concurrent(),
startup_timeout_ms: None,
}
}
}
fn default_provider_enabled() -> bool {
default_enabled()
}
fn default_provider_max_concurrent() -> usize {
3
}
fn default_mcp_protocol_version() -> String {
"2024-11-05".into()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mcp_provider_config_http_transport_from_toml() {
let toml_str = r#"
name = "deepwiki"
enabled = true
endpoint = "https://mcp.deepwiki.com/mcp"
protocol_version = "2024-11-05"
max_concurrent_requests = 3
[http_headers]
Authorization = "Bearer token"
"#;
let provider: McpProviderConfig = toml::from_str(toml_str).expect("http provider must parse");
assert_eq!(provider.name, "deepwiki");
assert!(provider.enabled);
assert_eq!(provider.max_concurrent_requests, 3);
match provider.transport {
McpTransportConfig::Http(http) => {
assert_eq!(http.endpoint, "https://mcp.deepwiki.com/mcp");
assert_eq!(http.protocol_version, "2024-11-05");
assert_eq!(http.http_headers.get("Authorization"), Some(&"Bearer token".to_string()));
}
McpTransportConfig::Stdio(_) => panic!("expected HTTP transport"),
}
}
#[test]
fn test_mcp_provider_config_stdio_transport_from_toml() {
let toml_str = r#"
name = "time"
command = "uvx"
args = ["mcp-server-time"]
working_directory = "/tmp"
"#;
let provider: McpProviderConfig = toml::from_str(toml_str).expect("stdio provider must parse");
match provider.transport {
McpTransportConfig::Stdio(stdio) => {
assert_eq!(stdio.command, "uvx");
assert_eq!(stdio.args, vec!["mcp-server-time"]);
assert_eq!(stdio.working_directory.as_deref(), Some("/tmp"));
}
McpTransportConfig::Http(_) => panic!("expected stdio transport"),
}
}
#[test]
fn test_mcp_provider_config_stdio_wins_when_both_transports_present() {
let toml_str = r#"
name = "mixed"
command = "uvx"
args = ["mcp-server-time"]
endpoint = "https://example.com/mcp"
"#;
let provider: McpProviderConfig = toml::from_str(toml_str).expect("mixed provider must parse");
assert!(matches!(provider.transport, McpTransportConfig::Stdio(_)));
}
#[test]
fn test_mcp_provider_config_http_fallback_when_command_lacks_args() {
let toml_str = r#"
name = "fallback"
command = "uvx"
endpoint = "https://example.com/mcp"
"#;
let provider: McpProviderConfig = toml::from_str(toml_str).expect("fallback provider must parse");
assert!(matches!(provider.transport, McpTransportConfig::Http(_)));
}
#[test]
fn test_mcp_provider_config_rejects_missing_transport() {
let toml_str = "name = \"bare\"";
let result: Result<McpProviderConfig, _> = toml::from_str(toml_str);
assert!(result.is_err(), "provider without command/endpoint must error");
}
#[test]
fn test_mcp_provider_config_rejects_malformed_known_field_on_non_selected_transport() {
let toml_str = r#"
name = "strict"
command = "uvx"
args = ["mcp-server-time"]
endpoint = 42
"#;
let result: Result<McpProviderConfig, _> = toml::from_str(toml_str);
assert!(result.is_err(), "malformed known field on non-selected transport must error under the flat wire");
}
}