use crate::{
cancellation::AgentCancellation,
config::{McpServerConfig, McpServersSettings, validate_mcp_server_name},
mcp::tool_schema::provider_tool_definition,
mcp::{CallToolResult, McpClient, McpError, McpResult, QualifiedMcpToolName, Tool},
};
use serde_json::Value;
use std::{collections::HashMap, path::Path, sync::Arc};
#[derive(Debug)]
pub(crate) struct McpManager {
servers: Vec<McpServerHandle>,
routes: HashMap<QualifiedMcpToolName, (usize, String)>,
statuses: HashMap<String, McpServerStatus>,
shutting_down: bool,
}
#[derive(Debug)]
pub(crate) struct McpServerHandle {
client: Arc<McpClient>,
pub(crate) tools: Vec<Tool>,
}
#[derive(Debug, Clone)]
pub(crate) struct ResolvedMcpToolCall {
client: Arc<McpClient>,
raw_tool_name: String,
}
impl ResolvedMcpToolCall {
pub(crate) fn call_tool_cancellable(
&self,
arguments: Option<Value>,
cancellation: &AgentCancellation,
) -> McpResult<CallToolResult> {
self.client
.call_tool_cancellable(&self.raw_tool_name, arguments, cancellation)
}
}
#[derive(Debug, Clone)]
pub(crate) enum McpServerStatus {
Connected { tool_count: usize },
Failed { error: String, phase: String },
}
impl McpManager {
pub(crate) fn from_settings_with_paths(
mcp_servers: &McpServersSettings,
mc_home: Option<&Path>,
) -> Self {
let mut manager = Self {
servers: Vec::new(),
routes: HashMap::new(),
statuses: HashMap::new(),
shutting_down: false,
};
for (server_name, config) in mcp_servers {
if !config.enabled() {
continue;
}
if let Err(error) = validate_mcp_server_name(server_name) {
manager.record_failure(server_name, "config", error);
continue;
}
match Self::connect_server(server_name, config, mc_home, None) {
Ok((client, tools)) => {
let server_index = manager.servers.len();
let route_result = manager.validate_routes(server_index, server_name, &tools);
match route_result {
Ok(routes) => {
manager.servers.push(McpServerHandle {
client: Arc::new(client),
tools,
});
manager.routes.extend(routes);
manager.statuses.insert(
server_name.clone(),
McpServerStatus::Connected {
tool_count: manager.servers[server_index].tools.len(),
},
);
}
Err(error) => manager.record_failure(server_name, "route", error),
}
}
Err((phase, error)) => manager.record_failure(server_name, phase, error),
}
}
manager
}
pub(crate) fn from_settings_strict_cancellable(
mcp_servers: &McpServersSettings,
mc_home: Option<&Path>,
cancellation: &AgentCancellation,
) -> McpResult<Self> {
Self::from_settings_strict_inner(mcp_servers, mc_home, Some(cancellation))
}
fn from_settings_strict_inner(
mcp_servers: &McpServersSettings,
mc_home: Option<&Path>,
cancellation: Option<&AgentCancellation>,
) -> McpResult<Self> {
let mut manager = Self {
servers: Vec::new(),
routes: HashMap::new(),
statuses: HashMap::new(),
shutting_down: false,
};
for (server_name, config) in mcp_servers {
if let Some(cancellation) = cancellation
&& let Err(error) = cancellation.check()
{
manager.shutdown();
return Err(McpError::transport(error));
}
if !config.enabled() {
continue;
}
if let Err(error) = validate_mcp_server_name(server_name) {
manager.shutdown();
return Err(McpError::Config(format!(
"MCP server '{server_name}' failed during config validation: {error}"
)));
}
let (client, tools) =
match Self::connect_server(server_name, config, mc_home, cancellation) {
Ok(connected) => connected,
Err((phase, error)) => {
manager.shutdown();
return Err(McpError::Config(format!(
"MCP server '{server_name}' failed during {phase}: {error}"
)));
}
};
let server_index = manager.servers.len();
let routes = match manager.validate_routes(server_index, server_name, &tools) {
Ok(routes) => routes,
Err(error) => {
manager.shutdown();
return Err(McpError::Config(format!(
"MCP server '{server_name}' failed during route validation: {error}"
)));
}
};
manager.servers.push(McpServerHandle {
client: Arc::new(client),
tools,
});
manager.routes.extend(routes);
manager.statuses.insert(
server_name.clone(),
McpServerStatus::Connected {
tool_count: manager.servers[server_index].tools.len(),
},
);
}
Ok(manager)
}
fn connect_server(
server_name: &str,
config: &McpServerConfig,
mc_home: Option<&Path>,
cancellation: Option<&AgentCancellation>,
) -> Result<(McpClient, Vec<Tool>), (&'static str, McpError)> {
let connect_phase = match config {
McpServerConfig::Stdio(_) => "spawn",
McpServerConfig::Http(_) => "connect",
};
if let Some(cancellation) = cancellation {
cancellation
.check()
.map_err(|error| (connect_phase, McpError::transport(error)))?;
}
let client = McpClient::connect_named(Some(server_name), config, mc_home)
.map_err(|error| (connect_phase, error))?;
if let Some(cancellation) = cancellation {
cancellation
.check()
.map_err(|error| (connect_phase, McpError::transport(error)))?;
}
match cancellation {
Some(cancellation) => client.initialize_cancellable(cancellation),
None => client.initialize(),
}
.map_err(|error| ("initialize", error))?;
if let Some(cancellation) = cancellation {
cancellation
.check()
.map_err(|error| ("initialize", McpError::transport(error)))?;
}
let tools = match cancellation {
Some(cancellation) => client.list_tools_cancellable(cancellation),
None => client.list_tools(),
}
.map_err(|error| ("list_tools", error))?;
if let Some(cancellation) = cancellation {
cancellation
.check()
.map_err(|error| ("list_tools", McpError::transport(error)))?;
}
Ok((client, tools))
}
fn validate_routes(
&self,
server_index: usize,
server_name: &str,
tools: &[Tool],
) -> McpResult<HashMap<QualifiedMcpToolName, (usize, String)>> {
let mut routes = HashMap::new();
for tool in tools {
let qualified = QualifiedMcpToolName::new(server_name, &tool.name)?;
if self.routes.contains_key(&qualified) || routes.contains_key(&qualified) {
return Err(McpError::Config(format!(
"duplicate MCP tool route '{qualified}'"
)));
}
routes.insert(qualified, (server_index, tool.name.clone()));
}
Ok(routes)
}
fn record_failure(
&mut self,
server_name: &str,
phase: impl Into<String>,
error: impl std::fmt::Display,
) {
self.statuses.insert(
server_name.to_string(),
McpServerStatus::Failed {
error: bounded_error(error.to_string()),
phase: phase.into(),
},
);
}
pub(crate) fn resolve(&self, qualified_name: &str) -> Option<(usize, String)> {
self.routes.get(qualified_name).cloned()
}
pub(crate) fn resolve_tool_call(&self, qualified_name: &str) -> McpResult<ResolvedMcpToolCall> {
if self.shutting_down {
return Err(McpError::Transport(
"MCP manager is shutting down".to_string(),
));
}
let (server_index, raw_tool_name) = self.resolve(qualified_name).ok_or_else(|| {
McpError::Config(format!("unknown MCP tool route '{qualified_name}'"))
})?;
let client = self
.servers
.get(server_index)
.ok_or_else(|| McpError::Config(format!("unknown MCP tool route '{qualified_name}'")))?
.client
.clone();
Ok(ResolvedMcpToolCall {
client,
raw_tool_name,
})
}
pub(crate) fn list_tool_definitions(&self) -> Vec<(QualifiedMcpToolName, Tool)> {
let mut definitions = self
.routes
.iter()
.filter_map(|(qualified, (server_index, _raw_tool_name))| {
let tool = self
.servers
.get(*server_index)?
.tools
.iter()
.find(|tool| tool.name == qualified.tool())?;
Some((qualified.clone(), tool.clone()))
})
.collect::<Vec<_>>();
definitions.sort_by(|left, right| left.0.cmp(&right.0));
definitions
}
pub(crate) fn provider_tool_definitions(&self) -> Vec<Value> {
let mut definitions = self
.list_tool_definitions()
.into_iter()
.map(|(qualified, tool)| provider_tool_definition(&qualified, &tool))
.collect::<Vec<_>>();
definitions.sort_by(|left, right| {
left.get("name")
.and_then(Value::as_str)
.cmp(&right.get("name").and_then(Value::as_str))
});
definitions
}
pub(crate) fn statuses(&self) -> &HashMap<String, McpServerStatus> {
&self.statuses
}
pub(crate) fn shutdown(&mut self) {
self.shutting_down = true;
for server in &self.servers {
server.client.request_shutdown();
}
for server in &self.servers {
server.client.shutdown_shared();
}
}
}
impl Drop for McpManager {
fn drop(&mut self) {
self.shutdown();
}
}
fn bounded_error(mut error: String) -> String {
const MAX_ERROR_BYTES: usize = 512;
if error.len() > MAX_ERROR_BYTES {
let truncate_at = error
.char_indices()
.map(|(index, _)| index)
.take_while(|index| *index <= MAX_ERROR_BYTES)
.last()
.unwrap_or(0);
error.truncate(truncate_at);
error.push_str("...");
}
error
}