Skip to main content

vtcode_mcp/
enhanced_config.rs

1//! Enhanced MCP Configuration
2//!
3//! This module provides enhanced configuration options for MCP with
4//! improved validation, security features, and better error handling.
5
6use serde::{Deserialize, Serialize};
7
8// Import canonical types from vtcode-config instead of defining locally
9pub use vtcode_config::mcp::{McpRateLimitConfig, McpValidationConfig};
10
11use tracing::{debug, warn};
12
13/// Enhanced security configuration for MCP
14#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
15#[derive(Debug, Clone, Deserialize, Serialize)]
16pub struct EnhancedMcpSecurityConfig {
17    /// Enable authentication for MCP server
18    #[serde(default = "default_auth_enabled")]
19    auth_enabled: bool,
20
21    /// API key environment variable name
22    #[serde(default)]
23    api_key_env: Option<String>,
24
25    /// Rate limiting configuration
26    #[serde(default)]
27    rate_limit: McpRateLimitConfig,
28
29    /// Tool call validation configuration
30    #[serde(default)]
31    validation: McpValidationConfig,
32}
33
34impl Default for EnhancedMcpSecurityConfig {
35    fn default() -> Self {
36        Self {
37            auth_enabled: default_auth_enabled(),
38            api_key_env: None,
39            rate_limit: McpRateLimitConfig::default(),
40            validation: McpValidationConfig::default(),
41        }
42    }
43}
44
45// Default functions
46fn default_auth_enabled() -> bool {
47    false
48}
49
50/// Enhanced MCP client configuration with validation
51#[derive(Debug, Clone)]
52pub struct ValidatedMcpClientConfig {
53    /// Original configuration
54    original: vtcode_config::mcp::McpClientConfig,
55    /// Enhanced security configuration
56    security: EnhancedMcpSecurityConfig,
57}
58
59impl ValidatedMcpClientConfig {
60    /// Create a new validated configuration from the original
61    fn new(original: vtcode_config::mcp::McpClientConfig) -> Self {
62        let security = EnhancedMcpSecurityConfig::default();
63        Self { original, security }
64    }
65
66    /// Validate the configuration and return any issues found
67    fn validate(&self) -> Vec<ValidationError> {
68        let mut errors = Vec::new();
69
70        // Validate server configuration if enabled
71        if self.original.server.enabled {
72            // Validate port range
73            if self.original.server.port == 0 {
74                errors.push(ValidationError::InvalidPort(self.original.server.port.into()));
75            }
76
77            // Validate bind address
78            if self.original.server.bind_address.is_empty() {
79                errors.push(ValidationError::EmptyBindAddress);
80            }
81
82            // Validate security settings if auth is enabled
83            if self.security.auth_enabled && self.security.api_key_env.is_none() {
84                errors.push(ValidationError::MissingApiKeyEnv);
85            }
86        }
87
88        // Validate timeouts
89        if let Some(startup_timeout) = self.original.startup_timeout_seconds
90            && startup_timeout > 300
91        {
92            // Max 5 minutes
93            errors.push(ValidationError::InvalidStartupTimeout(startup_timeout));
94        }
95
96        if let Some(tool_timeout) = self.original.tool_timeout_seconds
97            && tool_timeout > 3600
98        {
99            // Max 1 hour
100            errors.push(ValidationError::InvalidToolTimeout(tool_timeout));
101        }
102
103        // Validate provider configurations
104        for provider in &self.original.providers {
105            if provider.name.is_empty() {
106                errors.push(ValidationError::EmptyProviderName);
107            }
108
109            // Validate max_concurrent_requests
110            if provider.max_concurrent_requests == 0 {
111                errors.push(ValidationError::InvalidMaxConcurrentRequests(
112                    provider.name.clone(),
113                    provider.max_concurrent_requests,
114                ));
115            }
116        }
117
118        errors
119    }
120
121    /// Check if the configuration is valid
122    fn is_valid(&self) -> bool {
123        self.validate().is_empty()
124    }
125
126    /// Log any validation warnings
127    pub fn log_warnings(&self) {
128        let errors = self.validate();
129        if !errors.is_empty() {
130            warn!("MCP configuration validation issues found:");
131            for error in errors {
132                warn!("  - {}", error);
133            }
134        } else {
135            debug!("MCP configuration validation passed");
136        }
137    }
138}
139
140/// Validation error types
141#[derive(Debug, Clone, thiserror::Error)]
142pub enum ValidationError {
143    #[error("Invalid server port: {0}")]
144    InvalidPort(u64),
145    #[error("Server bind address cannot be empty")]
146    EmptyBindAddress,
147    #[error("API key environment variable must be set when auth is enabled")]
148    MissingApiKeyEnv,
149    #[error("Startup timeout cannot exceed 300 seconds: {0}")]
150    InvalidStartupTimeout(u64),
151    #[error("Tool timeout cannot exceed 3600 seconds: {0}")]
152    InvalidToolTimeout(u64),
153    #[error("MCP provider name cannot be empty")]
154    EmptyProviderName,
155    #[error("Max concurrent requests must be greater than 0 for provider '{0}': {1}")]
156    InvalidMaxConcurrentRequests(String, usize),
157}
158
159/// Enhanced tool configuration
160#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
161#[derive(Debug, Clone, Deserialize, Serialize)]
162pub struct EnhancedMcpToolConfig {
163    /// Name of the tool to expose
164    name: String,
165    /// Whether the tool is enabled
166    #[serde(default = "default_tool_enabled")]
167    enabled: bool,
168    /// Optional description override
169    description: Option<String>,
170    /// Rate limiting for this specific tool
171    #[serde(default)]
172    rate_limit: Option<McpRateLimitConfig>,
173    /// Validation rules specific to this tool
174    #[serde(default)]
175    validation: Option<McpValidationConfig>,
176}
177
178fn default_tool_enabled() -> bool {
179    true
180}
181
182#[cfg(test)]
183mod tests {
184    use super::*;
185    use hashbrown::HashMap;
186    use vtcode_config::mcp::{
187        McpClientConfig, McpProviderConfig, McpServerConfig, McpStdioServerConfig, McpTransportConfig,
188    };
189
190    fn create_test_config() -> McpClientConfig {
191        McpClientConfig {
192            enabled: true,
193            ui: Default::default(),
194            providers: vec![McpProviderConfig {
195                name: "test_provider".to_owned(),
196                transport: McpTransportConfig::Stdio(McpStdioServerConfig {
197                    command: "test_command".to_owned(),
198                    args: vec![],
199                    working_directory: None,
200                }),
201                env: HashMap::new(),
202                enabled: true,
203                max_concurrent_requests: 5,
204                startup_timeout_ms: None,
205            }],
206            server: McpServerConfig {
207                enabled: true,
208                bind_address: "127.0.0.1".to_owned(),
209                port: 3000,
210                transport: vtcode_config::mcp::McpServerTransport::Sse,
211                name: "test_server".to_owned(),
212                version: "1.0.0".to_owned(),
213                exposed_tools: vec![],
214            },
215            allowlist: Default::default(),
216            requirements: Default::default(),
217            max_concurrent_connections: 10,
218            request_timeout_seconds: 30,
219            retry_attempts: 3,
220            startup_timeout_seconds: Some(60),
221            tool_timeout_seconds: Some(300),
222            experimental_use_rmcp_client: false,
223            connection_pooling_enabled: true,
224            tool_cache_capacity: 128,
225            connection_timeout_seconds: 30,
226            security: Default::default(),
227            lifecycle: Default::default(),
228        }
229    }
230
231    #[test]
232    fn test_validated_config_creation() {
233        let original = create_test_config();
234        let validated = ValidatedMcpClientConfig::new(original);
235        assert!(validated.is_valid());
236    }
237
238    #[test]
239    fn test_invalid_port_validation() {
240        let mut original = create_test_config();
241        original.server.port = 65535; // Max valid port
242        let validated = ValidatedMcpClientConfig::new(original);
243        assert!(validated.is_valid());
244    }
245
246    #[test]
247    fn test_empty_bind_address_validation() {
248        let mut original = create_test_config();
249        original.server.bind_address = String::new(); // Empty bind address
250        let validated = ValidatedMcpClientConfig::new(original);
251        assert!(!validated.is_valid());
252    }
253
254    #[test]
255    fn test_timeout_validation() {
256        let mut original = create_test_config();
257        original.startup_timeout_seconds = Some(400); // Too long
258        let validated = ValidatedMcpClientConfig::new(original);
259        assert!(!validated.is_valid());
260    }
261
262    #[test]
263    fn test_zero_concurrent_requests_validation() {
264        let mut original = create_test_config();
265        original.providers[0].max_concurrent_requests = 0; // Invalid
266        let validated = ValidatedMcpClientConfig::new(original);
267        assert!(!validated.is_valid());
268    }
269}