use std::collections::HashMap;
use std::sync::Arc;
use super::bridge::{McpBridgedTool, McpToolClient, mcp_tool_name};
use super::client::{McpClientError, McpStdioClient};
use super::protocol::{McpServerConfig, McpToolDefinition};
use super::sse::client::{McpSseClient, McpSseError};
use super::sse::config::McpSseServerConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum McpServerStatus {
Disconnected,
Connecting,
Connected,
Error,
}
impl std::fmt::Display for McpServerStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Disconnected => write!(f, "disconnected"),
Self::Connecting => write!(f, "connecting"),
Self::Connected => write!(f, "connected"),
Self::Error => write!(f, "error"),
}
}
}
#[derive(Debug, Clone)]
pub struct McpServerSummary {
pub name: String,
pub status: McpServerStatus,
pub server_version: Option<String>,
pub tool_count: usize,
pub error: Option<String>,
}
enum TransportClient {
Stdio(Arc<McpStdioClient>),
Sse(Arc<McpSseClient>),
}
impl TransportClient {
fn server_version(&self) -> Option<String> {
match self {
Self::Stdio(client) => client.server_info().map(|info| info.version.clone()),
Self::Sse(client) => client.server_info().map(|info| info.version.clone()),
}
}
async fn shutdown(&self) {
match self {
Self::Stdio(client) => client.shutdown().await,
Self::Sse(client) => client.shutdown().await,
}
}
async fn call_tool(
&self,
tool_name: &str,
arguments: Option<serde_json::Value>,
) -> Result<super::protocol::McpToolCallResult, String> {
match self {
Self::Stdio(client) => McpToolClient::call_tool(&**client, tool_name, arguments).await,
Self::Sse(client) => McpToolClient::call_tool(&**client, tool_name, arguments).await,
}
}
fn bridge(&self, server_name: &str, tools: &[McpToolDefinition]) -> Vec<McpBridgedTool> {
tools
.iter()
.map(|tool| match self {
Self::Stdio(client) => {
McpBridgedTool::new(server_name.to_string(), tool.clone(), client.clone())
}
Self::Sse(client) => {
McpBridgedTool::new(server_name.to_string(), tool.clone(), client.clone())
}
})
.collect()
}
}
struct ConnectedServer {
client: TransportClient,
tools: Vec<McpToolDefinition>,
}
pub struct McpManager {
servers: HashMap<String, ConnectedServer>,
errors: HashMap<String, String>,
}
impl McpManager {
pub fn new() -> Self {
Self {
servers: HashMap::new(),
errors: HashMap::new(),
}
}
pub async fn connect(
&mut self,
config: &McpServerConfig,
) -> Result<Vec<McpBridgedTool>, McpClientError> {
self.disconnect(&config.name).await;
let client = McpStdioClient::connect(config).await.inspect_err(|e| {
self.errors.insert(config.name.clone(), e.to_string());
})?;
let tools = client.tools().to_vec();
let client = TransportClient::Stdio(Arc::new(client));
Ok(self.register(config.name.clone(), client, tools))
}
pub async fn connect_sse(
&mut self,
config: &McpSseServerConfig,
) -> Result<Vec<McpBridgedTool>, McpSseError> {
self.disconnect(&config.name).await;
let client = McpSseClient::connect(config).await.inspect_err(|error| {
self.errors.insert(config.name.clone(), error.to_string());
})?;
let tools = client.tools().to_vec();
let client = TransportClient::Sse(Arc::new(client));
Ok(self.register(config.name.clone(), client, tools))
}
fn register(
&mut self,
name: String,
client: TransportClient,
tools: Vec<McpToolDefinition>,
) -> Vec<McpBridgedTool> {
let bridged = client.bridge(&name, &tools);
self.errors.remove(&name);
self.servers.insert(name, ConnectedServer { client, tools });
bridged
}
pub async fn disconnect(&mut self, name: &str) {
if let Some(server) = self.servers.remove(name) {
server.client.shutdown().await;
}
}
pub async fn shutdown_all(&mut self) {
let names: Vec<String> = self.servers.keys().cloned().collect();
for name in names {
self.disconnect(&name).await;
}
}
pub fn list_servers(&self) -> Vec<McpServerSummary> {
let mut summaries: Vec<McpServerSummary> = self
.servers
.iter()
.map(|(name, server)| McpServerSummary {
name: name.clone(),
status: McpServerStatus::Connected,
server_version: server.client.server_version(),
tool_count: server.tools.len(),
error: None,
})
.collect();
for (name, error) in &self.errors {
if !self.servers.contains_key(name) {
summaries.push(McpServerSummary {
name: name.clone(),
status: McpServerStatus::Error,
server_version: None,
tool_count: 0,
error: Some(error.clone()),
});
}
}
summaries.sort_by(|a, b| a.name.cmp(&b.name));
summaries
}
pub fn all_tool_names(&self) -> Vec<String> {
self.servers
.iter()
.flat_map(|(name, server)| {
server
.tools
.iter()
.map(move |tool| mcp_tool_name(name, &tool.name))
})
.collect()
}
pub async fn call_tool(
&self,
server_name: &str,
tool_name: &str,
arguments: Option<serde_json::Value>,
) -> Result<super::protocol::McpToolCallResult, String> {
let server = self
.servers
.get(server_name)
.ok_or_else(|| format!("MCP server '{server_name}' not connected"))?;
server.client.call_tool(tool_name, arguments).await
}
pub fn is_connected(&self, name: &str) -> bool {
self.servers.contains_key(name)
}
pub fn connected_count(&self) -> usize {
self.servers.len()
}
}
impl Default for McpManager {
fn default() -> Self {
Self::new()
}
}