use crate::error::ToolError;
use crate::tool::{ToolDefinition, ToolDyn};
use rmcp::{
model::{CallToolRequestParams, Content, Tool},
service::RunningService,
transport::streamable_http_client::{
StreamableHttpClientTransportConfig, StreamableHttpClientWorker,
},
RoleClient, ServiceExt,
};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::RwLock;
type McpClient = RunningService<RoleClient, ()>;
#[derive(Clone, Debug)]
pub struct McpToolSource {
endpoint: String,
tools: Arc<RwLock<Vec<ToolDefinition>>>,
client: Arc<RwLock<Option<McpClient>>>,
}
impl McpToolSource {
pub async fn connect(endpoint: &str) -> Result<Self, ToolError> {
let config = StreamableHttpClientTransportConfig::with_uri(endpoint);
let transport = StreamableHttpClientWorker::new(reqwest::Client::default(), config);
let client: McpClient = ().serve(transport).await.map_err(|e| {
ToolError::McpError(format!(
"Failed to connect to MCP server at {}: {}",
endpoint, e
))
})?;
let source = Self {
endpoint: endpoint.to_string(),
tools: Arc::new(RwLock::new(Vec::new())),
client: Arc::new(RwLock::new(Some(client))),
};
Ok(source)
}
pub async fn is_connected(&self) -> bool {
self.client.read().await.is_some()
}
pub async fn discover(&self) -> Result<Vec<ToolDefinition>, ToolError> {
let client_guard = self.client.read().await;
let client = client_guard
.as_ref()
.ok_or_else(|| ToolError::McpError("Not connected to MCP server".into()))?;
let response = client
.list_tools(Default::default())
.await
.map_err(|e| ToolError::McpError(format!("Failed to list tools: {}", e)))?;
let definitions: Vec<ToolDefinition> = response
.tools
.into_iter()
.map(|tool| convert_mcp_tool_to_definition(&tool))
.collect();
drop(client_guard);
*self.tools.write().await = definitions.clone();
Ok(definitions)
}
pub async fn tools(&self) -> Vec<ToolDefinition> {
self.tools.read().await.clone()
}
pub async fn call_tool(
&self,
name: &str,
args: serde_json::Value,
) -> Result<String, ToolError> {
let client_guard = self.client.read().await;
let client = client_guard
.as_ref()
.ok_or_else(|| ToolError::McpError("Not connected to MCP server".into()))?;
let mut params = CallToolRequestParams::new(name.to_string());
if let Some(obj) = args.as_object().cloned() {
params = params.with_arguments(obj);
}
let result = client
.call_tool(params)
.await
.map_err(|e| ToolError::McpError(format!("Tool call failed: {}", e)))?;
if result.is_error.unwrap_or(false) {
let error_text = extract_text_from_content(&result.content);
return Err(ToolError::ExecutionFailed(error_text));
}
let output = extract_text_from_content(&result.content);
Ok(output)
}
pub fn endpoint(&self) -> &str {
&self.endpoint
}
pub async fn as_tools(&self) -> Vec<McpTool> {
let tools = self.tools.read().await;
tools
.iter()
.map(|def| McpTool {
definition: def.clone(),
source: self.clone(),
})
.collect()
}
pub async fn disconnect(&self) {
let mut client_guard = self.client.write().await;
*client_guard = None;
}
}
fn extract_text_from_content(content: &[Content]) -> String {
content
.iter()
.filter_map(|c| {
c.as_text()
.map(|text_content| text_content.text.to_string())
})
.collect::<Vec<_>>()
.join("\n")
}
fn convert_mcp_tool_to_definition(tool: &Tool) -> ToolDefinition {
let parameters = serde_json::to_value(tool.input_schema.as_ref()).unwrap_or_else(|_| {
serde_json::json!({
"type": "object",
"properties": {}
})
});
ToolDefinition {
name: tool.name.to_string(),
description: tool
.description
.clone()
.map(|s| s.to_string())
.unwrap_or_default(),
parameters,
}
}
#[derive(Clone)]
pub struct McpTool {
pub definition: ToolDefinition,
source: McpToolSource,
}
impl McpTool {
pub fn new(definition: ToolDefinition, source: McpToolSource) -> Self {
Self { definition, source }
}
}
impl ToolDyn for McpTool {
fn name(&self) -> &str {
&self.definition.name
}
fn definition<'a>(
&'a self,
_prompt: String,
) -> Pin<Box<dyn Future<Output = ToolDefinition> + Send + 'a>> {
Box::pin(async { self.definition.clone() })
}
fn call<'a>(
&'a self,
args: String,
) -> Pin<Box<dyn Future<Output = Result<String, ToolError>> + Send + 'a>> {
Box::pin(async move {
let parsed: serde_json::Value = serde_json::from_str(&args)?;
self.source.call_tool(&self.definition.name, parsed).await
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_convert_mcp_tool_to_definition() {
let schema = serde_json::json!({
"type": "object",
"properties": {
"input": {"type": "string"}
}
});
let mcp_tool = Tool::new(
"test_tool",
"A test tool",
Arc::new(schema.as_object().unwrap().clone()),
);
let def = convert_mcp_tool_to_definition(&mcp_tool);
assert_eq!(def.name, "test_tool");
assert_eq!(def.description, "A test tool");
assert!(def.parameters["properties"]["input"]["type"]
.as_str()
.is_some());
}
#[test]
fn test_convert_mcp_tool_no_description() {
let mcp_tool = Tool::new_with_raw("simple_tool", None, Arc::new(serde_json::Map::new()));
let def = convert_mcp_tool_to_definition(&mcp_tool);
assert_eq!(def.name, "simple_tool");
assert_eq!(def.description, "");
}
#[test]
fn test_mcp_tool_definition() {
let def = ToolDefinition::new(
"mcp_test",
"MCP test tool",
serde_json::json!({
"type": "object",
"properties": {}
}),
);
assert_eq!(def.name, "mcp_test");
assert_eq!(def.description, "MCP test tool");
}
#[test]
fn test_extract_text_from_content() {
let content = vec![Content::text("Hello"), Content::text("World")];
let result = extract_text_from_content(&content);
assert_eq!(result, "Hello\nWorld");
}
#[test]
fn test_extract_text_empty() {
let content: Vec<Content> = vec![];
let result = extract_text_from_content(&content);
assert_eq!(result, "");
}
}