use proc_macro2::TokenStream;
use quote::quote;
use syn::{Ident, Result, Token, Type, parse::Parse, parse::ParseStream};
struct ToolsList {
context_type: Option<Type>,
tools: Vec<Ident>,
}
impl Parse for ToolsList {
fn parse(input: ParseStream) -> Result<Self> {
let context_type = if input.peek2(Token![=>])
|| (input.peek(syn::Ident) && {
let fork = input.fork();
fork.parse::<Type>().is_ok() && fork.peek(Token![=>])
}) {
let ty = input.parse::<Type>()?;
input.parse::<Token![=>]>()?;
Some(ty)
} else {
None
};
let mut tools = Vec::new();
while !input.is_empty() {
tools.push(input.parse::<Ident>()?);
if input.peek(Token![,]) {
input.parse::<Token![,]>()?;
} else {
break;
}
}
Ok(ToolsList {
context_type,
tools,
})
}
}
pub fn tools_impl(input: TokenStream) -> Result<TokenStream> {
let tools_list = syn::parse2::<ToolsList>(input)?;
if tools_list.tools.is_empty() {
return Err(syn::Error::new(
proc_macro2::Span::call_site(),
"toolset! macro requires at least one tool function",
));
}
for (index, tool) in tools_list.tools.iter().enumerate() {
if tools_list.tools[..index].contains(tool) {
return Err(syn::Error::new(
tool.span(),
format!("tool '{}' is listed more than once", tool),
));
}
}
let wrapper_names: Vec<_> = tools_list
.tools
.iter()
.map(|tool_name| {
quote::format_ident!(
"{}Tool",
crate::common::to_pascal_case(&tool_name.to_string())
)
})
.collect();
let expanded = if let Some(ctx_type) = tools_list.context_type {
quote! {
{
use rsai::{ToolFunction, ToolSetBuilder};
let mut builder = ToolSetBuilder::<#ctx_type>::new();
#(
builder = builder.add_tool(std::sync::Arc::new(#wrapper_names));
)*
builder
}
}
} else {
quote! {
{
use rsai::{Tool, ToolChoice, ToolFunction, ToolRegistry, ToolSet};
let registry = ToolRegistry::new();
#(
registry.register(std::sync::Arc::new(#wrapper_names))
.expect("toolset! guarantees unique tool names");
)*
ToolSet {
registry,
}
}
}
};
Ok(expanded)
}