use std::sync::Arc;
use rmcp::ErrorData;
use rmcp::model::CallToolRequestParams;
use rmcp::model::CallToolResult;
use rmcp::model::JsonObject;
use rmcp::model::Tool;
use schemars::generate::SchemaSettings;
use super::HandlerContext;
use super::annotations::Annotation;
use super::handler::ErasedToolFn;
use super::json_response::ToolCallJsonResponse;
use super::name::ToolName;
use super::parameters::ParameterBuilder;
#[derive(Clone)]
pub struct ToolDef {
pub tool_name: ToolName,
pub annotations: Annotation,
pub handler: Arc<dyn ErasedToolFn>,
pub parameters: Option<fn() -> ParameterBuilder>,
}
impl ToolDef {
pub fn name(&self) -> &'static str { self.tool_name.into() }
pub async fn call_tool(
&self,
request: CallToolRequestParams,
) -> std::result::Result<CallToolResult, ErrorData> {
let handler_context = HandlerContext::new(self.clone(), request);
Ok(self.handler.call_erased(handler_context).await)
}
fn generate_output_schema() -> Arc<JsonObject> {
let mut schema_settings = SchemaSettings::default();
schema_settings.inline_subschemas = true;
let generator = schema_settings.into_generator();
let schema = generator.into_root_schema_for::<ToolCallJsonResponse>();
let Ok(schema_value) = serde_json::to_value(schema) else {
return Arc::new(rmcp::model::JsonObject::new());
};
let schema_object = schema_value
.as_object()
.map_or_else(rmcp::model::JsonObject::new, Clone::clone);
Arc::new(schema_object)
}
pub fn to_tool(&self) -> Tool {
let builder = self
.parameters
.map_or_else(ParameterBuilder::new, |builder_fn| builder_fn());
let enhanced_annotations = {
let mut enhanced = self.annotations.clone();
let category_prefix = enhanced.tool_category.as_ref();
let base_title = &enhanced.title;
let full_title = format!("{category_prefix}: {base_title}");
enhanced.title = full_title;
enhanced
};
rmcp::model::Tool::new(
<&'static str>::from(self.tool_name),
self.tool_name.description(),
builder.build(),
)
.with_title(self.tool_name.short_title())
.with_raw_output_schema(Self::generate_output_schema())
.with_annotations(enhanced_annotations.into())
}
}