use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::{quote, ToTokens};
use syn::{
parse_macro_input, Error, Expr, FnArg, Ident, ImplItem, ItemImpl, Lit, Meta, MetaNameValue, Pat
};
use std::collections::HashSet;
use heck::ToUpperCamelCase;
#[proc_macro_attribute]
pub fn toolbox(_attr: TokenStream, item: TokenStream) -> TokenStream {
let mut item_impl = parse_macro_input!(item as ItemImpl);
let struct_name = &item_impl.self_ty;
let struct_ident = match &**struct_name {
syn::Type::Path(type_path) => {
type_path.path.get_ident().expect("Expected an identifier for the struct")
}
_ => return Error::new(Span::call_site(), "toolbox! macro only supports impl blocks for structs").to_compile_error().into(),
};
let mut generated_code = TokenStream2::new();
let mut tool_definitions = TokenStream2::new();
let mut match_arms = TokenStream2::new();
let mut found_tools = HashSet::new();
for item in item_impl.items.iter_mut() {
if let ImplItem::Fn(ref mut method) = item {
if let Some(tool_attr) = method.attrs.clone().iter().find(|attr| attr.path().is_ident("tool")) {
method.attrs.retain(|attr| !attr.path().is_ident("tool"));
let fn_name_sig = &method.sig.ident;
let fn_name = fn_name_sig.to_string();
let mut tool_name = fn_name.clone();
let mut name_arg_found = false;
let parser = syn::punctuated::Punctuated::<Meta, syn::Token![,]>::parse_terminated;
if let Ok(args) = tool_attr.parse_args_with(parser) {
for arg_meta in args {
match arg_meta {
Meta::NameValue(name_value) if name_value.path.is_ident("name") => {
if name_arg_found {
return Error::new_spanned(name_value.to_token_stream(), "Duplicate 'name' argument in tool attribute").to_compile_error().into();
}
let Expr::Lit(expr_lit) = &name_value.value else {
return Error::new_spanned(name_value.value.to_token_stream(), "Expected literal value for tool name").to_compile_error().into();
};
let Lit::Str(lit_str) = &expr_lit.lit else {
return Error::new_spanned(expr_lit.to_token_stream(), "Expected string literal for tool name").to_compile_error().into();
};
tool_name = lit_str.value();
name_arg_found = true;
},
_ => {
return Error::new_spanned(arg_meta.to_token_stream(), "Expected name = \"...\" in tool attribute").to_compile_error().into();
}
};
}
}
if !found_tools.insert(tool_name.clone()) {
return Error::new_spanned(tool_attr.to_token_stream(), format!("Duplicate tool name found: {}", tool_name)).to_compile_error().into();
}
let description = method.attrs.iter()
.filter_map(|attr|
match attr.meta.clone() {
Meta::NameValue(MetaNameValue { path, value: Expr::Lit(expr_lit), .. }) if path.is_ident("doc") => {
match expr_lit.lit {
Lit::Str(lit_str) => {
Some(lit_str.value().trim().trim_start_matches(|c: char| c == '/' || c == '*' || c.is_whitespace()).to_string())
}
_ => None, }
},
_ => None, }
)
.collect::<Vec<String>>()
.join("\n");
let description_token = if description.trim().is_empty() {
quote! { None }
} else {
let desc = description.trim().to_string();
quote! { Some(#desc.to_string()) }
};
let params_struct_name = Ident::new(&format!("{}Params", fn_name.to_upper_camel_case()), fn_name_sig.span());
let mut param_fields = TokenStream2::new();
let mut param_assignments = TokenStream2::new();
for arg in method.sig.inputs.iter_mut() {
if let FnArg::Typed(ref mut pat_type) = arg {
let ty = pat_type.ty.clone();
let attrs = pat_type.attrs.clone();
pat_type.attrs.clear();
let Pat::Ident(ref pat_ident) = *pat_type.pat else {
return Error::new_spanned(pat_type.pat.to_token_stream(), "Tool function parameters must be simple identifiers").to_compile_error().into();
};
let arg_name = &pat_ident.ident;
param_fields.extend(quote! {
#(#attrs)* pub #arg_name: #ty,
});
param_assignments.extend(quote! {
params.#arg_name
});
}
}
if !param_fields.is_empty() {
generated_code.extend(quote! {
#[derive(serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
#[allow(dead_code)]
#[allow(clippy::all)]
struct #params_struct_name {
#param_fields
}
});
}
let schema_token = if param_fields.is_empty() {
quote! { None }
} else {
quote! {
Some({
let generator = ::schemars::generate::SchemaSettings::draft2020_12().with(|s| {
s.meta_schema = None;
}).into_generator();
generator.into_root_schema_for::<#params_struct_name>().into()
})
}
};
tool_definitions.extend(quote! {
Tool {
name: #tool_name.to_string(),
description: #description_token,
schema: #schema_token,
},
});
let mut method_call = TokenStream2::new();
if !param_fields.is_empty(){
method_call.extend(quote! {
let params: #params_struct_name = serde_json::from_value(parameters)
.map_err(|e| {
eprintln!("Tool parameter deserialization error for '{}': {:?}", #tool_name, e);
ToolError::ExecutionError
})?;
});
}
method_call.extend(quote! { self.#fn_name_sig(#param_assignments) });
if method.sig.asyncness.is_some() {
method_call.extend(quote! {.await});
}
method_call.extend(quote! { .map_err(|e| {
eprintln!("Tool execution error for '{}': {:?}", #tool_name, e);
ToolError::ExecutionError
}) });
match_arms.extend(quote! {
#tool_name => {
#method_call
},
});
}
}
}
if found_tools.is_empty() {
return Error::new(Span::call_site(), "No #[tool] definition in impl block").to_compile_error().into()
}
let toolbox_impl = quote! {
#[::async_trait::async_trait]
impl ToolBox for #struct_ident {
fn tools_definitions(&self) -> Result<Vec<Tool>, ToolError> {
Ok(vec![
#tool_definitions
])
}
async fn call_tool(&self, tool_name: String, parameters: serde_json::Value) -> Result<String, ToolError> {
match tool_name.as_str() {
#match_arms
_ => {
Err(ToolError::NoToolFound(tool_name))
}
}
}
}
};
let final_code = quote! {
#item_impl
#toolbox_impl
#generated_code
};
final_code.into()
}