use std::collections::HashMap;
use async_trait::async_trait;
use embacle::types::RunnerError;
use embacle::{FunctionDeclaration, McpServerConfig, McpToolExecutor};
use rmcp::model::{CallToolRequestParams, CallToolResult};
use rmcp::service::{RoleClient, RunningService};
use rmcp::transport::TokioChildProcess;
use rmcp::ServiceExt;
use serde_json::Value;
use tokio::process::Command;
use tracing::{info, warn};
struct ConnectedServer {
name: String,
service: RunningService<RoleClient, ()>,
}
pub struct McpClientPool {
servers: Vec<ConnectedServer>,
routing: HashMap<String, usize>,
declarations: Vec<FunctionDeclaration>,
}
impl McpClientPool {
pub async fn connect(configs: &[McpServerConfig]) -> Result<Self, RunnerError> {
let mut servers = Vec::with_capacity(configs.len());
let mut routing = HashMap::new();
let mut declarations = Vec::new();
for cfg in configs {
let mut command = Command::new(&cfg.command);
command.args(&cfg.args);
for (key, value) in &cfg.env {
command.env(key, value);
}
let transport = TokioChildProcess::new(command).map_err(|e| {
RunnerError::external_service(
"mcp",
format!("failed to spawn MCP server '{}': {e}", cfg.name),
)
})?;
let service = ().serve(transport).await.map_err(|e| {
RunnerError::external_service(
"mcp",
format!("failed to initialize MCP server '{}': {e}", cfg.name),
)
})?;
let tools = service.list_all_tools().await.map_err(|e| {
RunnerError::external_service(
"mcp",
format!("failed to list tools for MCP server '{}': {e}", cfg.name),
)
})?;
let server_idx = servers.len();
let mut registered = 0_usize;
for tool in tools {
let name = tool.name.to_string();
if routing.contains_key(&name) {
warn!(
tool = %name,
server = %cfg.name,
"Duplicate MCP tool name across servers; keeping first registration, dropping this one"
);
continue;
}
declarations.push(FunctionDeclaration {
name: name.clone(),
description: tool.description.map(|d| d.to_string()).unwrap_or_default(),
parameters: Some(Value::Object((*tool.input_schema).clone())),
});
routing.insert(name, server_idx);
registered += 1;
}
info!(
server = %cfg.name,
command = %cfg.command,
tools = registered,
"Connected to MCP tool server"
);
servers.push(ConnectedServer {
name: cfg.name.clone(),
service,
});
}
info!(
servers = servers.len(),
tools = declarations.len(),
"MCP client pool ready"
);
Ok(Self {
servers,
routing,
declarations,
})
}
pub fn declarations(&self) -> &[FunctionDeclaration] {
&self.declarations
}
pub fn is_empty(&self) -> bool {
self.servers.is_empty()
}
pub fn tool_count(&self) -> usize {
self.declarations.len()
}
}
#[async_trait]
impl McpToolExecutor for McpClientPool {
async fn execute(&self, tool_name: &str, arguments: &Value) -> Result<Value, RunnerError> {
let server_idx = *self.routing.get(tool_name).ok_or_else(|| {
RunnerError::internal(format!(
"no connected MCP server provides tool '{tool_name}'"
))
})?;
let server = &self.servers[server_idx];
let mut params = CallToolRequestParams::new(tool_name.to_owned());
if let Some(object) = arguments.as_object() {
params = params.with_arguments(object.clone());
}
let result = server.service.call_tool(params).await.map_err(|e| {
RunnerError::external_service(
"mcp",
format!("tool '{tool_name}' on server '{}' failed: {e}", server.name),
)
})?;
Ok(call_result_to_json(result))
}
}
fn call_result_to_json(result: CallToolResult) -> Value {
if let Some(structured) = result.structured_content {
return structured;
}
let text = result
.content
.iter()
.filter_map(|c| c.as_text().map(|t| t.text.clone()))
.collect::<Vec<_>>()
.join("\n");
if result.is_error == Some(true) {
return serde_json::json!({ "error": text });
}
serde_json::from_str::<Value>(&text).unwrap_or(Value::String(text))
}
#[cfg(test)]
mod tests {
use rmcp::model::Content;
use super::*;
#[test]
fn call_result_prefers_structured_content() {
let result = CallToolResult::structured(serde_json::json!({"temp": 72}));
let value = call_result_to_json(result);
assert_eq!(value["temp"], 72);
}
#[test]
fn call_result_parses_json_text() {
let result = CallToolResult::success(vec![Content::text(r#"{"ok":true}"#)]);
let value = call_result_to_json(result);
assert_eq!(value["ok"], true);
}
#[test]
fn call_result_falls_back_to_string() {
let result = CallToolResult::success(vec![Content::text("plain text answer")]);
let value = call_result_to_json(result);
assert_eq!(value, Value::String("plain text answer".to_owned()));
}
#[test]
fn call_result_wraps_errors() {
let result = CallToolResult::error(vec![Content::text("boom")]);
let value = call_result_to_json(result);
assert_eq!(value["error"], "boom");
}
}