diff --git a/crates/tui/src/core/engine.rs b/crates/tui/src/core/engine.rs index cece3975a4..9513da816e 100644 --- a/crates/tui/src/core/engine.rs +++ b/crates/tui/src/core/engine.rs @@ -6329,6 +6329,7 @@ impl Engine { } self.mcp_event_generation = generation; self.replace_mcp_boot_errors(&authority_errors, connection_errors); + self.session.pending_prefix_change_reason = Some("mcp-session-boot".to_string()); if let Ok(snapshot) = self.mcp_session_snapshot().await { let _ = self.tx_event.try_send(Event::McpSessionBoot { generation, @@ -6382,6 +6383,8 @@ impl Engine { } self.mcp_event_generation = generation; self.replace_mcp_boot_errors(&authority_errors, connection_errors); + self.session.pending_prefix_change_reason = + Some("mcp-session-boot".to_string()); } McpBootUpdate::Finished { generation, @@ -6484,33 +6487,54 @@ impl Engine { async move { let mut remaining: Vec = pending.iter().map(|(name, _)| name.clone()).collect(); - let results = McpPool::connect_pending_concurrently( + let mut connects = McpPool::spawn_pending_connects( pending, timeouts, network_policy, catalog_generation, - ) - .await; + ); let mut connection_errors = HashMap::new(); - { - let mut pool = pool_for_task.lock().await; - for (name, result) in results { - remaining.retain(|pending_name| pending_name != &name); - match result { - Ok(connection) => pool.store_ready_connection(name, connection), - Err(error) => { - pool.note_connect_failure(&name, &error); - connection_errors - .insert(name, crate::mcp::format_mcp_error_for_display(&error)); + while let Some(joined) = connects.join_next().await { + let (name, result) = joined + .unwrap_or_else(|error| ("connection task".to_string(), Err(error.into()))); + remaining.retain(|pending_name| pending_name != &name); + { + let mut pool = pool_for_task.lock().await; + // A turn may have reloaded the pool while these handshakes + // were in flight. Never let their old authority or failures + // overwrite the newly installed configuration. + let reload = pool.reload_if_config_changed().await; + if reload.is_err() + || pool.current_catalog_generation() != catalog_generation + { + connects.abort_all(); + connection_errors.clear(); + if let Err(error) = reload { + connection_errors.insert( + "configuration".to_string(), + crate::mcp::format_mcp_error_for_display(&error), + ); } + break; } - let _ = progress_tx.send(McpBootUpdate::Progress { - generation, - authority_errors: Arc::clone(&authority_errors), - connection_errors: connection_errors.clone(), - connecting: remaining.clone(), + let result = result.and_then(|connection| { + pool.store_ready_connection(name.clone(), connection) }); + if let Err(error) = result { + pool.note_connect_failure(&name, &error); + connection_errors + .insert(name, crate::mcp::format_mcp_error_for_display(&error)); + } } + let _ = progress_tx.send(McpBootUpdate::Progress { + generation, + authority_errors: Arc::clone(&authority_errors), + connection_errors: connection_errors.clone(), + connecting: remaining.clone(), + }); + } + { + let pool = pool_for_task.lock().await; let mut required = Vec::new(); pool.push_required_server_errors(&mut required); for (name, error) in required { @@ -6609,8 +6633,9 @@ impl Engine { if self.mcp_boot_in_flight { // Optional servers are still connecting in the background. Snapshot // currently-ready tools so the first LLM call is not serialized - // behind the slowest handshake. The catalog refreshes on a later - // turn once boot settles (KV-cache prefix re-pin: mcp-session-boot). + // behind the slowest handshake. Declare the refresh here as well + // as on Progress: a ready connection can precede its mailbox update. + self.session.pending_prefix_change_reason = Some("mcp-session-boot".to_string()); return pool.lock().await.to_api_tools(); } diff --git a/crates/tui/src/core/engine/tests.rs b/crates/tui/src/core/engine/tests.rs index d428feaa83..2241aa44cc 100644 --- a/crates/tui/src/core/engine/tests.rs +++ b/crates/tui/src/core/engine/tests.rs @@ -20754,6 +20754,148 @@ async fn reload_mcp_op_recovers_from_invalid_initial_config_in_process() { task.await.expect("engine task"); } +#[tokio::test] +async fn mcp_boot_reports_ready_server_before_stalled_server_finishes() { + assert_incremental_mcp_boot(false).await; +} + +#[tokio::test] +async fn mcp_boot_does_not_restore_servers_removed_during_handshake() { + assert_incremental_mcp_boot(true).await; +} + +async fn assert_incremental_mcp_boot(invalidate_config: bool) { + if std::process::Command::new("node") + .arg("--version") + .output() + .is_err() + { + tracing::warn!("skipping MCP stdio fixture because node is unavailable"); + return; + } + let tmp = tempdir().expect("tempdir"); + let server = tmp.path().join("server.mjs"); + let release = tmp.path().join("release-slow"); + std::fs::write( + &server, + r#"import fs from 'node:fs'; +import readline from 'node:readline'; +const lines = readline.createInterface({ input: process.stdin }); +lines.on('line', async (line) => { + const request = JSON.parse(line); + if (request.id === undefined) return; + if (process.argv[2] === 'slow' && request.method === 'initialize') { + while (!fs.existsSync(process.argv[3])) { + await new Promise(resolve => setTimeout(resolve, 10)); + } + } + const result = request.method === 'initialize' + ? { protocolVersion: '2024-11-05', capabilities: { tools: {} }, + serverInfo: { name: process.argv[2], version: '1' } } + : { tools: [{ name: 'ready', inputSchema: { type: 'object' } }] }; + process.stdout.write(JSON.stringify({ jsonrpc: '2.0', id: request.id, result }) + '\n'); +}); +"#, + ) + .expect("server fixture"); + let config_path = tmp.path().join("mcp.json"); + std::fs::write( + &config_path, + serde_json::to_vec(&serde_json::json!({ + "timeouts": { "connect_timeout": 30 }, + "servers": { + "fast": { "command": "node", "args": [server, "fast", release] }, + "slow": { "command": "node", "args": [server, "slow", release] } + } + })) + .expect("config JSON"), + ) + .expect("MCP config"); + let (mut engine, handle) = Engine::new( + EngineConfig { + workspace: tmp.path().to_path_buf(), + mcp_config_path: config_path.clone(), + ..Default::default() + }, + &Config::default(), + ); + let pool = engine.ensure_mcp_pool().await.expect("engine pool"); + let task = tokio::spawn(async move { engine.run().await }); + let mut events = handle.rx_event.write().await; + let progress = tokio::time::timeout(Duration::from_secs(10), async { + while let Some(event) = events.recv().await { + if let Event::McpSessionBoot { + snapshot, + connecting, + finished: false, + .. + } = event + && connecting == ["slow"] + { + return snapshot; + } + } + panic!("engine event channel closed"); + }) + .await; + let ready_tools = pool.lock().await.to_api_tools(); + if invalidate_config { + std::fs::write( + &config_path, + r#"{"servers":{"slow":{"command":"node","disabled":true}}}"#, + ) + .expect("remove servers"); + pool.lock() + .await + .reload_if_config_changed() + .await + .expect("reload config"); + } + // Release and shut down even when testing the old batch-buffered behavior. + std::fs::write(&release, "continue").expect("release stalled fixture"); + let finished = tokio::time::timeout(Duration::from_secs(10), async { + while let Some(event) = events.recv().await { + if let Event::McpSessionBoot { + snapshot, + finished: true, + .. + } = event + { + return snapshot; + } + } + panic!("engine event channel closed"); + }) + .await; + drop(events); + handle.send(Op::Shutdown).await.expect("shutdown"); + task.await.expect("engine task"); + let progress = progress.expect("fast server must be visible before slow server is released"); + assert!( + progress + .servers + .iter() + .any(|row| row.name == "fast" && row.connected) + ); + assert!( + progress + .servers + .iter() + .any(|row| row.name == "slow" && !row.connected) + ); + assert!(ready_tools.iter().any(|tool| tool.name == "mcp_fast_ready")); + assert!(!ready_tools.iter().any(|tool| tool.name == "mcp_slow_ready")); + let finished = finished.expect("finished boot"); + if invalidate_config { + assert_eq!(finished.servers.len(), 1); + assert!(!finished.servers[0].enabled); + assert!(!finished.servers[0].connected); + assert!(pool.lock().await.to_api_tools().is_empty()); + } else { + assert!(finished.servers.iter().all(|row| row.connected)); + } +} + #[tokio::test] async fn mcp_boot_updates_preserve_authority_errors_and_replace_ordinary_errors() { let tmp = tempdir().expect("tempdir"); @@ -20799,6 +20941,10 @@ async fn mcp_boot_updates_preserve_authority_errors_and_replace_ordinary_errors( ]) ); assert!(!engine.mcp_connection_errors.contains_key("stale-transport")); + assert_eq!( + engine.session.pending_prefix_change_reason.as_deref(), + Some("mcp-session-boot") + ); engine.mcp_connection_errors.insert( "stale-between-updates".to_string(), @@ -20829,6 +20975,42 @@ async fn mcp_boot_updates_preserve_authority_errors_and_replace_ordinary_errors( ); } +#[tokio::test] +async fn mcp_boot_catalog_refresh_declares_prefix_before_mailbox_delivery() { + let tmp = tempdir().expect("tempdir"); + let (mut engine, _handle) = Engine::new( + EngineConfig { + workspace: tmp.path().to_path_buf(), + ..Default::default() + }, + &Config::default(), + ); + engine.mcp_boot_generation = Some(1); + engine.mcp_boot_in_flight = true; + engine.session.pending_prefix_change_reason = None; + let _tools = engine.mcp_tools().await; + assert_eq!( + engine.session.pending_prefix_change_reason.as_deref(), + Some("mcp-session-boot") + ); + + engine.session.pending_prefix_change_reason = None; + let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); + engine.mcp_boot_rx = Some(rx); + tx.send(McpBootUpdate::Progress { + generation: 1, + authority_errors: Arc::new(HashMap::new()), + connection_errors: HashMap::new(), + connecting: vec!["slow".to_string()], + }) + .expect("queue progress"); + engine.drain_mcp_boot_updates().await; + assert_eq!( + engine.session.pending_prefix_change_reason.as_deref(), + Some("mcp-session-boot") + ); +} + #[tokio::test] async fn stale_boot_finished_does_not_clear_a_newer_receiver() { let tmp = tempdir().expect("tempdir"); diff --git a/crates/tui/src/mcp.rs b/crates/tui/src/mcp.rs index fbe356b722..63850d5fcb 100644 --- a/crates/tui/src/mcp.rs +++ b/crates/tui/src/mcp.rs @@ -16,6 +16,7 @@ use std::sync::atomic::{AtomicU64, Ordering}; use std::time::Duration; use anyhow::{Context, Result}; +use futures_util::FutureExt; use parking_lot::RwLock; use serde::{Deserialize, Serialize}; use sha2::Digest as _; @@ -2864,7 +2865,7 @@ impl McpPool { }; connection.catalog_generation = self.catalog_generation.load(Ordering::SeqCst); - self.store_ready_connection(server_name.to_string(), connection); + self.store_ready_connection(server_name.to_string(), connection)?; self.connections .get_mut(server_name) .ok_or_else(|| anyhow::anyhow!("Failed to store MCP connection for {server_name}")) @@ -2913,7 +2914,7 @@ impl McpPool { anyhow::bail!("Failed to connect MCP server '{server_name}': server is disabled"); } - let connection = match McpConnection::connect_with_policy( + let mut connection = match McpConnection::connect_with_policy( server_name.to_string(), server_config, &self.config.timeouts, @@ -2927,14 +2928,25 @@ impl McpPool { return Err(error); } }; - self.store_ready_connection(server_name.to_string(), connection); + connection.catalog_generation = self.current_catalog_generation(); + self.store_ready_connection(server_name.to_string(), connection)?; self.connections .get_mut(server_name) .ok_or_else(|| anyhow::anyhow!("Failed to store MCP connection for {server_name}")) } - pub(crate) fn store_ready_connection(&mut self, name: String, mut connection: McpConnection) { - connection.catalog_generation = self.catalog_generation.load(Ordering::SeqCst); + pub(crate) fn store_ready_connection( + &mut self, + name: String, + connection: McpConnection, + ) -> Result<()> { + anyhow::ensure!( + connection.catalog_generation == self.current_catalog_generation(), + "MCP configuration changed while connecting {name}; retry against the current config" + ); + if let Some(source) = connection.config().reviewed_plugin.as_ref() { + source.validate_before_use(&name, "use")?; + } // A successful connect settles the auth question for this server, // and the cooldown with it. self.connect_backoff.remove(&name); @@ -2942,6 +2954,7 @@ impl McpPool { self.needs_auth_generation = self.needs_auth_generation.wrapping_add(1); } self.connections.insert(name, connection); + Ok(()) } /// Record a connect failure's auth classification. When the failure looks @@ -3094,15 +3107,12 @@ impl McpPool { /// Handshake the pending servers concurrently without holding the pool /// lock. Callers insert results under a short lock so a live turn can /// snapshot ready tools while optional servers are still connecting. - pub(crate) async fn connect_pending_concurrently( + pub(crate) fn spawn_pending_connects( pending: Vec, timeouts: McpTimeouts, network_policy: Option, catalog_generation: u64, - ) -> Vec<(String, Result)> { - if pending.is_empty() { - return Vec::new(); - } + ) -> tokio::task::JoinSet<(String, Result)> { let semaphore = std::sync::Arc::new(tokio::sync::Semaphore::new(Self::CONNECT_CONCURRENCY)); let mut joins: tokio::task::JoinSet<(String, Result)> = tokio::task::JoinSet::new(); @@ -3110,36 +3120,28 @@ impl McpPool { let permit = semaphore.clone(); let network_policy = network_policy.clone(); joins.spawn(async move { - let _permit = permit.acquire_owned().await; - let connection = McpConnection::connect_with_policy( - name.clone(), - config, - &timeouts, - network_policy.as_ref(), - ) + let connection = std::panic::AssertUnwindSafe(async { + let _permit = permit.acquire_owned().await; + McpConnection::connect_with_policy( + name.clone(), + config, + &timeouts, + network_policy.as_ref(), + ) + .await + .map(|mut connection| { + connection.catalog_generation = catalog_generation; + connection + }) + }) + .catch_unwind() .await - .map(|mut connection| { - connection.catalog_generation = catalog_generation; - connection - }); + .unwrap_or_else(|_| Err(anyhow::anyhow!("MCP connection task panicked"))); (name, connection) }); } - let mut results = Vec::new(); - while let Some(joined) = joins.join_next().await { - match joined { - Ok(result) => results.push(result), - // A panicked connect task loses its server name in the - // JoinError; attribute generically. The sequential loop - // would have propagated the panic and taken the whole - // pool down with it, so this is strictly better. - Err(join_error) => { - results.push(("connection task".to_string(), Err(join_error.into()))); - } - } - } - results + joins } /// Connect to all enabled servers, returning errors for failed connections. @@ -3197,20 +3199,20 @@ impl McpPool { break; } - let results = Self::connect_pending_concurrently( + let mut connects = Self::spawn_pending_connects( pending, self.config.timeouts, self.network_policy.clone(), self.catalog_generation.load(Ordering::SeqCst), - ) - .await; - for (name, result) in results { - match result { - Ok(connection) => self.store_ready_connection(name, connection), - Err(error) => { - self.note_connect_failure(&name, &error); - errors.push((name, error)); - } + ); + while let Some(joined) = connects.join_next().await { + let (name, result) = joined + .unwrap_or_else(|error| ("connection task".to_string(), Err(error.into()))); + let result = result + .and_then(|connection| self.store_ready_connection(name.clone(), connection)); + if let Err(error) = result { + self.note_connect_failure(&name, &error); + errors.push((name, error)); } } diff --git a/crates/tui/src/mcp/tests.rs b/crates/tui/src/mcp/tests.rs index c65fdac0d6..18cb057d01 100644 --- a/crates/tui/src/mcp/tests.rs +++ b/crates/tui/src/mcp/tests.rs @@ -2913,6 +2913,27 @@ async fn reload_if_config_changed_swaps_config_on_content_change() { ); } +#[tokio::test] +async fn stale_handshake_cannot_be_restamped_after_config_reload() { + let dir = tempfile::tempdir().unwrap(); + let path = dir.path().join("mcp.json"); + std::fs::write(&path, r#"{"servers":{"local":{"command":"node"}}}"#).unwrap(); + let mut pool = McpPool::from_config_path(&path).unwrap(); + let drops = Arc::new(AtomicUsize::new(0)); + let mut connection = test_connection(Box::new(DropCountingTransport { + drops: drops.clone(), + })); + connection.catalog_generation = pool.current_catalog_generation(); + std::fs::write(&path, r#"{"servers":{}}"#).unwrap(); + pool.reload_from_config_sources(true).unwrap(); + let error = pool + .store_ready_connection("local".to_string(), connection) + .unwrap_err(); + assert!(error.to_string().contains("configuration changed")); + assert!(!pool.connections.contains_key("local")); + assert_eq!(drops.load(AtomicOrdering::SeqCst), 1); +} + #[tokio::test] async fn reload_if_config_changed_drops_live_connections() { let dir = tempfile::tempdir().unwrap();