diff --git a/src/openhuman/mcp_registry/boot.rs b/src/openhuman/mcp_registry/boot.rs index b34ee7eeb..7c3f7b174 100644 --- a/src/openhuman/mcp_registry/boot.rs +++ b/src/openhuman/mcp_registry/boot.rs @@ -10,10 +10,24 @@ //! to spawn. Once the `InstalledServer` model grows a remote-transport //! variant this function will skip them (or call a remote "warm-up" path). +use futures::stream::StreamExt; + use crate::openhuman::config::Config; +use super::types::InstalledServer; use super::{connections, store}; +/// How many installed MCP servers are brought up concurrently at boot. +/// +/// Each `connect` spawns a stdio subprocess and does the MCP `initialize` + +/// `tools/list` handshake, so the connects are independent network/process +/// round-trips — serial spawning made boot latency the *sum* of every +/// server's warmup. Bounded concurrency overlaps the handshakes while still +/// capping the subprocess spawn burst (a user with dozens of servers +/// shouldn't fork them all at once). The registry insert each `connect` +/// performs is internally synchronized, so there is no shared-state hazard. +const BOOT_SPAWN_CONCURRENCY: usize = 8; + /// Spawn every locally-installed MCP server. Per-server failures are logged /// and swallowed. pub async fn spawn_installed_servers(config: &Config) { @@ -35,29 +49,62 @@ pub async fn spawn_installed_servers(config: &Config) { servers.len() ); - for server in servers { + spawn_servers_concurrently(servers, |server| async move { + connections::connect(config, &server) + .await + .map(|tools| tools.len()) + }) + .await; +} + +/// Bring up `servers` with bounded concurrency, logging per-server outcomes. +/// +/// Disabled servers are filtered (and logged) before fan-out. `connect_fn` +/// takes the server **by value** (not `&InstalledServer`) so the returned +/// future owns its input — borrowing the argument would force rustc into a +/// higher-ranked `FnOnce` bound that fails to infer. Order is irrelevant: +/// each connect's effect (registry insert) is independent, so this uses +/// `for_each_concurrent` rather than an ordered combinator. +async fn spawn_servers_concurrently(servers: Vec, connect_fn: F) +where + F: Fn(InstalledServer) -> Fut, + Fut: std::future::Future>, +{ + let enabled = servers.into_iter().filter(|server| { if !server.enabled { tracing::info!( "[mcp-registry] boot: skipping disabled server_id={} qualified={}", server.server_id, server.qualified_name ); - continue; } - let server_id = server.server_id.clone(); - let qualified = server.qualified_name.clone(); - match connections::connect(config, &server).await { - Ok(tools) => tracing::info!( - "[mcp-registry] boot: connected server_id={} qualified={} tools={}", - server_id, - qualified, - tools.len() - ), - Err(err) => tracing::warn!( - "[mcp-registry] boot: connect failed server_id={} qualified={} err={err}", - server_id, - qualified - ), - } - } + server.enabled + }); + + futures::stream::iter(enabled) + .for_each_concurrent(BOOT_SPAWN_CONCURRENCY, |server| { + let connect_fn = &connect_fn; + let server_id = server.server_id.clone(); + let qualified = server.qualified_name.clone(); + async move { + match connect_fn(server).await { + Ok(tool_count) => tracing::info!( + "[mcp-registry] boot: connected server_id={} qualified={} tools={}", + server_id, + qualified, + tool_count + ), + Err(err) => tracing::warn!( + "[mcp-registry] boot: connect failed server_id={} qualified={} err={err}", + server_id, + qualified + ), + } + } + }) + .await; } + +#[cfg(test)] +#[path = "boot_tests.rs"] +mod tests; diff --git a/src/openhuman/mcp_registry/boot_tests.rs b/src/openhuman/mcp_registry/boot_tests.rs new file mode 100644 index 000000000..c78a1f62b --- /dev/null +++ b/src/openhuman/mcp_registry/boot_tests.rs @@ -0,0 +1,136 @@ +//! Tests for boot-time concurrent MCP server spawn. + +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; +use std::time::Duration; + +use super::spawn_servers_concurrently; +use super::BOOT_SPAWN_CONCURRENCY; +use crate::openhuman::mcp_registry::types::{CommandKind, InstalledServer, Transport}; + +fn sample_server(id: &str, enabled: bool) -> InstalledServer { + InstalledServer { + server_id: id.to_string(), + qualified_name: format!("@test/{id}"), + display_name: "Test Server".to_string(), + description: None, + icon_url: None, + command_kind: CommandKind::Node, + command: "npx".to_string(), + args: vec!["-y".to_string()], + env_keys: Vec::new(), + config: None, + installed_at: 1_700_000_000_000, + last_connected_at: None, + transport: Transport::Stdio, + enabled, + } +} + +/// Tracks how many `connect_fn` invocations overlap so the test can assert +/// real concurrency (peak in-flight > 1) bounded by `BOOT_SPAWN_CONCURRENCY`. +#[derive(Default)] +struct ConcurrencyProbe { + in_flight: AtomicUsize, + peak: AtomicUsize, + calls: AtomicUsize, +} + +impl ConcurrencyProbe { + async fn run(&self) { + self.calls.fetch_add(1, Ordering::SeqCst); + let now = self.in_flight.fetch_add(1, Ordering::SeqCst) + 1; + self.peak.fetch_max(now, Ordering::SeqCst); + // Hold the slot long enough that siblings pile up if run concurrently. + tokio::time::sleep(Duration::from_millis(40)).await; + self.in_flight.fetch_sub(1, Ordering::SeqCst); + } +} + +#[tokio::test] +async fn spawns_enabled_servers_concurrently() { + let probe = Arc::new(ConcurrencyProbe::default()); + let servers: Vec = (0..4) + .map(|i| sample_server(&format!("srv-{i}"), true)) + .collect(); + + let probe_for_fn = probe.clone(); + spawn_servers_concurrently(servers, move |_server| { + let probe = probe_for_fn.clone(); + async move { + probe.run().await; + Ok(3usize) + } + }) + .await; + + assert_eq!( + probe.calls.load(Ordering::SeqCst), + 4, + "all enabled servers connected" + ); + let peak = probe.peak.load(Ordering::SeqCst); + assert!( + peak >= 2, + "expected overlapping connects, peak in-flight was {peak}" + ); + assert!( + peak <= BOOT_SPAWN_CONCURRENCY, + "peak in-flight {peak} must not exceed BOOT_SPAWN_CONCURRENCY {BOOT_SPAWN_CONCURRENCY}" + ); +} + +#[tokio::test] +async fn skips_disabled_servers() { + let probe = Arc::new(ConcurrencyProbe::default()); + let servers = vec![ + sample_server("on-1", true), + sample_server("off-1", false), + sample_server("on-2", true), + sample_server("off-2", false), + ]; + + let probe_for_fn = probe.clone(); + spawn_servers_concurrently(servers, move |_server| { + let probe = probe_for_fn.clone(); + async move { + probe.run().await; + Ok(0usize) + } + }) + .await; + + assert_eq!( + probe.calls.load(Ordering::SeqCst), + 2, + "only the two enabled servers should be connected" + ); +} + +/// An error from one connect must not abort the others (boot is best-effort). +#[tokio::test] +async fn one_failure_does_not_abort_the_rest() { + let probe = Arc::new(ConcurrencyProbe::default()); + let servers: Vec = (0..3) + .map(|i| sample_server(&format!("srv-{i}"), true)) + .collect(); + + let probe_for_fn = probe.clone(); + spawn_servers_concurrently(servers, move |server| { + let probe = probe_for_fn.clone(); + async move { + probe.run().await; + if server.server_id == "srv-1" { + anyhow::bail!("boom"); + } + Ok(1usize) + } + }) + .await; + + assert_eq!( + probe.calls.load(Ordering::SeqCst), + 3, + "every server is attempted even when one fails" + ); +}