Skip to main content

a3s_boot_macros/
lib.rs

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}