use darling::ast::NestedMeta;
use darling::FromMeta;
use proc_macro2::{Ident, TokenStream};
use quote::quote;
use syn::parse::Parser;
use syn::{parse_quote, Attribute, ImplItem, ImplItemFn, ItemImpl, Visibility};
#[derive(Debug, Default, FromMeta)]
struct ToolRouterArgs {
#[darling(default = "default_router_name")]
router: String,
#[darling(default)]
vis: Option<String>,
}
fn default_router_name() -> String {
"tool_router".to_string()
}
struct ToolMethod {
name: Ident,
tool_name: String,
description: String,
is_async: bool,
}
pub fn expand_tool_router(args: TokenStream, mut input: ItemImpl) -> syn::Result<TokenStream> {
let nested_metas = if args.is_empty() {
vec![]
} else {
let parser = syn::punctuated::Punctuated::<NestedMeta, syn::Token![,]>::parse_terminated;
parser
.parse2(args)
.map(|p| p.into_iter().collect::<Vec<_>>())
.unwrap_or_default()
};
let args = ToolRouterArgs::from_list(&nested_metas).unwrap_or_default();
let tool_methods = collect_tool_methods(&input)?;
if tool_methods.is_empty() {
return Err(syn::Error::new_spanned(
&input,
"No methods marked with #[tool] found in impl block",
));
}
let _router_field = Ident::new(&args.router, proc_macro2::Span::call_site());
let vis = parse_visibility(args.vis.as_ref())?;
let tools_method = generate_tools_method(&tool_methods, &vis);
let handle_tool_method = generate_handle_tool_method(&tool_methods, &vis);
input.items.push(ImplItem::Fn(tools_method));
input.items.push(ImplItem::Fn(handle_tool_method));
Ok(quote! { #input })
}
fn collect_tool_methods(impl_block: &ItemImpl) -> syn::Result<Vec<ToolMethod>> {
let mut methods = Vec::new();
for item in &impl_block.items {
if let ImplItem::Fn(method) = item {
if let Some(tool_attr) = find_tool_attribute(&method.attrs) {
let tool_info = parse_tool_attribute(tool_attr)?;
let method_name = method.sig.ident.clone();
let tool_name = tool_info.name.unwrap_or_else(|| method_name.to_string());
methods.push(ToolMethod {
name: method_name,
tool_name,
description: tool_info.description,
is_async: method.sig.asyncness.is_some(),
});
}
}
}
Ok(methods)
}
struct ToolInfo {
name: Option<String>,
description: String,
}
fn find_tool_attribute(attrs: &[Attribute]) -> Option<&Attribute> {
attrs.iter().find(|attr| attr.path().is_ident("tool"))
}
fn parse_tool_attribute(attr: &Attribute) -> syn::Result<ToolInfo> {
let args_str = quote!(#attr).to_string();
let mut name = None;
let mut description = None;
if args_str.contains("description") {
if let Some(desc_start) = args_str.find("description = \"") {
let desc_start = desc_start + 15; if let Some(desc_end) = args_str[desc_start..].find('"') {
description = Some(args_str[desc_start..desc_start + desc_end].to_string());
}
}
}
if args_str.contains("name") && args_str.contains("name = \"") {
if let Some(name_start) = args_str.find("name = \"") {
let name_start = name_start + 8; if let Some(name_end) = args_str[name_start..].find('"') {
name = Some(args_str[name_start..name_start + name_end].to_string());
}
}
}
Ok(ToolInfo {
name,
description: description
.ok_or_else(|| syn::Error::new_spanned(attr, "Tool must have a description"))?,
})
}
fn generate_tools_method(methods: &[ToolMethod], vis: &Visibility) -> ImplItemFn {
let tool_definitions: Vec<_> = methods
.iter()
.map(|method| {
let name = &method.tool_name;
let description = &method.description;
quote! {
pmcp::types::ToolInfo::new(
#name,
Some(#description.to_string()),
serde_json::json!({
"type": "object",
"properties": {},
"required": []
}),
)
}
})
.collect();
parse_quote! {
#vis fn tools(&self) -> Vec<pmcp::types::ToolInfo> {
vec![
#(#tool_definitions),*
]
}
}
}
fn generate_handle_tool_method(methods: &[ToolMethod], vis: &Visibility) -> ImplItemFn {
let match_arms: Vec<_> = methods
.iter()
.map(|method| {
let tool_name = &method.tool_name;
let method_name = &method.name;
let await_token = if method.is_async {
quote!(.await)
} else {
quote!()
};
quote! {
#tool_name => {
let result = self.#method_name(args.clone())#await_token;
match result {
Ok(value) => Ok(serde_json::to_value(value)?),
Err(e) => Err(pmcp::Error::internal(format!("Tool error: {}", e))),
}
}
}
})
.collect();
parse_quote! {
#vis async fn handle_tool(
&self,
name: &str,
args: serde_json::Value,
_extra: pmcp::RequestHandlerExtra,
) -> pmcp::Result<serde_json::Value> {
match name {
#(#match_arms)*
_ => Err(pmcp::Error::method_not_found(format!("Unknown tool: {}", name))),
}
}
}
}
fn parse_visibility(vis_str: Option<&String>) -> syn::Result<Visibility> {
match vis_str {
None => Ok(parse_quote!(pub)),
Some(s) if s == "pub" => Ok(parse_quote!(pub)),
Some(s) if s == "pub(crate)" => Ok(parse_quote!(pub(crate))),
Some(s) if s == "pub(super)" => Ok(parse_quote!(pub(super))),
Some(s) if s.is_empty() => Ok(Visibility::Inherited),
Some(s) => Err(syn::Error::new(
proc_macro2::Span::call_site(),
format!("Invalid visibility: {}", s),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_visibility() {
assert!(parse_visibility(None).is_ok());
let pub_str = "pub".to_string();
assert!(parse_visibility(Some(&pub_str)).is_ok());
let crate_str = "pub(crate)".to_string();
assert!(parse_visibility(Some(&crate_str)).is_ok());
let invalid_str = "invalid".to_string();
assert!(parse_visibility(Some(&invalid_str)).is_err());
}
}