Skip to main content

vtcode_mcp/
connection_pool.rs

1//! MCP connection pool for efficient provider management
2//!
3//! This module provides connection pooling and parallel initialization
4//! for MCP providers to eliminate sequential connection bottlenecks.
5
6use hashbrown::HashMap;
7use std::sync::Arc;
8use std::time::Duration;
9use tokio::sync::{RwLock, Semaphore};
10use tracing::{error, info, warn};
11
12use super::{McpElicitationHandler, McpProvider, McpSandboxContext};
13use rmcp::model::{ClientCapabilities, InitializeRequestParams};
14use vtcode_commons::MultiErrors;
15use vtcode_config::mcp::{McpAllowListConfig, McpProviderConfig};
16
17/// MCP connection pool for efficient provider management
18pub(crate) struct McpConnectionPool {
19    /// Active provider connections
20    providers: Arc<RwLock<HashMap<String, Arc<McpProvider>>>>,
21    /// Connection semaphore to limit concurrent connections
22    connection_semaphore: Arc<Semaphore>,
23    /// Maximum connections allowed concurrently
24    max_concurrent_connections: usize,
25    /// Connection timeout
26    connection_timeout: Duration,
27}
28
29impl McpConnectionPool {
30    pub(crate) fn new(max_concurrent_connections: usize, connection_timeout_seconds: u64) -> Self {
31        Self {
32            providers: Arc::new(RwLock::new(HashMap::new())),
33            connection_semaphore: Arc::new(Semaphore::new(max_concurrent_connections)),
34            max_concurrent_connections,
35            connection_timeout: Duration::from_secs(connection_timeout_seconds),
36        }
37    }
38
39    /// Initialize multiple providers in parallel with controlled concurrency
40    pub(crate) async fn initialize_providers_parallel(
41        &self,
42        provider_configs: Vec<McpProviderConfig>,
43        elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
44        sandbox_context: Option<McpSandboxContext>,
45        tool_timeout: Option<Duration>,
46        allowlist_snapshot: &McpAllowListConfig,
47    ) -> Result<Vec<(String, Arc<McpProvider>)>, McpPoolError> {
48        use futures::future::join_all;
49
50        // Create initialization tasks for each provider
51        let tasks: Vec<_> = provider_configs
52            .into_iter()
53            .map(|config| {
54                let elicitation_handler = elicitation_handler.clone();
55                let sandbox_context = sandbox_context.clone();
56                let allowlist_snapshot = allowlist_snapshot.clone();
57
58                async move {
59                    self.initialize_provider(
60                        config,
61                        elicitation_handler,
62                        sandbox_context,
63                        tool_timeout.unwrap_or(Duration::from_secs(30)),
64                        allowlist_snapshot,
65                    )
66                    .await
67                }
68            })
69            .collect();
70
71        // Execute all tasks in parallel
72        let results = join_all(tasks).await;
73
74        // Collect successful connections, accumulating errors
75        let mut successful_providers = Vec::new();
76        let mut errors: MultiErrors<McpPoolError> = MultiErrors::new();
77
78        for result in results {
79            if let Some(provider) = errors.collect_result(result) {
80                successful_providers.push(provider);
81            }
82        }
83
84        if !errors.is_empty() {
85            warn!("Some MCP provider connections failed: {errors}");
86        }
87
88        Ok(successful_providers)
89    }
90
91    /// Initialize a single provider with connection pooling
92    async fn initialize_provider(
93        &self,
94        config: McpProviderConfig,
95        elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
96        sandbox_context: Option<McpSandboxContext>,
97        tool_timeout: Duration,
98        allowlist_snapshot: McpAllowListConfig,
99    ) -> Result<(String, Arc<McpProvider>), McpPoolError> {
100        // Acquire semaphore permit to limit concurrent connections
101        let _permit = self
102            .connection_semaphore
103            .acquire()
104            .await
105            .map_err(|e| McpPoolError::SemaphoreError(e.to_string()))?;
106
107        info!("Initializing MCP provider '{}'", config.name);
108
109        // Connect to provider with timeout
110        let provider = tokio::time::timeout(
111            self.connection_timeout,
112            McpProvider::connect(config.clone(), elicitation_handler, sandbox_context),
113        )
114        .await
115        .map_err(|_e| McpPoolError::ConnectionTimeout(config.name.clone()))?
116        .map_err(|e| McpPoolError::ConnectionError(config.name.clone(), e.to_string()))?;
117
118        // Initialize the provider with proper parameters
119        let provider_startup_timeout = self.resolve_startup_timeout(&config);
120        let initialize_params = build_pool_initialize_params(&provider);
121        let tool_timeout_opt = Some(tool_timeout);
122
123        if let Err(err) = provider
124            .initialize(initialize_params, provider_startup_timeout, tool_timeout_opt, &allowlist_snapshot)
125            .await
126        {
127            return Err(McpPoolError::InitializationError(config.name.clone(), err.to_string()));
128        }
129
130        // Refresh tools
131        if let Err(err) = provider
132            .cached_tools_or_refresh_shared(&allowlist_snapshot, tool_timeout_opt)
133            .await
134        {
135            warn!("Failed to fetch tools for provider '{}': {}", config.name, err);
136        }
137
138        info!("Successfully initialized MCP provider '{}'", config.name);
139
140        Ok((config.name.clone(), Arc::new(provider)))
141    }
142
143    /// Get a provider by name
144    async fn get_provider(&self, name: &str) -> Option<Arc<McpProvider>> {
145        let providers = self.providers.read().await;
146        providers.get(name).cloned()
147    }
148
149    /// Get all active providers
150    async fn get_all_providers(&self) -> Vec<Arc<McpProvider>> {
151        let providers = self.providers.read().await;
152        providers.values().cloned().collect()
153    }
154
155    /// Remove a provider from the pool
156    pub async fn remove_provider(&self, name: &str) -> Option<Arc<McpProvider>> {
157        let mut providers = self.providers.write().await;
158        providers.remove(name)
159    }
160
161    /// Check if a provider exists in the pool
162    async fn has_provider(&self, name: &str) -> bool {
163        let providers = self.providers.read().await;
164        providers.contains_key(name)
165    }
166
167    /// Get connection pool statistics
168    async fn stats(&self) -> ConnectionPoolStats {
169        let providers = self.providers.read().await;
170        let semaphore = self.connection_semaphore.available_permits();
171
172        ConnectionPoolStats {
173            active_connections: providers.len(),
174            available_permits: semaphore,
175            max_connections: self.max_concurrent_connections,
176        }
177    }
178
179    /// Shutdown all providers gracefully
180    pub async fn shutdown_all(&self) {
181        let providers: Vec<_> = {
182            let mut providers = self.providers.write().await;
183            providers.drain().collect()
184        };
185
186        for (name, provider) in providers {
187            if let Err(err) = provider.shutdown().await {
188                error!("Failed to shutdown MCP provider '{}': {}", name, err);
189            }
190        }
191    }
192
193    /// Check the health of all active providers.
194    ///
195    /// Returns a map of provider name → `true` (healthy) / `false` (unhealthy).
196    pub async fn health_check(&self) -> HashMap<String, bool> {
197        let providers: Vec<_> = {
198            let providers = self.providers.read().await;
199            providers
200                .iter()
201                .map(|(name, provider)| (name.clone(), Arc::clone(provider)))
202                .collect()
203        };
204        let mut results = HashMap::with_capacity(providers.len());
205        for (name, provider) in providers {
206            let _previous = results.insert(name, provider.is_healthy().await);
207        }
208        results
209    }
210
211    /// Attempt to reconnect any unhealthy providers.
212    ///
213    /// Returns the names of providers that were successfully reconnected.
214    pub async fn reconnect_unhealthy(
215        &self,
216        startup_timeout: Option<Duration>,
217        tool_timeout: Option<Duration>,
218        allowlist: &McpAllowListConfig,
219    ) -> Vec<String> {
220        let providers: Vec<_> = {
221            let providers = self.providers.read().await;
222            providers
223                .iter()
224                .map(|(name, provider)| (name.clone(), Arc::clone(provider)))
225                .collect()
226        };
227        let mut reconnected = Vec::new();
228        for (name, provider) in providers {
229            if !provider.is_healthy().await {
230                info!("Provider '{}' is unhealthy, attempting reconnect", name);
231                match provider.reconnect(startup_timeout, tool_timeout, allowlist).await {
232                    Ok(()) => {
233                        info!("Successfully reconnected MCP provider '{}'", name);
234                        reconnected.push(name);
235                    }
236                    Err(err) => {
237                        error!("Failed to reconnect MCP provider '{}': {}", name, err);
238                    }
239                }
240            }
241        }
242        reconnected
243    }
244
245    /// Resolve startup timeout based on provider configuration
246    fn resolve_startup_timeout(&self, config: &McpProviderConfig) -> Option<Duration> {
247        config.startup_timeout_ms.map(Duration::from_millis)
248    }
249}
250
251/// Connection pool statistics
252#[derive(Debug, Clone)]
253pub struct ConnectionPoolStats {
254    active_connections: usize,
255    available_permits: usize,
256    max_connections: usize,
257}
258
259/// Enhanced MCP manager with connection pooling
260pub(crate) struct PooledMcpManager {
261    /// Connection pool for providers
262    pool: Arc<McpConnectionPool>,
263    /// Tool discovery cache
264    tool_cache: Arc<super::tool_discovery_cache::ToolDiscoveryCache>,
265}
266
267impl PooledMcpManager {
268    fn new(max_concurrent_connections: usize, connection_timeout_seconds: u64, tool_cache_capacity: usize) -> Self {
269        Self {
270            pool: Arc::new(McpConnectionPool::new(max_concurrent_connections, connection_timeout_seconds)),
271            tool_cache: Arc::new(super::tool_discovery_cache::ToolDiscoveryCache::new(tool_cache_capacity)),
272        }
273    }
274
275    /// Initialize providers with pooling and caching
276    pub async fn initialize_providers(
277        &self,
278        provider_configs: Vec<McpProviderConfig>,
279        elicitation_handler: Option<Arc<dyn McpElicitationHandler>>,
280        sandbox_context: Option<McpSandboxContext>,
281        tool_timeout: Option<Duration>,
282        allowlist_snapshot: &McpAllowListConfig,
283    ) -> Result<Vec<(String, Arc<McpProvider>)>, McpPoolError> {
284        // Initialize providers in parallel
285        let providers = self
286            .pool
287            .initialize_providers_parallel(
288                provider_configs,
289                elicitation_handler,
290                sandbox_context,
291                tool_timeout,
292                allowlist_snapshot,
293            )
294            .await?;
295
296        // Add providers to the pool
297        let mut pool_providers = self.pool.providers.write().await;
298        for (name, provider) in &providers {
299            drop(pool_providers.insert(name.clone(), provider.clone()));
300        }
301
302        Ok(providers)
303    }
304
305    /// Execute a tool on a specific provider
306    pub async fn execute_tool(
307        &self,
308        provider_name: &str,
309        tool_name: &str,
310        arguments: serde_json::Value,
311        allowlist: &McpAllowListConfig,
312        tool_timeout: Option<Duration>,
313    ) -> Result<serde_json::Value, McpPoolError> {
314        let provider = self
315            .pool
316            .get_provider(provider_name)
317            .await
318            .ok_or_else(|| McpPoolError::ProviderNotFound(provider_name.to_string()))?;
319
320        // Convert arguments to proper format
321        let args_ref = &arguments;
322
323        // Execute the tool with correct signature
324        let result = provider
325            .call_tool(tool_name, args_ref, tool_timeout, allowlist)
326            .await
327            .map_err(|e| McpPoolError::ToolExecutionError(provider_name.to_string(), e.to_string()))?;
328
329        // Convert result to JSON value
330        Ok(serde_json::to_value(&result).unwrap_or(serde_json::Value::Null))
331    }
332
333    /// Check if a tool is read-only (safe to cache)
334    #[allow(dead_code, reason = "Intentional compatibility, platform, or test-only suppression.")]
335    fn is_read_only_tool(&self, tool_name: &str) -> bool {
336        // This is a simple heuristic - in practice, you might want to
337        // check tool metadata or maintain a list of read-only tools
338        matches!(
339            tool_name,
340            "read_file"
341                | "list_directory"
342                | "search_files"
343                | "get_file_info"
344                | "read_environment"
345                | "get_system_info"
346                | "search_code"
347                | "analyze_code"
348        )
349    }
350
351    /// Get pool statistics
352    async fn stats(&self) -> PooledMcpStats {
353        let pool_stats = self.pool.stats().await;
354        let tool_cache_stats = self.tool_cache.stats();
355
356        PooledMcpStats {
357            connection_pool: pool_stats,
358            tool_cache: tool_cache_stats,
359        }
360    }
361
362    /// Shutdown all providers gracefully
363    pub async fn shutdown(&self) {
364        self.pool.shutdown_all().await;
365    }
366}
367
368/// Pooled MCP manager statistics
369#[derive(Debug, Clone)]
370pub struct PooledMcpStats {
371    connection_pool: ConnectionPoolStats,
372    tool_cache: super::tool_discovery_cache::ToolCacheStats,
373}
374
375/// Build initialize params for an MCP provider
376fn build_pool_initialize_params(_provider: &McpProvider) -> InitializeRequestParams {
377    InitializeRequestParams::new(ClientCapabilities::default(), super::utils::build_client_implementation())
378        .with_protocol_version(super::rmcp_client::latest_protocol_version())
379}
380
381/// MCP connection pool errors
382#[derive(Debug, thiserror::Error)]
383pub enum McpPoolError {
384    #[error("Connection timeout for provider '{0}'")]
385    ConnectionTimeout(String),
386
387    #[error("Connection error for provider '{0}': {1}")]
388    ConnectionError(String, String),
389
390    #[error("Initialization timeout for provider '{0}'")]
391    InitializationTimeout(String),
392
393    #[error("Initialization error for provider '{0}': {1}")]
394    InitializationError(String, String),
395
396    #[error("Provider not found: {0}")]
397    ProviderNotFound(String),
398
399    #[error("Tool execution error for provider '{0}': {1}")]
400    ToolExecutionError(String, String),
401
402    #[error("Semaphore error: {0}")]
403    SemaphoreError(String),
404}
405
406#[cfg(test)]
407pub(crate) mod tests {
408    use super::*;
409
410    #[tokio::test]
411    async fn test_connection_pool_creation() {
412        let pool = McpConnectionPool::new(5, 30);
413        let stats = pool.stats().await;
414
415        assert_eq!(stats.active_connections, 0);
416        assert_eq!(stats.max_connections, 5);
417        assert_eq!(stats.available_permits, 5);
418    }
419
420    #[tokio::test]
421    async fn test_connection_pool_semaphore_limits() {
422        let pool = McpConnectionPool::new(3, 30);
423
424        // Acquire 3 permits
425        let permit1 = pool.connection_semaphore.acquire().await.unwrap();
426        let _permit2 = pool.connection_semaphore.acquire().await.unwrap();
427        let _permit3 = pool.connection_semaphore.acquire().await.unwrap();
428
429        let stats = pool.stats().await;
430        assert_eq!(stats.available_permits, 0);
431
432        // Try to acquire another (would block if not in test)
433        drop(permit1);
434        let _permit4 = pool.connection_semaphore.acquire().await.unwrap();
435
436        let stats = pool.stats().await;
437        assert_eq!(stats.available_permits, 0);
438    }
439
440    #[tokio::test]
441    async fn test_pooled_manager_creation() {
442        let manager = PooledMcpManager::new(10, 30, 100);
443        let stats = manager.stats().await;
444
445        assert_eq!(stats.connection_pool.max_connections, 10);
446        assert_eq!(stats.connection_pool.active_connections, 0);
447    }
448
449    #[tokio::test]
450    async fn test_read_only_tool_detection() {
451        let manager = PooledMcpManager::new(5, 30, 50);
452
453        assert!(manager.is_read_only_tool("read_file"));
454        assert!(manager.is_read_only_tool("search_files"));
455        assert!(manager.is_read_only_tool("get_system_info"));
456        assert!(manager.is_read_only_tool("get_file_info"));
457
458        assert!(!manager.is_read_only_tool("write_file"));
459        assert!(!manager.is_read_only_tool("edit_file"));
460        assert!(!manager.is_read_only_tool("execute_command"));
461        assert!(!manager.is_read_only_tool("delete_file"));
462    }
463
464    #[test]
465    fn test_connection_pool_error_display() {
466        let error = McpPoolError::ConnectionTimeout("test_provider".to_string());
467        assert!(error.to_string().contains("test_provider"));
468
469        let error = McpPoolError::InitializationError("auth".to_string(), "invalid credentials".to_string());
470        assert!(error.to_string().contains("auth"));
471        assert!(error.to_string().contains("invalid credentials"));
472    }
473
474    #[tokio::test]
475    async fn test_pool_provider_not_found() {
476        let pool = McpConnectionPool::new(5, 30);
477        let provider = pool.get_provider("nonexistent").await;
478        assert!(provider.is_none());
479    }
480
481    #[tokio::test]
482    async fn test_pool_has_provider() {
483        let pool = McpConnectionPool::new(5, 30);
484        assert!(!pool.has_provider("test").await);
485    }
486
487    #[tokio::test]
488    async fn test_pool_get_all_providers_empty() {
489        let pool = McpConnectionPool::new(5, 30);
490        let providers = pool.get_all_providers().await;
491        assert_eq!(providers.len(), 0);
492    }
493
494    #[tokio::test]
495    async fn test_pool_stats() {
496        let pool = McpConnectionPool::new(7, 60);
497        let stats = pool.stats().await;
498
499        assert_eq!(stats.max_connections, 7);
500        assert_eq!(stats.available_permits, 7);
501        assert_eq!(stats.active_connections, 0);
502    }
503}