use async_trait::async_trait;
use mofa_kernel::agent::components::mcp::{
McpClient, McpServerConfig, McpServerInfo, McpToolInfo, McpTransportConfig,
};
use mofa_kernel::agent::error::{AgentError, AgentResult};
use rmcp::model::{CallToolRequestParams, ClientCapabilities, ClientInfo, Implementation};
use rmcp::service::{RoleClient, RunningService};
use rmcp::transport::TokioChildProcess;
use rmcp::ServiceExt;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::process::Command;
use tokio::sync::RwLock;
use tracing;
struct McpConnection {
service: RunningService<RoleClient, ClientInfo>,
config: McpServerConfig,
}
pub struct McpClientManager {
connections: HashMap<String, McpConnection>,
}
impl McpClientManager {
pub fn new() -> Self {
Self {
connections: HashMap::new(),
}
}
fn get_connection(&self, server_name: &str) -> AgentResult<&McpConnection> {
self.connections
.get(server_name)
.ok_or_else(|| AgentError::ToolNotFound(
format!("MCP server '{}' not connected", server_name),
))
}
pub fn into_shared(self) -> Arc<RwLock<Self>> {
Arc::new(RwLock::new(self))
}
}
impl Default for McpClientManager {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl McpClient for McpClientManager {
async fn connect(&mut self, config: McpServerConfig) -> AgentResult<()> {
let server_name = config.name.clone();
if self.connections.contains_key(&server_name) {
return Err(AgentError::ConfigError(
format!("MCP server '{}' is already connected", server_name),
));
}
tracing::info!("Connecting to MCP server '{}'...", server_name);
let service = match &config.transport {
McpTransportConfig::Stdio { command, args, env } => {
let mut cmd = Command::new(command);
cmd.args(args);
for (key, value) in env {
cmd.env(key, value);
}
let transport = TokioChildProcess::new(cmd).map_err(|e| {
AgentError::InitializationFailed(
format!("Failed to start MCP server process '{}': {}", server_name, e),
)
})?;
let client_info = ClientInfo {
meta: None,
protocol_version: Default::default(),
capabilities: ClientCapabilities::default(),
client_info: Implementation {
name: "mofa-agent".to_string(),
version: "0.1.0".to_string(),
title: None,
description: None,
icons: None,
website_url: None,
},
};
client_info.serve(transport).await.map_err(|e| {
AgentError::InitializationFailed(
format!("Failed to initialize MCP session with '{}': {}", server_name, e),
)
})?
}
McpTransportConfig::Http { url: _ } => {
return Err(AgentError::ConfigError(
"HTTP transport is not yet supported. Use Stdio transport instead.".to_string(),
));
}
};
tracing::info!("Connected to MCP server '{}'", server_name);
self.connections.insert(
server_name,
McpConnection { service, config },
);
Ok(())
}
async fn disconnect(&mut self, server_name: &str) -> AgentResult<()> {
if let Some(connection) = self.connections.remove(server_name) {
tracing::info!("Disconnecting from MCP server '{}'...", server_name);
connection.service.cancel().await.map_err(|e| {
AgentError::ShutdownFailed(
format!("Failed to disconnect from MCP server '{}': {:?}", server_name, e),
)
})?;
tracing::info!("Disconnected from MCP server '{}'", server_name);
Ok(())
} else {
Err(AgentError::ToolNotFound(
format!("MCP server '{}' not connected", server_name),
))
}
}
async fn list_tools(&self, server_name: &str) -> AgentResult<Vec<McpToolInfo>> {
let connection = self.get_connection(server_name)?;
let result = connection
.service
.peer()
.list_tools(None)
.await
.map_err(|e| {
AgentError::ExecutionFailed(
format!("Failed to list tools from MCP server '{}': {}", server_name, e),
)
})?;
let tools = result
.tools
.into_iter()
.map(|tool| McpToolInfo {
name: tool.name.to_string(),
description: tool.description.unwrap_or_default().to_string(),
input_schema: serde_json::to_value(&tool.input_schema)
.unwrap_or(serde_json::json!({})),
})
.collect();
Ok(tools)
}
async fn call_tool(
&self,
server_name: &str,
tool_name: &str,
arguments: serde_json::Value,
) -> AgentResult<serde_json::Value> {
let connection = self.get_connection(server_name)?;
let params = CallToolRequestParams {
name: tool_name.to_string().into(),
arguments: if arguments.is_object() {
Some(arguments.as_object().unwrap().clone())
} else {
Some(serde_json::Map::new())
},
meta: None,
task: None,
};
let result = connection
.service
.peer()
.call_tool(params)
.await
.map_err(|e| {
AgentError::ExecutionFailed(
format!("MCP tool call '{}' on server '{}' failed: {}", tool_name, server_name, e),
)
})?;
let content_values: Vec<serde_json::Value> = result
.content
.iter()
.map(|content| {
serde_json::to_value(content).unwrap_or(serde_json::json!({"error": "serialization failed"}))
})
.collect();
if result.is_error.unwrap_or(false) {
let error_text = content_values
.first()
.and_then(|v| v.get("text").and_then(|t| t.as_str()))
.unwrap_or("Unknown MCP error")
.to_string();
return Err(AgentError::ExecutionFailed(error_text));
}
Ok(serde_json::json!({
"content": content_values,
}))
}
async fn server_info(&self, server_name: &str) -> AgentResult<McpServerInfo> {
let connection = self.get_connection(server_name)?;
let peer = connection.service.peer();
let server_info = peer.peer_info().ok_or_else(|| {
AgentError::ExecutionFailed(format!("No server info available for '{}'", server_name))
})?;
Ok(McpServerInfo {
name: server_info.server_info.name.clone(),
version: server_info.server_info.version.clone(),
instructions: server_info.instructions.clone(),
})
}
fn connected_servers(&self) -> Vec<String> {
self.connections.keys().cloned().collect()
}
fn is_connected(&self, server_name: &str) -> bool {
self.connections.contains_key(server_name)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_client_manager_new() {
let manager = McpClientManager::new();
assert!(manager.connected_servers().is_empty());
assert!(!manager.is_connected("nonexistent"));
}
#[test]
fn test_client_manager_default() {
let manager = McpClientManager::default();
assert!(manager.connected_servers().is_empty());
}
#[test]
fn test_into_shared() {
let manager = McpClientManager::new();
let _shared = manager.into_shared();
}
#[tokio::test]
async fn test_get_connection_missing() {
let manager = McpClientManager::new();
let result = manager.list_tools("nonexistent").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_disconnect_missing() {
let mut manager = McpClientManager::new();
let result = manager.disconnect("nonexistent").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_connect_duplicate() {
let mut manager = McpClientManager::new();
let config = McpServerConfig::stdio("test", "nonexistent-command-xyz", vec![]);
let _ = manager.connect(config).await;
let http_config = McpServerConfig::http("http-test", "http://localhost:9999");
let result = manager.connect(http_config).await;
assert!(result.is_err()); }
}