1use proc_macro::TokenStream;
2use quote::{format_ident, quote};
3use syn::parse::{Parse, ParseStream};
4use syn::{
5 parse_macro_input, Attribute, FnArg, Ident, ImplItem, ImplItemFn, Item, ItemImpl, LitInt,
6 LitStr, Pat, PatType, Result, Token,
7};
8
9#[proc_macro_attribute]
10pub fn injectable(attr: TokenStream, item: TokenStream) -> TokenStream {
11 if !attr.is_empty() {
12 return syn::Error::new(
13 proc_macro2::TokenStream::from(attr)
14 .into_iter()
15 .next()
16 .unwrap()
17 .span(),
18 "#[injectable] does not accept arguments",
19 )
20 .to_compile_error()
21 .into();
22 }
23
24 let item = parse_macro_input!(item as Item);
25 match item {
26 Item::Struct(item_struct) => expand_injectable(item_struct)
27 .unwrap_or_else(syn::Error::into_compile_error)
28 .into(),
29 item => syn::Error::new_spanned(item, "#[injectable] can only be used on structs")
30 .to_compile_error()
31 .into(),
32 }
33}
34
35#[proc_macro_attribute]
36pub fn controller(attr: TokenStream, item: TokenStream) -> TokenStream {
37 let prefix = parse_macro_input!(attr as LitStr);
38 let item_impl = parse_macro_input!(item as ItemImpl);
39
40 expand_controller(prefix, item_impl)
41 .unwrap_or_else(syn::Error::into_compile_error)
42 .into()
43}
44
45#[proc_macro_attribute]
46pub fn get(_attr: TokenStream, item: TokenStream) -> TokenStream {
47 route_attribute_outside_controller("get", item)
48}
49
50#[proc_macro_attribute]
51pub fn sse(_attr: TokenStream, item: TokenStream) -> TokenStream {
52 route_attribute_outside_controller("sse", item)
53}
54
55#[proc_macro_attribute]
56pub fn post(_attr: TokenStream, item: TokenStream) -> TokenStream {
57 route_attribute_outside_controller("post", item)
58}
59
60#[proc_macro_attribute]
61pub fn put(_attr: TokenStream, item: TokenStream) -> TokenStream {
62 route_attribute_outside_controller("put", item)
63}
64
65#[proc_macro_attribute]
66pub fn patch(_attr: TokenStream, item: TokenStream) -> TokenStream {
67 route_attribute_outside_controller("patch", item)
68}
69
70#[proc_macro_attribute]
71pub fn delete(_attr: TokenStream, item: TokenStream) -> TokenStream {
72 route_attribute_outside_controller("delete", item)
73}
74
75#[proc_macro_attribute]
76pub fn options(_attr: TokenStream, item: TokenStream) -> TokenStream {
77 route_attribute_outside_controller("options", item)
78}
79
80#[proc_macro_attribute]
81pub fn head(_attr: TokenStream, item: TokenStream) -> TokenStream {
82 route_attribute_outside_controller("head", item)
83}
84
85#[proc_macro_attribute]
86pub fn get_json(_attr: TokenStream, item: TokenStream) -> TokenStream {
87 route_attribute_outside_controller("get_json", item)
88}
89
90#[proc_macro_attribute]
91pub fn post_json(_attr: TokenStream, item: TokenStream) -> TokenStream {
92 route_attribute_outside_controller("post_json", item)
93}
94
95#[proc_macro_attribute]
96pub fn put_json(_attr: TokenStream, item: TokenStream) -> TokenStream {
97 route_attribute_outside_controller("put_json", item)
98}
99
100#[proc_macro_attribute]
101pub fn patch_json(_attr: TokenStream, item: TokenStream) -> TokenStream {
102 route_attribute_outside_controller("patch_json", item)
103}
104
105#[proc_macro_attribute]
106pub fn delete_json(_attr: TokenStream, item: TokenStream) -> TokenStream {
107 route_attribute_outside_controller("delete_json", item)
108}
109
110fn expand_injectable(item_struct: syn::ItemStruct) -> Result<proc_macro2::TokenStream> {
111 let ident = &item_struct.ident;
112 let (impl_generics, ty_generics, where_clause) = item_struct.generics.split_for_impl();
113
114 Ok(quote! {
115 #item_struct
116
117 impl #impl_generics #ident #ty_generics #where_clause {
118 pub fn into_provider(self) -> ::a3s_boot::ProviderDefinition
119 where
120 Self: ::std::marker::Send + ::std::marker::Sync + 'static,
121 {
122 ::a3s_boot::ProviderDefinition::singleton(self)
123 }
124
125 pub fn into_named_provider(self, token: impl Into<String>) -> ::a3s_boot::ProviderDefinition
126 where
127 Self: ::std::marker::Send + ::std::marker::Sync + 'static,
128 {
129 ::a3s_boot::ProviderDefinition::named_singleton(token, self)
130 }
131
132 pub fn from_arc_provider(value: ::std::sync::Arc<Self>) -> ::a3s_boot::ProviderDefinition
133 where
134 Self: ::std::marker::Send + ::std::marker::Sync + 'static,
135 {
136 ::a3s_boot::ProviderDefinition::from_arc(value)
137 }
138
139 pub fn from_named_arc_provider(
140 token: impl Into<String>,
141 value: ::std::sync::Arc<Self>,
142 ) -> ::a3s_boot::ProviderDefinition
143 where
144 Self: ::std::marker::Send + ::std::marker::Sync + 'static,
145 {
146 ::a3s_boot::ProviderDefinition::named_from_arc(token, value)
147 }
148 }
149 })
150}
151
152fn expand_controller(prefix: LitStr, mut item_impl: ItemImpl) -> Result<proc_macro2::TokenStream> {
153 if item_impl.trait_.is_some() {
154 return Err(syn::Error::new_spanned(
155 &item_impl,
156 "#[controller] can only be used on inherent impl blocks",
157 ));
158 }
159
160 let self_ty = item_impl.self_ty.clone();
161 let mut routes = Vec::new();
162 let mut errors: Option<syn::Error> = None;
163
164 for item in &mut item_impl.items {
165 let ImplItem::Fn(method) = item else {
166 continue;
167 };
168
169 let (clean_attrs, method_routes, route_errors) = take_route_attrs(&method.attrs);
170 method.attrs = clean_attrs;
171 for error in route_errors {
172 push_error(&mut errors, error);
173 }
174
175 for route in method_routes {
176 match route_registration(route, method) {
177 Ok(registration) => routes.push(registration),
178 Err(error) => push_error(&mut errors, error),
179 }
180 }
181 }
182
183 if let Some(error) = errors {
184 return Err(error);
185 }
186
187 Ok(quote! {
188 #item_impl
189
190 impl #self_ty {
191 pub fn controller(
192 self: ::std::sync::Arc<Self>,
193 ) -> ::a3s_boot::Result<::a3s_boot::ControllerDefinition> {
194 let mut __a3s_boot_controller =
195 ::a3s_boot::ControllerDefinition::new(#prefix)?;
196 #(
197 __a3s_boot_controller = #routes;
198 )*
199 Ok(__a3s_boot_controller)
200 }
201 }
202 })
203}
204
205fn take_route_attrs(attrs: &[Attribute]) -> (Vec<Attribute>, Vec<RouteSpec>, Vec<syn::Error>) {
206 let mut clean_attrs = Vec::new();
207 let mut routes = Vec::new();
208 let mut errors = Vec::new();
209
210 for attr in attrs {
211 let Some(kind) = RouteKind::from_attribute(attr) else {
212 clean_attrs.push(attr.clone());
213 continue;
214 };
215
216 match attr.parse_args::<RouteArgs>() {
217 Ok(args) => routes.push(RouteSpec { kind, args }),
218 Err(error) => errors.push(error),
219 }
220 }
221
222 (clean_attrs, routes, errors)
223}
224
225fn route_registration(route: RouteSpec, method: &ImplItemFn) -> Result<proc_macro2::TokenStream> {
226 if method.sig.asyncness.is_none() {
227 return Err(syn::Error::new_spanned(
228 &method.sig.fn_token,
229 "controller route methods must be async",
230 ));
231 }
232
233 let method_ident = &method.sig.ident;
234 let input = RouteMethodInput::from_method(method)?;
235 let status = route.args.status_value()?;
236 let path = route.args.path;
237
238 let raw = route.args.raw.is_some();
239 if raw && route.kind.is_explicit_json() {
240 return Err(syn::Error::new_spanned(
241 route.args.raw.unwrap(),
242 "raw is not supported on *_json route attributes",
243 ));
244 }
245
246 match route.kind.flavor(raw) {
247 RouteFlavor::Sse => {
248 if route.args.status.is_some() {
249 return Err(syn::Error::new_spanned(
250 route.args.status.unwrap(),
251 "status is not supported on SSE route attributes",
252 ));
253 }
254 if route.args.raw.is_some() {
255 return Err(syn::Error::new_spanned(
256 route.args.raw.unwrap(),
257 "raw is not supported on SSE route attributes",
258 ));
259 }
260 let handler = raw_or_json_request_handler(method_ident, input);
261 Ok(quote! {
262 __a3s_boot_controller.sse(#path, #handler)?
263 })
264 }
265 RouteFlavor::Raw => {
266 if route.args.status.is_some() {
267 return Err(syn::Error::new_spanned(
268 route.args.status.unwrap(),
269 "status is only supported on JSON route attributes",
270 ));
271 }
272 let builder = route.kind.raw_builder_ident();
273 let handler = raw_or_json_request_handler(method_ident, input);
274 Ok(quote! {
275 __a3s_boot_controller.#builder(#path, #handler)?
276 })
277 }
278 RouteFlavor::JsonRequest => {
279 let builder = route.kind.json_builder_ident().ok_or_else(|| {
280 syn::Error::new_spanned(
281 &method.sig.ident,
282 "this HTTP method does not support JSON route inference",
283 )
284 })?;
285 let handler = raw_or_json_request_handler(method_ident, input);
286 Ok(quote! {
287 __a3s_boot_controller.#builder(#path, #status, #handler)?
288 })
289 }
290 RouteFlavor::JsonBody => {
291 let Some(input) = input.arg else {
292 return Err(syn::Error::new_spanned(
293 &method.sig.ident,
294 "JSON body routes must accept one DTO argument after &self",
295 ));
296 };
297 let builder = route.kind.json_builder_ident().ok_or_else(|| {
298 syn::Error::new_spanned(
299 &method.sig.ident,
300 "this HTTP method does not support JSON route inference",
301 )
302 })?;
303 let handler = json_body_handler(method_ident, input);
304 Ok(quote! {
305 __a3s_boot_controller.#builder(#path, #status, #handler)?
306 })
307 }
308 }
309}
310
311fn raw_or_json_request_handler(
312 method_ident: &Ident,
313 input: RouteMethodInput,
314) -> proc_macro2::TokenStream {
315 let controller_name = format_ident!("__a3s_boot_{}", method_ident);
316 match input.arg {
317 Some(MethodArg { ident, ty }) => quote! {
318 {
319 let #controller_name = ::std::sync::Arc::clone(&self);
320 move |#ident: #ty| {
321 let #controller_name = ::std::sync::Arc::clone(&#controller_name);
322 async move { #controller_name.#method_ident(#ident).await }
323 }
324 }
325 },
326 None => quote! {
327 {
328 let #controller_name = ::std::sync::Arc::clone(&self);
329 move |_request: ::a3s_boot::BootRequest| {
330 let #controller_name = ::std::sync::Arc::clone(&#controller_name);
331 async move { #controller_name.#method_ident().await }
332 }
333 }
334 },
335 }
336}
337
338fn json_body_handler(method_ident: &Ident, input: MethodArg) -> proc_macro2::TokenStream {
339 let controller_name = format_ident!("__a3s_boot_{}", method_ident);
340 let MethodArg { ident, ty } = input;
341 quote! {
342 {
343 let #controller_name = ::std::sync::Arc::clone(&self);
344 move |#ident: #ty| {
345 let #controller_name = ::std::sync::Arc::clone(&#controller_name);
346 async move { #controller_name.#method_ident(#ident).await }
347 }
348 }
349 }
350}
351
352fn route_attribute_outside_controller(name: &str, item: TokenStream) -> TokenStream {
353 let item = proc_macro2::TokenStream::from(item);
354 let message =
355 format!("#[{name}] must be used inside an impl block annotated with #[controller]");
356 quote! {
357 compile_error!(#message);
358 #item
359 }
360 .into()
361}
362
363fn push_error(slot: &mut Option<syn::Error>, error: syn::Error) {
364 if let Some(existing) = slot {
365 existing.combine(error);
366 } else {
367 *slot = Some(error);
368 }
369}
370
371struct RouteArgs {
372 path: LitStr,
373 status: Option<LitInt>,
374 raw: Option<Ident>,
375}
376
377impl RouteArgs {
378 fn status_value(&self) -> Result<proc_macro2::TokenStream> {
379 let Some(status) = &self.status else {
380 return Ok(quote!(200));
381 };
382 let value = status.base10_parse::<u16>()?;
383 Ok(quote!(#value))
384 }
385}
386
387impl Parse for RouteArgs {
388 fn parse(input: ParseStream<'_>) -> Result<Self> {
389 let path = input.parse::<LitStr>()?;
390 let mut status = None;
391 let mut raw = None;
392
393 if !input.is_empty() {
394 while !input.is_empty() {
395 input.parse::<Token![,]>()?;
396 let name = input.parse::<Ident>()?;
397
398 if name == "status" {
399 if status.is_some() {
400 return Err(syn::Error::new_spanned(name, "duplicate `status` option"));
401 }
402 input.parse::<Token![=]>()?;
403 status = Some(input.parse::<LitInt>()?);
404 } else if name == "raw" {
405 if raw.is_some() {
406 return Err(syn::Error::new_spanned(name, "duplicate `raw` option"));
407 }
408 raw = Some(name);
409 } else {
410 return Err(syn::Error::new_spanned(
411 name,
412 "expected `status = <u16>` or `raw`",
413 ));
414 }
415 }
416 }
417
418 if !input.is_empty() {
419 return Err(input.error("unexpected route attribute arguments"));
420 }
421
422 Ok(Self { path, status, raw })
423 }
424}
425
426struct RouteSpec {
427 kind: RouteKind,
428 args: RouteArgs,
429}
430
431#[derive(Clone, Copy)]
432enum RouteKind {
433 Get,
434 Sse,
435 Post,
436 Put,
437 Patch,
438 Delete,
439 Options,
440 Head,
441 GetJson,
442 PostJson,
443 PutJson,
444 PatchJson,
445 DeleteJson,
446}
447
448impl RouteKind {
449 fn from_attribute(attr: &Attribute) -> Option<Self> {
450 let ident = attr.path().segments.last()?.ident.to_string();
451 match ident.as_str() {
452 "get" => Some(Self::Get),
453 "sse" => Some(Self::Sse),
454 "post" => Some(Self::Post),
455 "put" => Some(Self::Put),
456 "patch" => Some(Self::Patch),
457 "delete" => Some(Self::Delete),
458 "options" => Some(Self::Options),
459 "head" => Some(Self::Head),
460 "get_json" => Some(Self::GetJson),
461 "post_json" => Some(Self::PostJson),
462 "put_json" => Some(Self::PutJson),
463 "patch_json" => Some(Self::PatchJson),
464 "delete_json" => Some(Self::DeleteJson),
465 _ => None,
466 }
467 }
468
469 fn raw_builder_ident(self) -> Ident {
470 match self {
471 Self::Get => format_ident!("get"),
472 Self::Sse => format_ident!("get"),
473 Self::Post => format_ident!("post"),
474 Self::Put => format_ident!("put"),
475 Self::Patch => format_ident!("patch"),
476 Self::Delete => format_ident!("delete"),
477 Self::Options => format_ident!("options"),
478 Self::Head => format_ident!("head"),
479 Self::GetJson => format_ident!("get"),
480 Self::PostJson => format_ident!("post"),
481 Self::PutJson => format_ident!("put"),
482 Self::PatchJson => format_ident!("patch"),
483 Self::DeleteJson => format_ident!("delete"),
484 }
485 }
486
487 fn json_builder_ident(self) -> Option<Ident> {
488 match self {
489 Self::Get | Self::GetJson => Some(format_ident!("get_json_with_status")),
490 Self::Post | Self::PostJson => Some(format_ident!("post_json_with_status")),
491 Self::Put | Self::PutJson => Some(format_ident!("put_json_with_status")),
492 Self::Patch | Self::PatchJson => Some(format_ident!("patch_json_with_status")),
493 Self::Delete | Self::DeleteJson => Some(format_ident!("delete_json_with_status")),
494 Self::Sse | Self::Options | Self::Head => None,
495 }
496 }
497
498 fn is_explicit_json(self) -> bool {
499 matches!(
500 self,
501 Self::GetJson | Self::PostJson | Self::PutJson | Self::PatchJson | Self::DeleteJson
502 )
503 }
504
505 fn flavor(self, raw: bool) -> RouteFlavor {
506 if matches!(self, Self::Sse) {
507 return RouteFlavor::Sse;
508 }
509
510 if raw {
511 return RouteFlavor::Raw;
512 }
513
514 match self {
515 Self::Sse => RouteFlavor::Sse,
516 Self::Get | Self::GetJson | Self::Delete | Self::DeleteJson => RouteFlavor::JsonRequest,
517 Self::Post
518 | Self::PostJson
519 | Self::Put
520 | Self::PutJson
521 | Self::Patch
522 | Self::PatchJson => RouteFlavor::JsonBody,
523 Self::Options | Self::Head => RouteFlavor::Raw,
524 }
525 }
526}
527
528enum RouteFlavor {
529 Sse,
530 Raw,
531 JsonRequest,
532 JsonBody,
533}
534
535struct RouteMethodInput {
536 arg: Option<MethodArg>,
537}
538
539impl RouteMethodInput {
540 fn from_method(method: &ImplItemFn) -> Result<Self> {
541 let mut inputs = method.sig.inputs.iter();
542 let Some(FnArg::Receiver(receiver)) = inputs.next() else {
543 return Err(syn::Error::new_spanned(
544 &method.sig.ident,
545 "controller route methods must take &self as their first argument",
546 ));
547 };
548
549 if receiver.reference.is_none() || receiver.mutability.is_some() {
550 return Err(syn::Error::new_spanned(
551 receiver,
552 "controller route methods must use an immutable &self receiver",
553 ));
554 }
555
556 let args = inputs
557 .map(|input| match input {
558 FnArg::Typed(input) => MethodArg::from_pat_type(input),
559 FnArg::Receiver(receiver) => Err(syn::Error::new_spanned(
560 receiver,
561 "unexpected receiver argument",
562 )),
563 })
564 .collect::<Result<Vec<_>>>()?;
565
566 match args.len() {
567 0 => Ok(Self { arg: None }),
568 1 => Ok(Self {
569 arg: args.into_iter().next(),
570 }),
571 _ => Err(syn::Error::new_spanned(
572 &method.sig.inputs,
573 "controller route methods can accept at most one argument after &self",
574 )),
575 }
576 }
577}
578
579struct MethodArg {
580 ident: Ident,
581 ty: Box<syn::Type>,
582}
583
584impl MethodArg {
585 fn from_pat_type(input: &PatType) -> Result<Self> {
586 let Pat::Ident(ident) = input.pat.as_ref() else {
587 return Err(syn::Error::new_spanned(
588 &input.pat,
589 "controller route arguments must be simple identifiers",
590 ));
591 };
592
593 Ok(Self {
594 ident: ident.ident.clone(),
595 ty: input.ty.clone(),
596 })
597 }
598}