use std::sync::Arc;
use async_trait::async_trait;
use rmcp::model::{CallToolRequestParams, ErrorCode};
use rmcp::service::{RoleClient, RunningService, ServiceError};
use serde_json::Value;
use super::{Advertised, ToolClient, ToolError, ToolId};
#[derive(Debug)]
pub struct McpClient {
server: String,
service: Arc<RunningService<RoleClient, ()>>,
}
impl McpClient {
#[must_use]
pub fn new(server: impl Into<String>, service: Arc<RunningService<RoleClient, ()>>) -> Self {
Self {
server: server.into(),
service,
}
}
pub async fn discover(&self) -> Result<Vec<(ToolId, Advertised)>, ToolError> {
let listed = self
.service
.list_all_tools()
.await
.map_err(|e| Self::classify(&ToolId::new(&self.server, "tools/list"), &e))?;
Ok(listed
.into_iter()
.map(|t| {
let annotations = t.annotations.as_ref();
(
ToolId::new(&self.server, t.name.to_string()),
Advertised {
read_only: annotations.and_then(|a| a.read_only_hint),
destructive: annotations.and_then(|a| a.destructive_hint),
idempotent: annotations.and_then(|a| a.idempotent_hint),
},
)
})
.collect())
}
#[allow(clippy::match_same_arms)]
fn classify(tool: &ToolId, e: &ServiceError) -> ToolError {
let detail = e.to_string();
match e {
ServiceError::McpError(err)
if matches!(
err.code,
ErrorCode::METHOD_NOT_FOUND
| ErrorCode::INVALID_PARAMS
| ErrorCode::PARSE_ERROR
) =>
{
ToolError::Refused {
tool: tool.clone(),
detail,
}
}
ServiceError::McpError(_) => ToolError::TimedOut {
tool: tool.clone(),
detail,
},
ServiceError::Timeout { .. } => ToolError::TimedOut {
tool: tool.clone(),
detail,
},
ServiceError::Cancelled { .. } => ToolError::TimedOut {
tool: tool.clone(),
detail,
},
ServiceError::TransportSend(_) | ServiceError::TransportClosed => ToolError::TimedOut {
tool: tool.clone(),
detail,
},
ServiceError::UnexpectedResponse => ToolError::Malformed {
tool: tool.clone(),
detail,
},
_ => ToolError::TimedOut {
tool: tool.clone(),
detail,
},
}
}
}
#[async_trait]
impl ToolClient for McpClient {
async fn call(
&self,
tool: &ToolId,
arguments: &Value,
provenance: Option<&crate::core::Provenance>,
) -> Result<Value, ToolError> {
let object = match arguments {
Value::Object(map) => Some(map.clone()),
Value::Null => None,
other => {
return Err(ToolError::Refused {
tool: tool.clone(),
detail: format!(
"MCP tool arguments must be a JSON object, got {}",
kind_of(other)
),
});
}
};
let mut params = match object {
Some(args) => CallToolRequestParams::new(tool.tool.clone()).with_arguments(args),
None => CallToolRequestParams::new(tool.tool.clone()),
};
if let Some(p) = provenance {
use rmcp::model::RequestParamsMeta;
params.set_meta(rmcp::model::RequestMetaObject(rmcp::model::MetaObject(
p.to_meta(),
)));
}
let result = self
.service
.call_tool(params)
.await
.map_err(|e| Self::classify(tool, &e))?;
if result.is_error == Some(true) {
return Err(ToolError::ToolFailed {
tool: tool.clone(),
detail: render(&result),
});
}
if let Some(structured) = result.structured_content {
return Ok(structured);
}
serde_json::to_value(&result.content).map_err(|error| ToolError::Malformed {
tool: tool.clone(),
detail: format!("MCP tool result content could not be represented: {error}"),
})
}
}
fn kind_of(v: &Value) -> &'static str {
match v {
Value::Null => "null",
Value::Bool(_) => "a boolean",
Value::Number(_) => "a number",
Value::String(_) => "a string",
Value::Array(_) => "an array",
Value::Object(_) => "an object",
}
}
fn render(result: &rmcp::model::CallToolResult) -> String {
result
.content
.iter()
.filter_map(|c| c.as_text().map(|t| t.text.clone()))
.collect::<Vec<_>>()
.join("\n")
}