use darling::FromMeta;
use heck::ToUpperCamelCase;
use proc_macro2::{Ident, TokenStream};
use quote::quote;
use syn::parse::Parser;
use syn::{parse_quote, FnArg, ItemFn, Pat, PatType, ReturnType, Type};
#[derive(Debug, FromMeta)]
struct ToolArgs {
#[darling(default)]
name: Option<String>,
description: String,
#[darling(default)]
annotations: Option<ToolAnnotations>,
}
#[derive(Debug, Default, FromMeta)]
struct ToolAnnotations {
#[darling(default)]
#[allow(dead_code)]
category: Option<String>,
#[darling(default)]
#[allow(dead_code)]
complexity: Option<String>,
#[darling(default)]
read_only: Option<bool>,
#[darling(default)]
destructive: Option<bool>,
#[darling(default)]
idempotent: Option<bool>,
#[darling(default)]
open_world: Option<bool>,
#[darling(default)]
output_type: Option<String>,
}
pub fn expand_tool(args: TokenStream, input: &ItemFn) -> syn::Result<TokenStream> {
let nested_metas = if args.is_empty() {
vec![]
} else {
let parser = syn::punctuated::Punctuated::<darling::ast::NestedMeta, syn::Token![,]>::parse_terminated;
parser
.parse2(args)
.map(|p| p.into_iter().collect::<Vec<_>>())
.unwrap_or_default()
};
let args = ToolArgs::from_list(&nested_metas)
.map_err(|e| syn::Error::new_spanned(&input.sig.ident, e.to_string()))?;
let fn_name = &input.sig.ident;
let tool_name = args.name.unwrap_or_else(|| fn_name.to_string());
let description = args.description;
let params = extract_parameters(input);
let return_type = extract_return_type(input);
let wrapper_name = Ident::new(
&format!("{}ToolHandler", fn_name.to_string().to_upper_camel_case()),
fn_name.span(),
);
let is_async = input.sig.asyncness.is_some();
let await_token = if is_async { quote!(.await) } else { quote!() };
let param_extraction = generate_param_extraction(¶ms);
let param_names: Vec<_> = params.iter().map(|p| &p.name).collect();
let result_conversion = generate_result_conversion(&return_type);
let (annotations_code, definition_code) =
generate_definition_code(&tool_name, &description, args.annotations.as_ref());
let expanded = quote! {
#input
#[derive(Debug, Clone)]
pub struct #wrapper_name;
#[async_trait::async_trait]
impl pmcp::ToolHandler for #wrapper_name {
async fn handle(
&self,
args: serde_json::Value,
_extra: pmcp::RequestHandlerExtra,
) -> pmcp::Result<serde_json::Value> {
#param_extraction
let result = #fn_name(#(#param_names),*)#await_token;
#result_conversion
}
}
impl #wrapper_name {
#annotations_code
pub fn definition() -> pmcp::types::ToolInfo {
#definition_code
}
fn input_schema() -> serde_json::Value {
serde_json::json!({
"type": "object",
"properties": {},
"required": []
})
}
}
};
Ok(expanded)
}
fn generate_definition_code(
tool_name: &str,
description: &str,
annotations: Option<&ToolAnnotations>,
) -> (TokenStream, TokenStream) {
match annotations {
None => {
let definition = quote! {
pmcp::types::ToolInfo::new(
#tool_name,
Some(#description.to_string()),
Self::input_schema(),
)
};
(quote!(), definition)
},
Some(ann) => {
let mut annotation_chain = vec![quote!(pmcp::types::ToolAnnotations::new())];
if let Some(read_only) = ann.read_only {
annotation_chain.push(quote!(.with_read_only(#read_only)));
}
if let Some(destructive) = ann.destructive {
annotation_chain.push(quote!(.with_destructive(#destructive)));
}
if let Some(idempotent) = ann.idempotent {
annotation_chain.push(quote!(.with_idempotent(#idempotent)));
}
if let Some(open_world) = ann.open_world {
annotation_chain.push(quote!(.with_open_world(#open_world)));
}
let (output_schema_code, definition) = if let Some(ref output_type) = ann.output_type {
let output_type_ident = Ident::new(output_type, proc_macro2::Span::call_site());
let output_type_name = output_type.clone();
let output_schema_fn = quote! {
#[cfg(feature = "schema-generation")]
fn output_schema() -> serde_json::Value {
let schema = schemars::schema_for!(#output_type_ident);
serde_json::to_value(&schema).unwrap_or_else(|_| {
serde_json::json!({
"type": "object",
"additionalProperties": true
})
})
}
};
let annotations_with_output = quote! {
#(#annotation_chain)*
.with_output_type_name(#output_type_name)
};
let def = quote! {
#[cfg(feature = "schema-generation")]
{
let annotations = #annotations_with_output;
pmcp::types::ToolInfo::with_annotations(
#tool_name,
Some(#description.to_string()),
Self::input_schema(),
annotations,
)
.with_output_schema(Self::output_schema())
}
#[cfg(not(feature = "schema-generation"))]
{
let annotations = #(#annotation_chain)*
.with_output_type_name(#output_type_name);
pmcp::types::ToolInfo::with_annotations(
#tool_name,
Some(#description.to_string()),
Self::input_schema(),
annotations,
)
}
};
(output_schema_fn, def)
} else {
let annotations_build = quote! {
#(#annotation_chain)*
};
let def = quote! {
let annotations = #annotations_build;
pmcp::types::ToolInfo::with_annotations(
#tool_name,
Some(#description.to_string()),
Self::input_schema(),
annotations,
)
};
(quote!(), def)
};
(output_schema_code, definition)
},
}
}
struct ParamInfo {
name: Ident,
ty: Type,
optional: bool,
}
fn extract_parameters(func: &ItemFn) -> Vec<ParamInfo> {
let mut params = Vec::new();
for arg in &func.sig.inputs {
match arg {
FnArg::Receiver(_) => {
},
FnArg::Typed(PatType { pat, ty, .. }) => {
if let Pat::Ident(pat_ident) = pat.as_ref() {
let name = pat_ident.ident.clone();
let ty = ty.as_ref().clone();
let optional = crate::mcp_common::type_name_matches(&ty, "Option");
params.push(ParamInfo { name, ty, optional });
}
},
}
}
params
}
fn extract_return_type(func: &ItemFn) -> Type {
match &func.sig.output {
ReturnType::Default => parse_quote!(()),
ReturnType::Type(_, ty) => ty.as_ref().clone(),
}
}
fn generate_param_extraction(params: &[ParamInfo]) -> TokenStream {
let mut extractions = Vec::new();
for param in params {
let name = ¶m.name;
let name_str = name.to_string();
let ty = ¶m.ty;
if param.optional {
extractions.push(quote! {
let #name: #ty = args.get(#name_str)
.and_then(|v| serde_json::from_value(v.clone()).ok())
.unwrap_or_default();
});
} else {
extractions.push(quote! {
let #name: #ty = args.get(#name_str)
.and_then(|v| serde_json::from_value(v.clone()).ok())
.ok_or_else(|| pmcp::Error::invalid_params(
format!("Missing required parameter: {}", #name_str)
))?;
});
}
}
quote! {
#(#extractions)*
}
}
fn generate_result_conversion(return_type: &Type) -> TokenStream {
if crate::mcp_common::type_name_matches(return_type, "Result") {
quote! {
match result {
Ok(value) => {
let json_value = serde_json::to_value(value)
.map_err(|e| pmcp::Error::internal(e.to_string()))?;
Ok(json_value)
}
Err(e) => Err(pmcp::Error::internal(format!("Tool error: {}", e)))
}
}
} else {
quote! {
let json_value = serde_json::to_value(result)
.map_err(|e| pmcp::Error::internal(e.to_string()))?;
Ok(json_value)
}
}
}
#[cfg(test)]
mod tests {
use crate::mcp_common;
use syn::{parse_quote, Type};
#[test]
fn test_type_name_matches_option() {
let opt_type: Type = parse_quote!(Option<String>);
assert!(mcp_common::type_name_matches(&opt_type, "Option"));
let non_opt_type: Type = parse_quote!(String);
assert!(!mcp_common::type_name_matches(&non_opt_type, "Option"));
}
#[test]
fn test_type_name_matches_result() {
let result_type: Type = parse_quote!(Result<String, Error>);
assert!(mcp_common::type_name_matches(&result_type, "Result"));
let non_result_type: Type = parse_quote!(String);
assert!(!mcp_common::type_name_matches(&non_result_type, "Result"));
}
}