use crate::mcp::client::McpClient;
use crate::mcp::oauth;
use crate::mcp::protocol::{
CallToolResult, McpServerConfig, McpTool, McpTransportConfig, OAuthConfig,
};
use crate::mcp::transport::http_sse::HttpSseTransport;
use crate::mcp::transport::stdio::StdioTransport;
use crate::mcp::transport::streamable_http::StreamableHttpTransport;
use crate::mcp::transport::McpTransport;
use anyhow::{anyhow, Result};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
pub use crate::mcp::result::tool_result_to_string;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct McpServerStatus {
pub name: String,
pub connected: bool,
pub enabled: bool,
pub tool_count: usize,
pub error: Option<String>,
}
pub struct McpManager {
clients: RwLock<HashMap<String, Arc<McpClient>>>,
configs: RwLock<HashMap<String, McpServerConfig>>,
connect_errors: RwLock<HashMap<String, String>>,
last_used_at_ms: RwLock<HashMap<String, u64>>,
}
impl McpManager {
pub fn new() -> Self {
Self {
clients: RwLock::new(HashMap::new()),
configs: RwLock::new(HashMap::new()),
connect_errors: RwLock::new(HashMap::new()),
last_used_at_ms: RwLock::new(HashMap::new()),
}
}
pub async fn register_server(&self, config: McpServerConfig) {
let name = config.name.clone();
let mut configs = self.configs.write().await;
configs.insert(name.clone(), config);
tracing::info!("Registered MCP server: {}", name);
}
pub async fn connect(&self, name: &str) -> Result<()> {
let result = self.do_connect(name).await;
match &result {
Ok(_) => {
self.connect_errors.write().await.remove(name);
}
Err(e) => {
self.connect_errors
.write()
.await
.insert(name.to_string(), e.to_string());
}
}
result
}
async fn do_connect(&self, name: &str) -> Result<()> {
let config = {
let configs = self.configs.read().await;
configs
.get(name)
.cloned()
.ok_or_else(|| anyhow!("MCP server not found: {}", name))?
};
let (client, tools) = connect_ready_client(&config).await?;
tracing::info!("MCP server '{}' connected with {} tools", name, tools.len());
{
let mut clients = self.clients.write().await;
clients.insert(name.to_string(), client);
}
self.last_used_at_ms
.write()
.await
.insert(name.to_string(), now_epoch_ms());
Ok(())
}
pub async fn disconnect(&self, name: &str) -> Result<()> {
let client = {
let mut clients = self.clients.write().await;
clients.remove(name)
};
self.last_used_at_ms.write().await.remove(name);
if let Some(client) = client {
client.close().await?;
tracing::info!("MCP server '{}' disconnected", name);
}
Ok(())
}
pub async fn remove_server(&self, name: &str) -> Result<bool> {
let client = self.clients.write().await.remove(name);
let had_error = self.connect_errors.write().await.remove(name).is_some();
let had_timestamp = self.last_used_at_ms.write().await.remove(name).is_some();
let had_config = self.configs.write().await.remove(name).is_some();
let removed = client.is_some() || had_error || had_timestamp || had_config;
if let Some(client) = client {
client.close().await?;
tracing::info!("MCP server '{}' removed", name);
}
Ok(removed)
}
pub async fn contains_server(&self, name: &str) -> bool {
self.configs.read().await.contains_key(name)
}
#[cfg(test)]
pub(crate) async fn insert_client_for_test(&self, name: &str, client: Arc<McpClient>) {
self.clients.write().await.insert(name.to_string(), client);
self.last_used_at_ms
.write()
.await
.insert(name.to_string(), now_epoch_ms());
}
pub async fn last_used_at_ms(&self, name: &str) -> Option<u64> {
self.last_used_at_ms.read().await.get(name).copied()
}
pub async fn touch(&self, name: &str) {
self.last_used_at_ms
.write()
.await
.insert(name.to_string(), now_epoch_ms());
}
pub async fn disconnect_idle(&self, idle_threshold_ms: u64) -> Vec<String> {
let cutoff = now_epoch_ms().saturating_sub(idle_threshold_ms);
let candidates: Vec<String> = {
let clients = self.clients.read().await;
let last_used = self.last_used_at_ms.read().await;
clients
.keys()
.filter(|name| match last_used.get(*name) {
Some(ts) => *ts < cutoff,
None => true,
})
.cloned()
.collect()
};
let mut disconnected = Vec::with_capacity(candidates.len());
for name in candidates {
match self.disconnect(&name).await {
Ok(()) => disconnected.push(name),
Err(e) => tracing::warn!(
server = %name,
error = %e,
"MCP idle disconnect failed; entry already removed from registry"
),
}
}
{
let clients = self.clients.read().await;
self.last_used_at_ms
.write()
.await
.retain(|name, _| clients.contains_key(name));
}
disconnected
}
pub async fn all_configs(&self) -> Vec<McpServerConfig> {
self.configs.read().await.values().cloned().collect()
}
pub async fn get_all_tools(&self) -> Vec<(String, McpTool)> {
let clients = self.clients.read().await;
let mut all_tools = Vec::new();
for (server_name, client) in clients.iter() {
let tools = client.get_cached_tools().await;
for tool in tools {
all_tools.push((server_name.clone(), tool));
}
}
all_tools
}
pub async fn call_tool(
&self,
full_name: &str,
arguments: Option<serde_json::Value>,
) -> Result<CallToolResult> {
let (server_name, tool_name) = Self::parse_tool_name(full_name)?;
let client = {
let clients = self.clients.read().await;
clients
.get(&server_name)
.cloned()
.ok_or_else(|| anyhow!("MCP server not connected: {}", server_name))?
};
self.last_used_at_ms
.write()
.await
.insert(server_name.clone(), now_epoch_ms());
client.call_tool(&tool_name, arguments).await
}
async fn resolve_auth_header(oauth: Option<&OAuthConfig>) -> Result<Option<(String, String)>> {
let Some(oauth) = oauth else {
return Ok(None);
};
let token = if let Some(static_token) = &oauth.access_token {
static_token.clone()
} else {
oauth::exchange_client_credentials(
&oauth.token_url,
&oauth.client_id,
oauth.client_secret.as_deref().unwrap_or(""),
&oauth.scopes,
)
.await?
};
Ok(Some((
"Authorization".to_string(),
format!("Bearer {}", token),
)))
}
fn parse_tool_name(full_name: &str) -> Result<(String, String)> {
if !full_name.starts_with("mcp__") {
return Err(anyhow!("Invalid MCP tool name: {}", full_name));
}
let rest = &full_name[5..]; let parts: Vec<&str> = rest.splitn(2, "__").collect();
if parts.len() != 2 {
return Err(anyhow!("Invalid MCP tool name format: {}", full_name));
}
Ok((parts[0].to_string(), parts[1].to_string()))
}
pub async fn get_status(&self) -> HashMap<String, McpServerStatus> {
let configs = self.configs.read().await;
let clients = self.clients.read().await;
let errors = self.connect_errors.read().await;
let mut status = HashMap::new();
for (name, config) in configs.iter() {
let client = clients.get(name);
let (connected, tool_count) = if let Some(c) = client {
(c.is_connected(), c.get_cached_tools().await.len())
} else {
(false, 0)
};
status.insert(
name.clone(),
McpServerStatus {
name: name.clone(),
connected,
enabled: config.enabled,
tool_count,
error: errors.get(name).cloned(),
},
);
}
status
}
pub async fn get_client(&self, name: &str) -> Option<Arc<McpClient>> {
let clients = self.clients.read().await;
clients.get(name).cloned()
}
pub async fn is_connected(&self, name: &str) -> bool {
let clients = self.clients.read().await;
clients.get(name).map(|c| c.is_connected()).unwrap_or(false)
}
pub async fn list_connected(&self) -> Vec<String> {
let clients = self.clients.read().await;
clients.keys().cloned().collect()
}
pub async fn get_server_tools(&self, name: &str) -> Vec<McpTool> {
let clients = self.clients.read().await;
match clients.get(name) {
Some(client) => client.get_cached_tools().await,
None => Vec::new(),
}
}
}
pub(crate) async fn connect_ready_client(
config: &McpServerConfig,
) -> Result<(Arc<McpClient>, Vec<McpTool>)> {
if !config.enabled {
return Err(anyhow!("MCP server is disabled: {}", config.name));
}
let auth_header = McpManager::resolve_auth_header(config.oauth.as_ref()).await?;
let transport: Arc<dyn McpTransport> = match &config.transport {
McpTransportConfig::Stdio { command, args } => Arc::new(
StdioTransport::spawn_with_timeout(
command,
args,
&config.env,
config.tool_timeout_secs,
)
.await?,
),
McpTransportConfig::Http { url, headers } => {
let mut merged = headers.clone();
if let Some((key, value)) = &auth_header {
merged.insert(key.clone(), value.clone());
}
Arc::new(
HttpSseTransport::connect_with_timeout(url, merged, config.tool_timeout_secs)
.await?,
)
}
McpTransportConfig::StreamableHttp { url, headers } => {
let mut merged = headers.clone();
if let Some((key, value)) = &auth_header {
merged.insert(key.clone(), value.clone());
}
Arc::new(
StreamableHttpTransport::connect_with_timeout(
url,
merged,
config.tool_timeout_secs,
)
.await?,
)
}
};
let client = Arc::new(McpClient::new(config.name.clone(), transport));
if let Err(error) = client.initialize().await {
if let Err(close_error) = client.close().await {
tracing::warn!(
server = %config.name,
error = %close_error,
"Failed to close MCP transport after initialize failure"
);
}
return Err(error);
}
let tools = match client.list_tools().await {
Ok(tools) => tools,
Err(error) => {
if let Err(close_error) = client.close().await {
tracing::warn!(
server = %config.name,
error = %close_error,
"Failed to close MCP transport after tool discovery failure"
);
}
return Err(error);
}
};
Ok((client, tools))
}
impl Default for McpManager {
fn default() -> Self {
Self::new()
}
}
fn now_epoch_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
#[cfg(test)]
#[path = "manager/tests.rs"]
mod tests;