use async_trait::async_trait;
use indexmap::IndexMap;
#[cfg(feature = "derive")]
pub use schemars::JsonSchema;
use crate::Error;
#[cfg(any(feature = "derive", test))]
use crate::types::Tool;
use crate::types::{ToolBinaryResult, ToolInvocation, ToolResult, ToolResultExpanded};
#[cfg(feature = "derive")]
pub fn schema_for<T: schemars::JsonSchema>() -> serde_json::Value {
let schema = schemars::schema_for!(T);
let mut value = serde_json::to_value(schema).expect("JSON Schema serialization cannot fail");
if let Some(obj) = value.as_object_mut() {
obj.remove("$schema");
obj.remove("title");
}
value
}
pub fn tool_parameters(schema: serde_json::Value) -> IndexMap<String, serde_json::Value> {
try_tool_parameters(schema).expect("tool parameter schema must be a JSON object")
}
pub fn try_tool_parameters(
schema: serde_json::Value,
) -> Result<IndexMap<String, serde_json::Value>, serde_json::Error> {
serde_json::from_value(schema)
}
pub fn convert_mcp_call_tool_result(value: &serde_json::Value) -> Option<ToolResult> {
let content = value.get("content")?.as_array()?;
let mut text_parts = Vec::new();
let mut binary_results = Vec::new();
for block in content {
match block.get("type").and_then(serde_json::Value::as_str) {
Some("text") => {
if let Some(text) = block.get("text").and_then(serde_json::Value::as_str) {
text_parts.push(text.to_string());
}
}
Some("image") => {
let data = block
.get("data")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty());
let mime_type = block
.get("mimeType")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty());
if let (Some(data), Some(mime_type)) = (data, mime_type) {
binary_results.push(ToolBinaryResult {
data: data.to_string(),
mime_type: mime_type.to_string(),
r#type: "image".to_string(),
description: None,
});
}
}
Some("resource") => {
let Some(resource) = block.get("resource").and_then(serde_json::Value::as_object)
else {
continue;
};
if let Some(text) = resource
.get("text")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty())
{
text_parts.push(text.to_string());
}
if let Some(blob) = resource
.get("blob")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty())
{
let mime_type = resource
.get("mimeType")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty())
.unwrap_or("application/octet-stream");
let description = resource
.get("uri")
.and_then(serde_json::Value::as_str)
.filter(|s| !s.is_empty())
.map(ToString::to_string);
binary_results.push(ToolBinaryResult {
data: blob.to_string(),
mime_type: mime_type.to_string(),
r#type: "resource".to_string(),
description,
});
}
}
_ => {}
}
}
Some(ToolResult::Expanded(ToolResultExpanded {
text_result_for_llm: text_parts.join("\n"),
result_type: if value.get("isError").and_then(serde_json::Value::as_bool) == Some(true) {
"failure".to_string()
} else {
"success".to_string()
},
binary_results_for_llm: (!binary_results.is_empty()).then_some(binary_results),
session_log: None,
error: None,
tool_telemetry: None,
tool_references: None,
}))
}
#[async_trait]
pub trait ToolHandler: Send + Sync + 'static {
async fn call(&self, invocation: ToolInvocation) -> Result<ToolResult, Error>;
}
#[cfg(feature = "derive")]
pub fn define_tool<P, F, Fut>(
name: impl Into<String>,
description: impl Into<String>,
handler: F,
) -> Tool
where
P: schemars::JsonSchema + serde::de::DeserializeOwned + Send + 'static,
F: Fn(ToolInvocation, P) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<ToolResult, Error>> + Send + 'static,
{
struct FnHandler<P, F> {
handler: F,
_marker: std::marker::PhantomData<fn(P)>,
}
#[async_trait]
impl<P, F, Fut> ToolHandler for FnHandler<P, F>
where
P: schemars::JsonSchema + serde::de::DeserializeOwned + Send + 'static,
F: Fn(ToolInvocation, P) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = Result<ToolResult, Error>> + Send + 'static,
{
async fn call(&self, mut invocation: ToolInvocation) -> Result<ToolResult, Error> {
let arguments = std::mem::take(&mut invocation.arguments);
let params: P = serde_json::from_value(arguments)?;
(self.handler)(invocation, params).await
}
}
Tool {
name: name.into(),
description: description.into(),
parameters: tool_parameters(schema_for::<P>()),
..Default::default()
}
.with_handler(std::sync::Arc::new(FnHandler {
handler,
_marker: std::marker::PhantomData,
}))
}
#[cfg(feature = "derive")]
pub fn define_tool_declaration<P>(name: impl Into<String>, description: impl Into<String>) -> Tool
where
P: schemars::JsonSchema,
{
Tool {
name: name.into(),
description: description.into(),
parameters: tool_parameters(schema_for::<P>()),
..Default::default()
}
}
#[cfg(test)]
mod tests;