1use 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
17pub(crate) struct McpConnectionPool {
19 providers: Arc<RwLock<HashMap<String, Arc<McpProvider>>>>,
21 connection_semaphore: Arc<Semaphore>,
23 max_concurrent_connections: usize,
25 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 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 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 let results = join_all(tasks).await;
73
74 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 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 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 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 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 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 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 async fn get_all_providers(&self) -> Vec<Arc<McpProvider>> {
151 let providers = self.providers.read().await;
152 providers.values().cloned().collect()
153 }
154
155 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 async fn has_provider(&self, name: &str) -> bool {
163 let providers = self.providers.read().await;
164 providers.contains_key(name)
165 }
166
167 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 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 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 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 fn resolve_startup_timeout(&self, config: &McpProviderConfig) -> Option<Duration> {
247 config.startup_timeout_ms.map(Duration::from_millis)
248 }
249}
250
251#[derive(Debug, Clone)]
253pub struct ConnectionPoolStats {
254 active_connections: usize,
255 available_permits: usize,
256 max_connections: usize,
257}
258
259pub(crate) struct PooledMcpManager {
261 pool: Arc<McpConnectionPool>,
263 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 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 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 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 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 let args_ref = &arguments;
322
323 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 Ok(serde_json::to_value(&result).unwrap_or(serde_json::Value::Null))
331 }
332
333 #[allow(dead_code, reason = "Intentional compatibility, platform, or test-only suppression.")]
335 fn is_read_only_tool(&self, tool_name: &str) -> bool {
336 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 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 pub async fn shutdown(&self) {
364 self.pool.shutdown_all().await;
365 }
366}
367
368#[derive(Debug, Clone)]
370pub struct PooledMcpStats {
371 connection_pool: ConnectionPoolStats,
372 tool_cache: super::tool_discovery_cache::ToolCacheStats,
373}
374
375fn 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#[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 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 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}