feat: 更新配置管理,支持异步读取和写入;添加 MCP 服务器状态接口

This commit is contained in:
oudecheng 2026-07-03 08:37:02 +08:00
parent 4071dcc6ee
commit 1af0fab3ad
6 changed files with 271 additions and 44 deletions

View File

@ -62,7 +62,7 @@ pub struct SaveConfigResponse {
pub async fn get_config(
State(state): State<Arc<GatewayState>>,
) -> Json<Config> {
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,19 +72,21 @@ pub async fn save_config(
) -> Result<Json<SaveConfigResponse>, (StatusCode, String)> {
// Merge: preserve original api_keys if the submitted ones are masked
let mut new_config = req.config;
// 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) {
// Restore original api_key
if let Some(original) = state.config.providers.get(name) {
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_channel) = cfg.channels.get(name) {
if let Some(original_feishu) = original_channel.as_feishu() {
feishu.app_secret = original_feishu.app_secret.clone();
}
@ -92,6 +94,7 @@ pub async fn save_config(
}
}
}
} // 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<Arc<GatewayState>>,
) -> Json<crate::mcp::client::McpStatusResponse> {
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)
}

View File

@ -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<RwLock<Config>>,
pub session_manager: SessionManager,
pub channel_manager: ChannelManager,
pub bus: Arc<MessageBus>,
@ -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);

View File

@ -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();

View File

@ -61,6 +61,10 @@ pub struct McpClientManager {
clients: RwLock<HashMap<String, Arc<McpClient>>>,
/// Server information cache keyed by server key
server_info: RwLock<HashMap<String, McpServerInfo>>,
/// 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<HashMap<String, String>>,
}
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<String, crate::mcp::McpServerConfig>,
) -> 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<String>,
}
/// 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<McpServerStatus>,
}
impl Default for McpClientManager {

View File

@ -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};

View File

@ -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<string, string>
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<string, ProviderConfig>
models: Record<string, ModelConfig>
@ -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<McpStatusResponse | null>(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<McpServerConfig>) => update('mcpServers', { ...config.mcpServers, [name]: { ...config.mcpServers[name], ...patch } })
return (
<div className="space-y-4">
{entries.map(([name, s]) => (
{/* MCP Status Summary */}
{mcpStatus && mcpStatus.enabled && (
<div className="flex items-center gap-3 p-3 rounded-lg bg-[var(--bg-tertiary)] text-xs">
<div className="flex items-center gap-1.5">
<span className={`inline-block w-2 h-2 rounded-full ${mcpStatus.connected_servers > 0 ? 'bg-green-400' : 'bg-gray-400'}`} />
<span className="text-[var(--text-secondary)]">{mcpStatus.connected_servers}/{mcpStatus.total_servers} </span>
</div>
{mcpStatus.failed_servers > 0 && (
<span className="text-red-400">{mcpStatus.failed_servers} </span>
)}
<span className="text-[var(--text-muted)]">{mcpStatus.total_tools} </span>
<button onClick={fetchMcpStatus} className="ml-auto px-2 py-1 rounded text-[var(--text-muted)] hover:text-[var(--accent-cyan)] transition-colors" title="刷新状态">
<RefreshCw className="h-3 w-3" />
</button>
</div>
)}
{entries.map(([name, s]) => {
const st = statusFor(name)
return (
<div key={name} className="rounded-xl border border-[var(--border-color)] bg-[var(--bg-secondary)]/60 overflow-hidden">
<MapEntryHeader name={name} onDelete={() => delMcp(name)} />
<div className="flex items-center gap-2 p-3 border-b border-[var(--border-color)]">
{st ? (
st.connected
? <span className="inline-flex items-center gap-1 text-xs text-green-400"><span className="w-2 h-2 rounded-full bg-green-400" /> {st.tool_count} </span>
: st.error
? <span className="inline-flex items-center gap-1 text-xs text-red-400" title={st.error}><span className="w-2 h-2 rounded-full bg-red-400" /> </span>
: <span className="inline-flex items-center gap-1 text-xs text-gray-400"><span className="w-2 h-2 rounded-full bg-gray-400" /> </span>
) : null}
<span className="flex-1 text-sm font-medium text-[var(--text-primary)]">{name}</span>
<button onClick={() => delMcp(name)} className="p-1 rounded text-[var(--text-muted)] hover:text-red-400 transition-colors"><Trash2 className="h-3.5 w-3.5" /></button>
</div>
<div className="p-4 space-y-3">
<Field label="传输类型">
<select value={s.type} onChange={e => updMcp(name, { type: e.target.value as McpServerConfig['type'] })} className={selectCls}>
@ -722,7 +783,8 @@ export function ConfigPage({ onClose, onSaveConnection }: ConfigPageProps) {
)}
</div>
</div>
))}
)
})}
<button onClick={addMcp} className="flex items-center gap-2 px-4 py-2.5 rounded-xl border border-dashed border-[var(--border-color)] text-[var(--text-muted)] hover:text-[var(--accent-cyan)] hover:border-[var(--accent-cyan)]/30 transition-colors text-sm w-full justify-center">
<Plus className="h-4 w-4" /> MCP
</button>