use std::sync::Arc;
use async_trait::async_trait;
use serde_json::{Value, json};
use crate::tool::{
ParallelToolContext, RuntimeToolDescriptor, ToolApprovalCategory, ToolCapability,
ToolDefinition, ToolDurability, ToolExecutionCategory, ToolExecutor, ToolResult,
ToolSideEffectLevel,
};
use super::client::McpStdioClient;
use super::protocol::{McpToolCallResult, McpToolDefinition};
use super::sse::client::McpSseClient;
use super::streamable_http::client::McpStreamableHttpClient;
#[async_trait]
pub trait McpToolClient: sealed::Sealed + Send + Sync {
async fn call_tool(
&self,
tool_name: &str,
arguments: Option<Value>,
) -> Result<McpToolCallResult, String>;
async fn shutdown(&self);
}
mod sealed {
pub trait Sealed {}
impl Sealed for super::McpStdioClient {}
impl Sealed for super::McpSseClient {}
impl Sealed for super::McpStreamableHttpClient {}
#[cfg(test)]
impl Sealed for crate::mcp::tests::SuccessfulMcpClient {}
}
#[async_trait]
impl McpToolClient for McpStdioClient {
async fn call_tool(
&self,
tool_name: &str,
arguments: Option<Value>,
) -> Result<McpToolCallResult, String> {
McpStdioClient::call_tool(self, tool_name, arguments)
.await
.map_err(|error| error.to_string())
}
async fn shutdown(&self) {
McpStdioClient::shutdown(self).await
}
}
#[async_trait]
impl McpToolClient for McpSseClient {
async fn call_tool(
&self,
tool_name: &str,
arguments: Option<Value>,
) -> Result<McpToolCallResult, String> {
McpSseClient::call_tool(self, tool_name, arguments)
.await
.map_err(|error| error.to_string())
}
async fn shutdown(&self) {
McpSseClient::shutdown(self).await
}
}
#[async_trait]
impl McpToolClient for McpStreamableHttpClient {
async fn call_tool(
&self,
tool_name: &str,
arguments: Option<Value>,
) -> Result<McpToolCallResult, String> {
McpStreamableHttpClient::call_tool(self, tool_name, arguments)
.await
.map_err(|error| error.to_string())
}
async fn shutdown(&self) {
McpStreamableHttpClient::shutdown(self).await
}
}
const MCP_TOOL_PREFIX: &str = "mcp__";
pub fn mcp_tool_name(server_name: &str, tool_name: &str) -> String {
format!("{MCP_TOOL_PREFIX}{server_name}__{tool_name}")
}
pub fn parse_mcp_tool_name(name: &str) -> Option<(&str, &str)> {
let rest = name.strip_prefix(MCP_TOOL_PREFIX)?;
let (server, tool) = rest.split_once("__")?;
Some((server, tool))
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum McpServerNameError {
#[error("MCP server name must not be empty")]
Empty,
#[error(
"MCP server name {0:?} must not contain \"__\", which mcp_tool_name uses as the \
separator between the server name and the tool name"
)]
ContainsDoubleUnderscore(String),
#[error(
"MCP server name {0:?} must not end with \"_\", which would merge with the \"__\" \
separator mcp_tool_name inserts after it"
)]
EndsWithUnderscore(String),
}
pub fn validate_mcp_server_name(server_name: &str) -> Result<(), McpServerNameError> {
if server_name.is_empty() {
return Err(McpServerNameError::Empty);
}
if server_name.contains("__") {
return Err(McpServerNameError::ContainsDoubleUnderscore(
server_name.to_string(),
));
}
if server_name.ends_with('_') {
return Err(McpServerNameError::EndsWithUnderscore(
server_name.to_string(),
));
}
Ok(())
}
pub struct McpBridgedTool {
server_name: String,
tool_def: McpToolDefinition,
client: Arc<dyn McpToolClient>,
}
impl McpBridgedTool {
pub fn new(
server_name: String,
tool_def: McpToolDefinition,
client: Arc<dyn McpToolClient>,
) -> Self {
Self {
server_name,
tool_def,
client,
}
}
fn full_name(&self) -> String {
mcp_tool_name(&self.server_name, &self.tool_def.name)
}
}
impl std::fmt::Debug for McpBridgedTool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("McpBridgedTool")
.field("name", &self.full_name())
.finish_non_exhaustive()
}
}
impl ToolDefinition for McpBridgedTool {
fn descriptor(&self) -> RuntimeToolDescriptor {
let description = self.tool_def.description.clone().unwrap_or_default();
let input_schema = self
.tool_def
.input_schema
.clone()
.unwrap_or_else(|| json!({"type": "object", "properties": {}}));
RuntimeToolDescriptor::builder(self.full_name())
.description(description)
.input_schema(input_schema)
.capability(ToolCapability::Custom(format!("mcp:{}", self.server_name)))
.side_effect_level(ToolSideEffectLevel::External)
.durability(ToolDurability::Ephemeral)
.execution_category(ToolExecutionCategory::ExclusiveLocalMutation)
.approval_category(ToolApprovalCategory::Process)
.build()
}
}
#[async_trait]
impl ToolExecutor for McpBridgedTool {
async fn execute(&self, _ctx: ParallelToolContext, input: Value) -> ToolResult {
let arguments = if input.is_null()
|| (input.is_object() && input.as_object().is_none_or(|o| o.is_empty()))
{
None
} else {
Some(input)
};
let result = self
.client
.call_tool(&self.tool_def.name, arguments)
.await
.map_err(|error| format!("MCP tool call failed: {error}"))?;
let mut output = String::new();
for block in &result.content {
if let Some(text) = &block.text {
if !output.is_empty() {
output.push('\n');
}
output.push_str(text);
}
}
if result.is_error {
Err(output)
} else {
Ok(output)
}
}
}
#[cfg(test)]
mod name_tests {
use super::*;
#[test]
fn empty_server_name_is_rejected() {
assert_eq!(validate_mcp_server_name(""), Err(McpServerNameError::Empty));
}
#[test]
fn server_name_containing_double_underscore_is_rejected() {
assert_eq!(
validate_mcp_server_name("evil__foo"),
Err(McpServerNameError::ContainsDoubleUnderscore(
"evil__foo".to_string()
))
);
}
#[test]
fn server_name_ending_in_underscore_is_rejected() {
assert_eq!(
validate_mcp_server_name("evil_"),
Err(McpServerNameError::EndsWithUnderscore("evil_".to_string()))
);
}
#[test]
fn ordinary_server_name_is_accepted() {
assert_eq!(validate_mcp_server_name("evil"), Ok(()));
}
#[test]
fn colliding_names_from_the_issue_no_longer_collide() {
assert!(validate_mcp_server_name("evil__foo").is_err());
assert_eq!(
parse_mcp_tool_name(&mcp_tool_name("evil", "foo__real_tool")),
Some(("evil", "foo__real_tool"))
);
assert!(validate_mcp_server_name("evil_").is_err());
assert_eq!(
parse_mcp_tool_name(&mcp_tool_name("evil", "__thing")),
Some(("evil", "__thing"))
);
}
#[test]
fn round_trip_preserves_tool_names_containing_double_underscore_or_leading_underscore() {
for (server, tool) in [
("obs", "search__logs"),
("obs", "__internal"),
("obs", "logs"),
] {
validate_mcp_server_name(server).expect("valid server name");
let encoded = mcp_tool_name(server, tool);
assert_eq!(parse_mcp_tool_name(&encoded), Some((server, tool)));
}
}
}