use std::collections::HashMap;
use serde::Serialize;
use tokio::sync::RwLock;
use xz_mcp_core::{McpClient, McpError, McpServerConfig, McpTool, McpToolResult, McpTransportConfig};
use crate::stdio::StdioMcpClient;
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct ServerStatus {
pub name: String,
pub transport: String,
pub connected: bool,
pub tool_count: usize,
pub error: Option<String>,
}
pub struct McpManager {
clients: RwLock<HashMap<String, Box<dyn McpClient>>>,
tools: RwLock<HashMap<String, (String, McpTool)>>,
statuses: RwLock<HashMap<String, ServerStatus>>,
configs: RwLock<HashMap<String, McpServerConfig>>,
}
impl Default for McpManager {
fn default() -> Self {
Self::new()
}
}
impl McpManager {
pub fn new() -> Self {
Self {
clients: RwLock::new(HashMap::new()),
tools: RwLock::new(HashMap::new()),
statuses: RwLock::new(HashMap::new()),
configs: RwLock::new(HashMap::new()),
}
}
pub async fn connect_all(&self, configs: &[McpServerConfig]) -> (usize, usize, Vec<String>) {
for cfg in configs.iter().filter(|c| c.enabled) {
self.configs.write().await.insert(cfg.name.clone(), cfg.clone());
}
let mut ok = 0;
let mut fail = 0;
let mut messages = Vec::new();
for cfg in configs.iter().filter(|c| c.enabled) {
match self.connect_one(cfg).await {
Ok(tool_count) => {
ok += 1;
messages.push(format!("✅ {} — {} 个工具", cfg.name, tool_count));
}
Err(e) => {
fail += 1;
messages.push(format!("❌ {} — {}", cfg.name, e));
}
}
}
(ok, fail, messages)
}
pub async fn connect_one(&self, cfg: &McpServerConfig) -> Result<usize, McpError> {
let transport_label = match &cfg.transport {
McpTransportConfig::Stdio { command, .. } => format!("stdio({command})"),
McpTransportConfig::Http { url, .. } => format!("http({url})"),
};
let mut status = ServerStatus {
name: cfg.name.clone(),
transport: transport_label.clone(),
connected: false,
tool_count: 0,
error: None,
};
self.statuses.write().await.insert(cfg.name.clone(), status.clone());
let mut client: Box<dyn McpClient> = match &cfg.transport {
McpTransportConfig::Stdio { command, args, env } => {
Box::new(StdioMcpClient::new(command, args.clone(), env.clone()))
}
McpTransportConfig::Http { url, headers } => {
Box::new(crate::HttpMcpClient::new(url, headers.clone()))
}
};
client.connect().await?;
let tools = client.list_tools().await?;
let name = cfg.name.clone();
for tool in &tools {
self.tools.write().await.insert(tool.name.clone(), (name.clone(), tool.clone()));
}
self.clients.write().await.insert(name.clone(), client);
status.connected = true;
status.tool_count = tools.len();
self.statuses.write().await.insert(name.clone(), status);
tracing::info!(server = %cfg.name, tool_count = tools.len(), "MCP server connected");
Ok(tools.len())
}
pub async fn all_tools(&self) -> Vec<McpTool> {
self.tools.read().await.values().map(|(_, t)| t.clone()).collect()
}
pub async fn call_tool(&self, name: &str, args: serde_json::Value) -> Result<McpToolResult, McpError> {
let (server_name, _tool) = self.tools.read().await
.get(name)
.cloned()
.ok_or_else(|| McpError::ToolNotFound(name.into()))?;
let clients = self.clients.read().await;
let client = clients.get(&server_name)
.ok_or_else(|| McpError::Connection(format!("server '{server_name}' disconnected")))?;
client.call_tool(name, args).await
}
pub async fn list_servers(&self) -> Vec<ServerStatus> {
self.statuses.read().await.values().cloned().collect()
}
pub async fn reconnect_server(&self, name: &str) -> Result<usize, McpError> {
let cfg = self.configs.read().await.get(name).cloned()
.ok_or_else(|| McpError::Other(format!("no config for '{name}'")))?;
self.connect_one(&cfg).await
}
pub async fn connect_server(&self, cfg: &McpServerConfig) -> Result<usize, McpError> {
self.disconnect_server(&cfg.name).await;
self.configs.write().await.insert(cfg.name.clone(), cfg.clone());
self.connect_one(cfg).await
}
pub async fn remove_server(&self, name: &str) {
self.disconnect_server(name).await;
self.configs.write().await.remove(name);
}
pub async fn disconnect_server(&self, name: &str) {
self.clients.write().await.remove(name);
self.tools.write().await.retain(|_, (s, _)| s != name);
if let Some(s) = self.statuses.write().await.get_mut(name) {
s.connected = false;
s.tool_count = 0;
}
}
}