use darling::FromMeta;
use proc_macro2::TokenStream;
use quote::{format_ident, quote};
use syn::{parse2, FnArg, ImplItemFn, ItemFn, Lit, Meta, Pat, Type};
use crate::schema::{is_option_type, type_to_schema};
pub fn rust_type_to_json_type(ty: &Type) -> &'static str {
match ty {
Type::Path(type_path) => {
if let Some(segment) = type_path.path.segments.last() {
let ident = segment.ident.to_string();
match ident.as_str() {
"String" | "str" => "string",
"i8" | "i16" | "i32" | "i64" | "i128" | "isize" | "u8" | "u16" | "u32"
| "u64" | "u128" | "usize" => "integer",
"f32" | "f64" => "number",
"bool" => "boolean",
"Vec" => "array",
"Option" => {
if let syn::PathArguments::AngleBracketed(args) = &segment.arguments {
if let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() {
return rust_type_to_json_type(inner_ty);
}
}
"object"
}
_ => "object",
}
} else {
"object"
}
}
Type::Reference(type_ref) => rust_type_to_json_type(&type_ref.elem),
_ => "object",
}
}
#[derive(Debug, FromMeta)]
pub struct ToolArgs {
pub description: String,
#[darling(default)]
pub name: Option<String>,
#[darling(default)]
pub group: Option<String>,
}
#[derive(Debug, Default, FromMeta)]
pub struct ToolParamArgs {
#[darling(default)]
pub description: Option<String>,
#[darling(default)]
pub name: Option<String>,
#[darling(default)]
pub required: Option<bool>,
}
pub fn parse_param_attrs(attrs: &[syn::Attribute]) -> Option<ToolParamArgs> {
for attr in attrs {
if attr.path().is_ident("param") {
if let Ok(Lit::Str(s)) = attr.parse_args::<Lit>() {
return Some(ToolParamArgs {
description: Some(s.value()),
name: None,
required: None,
});
}
if let Ok(args) = ToolParamArgs::from_meta(&attr.meta) {
return Some(args);
}
return Some(ToolParamArgs::default());
}
}
None
}
pub fn has_param_attr(attrs: &[syn::Attribute]) -> bool {
attrs.iter().any(|a| a.path().is_ident("param"))
}
pub fn strip_param_attrs(attrs: &mut Vec<syn::Attribute>) {
attrs.retain(|a| !a.path().is_ident("param"))
}
#[derive(Debug, Clone)]
pub struct ParamInfo {
pub name: String,
pub ty: syn::Type,
pub description: Option<String>,
pub required: bool,
}
pub fn mcp_tool_impl(attr: TokenStream, item: TokenStream) -> TokenStream {
let attr_args = match parse_tool_args(attr) {
Ok(args) => args,
Err(e) => return e.to_compile_error(),
};
if let Ok(func) = parse2::<ItemFn>(item.clone()) {
process_standalone_function(func, attr_args)
} else if let Ok(method) = parse2::<ImplItemFn>(item.clone()) {
process_impl_method(method, attr_args)
} else {
syn::Error::new_spanned(item, "mcp_tool can only be applied to functions or methods")
.to_compile_error()
}
}
fn parse_tool_args(attr: TokenStream) -> Result<ToolArgs, syn::Error> {
if attr.is_empty() {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"mcp_tool requires a description: #[mcp_tool(\"description\")] or #[mcp_tool(description = \"...\")]",
));
}
if let Ok(Lit::Str(s)) = parse2::<Lit>(attr.clone()) {
return Ok(ToolArgs {
description: s.value(),
name: None,
group: None,
});
}
let meta: Meta = parse2(quote! { mcp_tool(#attr) })?;
ToolArgs::from_meta(&meta).map_err(|e| syn::Error::new(proc_macro2::Span::call_site(), e))
}
fn process_impl_method(mut method: ImplItemFn, args: ToolArgs) -> TokenStream {
let params = extract_params(&method.sig.inputs.iter().collect::<Vec<_>>());
let tool_name = args.name.unwrap_or_else(|| method.sig.ident.to_string());
for input in &mut method.sig.inputs {
if let FnArg::Typed(pat_type) = input {
strip_param_attrs(&mut pat_type.attrs);
}
}
let description = &args.description;
let param_tokens = generate_param_metadata(¶ms);
quote! {
#[doc(hidden)]
#[allow(dead_code)]
const _: () = {
};
#[doc = #description]
#[mcp_tool_meta(name = #tool_name, description = #description, params = [#param_tokens])]
#method
}
}
fn process_standalone_function(mut func: ItemFn, args: ToolArgs) -> TokenStream {
let func_name = &func.sig.ident;
let tool_name = args.name.unwrap_or_else(|| func_name.to_string());
let description = &args.description;
let group_code = match &args.group {
Some(g) => quote! { Some(#g.to_string()) },
None => quote! { None },
};
let struct_name = format_ident!("{}Tool", to_pascal_case(&func_name.to_string()));
let params = extract_params(&func.sig.inputs.iter().collect::<Vec<_>>());
let is_async = func.sig.asyncness.is_some();
for input in &mut func.sig.inputs {
if let FnArg::Typed(pat_type) = input {
strip_param_attrs(&mut pat_type.attrs);
}
}
let properties = generate_json_properties(¶ms);
let required: Vec<&str> = params
.iter()
.filter(|p| p.required)
.map(|p| p.name.as_str())
.collect();
let param_extractions: Vec<TokenStream> = params
.iter()
.map(|p| {
let param_name = &p.name;
let param_ident = syn::Ident::new(&p.name, proc_macro2::Span::call_site());
let ty = &p.ty;
if is_option_type(ty) {
quote! {
let #param_ident: #ty = __args
.get(#param_name)
.and_then(|v| serde_json::from_value(v.clone()).ok());
}
} else {
quote! {
let #param_ident: #ty = {
let __raw = __args
.get(#param_name)
.ok_or_else(|| format!("Missing required parameter: {}", #param_name))?
.clone();
serde_json::from_value(__raw)
.map_err(|e| format!("Invalid parameter '{}': {}", #param_name, e))?
};
}
}
})
.collect();
let param_names: Vec<syn::Ident> = params
.iter()
.map(|p| syn::Ident::new(&p.name, proc_macro2::Span::call_site()))
.collect();
let call_expr = if is_async {
quote! { #func_name(#(#param_names),*).await }
} else {
quote! { #func_name(#(#param_names),*) }
};
let result_handling = generate_result_handling(&func.sig.output, call_expr);
let inventory_group = match &args.group {
Some(g) => quote! { Some(#g) },
None => quote! { None },
};
quote! {
#[doc = #description]
#func
#[derive(Clone, Copy, Default)]
pub struct #struct_name;
impl model_context_protocol::McpTool for #struct_name {
fn definition(&self) -> model_context_protocol::McpToolDefinition {
model_context_protocol::McpToolDefinition {
name: #tool_name.to_string(),
description: Some(#description.to_string()),
group: #group_code,
input_schema: serde_json::json!({
"type": "object",
"properties": { #properties },
"required": [#(#required),*]
}),
output_schema: None,
annotations: None,
execution: None,
title: None,
icons: None,
meta: None,
}
}
fn call<'a>(&'a self, __args: serde_json::Value) -> model_context_protocol::BoxFuture<'a, model_context_protocol::ToolCallResult> {
Box::pin(async move {
let __args = __args.as_object().cloned().unwrap_or_default();
#(#param_extractions)*
#result_handling
})
}
}
model_context_protocol::inventory::submit! {
model_context_protocol::ToolEntry::new(
|| std::sync::Arc::new(#struct_name) as model_context_protocol::DynTool,
#inventory_group
)
}
}
}
fn to_pascal_case(s: &str) -> String {
s.split('_')
.map(|word| {
let mut chars = word.chars();
match chars.next() {
None => String::new(),
Some(first) => first.to_uppercase().chain(chars).collect(),
}
})
.collect()
}
fn generate_result_handling(output: &syn::ReturnType, call_expr: TokenStream) -> TokenStream {
match output {
syn::ReturnType::Default => {
quote! {
#call_expr;
Ok(vec![model_context_protocol::ToolContent::text("ok")])
}
}
syn::ReturnType::Type(_, ty) => {
if is_result_type(ty) {
quote! {
match #call_expr {
Ok(value) => {
let text = serde_json::to_string(&value)
.unwrap_or_else(|_| format!("{:?}", value));
Ok(vec![model_context_protocol::ToolContent::text(text)])
}
Err(e) => Err(format!("{}", e)),
}
}
} else {
quote! {
let __result = #call_expr;
let text = serde_json::to_string(&__result)
.unwrap_or_else(|_| format!("{:?}", __result));
Ok(vec![model_context_protocol::ToolContent::text(text)])
}
}
}
}
}
fn is_result_type(ty: &Type) -> bool {
if let Type::Path(type_path) = ty {
if let Some(segment) = type_path.path.segments.last() {
return segment.ident == "Result";
}
}
false
}
fn extract_params(inputs: &[&FnArg]) -> Vec<ParamInfo> {
let mut params = Vec::new();
for input in inputs {
if let FnArg::Typed(pat_type) = input {
if let Pat::Ident(pat_ident) = pat_type.pat.as_ref() {
let name = pat_ident.ident.to_string();
if name == "self" {
continue;
}
if !has_param_attr(&pat_type.attrs) {
continue;
}
let param_args = parse_param_attrs(&pat_type.attrs).unwrap_or_default();
let description = param_args
.description
.or_else(|| extract_doc_comment(&pat_type.attrs));
let param_name = param_args.name.unwrap_or(name);
let required = param_args
.required
.unwrap_or_else(|| !is_option_type(&pat_type.ty));
params.push(ParamInfo {
name: param_name,
ty: (*pat_type.ty).clone(),
description,
required,
});
}
}
}
params
}
fn extract_doc_comment(attrs: &[syn::Attribute]) -> Option<String> {
for attr in attrs {
if attr.path().is_ident("doc") {
if let Meta::NameValue(meta) = &attr.meta {
if let syn::Expr::Lit(expr_lit) = &meta.value {
if let Lit::Str(lit_str) = &expr_lit.lit {
return Some(lit_str.value().trim().to_string());
}
}
}
}
}
None
}
fn generate_param_metadata(params: &[ParamInfo]) -> TokenStream {
let param_tokens: Vec<TokenStream> = params
.iter()
.map(|p| {
let name = &p.name;
let ty = &p.ty;
let desc = p.description.as_deref().unwrap_or("");
let required = p.required;
let schema = type_to_schema(ty);
quote! {
McpParamMeta {
name: #name,
description: #desc,
required: #required,
schema: #schema,
}
}
})
.collect();
quote! { #(#param_tokens),* }
}
fn generate_json_properties(params: &[ParamInfo]) -> TokenStream {
let props: Vec<TokenStream> = params
.iter()
.map(|p| {
let name = &p.name;
let ty_str = rust_type_to_json_type(&p.ty);
let desc = p.description.as_deref().unwrap_or("");
if desc.is_empty() {
quote! { #name: { "type": #ty_str } }
} else {
quote! { #name: { "type": #ty_str, "description": #desc } }
}
})
.collect();
quote! { #(#props),* }
}
#[derive(Debug, Clone)]
pub struct CollectedTool {
pub name: String,
pub description: String,
pub params: Vec<ParamInfo>,
pub method_ident: syn::Ident,
}
impl CollectedTool {
pub fn generate_mcp_tool_def(&self) -> TokenStream {
let name = &self.name;
let description = &self.description;
let properties = self.generate_json_properties();
let required: Vec<&str> = self
.params
.iter()
.filter(|p| p.required)
.map(|p| p.name.as_str())
.collect();
quote! {
model_context_protocol::McpToolDefinition {
name: #name.to_string(),
description: Some(#description.to_string()),
group: None,
input_schema: serde_json::json!({
"type": "object",
"properties": { #properties },
"required": [#(#required),*]
}),
output_schema: None,
annotations: None,
execution: None,
title: None,
icons: None,
meta: None,
}
}
}
fn generate_json_properties(&self) -> TokenStream {
let props: Vec<TokenStream> = self
.params
.iter()
.map(|p| {
let name = &p.name;
let ty_str = rust_type_to_json_type(&p.ty);
let desc = p.description.as_deref().unwrap_or("");
if desc.is_empty() {
quote! { #name: { "type": #ty_str } }
} else {
quote! { #name: { "type": #ty_str, "description": #desc } }
}
})
.collect();
quote! { #(#props),* }
}
pub fn generate_call_arm(&self) -> TokenStream {
let name = &self.name;
let method = &self.method_ident;
let param_extractions: Vec<TokenStream> = self
.params
.iter()
.map(|p| {
let param_name = &p.name;
let param_ident = syn::Ident::new(&p.name, proc_macro2::Span::call_site());
let ty = &p.ty;
if is_option_type(ty) {
quote! {
let #param_ident: #ty = args
.get(#param_name)
.and_then(|v| serde_json::from_value(v.clone()).ok());
}
} else {
quote! {
let #param_ident: #ty = {
let __raw = args
.get(#param_name)
.ok_or_else(|| format!("Missing required parameter: {}", #param_name))?
.clone();
serde_json::from_value(__raw)
.map_err(|e| format!("Invalid parameter '{}': {}", #param_name, e))?
};
}
}
})
.collect();
let param_names: Vec<syn::Ident> = self
.params
.iter()
.map(|p| syn::Ident::new(&p.name, proc_macro2::Span::call_site()))
.collect();
quote! {
#name => {
#(#param_extractions)*
let result = self.#method(#(#param_names),*);
match serde_json::to_string(&result) {
Ok(json) => Ok(vec![model_context_protocol::ToolContent::text(json)]),
Err(e) => Err(format!("Serialization error: {}", e)),
}
}
}
}
}