#![cfg(feature = "mcp")]
use std::collections::HashMap;
use std::path::Path;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, OnceLock};
use modular_agent_core::{AgentContext, AgentError, AgentValue, async_trait};
use rmcp::{
model::{CallToolRequestParam, CallToolResult},
service::ServiceExt,
transport::{ConfigureCommandExt, TokioChildProcess},
};
use serde::Deserialize;
use tokio::process::Command;
use tokio::sync::Mutex as AsyncMutex;
use crate::tool::{Tool, ToolInfo, register_tool};
struct MCPTool {
server_name: String,
server_config: MCPServerConfig,
tool: rmcp::model::Tool,
info: ToolInfo,
}
impl MCPTool {
fn new(
name: String,
server_name: String,
server_config: MCPServerConfig,
tool: rmcp::model::Tool,
) -> Self {
let info = ToolInfo::new(
name,
tool.description.clone().unwrap_or_default().into_owned(),
serde_json::to_value(&tool.input_schema).ok(),
);
Self {
server_name,
server_config,
tool,
info,
}
}
async fn tool_call(
&self,
_ctx: AgentContext,
value: AgentValue,
) -> Result<AgentValue, AgentError> {
let arguments = value.as_object().map(|obj| {
obj.iter()
.map(|(k, v)| {
(
k.clone(),
serde_json::to_value(v).unwrap_or(serde_json::Value::Null),
)
})
.collect::<serde_json::Map<String, serde_json::Value>>()
});
let entry = {
let mut pool = connection_pool().lock().await;
pool.get_or_create(&self.server_name, &self.server_config)
.await?
};
let tool_result = match self.call_once(&entry, arguments.clone()).await {
Ok(result) => result,
Err(e) => {
log::warn!(
"MCP tool call '{}' failed ({}); reconnecting to server '{}' and retrying",
self.tool.name,
e,
self.server_name
);
entry.dead.store(true, Ordering::Release);
let entry = {
let mut pool = connection_pool().lock().await;
pool.invalidate(&self.server_name);
pool.get_or_create(&self.server_name, &self.server_config)
.await?
};
self.call_once(&entry, arguments)
.await
.inspect_err(|_| entry.dead.store(true, Ordering::Release))?
}
};
call_tool_result_to_agent_value(tool_result)
}
async fn call_once(
&self,
entry: &PoolEntry,
arguments: Option<serde_json::Map<String, serde_json::Value>>,
) -> Result<CallToolResult, AgentError> {
let connection = entry.conn.lock().await;
let service = connection.service.as_ref().ok_or_else(|| {
AgentError::Other(format!(
"MCP service for '{}' is not available (tool '{}')",
self.server_name, self.info.name
))
})?;
service
.call_tool(CallToolRequestParam {
name: self.tool.name.clone(),
arguments,
task: None,
})
.await
.map_err(|e| {
AgentError::Other(format!("Failed to call MCP tool '{}': {e}", self.info.name))
})
}
}
#[async_trait]
impl Tool for MCPTool {
fn info(&self) -> &ToolInfo {
&self.info
}
async fn call(&self, ctx: AgentContext, args: AgentValue) -> Result<AgentValue, AgentError> {
self.tool_call(ctx, args).await
}
}
#[derive(Debug, Deserialize)]
pub struct MCPConfig {
#[serde(rename = "mcpServers")]
pub mcp_servers: HashMap<String, MCPServerConfig>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct MCPServerConfig {
pub command: String,
pub args: Vec<String>,
#[serde(default)]
pub env: Option<HashMap<String, String>>,
}
type MCPService = rmcp::service::RunningService<rmcp::service::RoleClient, ()>;
struct MCPConnection {
service: Option<MCPService>,
}
#[derive(Clone)]
struct PoolEntry {
conn: Arc<AsyncMutex<MCPConnection>>,
dead: Arc<AtomicBool>,
}
struct MCPConnectionPool {
connections: HashMap<String, PoolEntry>,
}
impl MCPConnectionPool {
fn new() -> Self {
Self {
connections: HashMap::new(),
}
}
async fn get_or_create(
&mut self,
server_name: &str,
config: &MCPServerConfig,
) -> Result<PoolEntry, AgentError> {
if let Some(entry) = self.connections.get(server_name) {
let service_gone = entry
.conn
.try_lock()
.map(|c| c.service.is_none())
.unwrap_or(false);
if !entry.dead.load(Ordering::Acquire) && !service_gone {
log::debug!("Reusing existing MCP connection for '{}'", server_name);
return Ok(entry.clone());
}
log::info!(
"Discarding dead MCP connection for '{}', creating a new one",
server_name
);
if let Some(stale) = self.connections.remove(server_name) {
cancel_in_background(stale, server_name.to_string());
}
}
log::info!(
"Starting MCP server '{}' (command: {})",
server_name,
config.command
);
let service = ()
.serve(
TokioChildProcess::new(Command::new(&config.command).configure(|cmd| {
for arg in &config.args {
cmd.arg(arg);
}
if let Some(env) = &config.env {
for (key, value) in env {
cmd.env(key, value);
}
}
}))
.map_err(|e| {
log::error!("Failed to start MCP process for '{}': {}", server_name, e);
AgentError::Other(format!(
"Failed to start MCP process for '{}': {e}",
server_name
))
})?,
)
.await
.map_err(|e| {
log::error!("Failed to start MCP service for '{}': {}", server_name, e);
AgentError::Other(format!(
"Failed to start MCP service for '{}': {e}",
server_name
))
})?;
log::info!("Successfully started MCP server '{}'", server_name);
let entry = PoolEntry {
conn: Arc::new(AsyncMutex::new(MCPConnection {
service: Some(service),
})),
dead: Arc::new(AtomicBool::new(false)),
};
self.connections
.insert(server_name.to_string(), entry.clone());
Ok(entry)
}
fn invalidate(&mut self, server_name: &str) {
let is_dead = self
.connections
.get(server_name)
.is_some_and(|e| e.dead.load(Ordering::Acquire));
if is_dead && let Some(entry) = self.connections.remove(server_name) {
log::info!("Invalidating MCP connection for '{}'", server_name);
cancel_in_background(entry, server_name.to_string());
}
}
fn take_all(&mut self) -> Vec<(String, PoolEntry)> {
self.connections.drain().collect()
}
}
fn cancel_in_background(entry: PoolEntry, server_name: String) {
entry.dead.store(true, Ordering::Release);
tokio::spawn(async move {
let mut connection = entry.conn.lock().await;
if let Some(service) = connection.service.take()
&& let Err(e) = service.cancel().await
{
log::warn!("Failed to cancel MCP service '{}': {}", server_name, e);
}
});
}
static CONNECTION_POOL: OnceLock<AsyncMutex<MCPConnectionPool>> = OnceLock::new();
fn connection_pool() -> &'static AsyncMutex<MCPConnectionPool> {
CONNECTION_POOL.get_or_init(|| AsyncMutex::new(MCPConnectionPool::new()))
}
pub async fn shutdown_all_mcp_connections() -> Result<(), AgentError> {
log::info!("Shutting down all MCP server connections");
let entries = { connection_pool().lock().await.take_all() };
for (name, entry) in entries {
entry.dead.store(true, Ordering::Release);
let conn = entry.conn.clone();
match tokio::time::timeout(std::time::Duration::from_secs(5), conn.lock()).await {
Ok(mut connection) => {
if let Some(service) = connection.service.take() {
if let Err(e) = service.cancel().await {
log::error!("Failed to cancel MCP service '{}': {}", name, e);
} else {
log::debug!("Successfully shut down MCP server '{}'", name);
}
}
}
Err(_) => {
log::warn!(
"MCP connection '{}' busy during shutdown; cancelling in background",
name
);
cancel_in_background(entry, name);
}
}
}
log::info!("All MCP server connections shut down");
Ok(())
}
async fn register_tools_from_server(
server_name: String,
server_config: MCPServerConfig,
) -> Result<Vec<String>, AgentError> {
log::debug!("Registering tools from MCP server '{}'", server_name);
let entry = {
let mut pool = connection_pool().lock().await;
pool.get_or_create(&server_name, &server_config).await?
};
log::debug!("Listing tools from MCP server '{}'", server_name);
let tools_list = {
let connection = entry.conn.lock().await;
let service = connection.service.as_ref().ok_or_else(|| {
log::error!("MCP service for '{}' is not available", server_name);
AgentError::Other(format!(
"MCP service for '{}' is not available",
server_name
))
})?;
service.list_tools(Default::default()).await.map_err(|e| {
log::error!("Failed to list MCP tools for '{}': {}", server_name, e);
AgentError::Other(format!(
"Failed to list MCP tools for '{}': {e}",
server_name
))
})?
};
let mut registered_tool_names = Vec::new();
for tool_info in tools_list.tools {
let mcp_tool_name = format!("{}::{}", server_name, tool_info.name);
registered_tool_names.push(mcp_tool_name.clone());
register_tool(MCPTool::new(
mcp_tool_name.clone(),
server_name.clone(),
server_config.clone(),
tool_info,
));
log::debug!("Registered MCP tool '{}'", mcp_tool_name);
}
log::info!(
"Registered {} tools from MCP server '{}'",
registered_tool_names.len(),
server_name
);
Ok(registered_tool_names)
}
pub async fn register_tools_from_mcp_json<P: AsRef<Path>>(
json_path: P,
) -> Result<Vec<String>, AgentError> {
let path = json_path.as_ref();
log::info!("Loading MCP configuration from: {}", path.display());
let json_content = std::fs::read_to_string(path).map_err(|e| {
log::error!("Failed to read MCP config file '{}': {}", path.display(), e);
AgentError::Other(format!("Failed to read MCP config file: {e}"))
})?;
let config: MCPConfig = serde_json::from_str(&json_content).map_err(|e| {
log::error!("Failed to parse MCP config JSON: {}", e);
AgentError::Other(format!("Failed to parse MCP config JSON: {e}"))
})?;
log::info!("Found {} MCP servers in config", config.mcp_servers.len());
let mut registered_tool_names = Vec::new();
for (server_name, server_config) in config.mcp_servers {
let tools = register_tools_from_server(server_name, server_config).await?;
registered_tool_names.extend(tools);
}
log::info!(
"Successfully registered {} MCP tools total",
registered_tool_names.len()
);
Ok(registered_tool_names)
}
fn call_tool_result_to_agent_value(result: CallToolResult) -> Result<AgentValue, AgentError> {
let mut contents = Vec::new();
for c in result.content.iter() {
match &c.raw {
rmcp::model::RawContent::Text(text) => {
contents.push(AgentValue::string(text.text.clone()));
}
_ => {
}
}
}
let data = AgentValue::array(contents.into());
if result.is_error == Some(true) {
return Err(AgentError::Other(
serde_json::to_string(&data).map_err(|e| AgentError::InvalidValue(e.to_string()))?,
));
}
Ok(data)
}
#[cfg(test)]
mod tests {
use super::*;
fn test_entry(dead: bool) -> PoolEntry {
PoolEntry {
conn: Arc::new(AsyncMutex::new(MCPConnection { service: None })),
dead: Arc::new(AtomicBool::new(dead)),
}
}
fn bogus_config() -> MCPServerConfig {
MCPServerConfig {
command: "modular-agent-test-nonexistent-command".to_string(),
args: Vec::new(),
env: None,
}
}
#[tokio::test]
async fn invalidate_removes_dead_entry() {
let mut pool = MCPConnectionPool::new();
pool.connections.insert("s".to_string(), test_entry(true));
pool.invalidate("s");
assert!(pool.connections.is_empty());
}
#[tokio::test]
async fn invalidate_keeps_entry_not_marked_dead() {
let mut pool = MCPConnectionPool::new();
pool.connections.insert("s".to_string(), test_entry(false));
pool.invalidate("s");
assert!(pool.connections.contains_key("s"));
}
#[tokio::test]
async fn get_or_create_reuses_busy_connection() {
let mut pool = MCPConnectionPool::new();
let entry = test_entry(false);
pool.connections.insert("s".to_string(), entry.clone());
let _guard = entry.conn.try_lock().unwrap();
let got = pool.get_or_create("s", &bogus_config()).await.unwrap();
assert!(Arc::ptr_eq(&got.conn, &entry.conn));
}
#[tokio::test]
async fn get_or_create_discards_dead_entry() {
let mut pool = MCPConnectionPool::new();
pool.connections.insert("s".to_string(), test_entry(true));
assert!(pool.get_or_create("s", &bogus_config()).await.is_err());
assert!(pool.connections.is_empty());
}
#[tokio::test]
async fn get_or_create_discards_entry_with_missing_service() {
let mut pool = MCPConnectionPool::new();
pool.connections.insert("s".to_string(), test_entry(false));
assert!(pool.get_or_create("s", &bogus_config()).await.is_err());
assert!(pool.connections.is_empty());
}
}