From 1af0fab3add02fcf45e7dc1b939527bf8ea0b082 Mon Sep 17 00:00:00 2001 From: oudecheng <13802883547@139.com> Date: Fri, 3 Jul 2026 08:37:02 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=9B=B4=E6=96=B0=E9=85=8D=E7=BD=AE?= =?UTF-8?q?=E7=AE=A1=E7=90=86=EF=BC=8C=E6=94=AF=E6=8C=81=E5=BC=82=E6=AD=A5?= =?UTF-8?q?=E8=AF=BB=E5=8F=96=E5=92=8C=E5=86=99=E5=85=A5=EF=BC=9B=E6=B7=BB?= =?UTF-8?q?=E5=8A=A0=20MCP=20=E6=9C=8D=E5=8A=A1=E5=99=A8=E7=8A=B6=E6=80=81?= =?UTF-8?q?=E6=8E=A5=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/gateway/http.rs | 65 ++++++++--- src/gateway/mod.rs | 52 +++++---- src/gateway/ws.rs | 2 +- src/mcp/client.rs | 120 +++++++++++++++++++++ src/mcp/mod.rs | 2 +- web/src/components/Settings/ConfigPage.tsx | 74 +++++++++++-- 6 files changed, 271 insertions(+), 44 deletions(-) diff --git a/src/gateway/http.rs b/src/gateway/http.rs index 1a35c59..a3bf39a 100644 --- a/src/gateway/http.rs +++ b/src/gateway/http.rs @@ -62,7 +62,7 @@ pub struct SaveConfigResponse { pub async fn get_config( State(state): State>, ) -> Json { - Json(mask_config(&state.config)) + Json(mask_config(&*state.config.read().await)) } /// PUT /api/config — Save config to file, preserving original api_keys if masked @@ -72,26 +72,29 @@ pub async fn save_config( ) -> Result, (StatusCode, String)> { // Merge: preserve original api_keys if the submitted ones are masked let mut new_config = req.config; - for (name, provider) in new_config.providers.iter_mut() { - if is_masked_key(&provider.api_key) { - // Restore original api_key - if let Some(original) = state.config.providers.get(name) { - provider.api_key = original.api_key.clone(); + + // Read old values under read lock (not held across disk I/O) + { + let cfg = state.config.read().await; + for (name, provider) in new_config.providers.iter_mut() { + if is_masked_key(&provider.api_key) { + if let Some(original) = cfg.providers.get(name) { + provider.api_key = original.api_key.clone(); + } } } - } - // Merge: preserve original app_secrets if the submitted ones are masked - for (name, channel) in new_config.channels.iter_mut() { - if let Some(feishu) = channel.as_feishu_mut() { - if is_masked_key(&feishu.app_secret) { - if let Some(original_channel) = state.config.channels.get(name) { - if let Some(original_feishu) = original_channel.as_feishu() { - feishu.app_secret = original_feishu.app_secret.clone(); + for (name, channel) in new_config.channels.iter_mut() { + if let Some(feishu) = channel.as_feishu_mut() { + if is_masked_key(&feishu.app_secret) { + if let Some(original_channel) = cfg.channels.get(name) { + if let Some(original_feishu) = original_channel.as_feishu() { + feishu.app_secret = original_feishu.app_secret.clone(); + } } } } } - } + } // read lock released here // Validate timezone if let Err(e) = new_config.time.parse_timezone() { @@ -103,13 +106,19 @@ pub async fn save_config( .map(std::path::PathBuf::from) .unwrap_or_else(|_| get_default_config_path()); - // Serialize and write + // Serialize and write to disk (no lock held) let json = serde_json::to_string_pretty(&new_config) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("Serialize error: {}", e)))?; - std::fs::write(&config_path, json) + std::fs::write(&config_path, &json) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("Write error: {}", e)))?; + // Update in-memory config (write lock, held only for assignment) + { + let mut cfg = state.config.write().await; + *cfg = new_config.clone(); + } + tracing::info!(path = %config_path.display(), "Config saved via API"); Ok(Json(SaveConfigResponse { @@ -151,3 +160,25 @@ pub async fn restart( message: "服务正在重启...".to_string(), })) } + +/// GET /api/mcp/status — Return MCP server connection status +pub async fn mcp_status( + State(state): State>, +) -> Json { + let status = match &state.mcp_manager { + Some(manager) => { + // Clone mcp_servers before await to avoid holding read lock across it + let mcp_servers = state.config.read().await.mcp_servers.clone(); + manager.get_status(&mcp_servers).await + } + None => crate::mcp::client::McpStatusResponse { + enabled: false, + total_servers: 0, + connected_servers: 0, + failed_servers: 0, + total_tools: 0, + servers: vec![], + }, + }; + Json(status) +} diff --git a/src/gateway/mod.rs b/src/gateway/mod.rs index 80ce060..b127707 100644 --- a/src/gateway/mod.rs +++ b/src/gateway/mod.rs @@ -52,10 +52,10 @@ use session_message_sender::BusSessionMessageSender; use session::SessionManager; use static_files::static_handler; -use tokio::sync::watch; +use tokio::sync::{watch, RwLock}; pub struct GatewayState { - pub config: Config, + pub config: Arc>, pub session_manager: SessionManager, pub channel_manager: ChannelManager, pub bus: Arc, @@ -107,7 +107,7 @@ impl GatewayState { let cancel_manager = CancelManager::new(); Ok(Self { - config, + config: Arc::new(RwLock::new(config)), session_manager, channel_manager, bus, @@ -122,18 +122,19 @@ impl GatewayState { pub async fn start_message_processing(&self) { let bus_for_outbound = self.bus.clone(); - // Create semaphore for controlling concurrent requests - let max_concurrent = self.config.gateway.max_concurrent_requests; - let semaphore = Arc::new(Semaphore::new(max_concurrent)); - - // Spawn inbound processor with semaphore-controlled concurrency - let provider_config = match self.config.get_provider_config("default") { + // Read config under read lock + let cfg = self.config.read().await; + let max_concurrent = cfg.gateway.max_concurrent_requests; + let provider_config = match cfg.get_provider_config("default") { Ok(config) => config, Err(e) => { tracing::error!(error = %e, "Failed to get provider config"); return; } }; + drop(cfg); // release read lock before spawning long-running tasks + + let semaphore = Arc::new(Semaphore::new(max_concurrent)); let inbound_processor = InboundProcessor::new(self.bus.clone(), self.session_manager.clone(), semaphore, provider_config, self.cancel_manager.clone()); tokio::spawn(inbound_processor.run()); @@ -170,23 +171,30 @@ pub async fn run( let state = Arc::new(GatewayState::from_config(config, restart_tx)?); // Get provider config for channels - let provider_config = state.config.get_provider_config("default")?; + let cfg = state.config.read().await; + let provider_config = cfg.get_provider_config("default")?; // Initialize and start channels state .channel_manager - .init(&state.config, provider_config.clone()) + .init(&*cfg, provider_config.clone()) .await?; + drop(cfg); state.channel_manager.start_all().await?; // Start message processing (inbound processor + outbound dispatcher) state.start_message_processing().await; let (scheduler_shutdown_tx, scheduler_shutdown_rx) = tokio::sync::watch::channel(false); - if state.config.scheduler.enabled { + let scheduler_enabled = { + let cfg = state.config.read().await; + cfg.scheduler.enabled + }; + if scheduler_enabled { + let scheduler_cfg = state.config.read().await.scheduler.clone(); let scheduler = Scheduler::new( state.bus.clone(), - state.config.scheduler.clone(), + scheduler_cfg, timezone, state.session_manager.store(), AgentTaskExecutor::new(state.session_manager.clone()), @@ -199,8 +207,12 @@ pub async fn run( } // CLI args override config file values - let bind_host = host.unwrap_or_else(|| state.config.gateway.host.clone()); - let bind_port = port.unwrap_or(state.config.gateway.port); + let (bind_host, bind_port) = { + let cfg = state.config.read().await; + let h = host.unwrap_or_else(|| cfg.gateway.host.clone()); + let p = port.unwrap_or(cfg.gateway.port); + (h, p) + }; // 使用嵌入的静态文件(编译时打包进二进制) // 开发模式下可通过 STATIC_DIR 环境变量使用磁盘文件 @@ -211,6 +223,7 @@ pub async fn run( .route("/health", routing::get(http::health)) .route("/api/config", routing::get(http::get_config).put(http::save_config)) .route("/api/restart", routing::post(http::restart)) + .route("/api/mcp/status", routing::get(http::mcp_status)) .route("/ws", routing::get(ws::ws_handler)) .fallback(static_handler) .with_state(state.clone()) @@ -220,6 +233,7 @@ pub async fn run( .route("/health", routing::get(http::health)) .route("/api/config", routing::get(http::get_config).put(http::save_config)) .route("/api/restart", routing::post(http::restart)) + .route("/api/mcp/status", routing::get(http::mcp_status)) .route("/ws", routing::get(ws::ws_handler)) .fallback_service(ServeDir::new(&static_dir)) .with_state(state.clone()) @@ -253,8 +267,8 @@ pub async fn run( cancel_manager.cancel_all().await; let _ = scheduler_shutdown_tx.send(true); if let Some(ref mgr) = mcp_manager { - tracing::info!("Disconnecting MCP servers before shutdown"); - let _ = mgr.disconnect_all().await; + tracing::info!("Shutting down MCP servers before shutdown"); + let _ = mgr.shutdown_all().await; } let _ = channel_manager.stop_all().await; let _ = result_tx.send(false); @@ -266,8 +280,8 @@ pub async fn run( cancel_manager.cancel_all().await; let _ = scheduler_shutdown_tx.send(true); if let Some(ref mgr) = mcp_manager { - tracing::info!("Disconnecting MCP servers before restart"); - let _ = mgr.disconnect_all().await; + tracing::info!("Shutting down MCP servers before restart"); + let _ = mgr.shutdown_all().await; } let _ = channel_manager.stop_all().await; let _ = result_tx.send(true); diff --git a/src/gateway/ws.rs b/src/gateway/ws.rs index 0cb8d85..596daa6 100644 --- a/src/gateway/ws.rs +++ b/src/gateway/ws.rs @@ -368,7 +368,7 @@ async fn handle_inbound( let store = state.session_manager.store(); let skills = state.session_manager.skills(); let skills_for_handler = skills.clone(); - let provider_config = state.config.get_provider_config("default") + let provider_config = state.config.read().await.get_provider_config("default") .map_err(|e| AgentError::Other(e.to_string()))?; let prompt_repository = state.session_manager.store().clone(); diff --git a/src/mcp/client.rs b/src/mcp/client.rs index 6d27db1..6dfd920 100644 --- a/src/mcp/client.rs +++ b/src/mcp/client.rs @@ -61,6 +61,10 @@ pub struct McpClientManager { clients: RwLock>>, /// Server information cache keyed by server key server_info: RwLock>, + /// Count of active stdio (child process) connections + stdio_client_count: std::sync::atomic::AtomicUsize, + /// Connection errors per server key (last error message) + connection_errors: RwLock>, } impl McpClientManager { @@ -69,6 +73,8 @@ impl McpClientManager { Self { clients: RwLock::new(HashMap::new()), server_info: RwLock::new(HashMap::new()), + stdio_client_count: std::sync::atomic::AtomicUsize::new(0), + connection_errors: RwLock::new(HashMap::new()), } } @@ -139,7 +145,12 @@ impl McpClientManager { attempts = MAX_RETRIES, "Failed to connect to MCP server after all retries" ); + // Record error for status reporting + self.connection_errors.write().await.insert(key.clone(), e.to_string()); failed += 1; + } else { + // Clear any previous error on successful connection + self.connection_errors.write().await.remove(&key); } } @@ -225,6 +236,9 @@ impl McpClientManager { // Use default client handler (empty tuple) let client = ().serve(transport).await?; + // Track that we have a stdio (child process) connection + self.stdio_client_count.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Ok(client) } @@ -361,10 +375,116 @@ impl McpClientManager { Ok(()) } + /// Shut down all MCP connections with proper cleanup + /// + /// Drops all client connections and waits for child processes to terminate + /// if stdio transports were in use. This prevents race conditions during + /// gateway restart where old MCP processes may still be running when + /// new ones start. + pub async fn shutdown_all(&self) -> anyhow::Result<()> { + let stdio_count = self.stdio_client_count.load(std::sync::atomic::Ordering::SeqCst); + + // Drop all clients (triggers cancellation + graceful shutdown in rmcp) + self.disconnect_all().await?; + + // If stdio connections were active, wait for child processes to be killed. + // rmcp's RunningService::drop() triggers async cancellation with: + // - 2 second graceful drain period + // - 3 second process kill timeout + // Total: ~5 seconds. We add 1 second buffer. + if stdio_count > 0 { + tracing::info!( + stdio_count, + "Waiting for MCP child processes to terminate (up to 6s)..." + ); + tokio::time::sleep(std::time::Duration::from_secs(6)).await; + tracing::info!("MCP child process cleanup wait complete"); + self.stdio_client_count.store(0, std::sync::atomic::Ordering::SeqCst); + } + + Ok(()) + } + /// Check if any servers are connected pub async fn has_connections(&self) -> bool { !self.clients.read().await.is_empty() } + + /// Get current MCP connection status for all configured servers + pub async fn get_status( + &self, + configured_servers: &HashMap, + ) -> McpStatusResponse { + let info_map = self.server_info.read().await; + let errors = self.connection_errors.read().await; + let clients = self.clients.read().await; + + let mut total_servers = 0usize; + let mut connected_servers = 0usize; + let mut failed_servers = 0usize; + let mut total_tools = 0usize; + let mut servers = Vec::new(); + + for (key, config) in configured_servers { + total_servers += 1; + let name = config.effective_name(key); + let transport_type = config.transport_type.clone(); + let is_active = config.is_active; + let connected = clients.contains_key(key); + let tool_count = info_map.get(key).map(|info| info.tools.len()).unwrap_or(0); + let error = errors.get(key).cloned(); + + if connected { + connected_servers += 1; + } else if is_active && error.is_some() { + failed_servers += 1; + } + + total_tools += tool_count; + + servers.push(McpServerStatus { + key: key.clone(), + name, + transport_type, + is_active, + connected, + tool_count, + error, + }); + } + + McpStatusResponse { + enabled: !configured_servers.is_empty(), + total_servers, + connected_servers, + failed_servers, + total_tools, + servers, + } + } +} + +/// Status of a single MCP server connection +#[derive(Debug, Clone, serde::Serialize)] +pub struct McpServerStatus { + pub key: String, + pub name: String, + pub transport_type: String, + pub is_active: bool, + pub connected: bool, + pub tool_count: usize, + pub error: Option, +} + +/// Overall MCP status response +#[derive(Debug, Clone, Default, serde::Serialize)] +pub struct McpStatusResponse { + pub enabled: bool, + pub total_servers: usize, + pub connected_servers: usize, + pub failed_servers: usize, + pub total_tools: usize, + pub servers: Vec, } impl Default for McpClientManager { diff --git a/src/mcp/mod.rs b/src/mcp/mod.rs index d3d9215..95cb165 100644 --- a/src/mcp/mod.rs +++ b/src/mcp/mod.rs @@ -16,5 +16,5 @@ pub mod client; pub mod tool_adapter; pub use config::{McpConfig, McpServerConfig, McpTransportConfig}; -pub use client::{McpClientManager, McpClient, McpServerInfo, McpInitializer}; +pub use client::{McpClientManager, McpClient, McpServerInfo, McpInitializer, McpServerStatus, McpStatusResponse}; pub use tool_adapter::{McpToolWrapper, register_mcp_tools}; \ No newline at end of file diff --git a/web/src/components/Settings/ConfigPage.tsx b/web/src/components/Settings/ConfigPage.tsx index 3489c9f..b5f0860 100644 --- a/web/src/components/Settings/ConfigPage.tsx +++ b/web/src/components/Settings/ConfigPage.tsx @@ -20,6 +20,7 @@ interface ImageContextConfig { max_images_in_context: number; max_image_age_roun interface SubagentsConfig { enabled: boolean; sources: string[] } interface ClientConfig { gateway_url: string } interface McpServerConfig { + name?: string type: 'stdio' | 'streamableHttp' | 'http' is_active: boolean command?: string @@ -29,6 +30,25 @@ interface McpServerConfig { headers?: Record description?: string } + +interface McpServerStatus { + key: string + name: string + transport_type: string + is_active: boolean + connected: boolean + tool_count: number + error?: string +} + +interface McpStatusResponse { + enabled: boolean + total_servers: number + connected_servers: number + failed_servers: number + total_tools: number + servers: McpServerStatus[] +} interface AppConfig { providers: Record models: Record @@ -284,6 +304,14 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) { const [dirty, setDirty] = useState(false) const [showRestartDialog, setShowRestartDialog] = useState(false) const [restarting, setRestarting] = useState(false) + const [mcpStatus, setMcpStatus] = useState(null) + + const fetchMcpStatus = useCallback(async () => { + try { + const resp = await fetch('/api/mcp/status') + if (resp.ok) setMcpStatus(await resp.json()) + } catch { /* ignore fetch errors */ } + }, []) const handleClose = useCallback(() => { if (dirty && !confirm('有未保存的更改,确定要关闭吗?')) return @@ -298,6 +326,11 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) { }).catch(e => { setError('加载配置失败: ' + e.message); setLoading(false) }) }, []) + // Fetch MCP status when MCP tab is selected + useEffect(() => { + if (activeTab === 'mcp') fetchMcpStatus() + }, [activeTab, fetchMcpStatus]) + // ESC to close useEffect(() => { const h = (e: KeyboardEvent) => { if (e.key === 'Escape') handleClose() } @@ -321,9 +354,8 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) { }) const data = await resp.json() if (!resp.ok) throw new Error(data.message || data.error || '保存失败') - // Reload config from server to get masked values - const refreshed = await fetch('/api/config').then(r => r.json()) - setConfig(refreshed) + // Config is now synced to both disk and in-memory state, + // so the local state is already correct. No need to re-fetch. setDirty(false) // Show restart confirmation dialog setShowRestartDialog(true) @@ -661,6 +693,7 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) { const renderMcp = () => { const entries = Object.entries(config.mcpServers) + const statusFor = (key: string) => mcpStatus?.servers?.find(s => s.key === key) const addMcp = () => { const name = prompt('MCP 服务器名称:')?.trim() if (name && !config.mcpServers[name]) { @@ -671,9 +704,37 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) { const updMcp = (name: string, patch: Partial) => update('mcpServers', { ...config.mcpServers, [name]: { ...config.mcpServers[name], ...patch } }) return (
- {entries.map(([name, s]) => ( + {/* MCP Status Summary */} + {mcpStatus && mcpStatus.enabled && ( +
+
+ 0 ? 'bg-green-400' : 'bg-gray-400'}`} /> + {mcpStatus.connected_servers}/{mcpStatus.total_servers} 已连接 +
+ {mcpStatus.failed_servers > 0 && ( + {mcpStatus.failed_servers} 失败 + )} + {mcpStatus.total_tools} 个工具 + +
+ )} + {entries.map(([name, s]) => { + const st = statusFor(name) + return (
- delMcp(name)} /> +
+ {st ? ( + st.connected + ? {st.tool_count} 工具 + : st.error + ? 错误 + : 未连接 + ) : null} + {name} + +