#![warn(missing_docs)]
use proc_macro::TokenStream;
use proc_macro2::TokenStream as TokenStream2;
use quote::{format_ident, quote};
use syn::{
parse_macro_input, Attribute, Expr, ExprLit, FnArg, Ident, ItemFn, Lit, Meta, MetaNameValue,
Pat, PatType, Result, Signature, Type,
};
const PARAM_ATTR: &str = "param";
#[proc_macro_attribute]
pub fn tool(attr: TokenStream, item: TokenStream) -> TokenStream {
let func = parse_macro_input!(item as ItemFn);
let description = match parse_tool_attr(attr) {
Ok(d) => d,
Err(err) => return err.to_compile_error().into(),
};
match tool_impl(description, func) {
Ok(tokens) => tokens.into(),
Err(err) => err.to_compile_error().into(),
}
}
fn parse_tool_attr(attr: TokenStream) -> Result<String> {
if attr.is_empty() {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"#[tool(description = \"...\")] is required",
));
}
let meta: Meta = syn::parse(attr)?;
if let Meta::NameValue(MetaNameValue { path, value, .. }) = &meta {
if path.is_ident("description") {
if let Expr::Lit(ExprLit {
lit: Lit::Str(lit), ..
}) = value
{
return Ok(lit.value());
}
}
}
Err(syn::Error::new(
proc_macro2::Span::call_site(),
"expected #[tool(description = \"...\")]",
))
}
fn tool_impl(description: String, mut func: ItemFn) -> Result<TokenStream2> {
let func_name_str = func.sig.ident.to_string();
let pascal = to_pascal_case(&func_name_str);
if !is_valid_ident_seed(&pascal) {
return Err(syn::Error::new_spanned(
&func.sig.ident,
format!(
"cannot derive `{pascal}Tool`/`{pascal}Input` from function name `{func_name_str}`: \
generated identifiers must start with an alphabetic or underscore character"
),
));
}
let tool_struct_name = format_ident!("{}Tool", pascal);
let input_struct_name = format_ident!("{}Input", pascal);
let func_name = func.sig.ident.clone();
if func.sig.asyncness.is_some() {
return Err(syn::Error::new_spanned(
&func.sig.ident,
"async tool functions are not supported by #[tool]: make the function synchronous \
(the derived Tool::invoke / BaseTool::run are already async)",
));
}
let params = extract_params(&func.sig)?;
let field_names: Vec<Ident> = params.iter().map(|p| p.name.clone()).collect();
let output_type = match &func.sig.output {
syn::ReturnType::Default => quote! { () },
syn::ReturnType::Type(_, ty) => {
if let Some(inner) = extract_result_ok(&func.sig.output) {
quote! { #inner }
} else {
quote! { #ty }
}
}
};
let invoke_body = if return_type_is_tool_error(&func.sig.output) {
quote! { #func_name(#(#field_names),*) }
} else {
quote! { #func_name(#(#field_names),*).map_err(|e| ::lc_core::tools::ToolError::ExecutionFailed(e.to_string())) }
};
let input_fields = generate_input_fields(¶ms);
let input_field_attrs = generate_field_attrs(¶ms);
strip_param_attrs(&mut func);
let expanded = quote! {
#func
#[derive(Debug, Clone)]
pub struct #tool_struct_name;
impl ::std::default::Default for #tool_struct_name {
fn default() -> Self {
Self
}
}
impl #tool_struct_name {
pub fn new() -> Self {
Self
}
}
#[derive(serde::Deserialize, schemars::JsonSchema)]
pub struct #input_struct_name {
#(#input_field_attrs)*
#(#input_fields)*
}
#[::async_trait::async_trait]
impl ::lc_core::tools::Tool for #tool_struct_name {
type Input = #input_struct_name;
type Output = #output_type;
async fn invoke(&self, input: Self::Input) -> ::std::result::Result<Self::Output, ::lc_core::tools::ToolError> {
let #input_struct_name { #(#field_names),* } = input;
#invoke_body
}
}
#[::async_trait::async_trait]
impl ::lc_core::tools::BaseTool for #tool_struct_name {
fn name(&self) -> &str {
#func_name_str
}
fn description(&self) -> &str {
#description
}
async fn run(&self, input: ::std::string::String) -> ::std::result::Result<::std::string::String, ::lc_core::tools::ToolError> {
let parsed: #input_struct_name = ::serde_json::from_str(&input)
.map_err(|e| ::lc_core::tools::ToolError::InvalidInput(format!("JSON parse error: {}", e)))?;
let #input_struct_name { #(#field_names),* } = parsed;
let result = #func_name(#(#field_names),*)
.map_err(|e| ::lc_core::tools::ToolError::ExecutionFailed(e.to_string()))?;
let serialized = ::serde_json::to_string(&result)
.map_err(|e| ::lc_core::tools::ToolError::ExecutionFailed(format!("Failed to serialize tool output: {}", e)))?;
Ok(serialized)
}
fn args_schema(&self) -> ::std::option::Option<::serde_json::Value> {
use ::schemars::schema_for;
Some(
::serde_json::to_value(schema_for!(#input_struct_name)).expect(
"[lc-tools-derive] internal error: generated Input schema failed to serialize \
(Input must derive schemars::JsonSchema)",
),
)
}
}
};
Ok(expanded)
}
struct ParamInfo {
name: Ident,
ty: Type,
desc: Option<String>,
}
fn extract_params(sig: &Signature) -> Result<Vec<ParamInfo>> {
let mut params = Vec::new();
for arg in &sig.inputs {
if let FnArg::Receiver(_) = arg {
continue;
}
if let FnArg::Typed(PatType { pat, ty, attrs, .. }) = arg {
let name = match pat.as_ref() {
Pat::Ident(ident) => ident.ident.clone(),
other => {
return Err(syn::Error::new_spanned(
other,
"tool parameters must be plain identifiers \
(ascription / tuple / wildcard patterns are not supported)",
));
}
};
let desc = extract_param_desc(attrs)?;
params.push(ParamInfo {
name,
ty: (*(*ty)).clone(),
desc,
});
}
}
Ok(params)
}
fn result_generics(ret: &syn::ReturnType) -> Option<(Type, Type)> {
let syn::ReturnType::Type(_, ty) = ret else {
return None;
};
let Type::Path(type_path) = &**ty else {
return None;
};
let seg = type_path.path.segments.last()?;
if seg.ident != "Result" {
return None;
}
let syn::PathArguments::AngleBracketed(args) = &seg.arguments else {
return None;
};
let mut generic = args.args.iter().filter_map(|a| match a {
syn::GenericArgument::Type(t) => Some(t),
_ => None,
});
Some((generic.next()?.clone(), generic.next()?.clone()))
}
fn extract_result_ok(ret: &syn::ReturnType) -> Option<Type> {
result_generics(ret).map(|(ok, _)| ok)
}
fn return_type_is_tool_error(ret: &syn::ReturnType) -> bool {
let Some((_, err)) = result_generics(ret) else {
return false;
};
let Type::Path(err_path) = err else {
return false;
};
let segs = err_path.path.segments;
let Some(last) = segs.last() else {
return false;
};
if last.ident != "ToolError" {
return false;
}
segs.len() == 1
|| segs
.get(segs.len().saturating_sub(2))
.is_some_and(|s| s.ident == "tools")
}
fn is_valid_ident_seed(s: &str) -> bool {
match s.chars().next() {
Some(c) => c == '_' || c.is_ascii_alphabetic(),
None => false,
}
}
fn extract_param_desc(attrs: &[Attribute]) -> Result<Option<String>> {
for attr in attrs {
if attr.path().is_ident(PARAM_ATTR) {
let meta: Meta = attr.parse_args().map_err(|e| {
syn::Error::new_spanned(
attr,
format!("failed to parse `#[{PARAM_ATTR}(...)]`: {e}"),
)
})?;
if let Meta::NameValue(MetaNameValue { path, value, .. }) = &meta {
if path.is_ident("desc") {
if let Expr::Lit(ExprLit {
lit: Lit::Str(lit), ..
}) = value
{
return Ok(Some(lit.value()));
}
}
}
}
}
Ok(None)
}
fn generate_input_fields(params: &[ParamInfo]) -> Vec<TokenStream2> {
params
.iter()
.map(|p| {
let name = &p.name;
let ty = &p.ty;
quote! {
pub #name: #ty,
}
})
.collect()
}
fn generate_field_attrs(params: &[ParamInfo]) -> Vec<TokenStream2> {
params
.iter()
.map(|p| {
if let Some(desc) = &p.desc {
quote! {
#[doc = #desc]
}
} else {
quote! {}
}
})
.collect()
}
fn to_pascal_case(s: &str) -> String {
s.split('_')
.map(|word| {
let mut chars = word.chars();
match chars.next() {
None => String::new(),
Some(c) => c.to_uppercase().collect::<String>() + chars.as_str(),
}
})
.collect()
}
fn strip_param_attrs(func: &mut ItemFn) {
for arg in &mut func.sig.inputs {
if let FnArg::Typed(pat_type) = arg {
pat_type
.attrs
.retain(|attr| !attr.path().is_ident(PARAM_ATTR));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use quote::quote;
use syn::ReturnType;
fn rt(src: &str) -> ReturnType {
syn::parse_str(src).unwrap()
}
fn ty_str(ret: &ReturnType) -> Option<String> {
extract_result_ok(ret).map(|t| quote! { #t }.to_string())
}
#[test]
fn result_generics_bare_and_qualified_agree() {
assert_eq!(ty_str(&rt("-> Result<f64, String>")), Some("f64".into()));
assert_eq!(
ty_str(&rt("-> std::result::Result<f64, String>")),
Some("f64".into())
);
assert_eq!(ty_str(&rt("-> f64")), None);
assert_eq!(ty_str(&rt("-> Result<f64>")), None); }
#[test]
fn return_type_is_tool_error_qualified_forms_only() {
assert!(return_type_is_tool_error(&rt("-> Result<String, ToolError>")));
assert!(return_type_is_tool_error(&rt("-> Result<String, lc_core::tools::ToolError>")));
assert!(return_type_is_tool_error(&rt("-> Result<String, tools::ToolError>")));
assert!(!return_type_is_tool_error(&rt("-> Result<String, MyToolError>")));
assert!(!return_type_is_tool_error(&rt("-> Result<String, other::ToolError>")));
assert!(!return_type_is_tool_error(&rt("-> Result<String, anyhow::Error>")));
assert!(!return_type_is_tool_error(&rt("-> String")));
}
#[test]
fn ident_seed_validation() {
assert!(is_valid_ident_seed("Calculator"));
assert!(is_valid_ident_seed("_private"));
assert!(!is_valid_ident_seed("9lives"));
assert!(!is_valid_ident_seed(""));
}
}