use crate::client::McpClient;
use crate::transport::McpConnection;
use anyhow::{Result, anyhow};
use async_trait::async_trait;
use everruns_core::mcp_server::sanitize_mcp_server_name;
use everruns_core::{AgentLoopError, McpToolInvoker, ToolCall, ToolResult, parse_mcp_tool_name};
use std::collections::HashMap;
use std::sync::Arc;
#[async_trait]
pub trait McpConnectionResolver: Send + Sync {
async fn resolve(&self, server_prefix: &str) -> Result<Option<McpConnection>>;
}
#[derive(Default)]
pub struct StaticConnectionResolver {
connections: HashMap<String, McpConnection>,
}
impl StaticConnectionResolver {
pub fn new() -> Self {
Self::default()
}
pub fn insert(&mut self, connection: McpConnection) {
let key = sanitize_mcp_server_name(&connection.name);
self.connections.insert(key, connection);
}
pub fn with(mut self, connection: McpConnection) -> Self {
self.insert(connection);
self
}
pub fn from_connections(connections: impl IntoIterator<Item = McpConnection>) -> Self {
let mut resolver = Self::new();
for connection in connections {
resolver.insert(connection);
}
resolver
}
pub fn is_empty(&self) -> bool {
self.connections.is_empty()
}
}
#[async_trait]
impl McpConnectionResolver for StaticConnectionResolver {
async fn resolve(&self, server_prefix: &str) -> Result<Option<McpConnection>> {
Ok(self.connections.get(server_prefix).cloned())
}
}
pub struct McpExecutor {
client: Arc<McpClient>,
resolver: Arc<dyn McpConnectionResolver>,
}
impl McpExecutor {
pub fn new(client: Arc<McpClient>, resolver: Arc<dyn McpConnectionResolver>) -> Self {
Self { client, resolver }
}
pub async fn execute_mcp_tool(&self, tool_call: &ToolCall) -> Result<ToolResult> {
let (server_prefix, original_tool_name) = parse_mcp_tool_name(&tool_call.name)
.ok_or_else(|| anyhow!("Invalid MCP tool name: {}", tool_call.name))?;
let connection = self
.resolver
.resolve(&server_prefix)
.await?
.ok_or_else(|| anyhow!("MCP server not found for prefix: {server_prefix}"))?;
if let Some(provider) = &connection.pending_oauth_provider {
return Ok(ToolResult {
tool_call_id: tool_call.id.clone(),
result: None,
images: None,
error: Some(format!(
"MCP server '{}' requires an OAuth connection. \
Ask the user to connect provider '{provider}'.",
connection.name
)),
connection_required: Some(provider.clone()),
raw_output: None,
});
}
self.client
.call_as_tool_result(
&connection,
tool_call.id.clone(),
&original_tool_name,
tool_call.arguments.clone(),
)
.await
}
}
#[async_trait]
impl McpToolInvoker for McpExecutor {
async fn invoke(&self, tool_call: &ToolCall) -> everruns_core::Result<ToolResult> {
self.execute_mcp_tool(tool_call).await.map_err(|e| {
tracing::error!(error = %e, "MCP tool execution failed");
AgentLoopError::tool(e.to_string())
})
}
}