Skip to main content

vtcode_mcp/
lib.rs

1#![allow(
2    clippy::large_futures,
3    missing_docs,
4    dead_code,
5    unused_imports,
6    reason = "Intentional compatibility, platform, or test-only suppression."
7)]
8//! Model Context Protocol (MCP) client management.
9//!
10//! This crate adapts reference MCP client, server, and type definitions
11//! from [openai/codex] to integrate them with VT Code's multi-provider
12//! configuration model (Apache-2.0). Copyright 2025 OpenAI. See the
13//! repository `THIRD-PARTY-NOTICES` file for the full license text and
14//! attribution.
15//!
16//! The VT Code-specific surface (allow lists, tool indexing, status
17//! reporting) is original; only the protocol interaction layer mirrors
18//! Codex' `mcp-client` crate.
19//!
20//! [openai/codex]: https://github.com/openai/codex
21
22mod client;
23pub mod connection_pool;
24pub(crate) mod conversion;
25pub mod enhanced_config;
26pub mod errors;
27mod provider;
28mod rmcp_client;
29pub mod rmcp_transport;
30mod sandbox_context;
31pub(crate) mod schema;
32pub mod tool_discovery;
33pub mod tool_discovery_cache;
34pub mod traits;
35pub mod trust;
36pub mod types;
37pub(crate) mod utils;
38
39pub use client::McpClient;
40
41pub(crate) use connection_pool::{
42    ConnectionPoolStats, McpConnectionPool, McpPoolError, PooledMcpManager, PooledMcpStats,
43};
44pub use errors::{
45    ErrorCode, McpResult, configuration_error, initialization_timeout, provider_not_found, provider_unavailable,
46    schema_invalid, tool_invocation_failed, tool_not_found,
47};
48pub use provider::McpProvider;
49pub(crate) use rmcp_client::RmcpClient;
50pub use rmcp_transport::{
51    HttpTransport, create_http_transport, create_stdio_transport, create_stdio_transport_with_stderr,
52};
53pub use sandbox_context::McpSandboxContext;
54pub(crate) use schema::{validate_against_schema, validate_tool_input};
55pub use tool_discovery::{DetailLevel, ToolDiscovery, ToolDiscoveryResult};
56pub use traits::{McpElicitationHandler, McpToolExecutor};
57pub use trust::{MAX_UNTRUSTED_MCP_DESCRIPTION_BYTES, render_untrusted_mcp_description};
58pub use types::{
59    FileParamSchemaEntry, FileUploadResult, McpClientStatus, McpElicitationRequest, McpElicitationResponse,
60    McpPromptDetail, McpPromptInfo, McpResourceData, McpResourceInfo, McpToolInfo, OPENAI_FILE_PARAMS_META_KEY,
61    OPENAI_FILE_PARAMS_VALUE, ProvidedFilePayload,
62};
63pub(crate) use utils::{
64    LOCAL_TIMEZONE_ENV_VAR, TIMEZONE_ARGUMENT, TZ_ENV_VAR, build_headers, detect_local_timezone,
65    ensure_timezone_argument, schema_requires_field,
66};
67
68use anyhow::{Result, anyhow};
69use hashbrown::HashMap;
70pub use rmcp::model::ElicitationAction;
71use std::ffi::OsString;
72use std::fmt::Write;
73
74/// MCP protocol version constants
75pub(crate) const LATEST_PROTOCOL_VERSION: &str = "2026-07-28";
76pub(crate) const SUPPORTED_PROTOCOL_VERSIONS: &[&str] =
77    &["2026-07-28", "2025-11-25", "2025-06-18", "2025-03-26", "2024-11-05"];
78
79/// Convert any serializable type to rmcp model type via JSON serialization
80pub(crate) fn convert_to_rmcp<T, U>(value: T) -> Result<U>
81where
82    T: serde::Serialize,
83    U: serde::de::DeserializeOwned,
84{
85    let json = serde_json::to_value(value)?;
86    serde_json::from_value(json).map_err(|err| anyhow!(err))
87}
88
89pub(crate) fn create_env_for_mcp_server(extra_env: Option<HashMap<OsString, OsString>>) -> HashMap<OsString, OsString> {
90    DEFAULT_ENV_VARS
91        .iter()
92        .filter_map(|var| std::env::var_os(var).map(|value| (OsString::from(*var), value)))
93        .chain(extra_env.unwrap_or_default())
94        .collect()
95}
96
97/// Validate MCP configuration settings
98pub fn validate_mcp_config(config: &vtcode_config::mcp::McpClientConfig) -> Result<()> {
99    // Validate server configuration if enabled
100    if config.server.enabled {
101        // Validate port range
102        if config.server.port == 0 {
103            return Err(anyhow::anyhow!("Invalid server port: {}", config.server.port));
104        }
105
106        // Validate bind address
107        if config.server.bind_address.is_empty() {
108            return Err(anyhow::anyhow!("Server bind address cannot be empty"));
109        }
110
111        // Validate security settings if auth is enabled
112        if config.security.auth_enabled && config.security.api_key_env.is_none() {
113            return Err(anyhow::anyhow!("API key environment variable must be set when auth is enabled"));
114        }
115    }
116
117    // Validate timeouts
118    if let Some(startup_timeout) = config.startup_timeout_seconds
119        && startup_timeout > 300
120    {
121        // Max 5 minutes
122        return Err(anyhow::anyhow!("Startup timeout cannot exceed 300 seconds"));
123    }
124
125    if let Some(tool_timeout) = config.tool_timeout_seconds
126        && tool_timeout > 3600
127    {
128        // Max 1 hour
129        return Err(anyhow::anyhow!("Tool timeout cannot exceed 3600 seconds"));
130    }
131
132    // Validate provider configurations
133    for provider in &config.providers {
134        if provider.name.is_empty() {
135            return Err(anyhow::anyhow!("MCP provider name cannot be empty"));
136        }
137
138        // Validate max_concurrent_requests
139        if provider.max_concurrent_requests == 0 {
140            return Err(anyhow::anyhow!(
141                "Max concurrent requests must be greater than 0 for provider '{}'",
142                provider.name
143            ));
144        }
145    }
146
147    Ok(())
148}
149
150#[cfg(unix)]
151const DEFAULT_ENV_VARS: &[&str] = &[
152    "HOME",
153    "LOGNAME",
154    "PATH",
155    "SHELL",
156    "USER",
157    "__CF_USER_TEXT_ENCODING",
158    "LANG",
159    "LC_ALL",
160    "TERM",
161    "TMPDIR",
162    "TZ",
163];
164
165#[cfg(windows)]
166const DEFAULT_ENV_VARS: &[&str] = &[
167    // Core path resolution
168    "PATH",
169    "PATHEXT",
170    // Shell and system roots
171    "COMSPEC",
172    "SYSTEMROOT",
173    "SYSTEMDRIVE",
174    // User context and profiles
175    "USERNAME",
176    "USERDOMAIN",
177    "USERPROFILE",
178    "HOMEDRIVE",
179    "HOMEPATH",
180    // Program locations
181    "PROGRAMFILES",
182    "PROGRAMFILES(X86)",
183    "PROGRAMW6432",
184    "PROGRAMDATA",
185    // App data and caches
186    "LOCALAPPDATA",
187    "APPDATA",
188    // Temp locations
189    "TEMP",
190    "TMP",
191    // Common shells/pwsh hints
192    "POWERSHELL",
193    "PWSH",
194];
195
196// Helper functions for file-based tool discovery
197
198/// Sanitize a string for use in a filename
199pub(crate) fn sanitize_filename(name: &str) -> String {
200    name.chars()
201        .map(|c| {
202            if c.is_alphanumeric() || c == '_' || c == '-' {
203                c
204            } else {
205                '_'
206            }
207        })
208        .collect()
209}
210
211/// Format a tool description as Markdown
212pub(crate) fn format_tool_markdown(tool: &McpToolInfo) -> String {
213    let mut content = String::new();
214    let _format_result = write!(content, "# {}\n\n", tool.name);
215    let _format_result = write!(content, "**Provider**: {}\n\n", tool.provider);
216    content.push_str("## Description\n\n");
217    content.push_str(&tool.description);
218    content.push_str("\n\n");
219
220    content.push_str("## Input Schema\n\n");
221    content.push_str("```json\n");
222    content
223        .push_str(&serde_json::to_string_pretty(&tool.input_schema).unwrap_or_else(|_| tool.input_schema.to_string()));
224    content.push_str("\n```\n\n");
225
226    if let Some(output_schema) = tool.output_schema.as_ref() {
227        content.push_str("## Output Schema\n\n");
228        content.push_str("```json\n");
229        content.push_str(&serde_json::to_string_pretty(output_schema).unwrap_or_else(|_| output_schema.to_string()));
230        content.push_str("\n```\n\n");
231    }
232
233    // Extract required fields if present
234    if let Some(obj) = tool.input_schema.as_object() {
235        if let Some(required) = obj.get("required").and_then(|v| v.as_array())
236            && !required.is_empty()
237        {
238            content.push_str("## Required Parameters\n\n");
239            for req in required {
240                if let Some(name) = req.as_str() {
241                    let _format_result = writeln!(content, "- `{name}`");
242                }
243            }
244            content.push('\n');
245        }
246
247        // Extract properties descriptions
248        if let Some(props) = obj.get("properties").and_then(|v| v.as_object())
249            && !props.is_empty()
250        {
251            content.push_str("## Parameters\n\n");
252            for (param_name, param_schema) in props {
253                let param_type = param_schema.get("type").and_then(|t| t.as_str()).unwrap_or("any");
254                let param_desc = param_schema.get("description").and_then(|d| d.as_str()).unwrap_or("");
255                let _format_result = write!(content, "### `{param_name}`\n\n");
256                let _format_result = writeln!(content, "- **Type**: {param_type}");
257                if !param_desc.is_empty() {
258                    let _format_result = writeln!(content, "- **Description**: {param_desc}");
259                }
260                content.push('\n');
261            }
262        }
263    }
264
265    content.push_str("---\n");
266    content.push_str("*Generated automatically for dynamic context discovery.*\n");
267
268    content
269}
270
271#[cfg(test)]
272mod tests {
273    use super::*;
274    use crate::utils::{LOCAL_TIMEZONE_ENV_VAR, TIMEZONE_ARGUMENT, clear_test_env_override, set_test_env_override};
275    use hashbrown::HashMap;
276    use rmcp::model::ElicitationAction;
277    use serde_json::{Map, Value, json};
278
279    #[cfg(unix)]
280    use serial_test::serial;
281
282    #[cfg(unix)]
283    use std::os::unix::ffi::OsStringExt;
284
285    struct EnvGuard {
286        key: &'static str,
287    }
288
289    impl EnvGuard {
290        fn set(key: &'static str, value: &str) -> Self {
291            set_test_env_override(key, Some(value));
292            Self { key }
293        }
294    }
295
296    impl Drop for EnvGuard {
297        fn drop(&mut self) {
298            clear_test_env_override(self.key);
299        }
300    }
301
302    #[test]
303    fn schema_detection_handles_required_entries() {
304        let schema = json!({
305            "type": "object",
306            "required": [TIMEZONE_ARGUMENT],
307            "properties": {
308                TIMEZONE_ARGUMENT: { "type": "string" }
309            }
310        });
311
312        assert!(schema_requires_field(&schema, TIMEZONE_ARGUMENT));
313        assert!(!schema_requires_field(&schema, "location"));
314    }
315
316    #[test]
317    fn ensure_timezone_injects_from_override_env() {
318        let _guard = EnvGuard::set(LOCAL_TIMEZONE_ENV_VAR, "Etc/UTC");
319        let mut arguments = Map::new();
320
321        ensure_timezone_argument(&mut arguments, true).unwrap();
322
323        assert_eq!(arguments.get(TIMEZONE_ARGUMENT).and_then(Value::as_str), Some("Etc/UTC"));
324    }
325
326    #[test]
327    fn ensure_timezone_does_not_override_existing_value() {
328        let mut arguments = Map::new();
329        drop(arguments.insert(TIMEZONE_ARGUMENT.to_string(), Value::String("America/New_York".to_owned())));
330
331        ensure_timezone_argument(&mut arguments, true).unwrap();
332
333        assert_eq!(arguments.get(TIMEZONE_ARGUMENT).and_then(Value::as_str), Some("America/New_York"));
334    }
335
336    #[test]
337    fn create_env_merges_configured_values() {
338        let mut extra_env = HashMap::new();
339        drop(extra_env.insert(OsString::from("A"), OsString::from("1")));
340        drop(extra_env.insert(OsString::from("B"), OsString::from("2")));
341
342        let env = create_env_for_mcp_server(Some(extra_env));
343
344        assert_eq!(env.get(&OsString::from("A")), Some(&OsString::from("1")));
345        assert_eq!(env.get(&OsString::from("B")), Some(&OsString::from("2")));
346    }
347
348    #[test]
349    #[cfg(unix)]
350    #[serial]
351    fn create_env_preserves_non_utf8_path() {
352        let env_guard = vtcode_commons::env_lock::lock();
353        let original_path = std::env::var_os("PATH");
354        let non_utf8_path = OsString::from_vec(b"/tmp/alpha:\xFFbeta".to_vec());
355
356        env_guard.set_var("PATH", &non_utf8_path);
357
358        let env = create_env_for_mcp_server(None);
359
360        env_guard.restore_var("PATH", original_path);
361
362        assert_eq!(env.get(&OsString::from("PATH")), Some(&non_utf8_path));
363    }
364
365    #[tokio::test]
366    async fn convert_to_rmcp_round_trip() {
367        use rmcp::model::{ClientCapabilities, Implementation, InitializeRequestParams, RootsCapabilities};
368
369        let mut capabilities = ClientCapabilities::default();
370        let mut roots = RootsCapabilities::default();
371        roots.list_changed = Some(true);
372        capabilities.roots = Some(roots);
373        let params = InitializeRequestParams::new(capabilities, Implementation::new("vtcode", "1.0"))
374            .with_protocol_version(rmcp::model::ProtocolVersion::V_2026_07_28);
375
376        let converted: InitializeRequestParams = convert_to_rmcp(params.clone()).unwrap();
377        // Verify the conversion succeeded by checking the name
378        assert_eq!(converted.client_info.name, "vtcode");
379        assert_eq!(converted.client_info.version, "1.0");
380    }
381
382    #[test]
383    fn supported_protocol_versions_cover_known_rmcp_versions_newest_first() {
384        let known: Vec<String> = rmcp::model::ProtocolVersion::KNOWN_VERSIONS
385            .iter()
386            .map(|version| version.to_string())
387            .collect();
388        let mut known_newest_first = known.clone();
389        known_newest_first.reverse();
390
391        assert_eq!(
392            SUPPORTED_PROTOCOL_VERSIONS,
393            known_newest_first.as_slice(),
394            "supported MCP versions must track rmcp's known versions, newest first"
395        );
396    }
397
398    #[test]
399    fn validate_elicitation_payload_rejects_invalid_content() {
400        use crate::rmcp_client::{build_elicitation_validator, validate_elicitation_payload};
401
402        let schema = json!({
403            "type": "object",
404            "properties": {
405                "name": { "type": "string" }
406            },
407            "required": ["name"]
408        });
409        let validator = build_elicitation_validator("test", &schema).expect("schema should compile");
410
411        let result = validate_elicitation_payload(
412            "test",
413            Some(&validator),
414            &ElicitationAction::Accept,
415            Some(&json!({ "name": 42 })),
416        );
417
418        assert!(result.is_err());
419    }
420
421    #[test]
422    fn validate_elicitation_payload_accepts_valid_content() {
423        use crate::rmcp_client::{build_elicitation_validator, validate_elicitation_payload};
424
425        let schema = json!({
426            "type": "object",
427            "properties": {
428                "email": { "type": "string", "format": "email" }
429            },
430            "required": ["email"]
431        });
432        let validator = build_elicitation_validator("test", &schema).expect("schema should compile");
433
434        let result = validate_elicitation_payload(
435            "test",
436            Some(&validator),
437            &ElicitationAction::Accept,
438            Some(&json!({ "email": "user@example.com" })),
439        );
440
441        result.unwrap();
442    }
443
444    #[tokio::test]
445    async fn provider_max_concurrency_defaults_to_one() {
446        use crate::provider::McpProvider;
447        use vtcode_config::mcp::{McpProviderConfig, McpStdioServerConfig, McpTransportConfig};
448
449        let config = McpProviderConfig {
450            name: "test".into(),
451            transport: McpTransportConfig::Stdio(McpStdioServerConfig {
452                command: "cat".into(),
453                args: vec![],
454                working_directory: None,
455            }),
456            env: HashMap::new(),
457            enabled: true,
458            max_concurrent_requests: 0,
459            startup_timeout_ms: None,
460        };
461
462        let provider = McpProvider::connect(config, None, None).await.unwrap();
463        assert_eq!(provider.semaphore.available_permits(), 1);
464    }
465
466    #[test]
467    fn directory_to_file_uri_generates_file_scheme() {
468        use crate::rmcp_client::directory_to_file_uri;
469
470        let temp_dir = std::env::temp_dir();
471        let uri = directory_to_file_uri(temp_dir.as_path()).expect("should create file uri for temp directory");
472        assert!(uri.starts_with("file://"));
473    }
474}