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}