Skip to main content

stasis_macros/
lib.rs

1use proc_macro::TokenStream;
2
3use quote::{format_ident, quote};
4use syn::parse::Parser;
5use syn::{
6    Expr, ExprLit, FnArg, ItemFn, Lit, LitStr, MetaNameValue, PathArguments, ReturnType, Type,
7    punctuated::Punctuated,
8};
9
10#[proc_macro_attribute]
11pub fn stasis_tool(attr: TokenStream, item: TokenStream) -> TokenStream {
12    let parser = Punctuated::<MetaNameValue, syn::Token![,]>::parse_terminated;
13    let args = match parser.parse(attr) {
14        Ok(args) => args,
15        Err(err) => return err.to_compile_error().into(),
16    };
17
18    let mut tool_name: Option<LitStr> = None;
19    let mut description: Option<LitStr> = None;
20    let mut crate_path_literal: Option<LitStr> = None;
21    let mut output_schema_enabled = false;
22
23    for arg in args {
24        if arg.path.is_ident("name") {
25            match arg.value {
26                Expr::Lit(ExprLit {
27                    lit: Lit::Str(value),
28                    ..
29                }) => {
30                    tool_name = Some(value);
31                }
32                _ => {
33                    return syn::Error::new_spanned(arg.value, "name must be a string literal")
34                        .to_compile_error()
35                        .into();
36                }
37            }
38            continue;
39        }
40
41        if arg.path.is_ident("description") {
42            match arg.value {
43                Expr::Lit(ExprLit {
44                    lit: Lit::Str(value),
45                    ..
46                }) => {
47                    description = Some(value);
48                }
49                _ => {
50                    return syn::Error::new_spanned(
51                        arg.value,
52                        "description must be a string literal",
53                    )
54                    .to_compile_error()
55                    .into();
56                }
57            }
58            continue;
59        }
60
61        if arg.path.is_ident("crate_path") {
62            match arg.value {
63                Expr::Lit(ExprLit {
64                    lit: Lit::Str(value),
65                    ..
66                }) => {
67                    crate_path_literal = Some(value);
68                }
69                _ => {
70                    return syn::Error::new_spanned(
71                        arg.value,
72                        "crate_path must be a string literal",
73                    )
74                    .to_compile_error()
75                    .into();
76                }
77            }
78            continue;
79        }
80
81        if arg.path.is_ident("output_schema") {
82            match arg.value {
83                Expr::Lit(ExprLit {
84                    lit: Lit::Bool(value),
85                    ..
86                }) => {
87                    output_schema_enabled = value.value;
88                }
89                _ => {
90                    return syn::Error::new_spanned(
91                        arg.value,
92                        "output_schema must be a bool literal",
93                    )
94                    .to_compile_error()
95                    .into();
96                }
97            }
98            continue;
99        }
100
101        return syn::Error::new_spanned(
102            arg.path,
103            "unsupported attribute key (expected: name, description, crate_path, output_schema)",
104        )
105        .to_compile_error()
106        .into();
107    }
108
109    let item_fn = syn::parse_macro_input!(item as ItemFn);
110    let fn_ident = item_fn.sig.ident.clone();
111
112    let tool_name = match tool_name {
113        Some(name) => name,
114        None => {
115            return syn::Error::new_spanned(
116                item_fn.sig.ident,
117                "missing required attribute argument: name = \"...\"",
118            )
119            .to_compile_error()
120            .into();
121        }
122    };
123
124    if item_fn.sig.asyncness.is_none() {
125        return syn::Error::new_spanned(item_fn.sig.fn_token, "stasis_tool function must be async")
126            .to_compile_error()
127            .into();
128    }
129
130    if !item_fn.sig.generics.params.is_empty() || item_fn.sig.generics.where_clause.is_some() {
131        return syn::Error::new_spanned(
132            item_fn.sig.generics,
133            "stasis_tool functions must not use generics",
134        )
135        .to_compile_error()
136        .into();
137    }
138
139    if item_fn.sig.inputs.len() != 1 {
140        return syn::Error::new_spanned(
141            item_fn.sig.inputs,
142            "stasis_tool function must accept exactly one typed argument",
143        )
144        .to_compile_error()
145        .into();
146    }
147
148    let input_ty = match item_fn.sig.inputs.first() {
149        Some(FnArg::Typed(pat_ty)) => pat_ty.ty.clone(),
150        Some(FnArg::Receiver(receiver)) => {
151            return syn::Error::new_spanned(
152                receiver,
153                "stasis_tool function must not use a self receiver",
154            )
155            .to_compile_error()
156            .into();
157        }
158        None => unreachable!(),
159    };
160
161    let output_ty = match extract_result_output_type(&item_fn.sig.output) {
162        Ok(output) => output,
163        Err(err) => return err.to_compile_error().into(),
164    };
165
166    let struct_name = format_ident!("{}Tool", to_pascal_case(&fn_ident.to_string()));
167    let ctor_name = format_ident!("{}_tool", fn_ident);
168
169    let crate_path_lit = crate_path_literal.unwrap_or_else(|| LitStr::new("stasis", fn_ident.span()));
170    let crate_path = match syn::parse_str::<syn::Path>(&crate_path_lit.value()) {
171        Ok(path) => path,
172        Err(err) => return err.to_compile_error().into(),
173    };
174
175    let description_expr = match description {
176        Some(value) => quote! { ::core::option::Option::Some(#value) },
177        None => quote! { ::core::option::Option::None },
178    };
179
180    let output_schema_impl = if output_schema_enabled {
181        quote! {
182            fn output_schema(&self) -> ::core::option::Option<#crate_path::macro_support::serde_json::Value> {
183                fn __assert_output_schema_traits<T: #crate_path::macro_support::schemars::JsonSchema>() {}
184                __assert_output_schema_traits::<#output_ty>();
185
186                let schema = #crate_path::macro_support::schemars::schema_for!(#output_ty);
187                #crate_path::macro_support::serde_json::to_value(schema.schema).ok()
188            }
189        }
190    } else {
191        quote! {}
192    };
193
194    let expanded = quote! {
195        #item_fn
196
197        #[derive(Clone, Copy, Debug, Default)]
198        pub struct #struct_name;
199
200        #[#crate_path::macro_support::async_trait::async_trait]
201        impl #crate_path::application::orchestration::tool_registry::StasisTool for #struct_name {
202            fn name(&self) -> &'static str {
203                #tool_name
204            }
205
206            fn description(&self) -> ::core::option::Option<&'static str> {
207                #description_expr
208            }
209
210            fn input_schema(&self) -> ::core::option::Option<#crate_path::macro_support::serde_json::Value> {
211                fn __assert_input_traits<T: #crate_path::macro_support::schemars::JsonSchema + #crate_path::macro_support::serde::de::DeserializeOwned>() {}
212                __assert_input_traits::<#input_ty>();
213
214                let schema = #crate_path::macro_support::schemars::schema_for!(#input_ty);
215                #crate_path::macro_support::serde_json::to_value(schema.schema).ok()
216            }
217
218            #output_schema_impl
219
220            async fn invoke(
221                &self,
222                input: #crate_path::macro_support::serde_json::Value,
223            ) -> #crate_path::domain::errors::Result<#crate_path::macro_support::serde_json::Value> {
224                fn __assert_output_traits<T: #crate_path::macro_support::serde::Serialize>() {}
225                __assert_output_traits::<#output_ty>();
226
227                let parsed_input: #input_ty = #crate_path::macro_support::serde_json::from_value(input).map_err(|err| {
228                    #crate_path::domain::errors::StasisError::PortFailure(
229                        format!("invalid input for tool '{}': {}", #tool_name, err)
230                    )
231                })?;
232
233                let output: #output_ty = #fn_ident(parsed_input).await?;
234
235                #crate_path::macro_support::serde_json::to_value(output).map_err(|err| {
236                    #crate_path::domain::errors::StasisError::PortFailure(
237                        format!("failed to serialize output for tool '{}': {}", #tool_name, err)
238                    )
239                })
240            }
241        }
242
243        pub fn #ctor_name() -> #struct_name {
244            #struct_name
245        }
246    };
247
248    expanded.into()
249}
250
251fn extract_result_output_type(output: &ReturnType) -> syn::Result<Type> {
252    let ReturnType::Type(_, ty) = output else {
253        return Err(syn::Error::new_spanned(
254            output,
255            "stasis_tool function must return Result<OutputType>",
256        ));
257    };
258
259    let Type::Path(type_path) = ty.as_ref() else {
260        return Err(syn::Error::new_spanned(
261            ty,
262            "stasis_tool function must return Result<OutputType>",
263        ));
264    };
265
266    let Some(segment) = type_path.path.segments.last() else {
267        return Err(syn::Error::new_spanned(
268            type_path,
269            "unable to parse function return type",
270        ));
271    };
272
273    if segment.ident != "Result" {
274        return Err(syn::Error::new_spanned(
275            segment,
276            "stasis_tool function must return Result<OutputType>",
277        ));
278    }
279
280    let PathArguments::AngleBracketed(args) = &segment.arguments else {
281        return Err(syn::Error::new_spanned(
282            segment,
283            "stasis_tool function must return Result<OutputType>",
284        ));
285    };
286
287    let Some(first_arg) = args.args.first() else {
288        return Err(syn::Error::new_spanned(
289            args,
290            "stasis_tool function must return Result<OutputType>",
291        ));
292    };
293
294    let syn::GenericArgument::Type(output_ty) = first_arg else {
295        return Err(syn::Error::new_spanned(
296            first_arg,
297            "stasis_tool function must return Result<OutputType>",
298        ));
299    };
300
301    Ok(output_ty.clone())
302}
303
304fn to_pascal_case(value: &str) -> String {
305    let mut output = String::new();
306
307    for part in value
308        .split(|ch: char| !ch.is_ascii_alphanumeric())
309        .filter(|part| !part.is_empty())
310    {
311        let mut chars = part.chars();
312        if let Some(first) = chars.next() {
313            output.push(first.to_ascii_uppercase());
314            output.push_str(chars.as_str());
315        }
316    }
317
318    if output.is_empty() {
319        "StasisToolGenerated".to_string()
320    } else {
321        output
322    }
323}