use proc_macro2::TokenStream;
use quote::quote;
use syn::{parse::Parse, parse::ParseStream, Ident, LitStr, Token};
pub struct ToolArgs {
pub description: String,
}
impl Parse for ToolArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
let mut description = None;
while !input.is_empty() {
let key: Ident = input.parse()?;
let _: Token![=] = input.parse()?;
match key.to_string().as_str() {
"name" => {
return Err(syn::Error::new(
key.span(),
"The 'name' attribute is not supported. Use the function name as the tool name instead.",
));
}
"description" => {
let value: LitStr = input.parse()?;
description = Some(value.value());
}
_ => {
return Err(syn::Error::new(
key.span(),
format!("Unknown tool attribute: {key}"),
));
}
}
if input.peek(Token![,]) {
let _: Token![,] = input.parse()?;
}
}
let description = description
.ok_or_else(|| syn::Error::new(input.span(), "Missing required field: description"))?;
Ok(Self { description })
}
}
#[allow(clippy::needless_pass_by_value)] pub fn generate_tool_impl(args: ToolArgs, item: TokenStream) -> TokenStream {
let func: syn::ItemFn = match syn::parse2(item) {
Ok(f) => f,
Err(e) => return e.to_compile_error(),
};
if func.sig.asyncness.is_none() {
return syn::Error::new_spanned(&func.sig, "Tool function must be async")
.to_compile_error();
}
let func_name = &func.sig.ident;
let func_name_str = func_name.to_string();
let tool_name = &func_name_str;
let description = &args.description;
let body = &func.block;
let vis = &func.vis;
let attrs = &func.attrs;
let generics = &func.sig.generics;
let where_clause = &func.sig.generics.where_clause;
let (args_param, ctx_param) = match parse_parameters(&func.sig) {
Ok(params) => params,
Err(e) => return e.to_compile_error(),
};
let args_type = match extract_type(&args_param) {
Ok(ty) => ty,
Err(e) => return e.to_compile_error(),
};
let param_bindings = generate_param_bindings(&args_param, ctx_param.as_ref());
let schema_gen = quote! {
::schemars::schema_for!(#args_type).to_value()
};
quote! {
#[derive(Clone, Copy, Debug)]
#[allow(non_camel_case_types)]
#(#attrs)*
#vis struct #func_name #generics #where_clause;
#[cfg_attr(all(target_os = "wasi", target_env = "p1"), ::async_trait::async_trait(?Send))]
#[cfg_attr(not(all(target_os = "wasi", target_env = "p1")), ::async_trait::async_trait)]
impl #generics ::radkit::tools::BaseTool for #func_name #generics #where_clause {
fn name(&self) -> &str {
#tool_name
}
fn description(&self) -> &str {
#description
}
fn declaration(&self) -> ::radkit::tools::FunctionDeclaration {
::radkit::tools::FunctionDeclaration::new(
#tool_name,
#description,
#schema_gen
)
}
async fn run_async(
&self,
__args: ::std::collections::HashMap<::std::string::String, ::serde_json::Value>,
__ctx: &::radkit::tools::ToolContext<'_>,
) -> ::radkit::tools::ToolResult {
#param_bindings
#body
}
}
}
}
fn parse_parameters(sig: &syn::Signature) -> syn::Result<(syn::FnArg, Option<syn::FnArg>)> {
let mut params = sig.inputs.iter();
let args_param = params
.next()
.ok_or_else(|| {
syn::Error::new_spanned(
sig,
"Tool function must have at least one parameter (args struct)",
)
})?
.clone();
let ctx_param = params.next().cloned();
if let Some(ref param) = ctx_param {
if !is_tool_context_type(param) {
return Err(syn::Error::new_spanned(
param,
"Second parameter must be &ToolContext<'_>",
));
}
}
if params.next().is_some() {
return Err(syn::Error::new_spanned(
sig,
"Tool function can have at most 2 parameters (args struct + optional ToolContext)",
));
}
Ok((args_param, ctx_param))
}
fn is_tool_context_type(param: &syn::FnArg) -> bool {
if let syn::FnArg::Typed(pat_type) = param {
if let syn::Type::Reference(type_ref) = &*pat_type.ty {
if let syn::Type::Path(type_path) = &*type_ref.elem {
let path_str = type_path
.path
.segments
.iter()
.map(|s| s.ident.to_string())
.collect::<Vec<_>>()
.join("::");
return path_str == "ToolContext"
|| path_str == "radkit::tools::ToolContext"
|| path_str == "tools::ToolContext";
}
}
}
false
}
fn extract_type(param: &syn::FnArg) -> syn::Result<&syn::Type> {
if let syn::FnArg::Typed(pat_type) = param {
Ok(&pat_type.ty)
} else {
Err(syn::Error::new_spanned(param, "Expected typed parameter"))
}
}
fn generate_param_bindings(args_param: &syn::FnArg, ctx_param: Option<&syn::FnArg>) -> TokenStream {
let args_binding = if let syn::FnArg::Typed(pat_type) = args_param {
let pat = &pat_type.pat;
let ty = &pat_type.ty;
quote! {
let #pat: #ty = match ::serde_json::from_value(
::serde_json::Value::Object(__args.into_iter().collect())
) {
Ok(val) => val,
Err(e) => {
return ::radkit::tools::ToolResult::error(
format!("Invalid arguments: {}", e)
);
}
};
}
} else {
quote! {}
};
let ctx_binding = if let Some(syn::FnArg::Typed(pat_type)) = ctx_param {
let pat = &pat_type.pat;
quote! {
let #pat = __ctx;
}
} else {
quote! {}
};
quote! {
#args_binding
#ctx_binding
}
}