extern crate proc_macro;
use proc_macro::TokenStream;
use proc_macro2::{Span, TokenStream as TokenStream2};
use quote::{quote, ToTokens};
use syn::{
parse::Parser, parse_macro_input, punctuated::Punctuated, Attribute, Data, DeriveInput, Expr,
ExprLit, FnArg, Ident, ItemFn, Lit, Meta, Pat, Token,
};
fn optional_string(s: &str) -> TokenStream2 {
if s.is_empty() {
quote! { None }
} else {
quote! { Some(String::from(#s)) }
}
}
fn optional_static_str(s: &str) -> TokenStream2 {
if s.is_empty() {
quote! { None }
} else {
quote! { Some(#s) }
}
}
fn is_option_type(ty: &syn::Type) -> bool {
let syn::Type::Path(tp) = ty else {
return false;
};
tp.path
.segments
.last()
.map(|s| s.ident == "Option")
.unwrap_or(false)
}
fn is_result_type(ty: &syn::Type) -> bool {
let syn::Type::Path(tp) = ty else {
return false;
};
tp.path
.segments
.last()
.map(|s| s.ident == "Result")
.unwrap_or(false)
}
fn type_token_string(ty: &syn::Type) -> String {
ty.to_token_stream()
.to_string()
.replace(" < ", "<")
.replace(" > ", ">")
.replace(" >", ">")
.replace(" ,", ",")
.replace(" ::", "::")
.replace(":: ", "::")
}
fn extract_parameter_description(attrs: &[Attribute]) -> Option<String> {
attrs.iter().find_map(|attr| {
if !attr.path().is_ident("function_call_description") {
return None;
}
let Meta::NameValue(nv) = &attr.meta else {
return None;
};
if let Expr::Lit(ExprLit {
lit: Lit::Str(s), ..
}) = &nv.value
{
Some(s.value())
} else {
None
}
})
}
fn collect_doc_comment(attrs: &[Attribute]) -> Option<String> {
let mut lines: Vec<String> = Vec::new();
for attr in attrs {
if !attr.path().is_ident("doc") {
continue;
}
let Meta::NameValue(nv) = &attr.meta else {
continue;
};
if let Expr::Lit(ExprLit {
lit: Lit::Str(s), ..
}) = &nv.value
{
lines.push(s.value().trim().to_string());
}
}
if lines.is_empty() {
None
} else {
Some(lines.join(" ").trim().to_string())
}
}
#[derive(Default)]
struct FieldOpts {
skip: bool,
rename: Option<String>,
description: Option<String>,
}
struct ContainerOpts {
crate_path: String,
}
impl Default for ContainerOpts {
fn default() -> Self {
Self {
crate_path: "::open_ai_rust::logoi::input::tool::raw_macro".to_string(),
}
}
}
fn parse_container_opts(attrs: &[Attribute]) -> syn::Result<ContainerOpts> {
let mut opts = ContainerOpts::default();
for attr in attrs {
if !attr.path().is_ident("function_call") {
continue;
}
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("crate") {
opts.crate_path = meta.value()?.parse::<syn::LitStr>()?.value();
return Ok(());
}
Err(meta.error("unknown function_call container attribute (expected: crate)"))
})?;
}
Ok(opts)
}
fn parse_field_opts(attrs: &[Attribute]) -> syn::Result<FieldOpts> {
let mut opts = FieldOpts::default();
for attr in attrs {
if !attr.path().is_ident("function_call") {
continue;
}
attr.parse_nested_meta(|meta| {
if meta.path.is_ident("skip") {
opts.skip = true;
return Ok(());
}
if meta.path.is_ident("rename") {
opts.rename = Some(meta.value()?.parse::<syn::LitStr>()?.value());
return Ok(());
}
if meta.path.is_ident("description") {
opts.description = Some(meta.value()?.parse::<syn::LitStr>()?.value());
return Ok(());
}
Err(meta.error("unknown function_call attribute (expected: skip, rename, description)"))
})?;
}
Ok(opts)
}
#[proc_macro_derive(FunctionCall, attributes(function_call))]
pub fn turn_type_to_function_call(input: TokenStream) -> TokenStream {
let input = parse_macro_input!(input as DeriveInput);
match derive_impl(input) {
Ok(ts) => ts.into(),
Err(e) => e.to_compile_error().into(),
}
}
struct FieldInfo {
ty: syn::Type,
name_out: String,
description: String,
required: bool,
}
fn collect_fields(s: &syn::DataStruct) -> syn::Result<Vec<FieldInfo>> {
let mut out = Vec::new();
for (i, f) in s.fields.iter().enumerate() {
let opts = parse_field_opts(&f.attrs)?;
if opts.skip {
continue;
}
let default_name = match &f.ident {
Some(id) => id.to_string(),
None => format!("_{i}"),
};
let name_out = opts.rename.unwrap_or(default_name);
let description = opts
.description
.or_else(|| collect_doc_comment(&f.attrs))
.unwrap_or_default();
let required = !is_option_type(&f.ty);
out.push(FieldInfo {
ty: f.ty.clone(),
name_out,
description,
required,
});
}
Ok(out)
}
fn emit_param_vec<'a>(
fields: &'a [FieldInfo],
trait_path: &syn::Path,
) -> impl Iterator<Item = TokenStream2> + 'a {
let tp = trait_path.clone();
fields.iter().map(move |f| {
let name = &f.name_out;
let desc = optional_string(&f.description);
let req = f.required;
let ty = &f.ty;
quote! {
FunctionParameter {
name: String::from(#name),
_type: <#ty as #tp ::FunctionCallable>::schema_type(),
description: #desc,
required: #req,
}
}
})
}
fn derive_impl(input: DeriveInput) -> syn::Result<TokenStream2> {
let name = &input.ident;
let struct_description = collect_doc_comment(&input.attrs).unwrap_or_default();
let (impl_generics, ty_generics, where_clause) = input.generics.split_for_impl();
let container = parse_container_opts(&input.attrs)?;
let trait_path: syn::Path = syn::parse_str(&container.crate_path)?;
let struct_desc_tokens = optional_string(&struct_description);
match &input.data {
Data::Struct(s) => {
let fields = collect_fields(s)?;
let params: Vec<_> = emit_param_vec(&fields, &trait_path).collect();
let params2 = params.clone();
Ok(quote! {
impl #impl_generics #trait_path ::FunctionCallable for #name #ty_generics #where_clause {
fn schema_type() -> FunctionType where Self: ::std::marker::Sized {
FunctionType::Object(vec![ #( #params, )* ])
}
fn fn_schema() -> FunctionCall where Self: ::std::marker::Sized {
FunctionCall {
name: String::from(stringify!(#name)),
description: #struct_desc_tokens,
parameters: vec![ #( #params2, )* ],
}
}
}
})
}
Data::Enum(e) => {
let variants: Vec<&Ident> = e
.variants
.iter()
.map(|v| {
if !matches!(v.fields, syn::Fields::Unit) {
return Err(syn::Error::new_spanned(
v,
"FunctionCall derive currently only supports unit enum variants; \
data-variant enums (oneOf) are planned for a future release",
));
}
Ok(&v.ident)
})
.collect::<syn::Result<_>>()?;
let variant_strings: Vec<String> = variants.iter().map(|v| v.to_string()).collect();
Ok(quote! {
impl #impl_generics #trait_path ::FunctionCallable for #name #ty_generics #where_clause {
fn schema_type() -> FunctionType where Self: ::std::marker::Sized {
FunctionType::Enum(EnumValues::String(vec![
#( String::from(#variant_strings) ),*
]))
}
fn fn_schema() -> FunctionCall where Self: ::std::marker::Sized {
FunctionCall {
name: String::from(stringify!(#name)),
description: #struct_desc_tokens,
parameters: vec![],
}
}
}
})
}
Data::Union(u) => Err(syn::Error::new_spanned(
u.union_token,
"FunctionCall derive does not support unions",
)),
}
}
struct ParamMeta {
ident: Ident,
name: String,
ty: syn::Type,
ty_string: String,
description: String,
}
fn extract_param_metadata(input: &ItemFn) -> Vec<ParamMeta> {
let mut out = Vec::new();
for arg in input.sig.inputs.iter() {
let FnArg::Typed(pat_type) = arg else {
continue;
};
let Pat::Ident(pi) = &*pat_type.pat else {
continue;
};
let ident = pi.ident.clone();
let name = ident.to_string();
let ty = (*pat_type.ty).clone();
let ty_string = type_token_string(&pat_type.ty);
let description = extract_parameter_description(&pat_type.attrs).unwrap_or_default();
out.push(ParamMeta {
ident,
name,
ty,
ty_string,
description,
});
}
out
}
fn strip_helper_attrs(input: &mut ItemFn) {
for arg in input.sig.inputs.iter_mut() {
if let FnArg::Typed(pat_type) = arg {
pat_type
.attrs
.retain(|a| !a.path().is_ident("function_call_description"));
}
}
}
fn build_schema_const(fn_name: &str, description: &str, params: &[ParamMeta]) -> TokenStream2 {
let fn_name_uppercase = Ident::new(&fn_name.to_uppercase(), Span::call_site());
let param_lits = params.iter().map(|p| {
let n = &p.name;
let t = &p.ty_string;
let d = &p.description;
quote! { FunctionParamRaw { name: #n, ty: #t, description: #d } }
});
quote! {
const #fn_name_uppercase: FunctionCallRaw<'static> = FunctionCallRaw {
name: #fn_name,
description: #description,
parameters: &[ #( #param_lits ),* ],
};
}
}
fn parse_attr_description(attr: TokenStream2, fn_attrs: &[Attribute]) -> syn::Result<String> {
let parser = Punctuated::<Expr, Token![,]>::parse_terminated;
let attr_args = parser.parse2(attr)?;
Ok(attr_args
.iter()
.find_map(|e| match e {
Expr::Lit(ExprLit {
lit: Lit::Str(s), ..
}) => Some(s.value()),
_ => None,
})
.or_else(|| collect_doc_comment(fn_attrs))
.unwrap_or_default())
}
#[proc_macro_attribute]
pub fn function_call(attr: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemFn);
match function_call_impl(attr.into(), input) {
Ok(ts) => ts.into(),
Err(e) => e.to_compile_error().into(),
}
}
fn function_call_impl(attr: TokenStream2, mut input: ItemFn) -> syn::Result<TokenStream2> {
let fn_name = input.sig.ident.to_string();
let description = parse_attr_description(attr, &input.attrs)?;
let params = extract_param_metadata(&input);
strip_helper_attrs(&mut input);
let const_body = build_schema_const(&fn_name, &description, ¶ms);
Ok(quote! { #const_body #input })
}
#[proc_macro_attribute]
pub fn tool(attr: TokenStream, item: TokenStream) -> TokenStream {
let input = parse_macro_input!(item as ItemFn);
match tool_impl(attr.into(), input) {
Ok(ts) => ts.into(),
Err(e) => e.to_compile_error().into(),
}
}
fn tool_impl(attr: TokenStream2, mut input: ItemFn) -> syn::Result<TokenStream2> {
let fn_name = input.sig.ident.to_string();
let fn_ident = input.sig.ident.clone();
let description = parse_attr_description(attr, &input.attrs)?;
let params = extract_param_metadata(&input);
strip_helper_attrs(&mut input);
let const_body = build_schema_const(&fn_name, &description, ¶ms);
let is_async = input.sig.asyncness.is_some();
let vis = &input.vis;
let dispatch_ident = Ident::new(&format!("{fn_name}_dispatch"), Span::call_site());
let arg_extractions: Vec<TokenStream2> = params
.iter()
.map(|p| {
let ident = &p.ident;
let name_lit = &p.name;
let ty = &p.ty;
quote! {
let #ident: #ty = ::serde_json::from_value(
args.get(#name_lit).cloned().unwrap_or(::serde_json::Value::Null)
).map_err(|e| e.to_string())?;
}
})
.collect();
let arg_idents: Vec<&Ident> = params.iter().map(|p| &p.ident).collect();
let ret_is_result = match &input.sig.output {
syn::ReturnType::Type(_, ty) => is_result_type(ty),
_ => false,
};
let raw_call = if is_async {
quote! { #fn_ident( #( #arg_idents ),* ).await }
} else {
quote! { #fn_ident( #( #arg_idents ),* ) }
};
let serialize_expr = if ret_is_result {
quote! {
let __r = #raw_call.map_err(|e| format!("{e}"))?;
::serde_json::to_value(__r).map_err(|e| e.to_string())
}
} else {
quote! {
let __r = #raw_call;
::serde_json::to_value(__r).map_err(|e| e.to_string())
}
};
let dispatch_fn = quote! {
#vis fn #dispatch_ident(
args: ::serde_json::Value,
) -> ::std::pin::Pin<
::std::boxed::Box<
dyn ::std::future::Future<
Output = ::std::result::Result<::serde_json::Value, String>
> + Send + 'static
>
>
{
::std::boxed::Box::pin(async move {
#( #arg_extractions )*
#serialize_expr
})
}
};
let tool_registry_entry = if cfg!(feature = "tool_registry") {
let desc_tokens = optional_static_str(&description);
let schema_helper = Ident::new(&format!("__oa_tool_schema_{fn_name}"), Span::call_site());
let dispatch_helper =
Ident::new(&format!("__oa_tool_dispatch_{fn_name}"), Span::call_site());
let static_ident = Ident::new(
&format!("__OA_TOOL_{}", fn_name.to_uppercase()),
Span::call_site(),
);
let schema_params = params.iter().map(|p| {
let n = &p.name;
let d = optional_string(&p.description);
let ty = &p.ty;
let req = !is_option_type(&p.ty);
quote! {
::open_ai_rust::logoi::input::tool::FunctionParameter {
name: String::from(#n),
_type: <#ty as ::open_ai_rust::logoi::input::tool::raw_macro::FunctionCallable>::schema_type(),
description: #d,
required: #req,
}
}
});
let fn_description = optional_string(&description);
quote! {
fn #schema_helper() -> ::open_ai_rust::logoi::input::tool::FunctionCall {
::open_ai_rust::logoi::input::tool::FunctionCall {
name: String::from(#fn_name),
description: #fn_description,
parameters: vec![ #( #schema_params, )* ],
}
}
fn #dispatch_helper(
args: ::serde_json::Value,
) -> ::open_ai_rust::tool_registry::DispatchFuture {
#dispatch_ident(args)
}
#[::open_ai_rust::__macro_support::linkme::distributed_slice(
::open_ai_rust::tool_registry::TOOLS
)]
#[linkme(crate = ::open_ai_rust::__macro_support::linkme)]
static #static_ident: ::open_ai_rust::tool_registry::ToolEntry =
::open_ai_rust::tool_registry::ToolEntry {
name: #fn_name,
description: #desc_tokens,
schema: #schema_helper,
dispatch: #dispatch_helper,
};
}
} else {
quote! {}
};
Ok(quote! {
#const_body
#dispatch_fn
#tool_registry_entry
#input
})
}