Skip to main content

xz_mcp_engine/
manager.rs

1use std::collections::HashMap;
2
3use serde::Serialize;
4use tokio::sync::RwLock;
5use xz_mcp_core::{McpClient, McpError, McpServerConfig, McpTool, McpToolResult, McpTransportConfig};
6
7use crate::stdio::StdioMcpClient;
8
9#[derive(Debug, Clone, Serialize)]
10#[serde(rename_all = "camelCase")]
11pub struct ServerStatus {
12    pub name: String,
13    pub transport: String,
14    pub connected: bool,
15    pub tool_count: usize,
16    pub error: Option<String>,
17}
18
19/// Multi-server MCP connection manager (stdio / future transports).
20pub struct McpManager {
21    clients: RwLock<HashMap<String, Box<dyn McpClient>>>,
22    tools: RwLock<HashMap<String, (String, McpTool)>>,
23    statuses: RwLock<HashMap<String, ServerStatus>>,
24    configs: RwLock<HashMap<String, McpServerConfig>>,
25}
26
27impl Default for McpManager {
28    fn default() -> Self {
29        Self::new()
30    }
31}
32
33impl McpManager {
34    /// Create an empty manager with no connected servers.
35    pub fn new() -> Self {
36        Self {
37            clients: RwLock::new(HashMap::new()),
38            tools: RwLock::new(HashMap::new()),
39            statuses: RwLock::new(HashMap::new()),
40            configs: RwLock::new(HashMap::new()),
41        }
42    }
43
44    /// Connect to all enabled servers. Returns (connected, failed) counts.
45    pub async fn connect_all(&self, configs: &[McpServerConfig]) -> (usize, usize, Vec<String>) {
46        for cfg in configs.iter().filter(|c| c.enabled) {
47            self.configs.write().await.insert(cfg.name.clone(), cfg.clone());
48        }
49        let mut ok = 0;
50        let mut fail = 0;
51        let mut messages = Vec::new();
52
53        for cfg in configs.iter().filter(|c| c.enabled) {
54            match self.connect_one(cfg).await {
55                Ok(tool_count) => {
56                    ok += 1;
57                    messages.push(format!("✅ {} — {} 个工具", cfg.name, tool_count));
58                }
59                Err(e) => {
60                    fail += 1;
61                    messages.push(format!("❌ {} — {}", cfg.name, e));
62                }
63            }
64        }
65        (ok, fail, messages)
66    }
67
68    /// Connect a single server. Errors are returned to the caller for UI feedback.
69    pub async fn connect_one(&self, cfg: &McpServerConfig) -> Result<usize, McpError> {
70        let transport_label = match &cfg.transport {
71            McpTransportConfig::Stdio { command, .. } => format!("stdio({command})"),
72            McpTransportConfig::Http { url, .. } => format!("http({url})"),
73        };
74
75        let mut status = ServerStatus {
76            name: cfg.name.clone(),
77            transport: transport_label.clone(),
78            connected: false,
79            tool_count: 0,
80            error: None,
81        };
82        self.statuses.write().await.insert(cfg.name.clone(), status.clone());
83
84        let mut client: Box<dyn McpClient> = match &cfg.transport {
85            McpTransportConfig::Stdio { command, args, env } => {
86                Box::new(StdioMcpClient::new(command, args.clone(), env.clone()))
87            }
88            McpTransportConfig::Http { url, headers } => {
89                Box::new(crate::HttpMcpClient::new(url, headers.clone()))
90            }
91        };
92
93        client.connect().await?;
94        let tools = client.list_tools().await?;
95
96        let name = cfg.name.clone();
97        for tool in &tools {
98            self.tools.write().await.insert(tool.name.clone(), (name.clone(), tool.clone()));
99        }
100        self.clients.write().await.insert(name.clone(), client);
101
102        status.connected = true;
103        status.tool_count = tools.len();
104        self.statuses.write().await.insert(name.clone(), status);
105
106        tracing::info!(server = %cfg.name, tool_count = tools.len(), "MCP server connected");
107        Ok(tools.len())
108    }
109
110    pub async fn all_tools(&self) -> Vec<McpTool> {
111        self.tools.read().await.values().map(|(_, t)| t.clone()).collect()
112    }
113
114    pub async fn call_tool(&self, name: &str, args: serde_json::Value) -> Result<McpToolResult, McpError> {
115        let (server_name, _tool) = self.tools.read().await
116            .get(name)
117            .cloned()
118            .ok_or_else(|| McpError::ToolNotFound(name.into()))?;
119
120        let clients = self.clients.read().await;
121        let client = clients.get(&server_name)
122            .ok_or_else(|| McpError::Connection(format!("server '{server_name}' disconnected")))?;
123
124        client.call_tool(name, args).await
125    }
126
127    /// Get status of all servers (both configured and connected).
128    pub async fn list_servers(&self) -> Vec<ServerStatus> {
129        self.statuses.read().await.values().cloned().collect()
130    }
131
132    /// Reconnect a previously disconnected server using its stored config.
133    pub async fn reconnect_server(&self, name: &str) -> Result<usize, McpError> {
134        let cfg = self.configs.read().await.get(name).cloned()
135            .ok_or_else(|| McpError::Other(format!("no config for '{name}'")))?;
136        self.connect_one(&cfg).await
137    }
138
139    /// Connect a server with the given config, storing it for future reconnects.
140    /// If a server with the same name already exists, it is disconnected first.
141    pub async fn connect_server(&self, cfg: &McpServerConfig) -> Result<usize, McpError> {
142        self.disconnect_server(&cfg.name).await;
143        self.configs.write().await.insert(cfg.name.clone(), cfg.clone());
144        self.connect_one(cfg).await
145    }
146
147    /// Fully remove a server: disconnect and clear its stored config.
148    pub async fn remove_server(&self, name: &str) {
149        self.disconnect_server(name).await;
150        self.configs.write().await.remove(name);
151    }
152
153    /// Disconnect a specific server.
154    pub async fn disconnect_server(&self, name: &str) {
155        self.clients.write().await.remove(name);
156        self.tools.write().await.retain(|_, (s, _)| s != name);
157        if let Some(s) = self.statuses.write().await.get_mut(name) {
158            s.connected = false;
159            s.tool_count = 0;
160        }
161    }
162}