Skip to main content

ukiapi_macros/
lib.rs

1use proc_macro::TokenStream;
2use quote::quote;
3use syn::{parse_macro_input, GenericArgument, ItemFn, PathArguments, Type};
4
5fn extractor_inner_type(ty: &Type) -> Option<(String, &Type)> {
6    if let Type::Path(type_path) = ty {
7        if let Some(segment) = type_path.path.segments.last() {
8            let name = segment.ident.to_string();
9            if matches!(name.as_str(), "Json" | "ValidatedJson" | "Path" | "Query") {
10                if let PathArguments::AngleBracketed(ref args) = segment.arguments {
11                    if let Some(GenericArgument::Type(inner)) = args.args.first() {
12                        return Some((name, inner));
13                    }
14                }
15            }
16        }
17    }
18    None
19}
20
21struct RouteArgs {
22    path: syn::LitStr,
23    registry: Option<syn::Path>,
24}
25
26impl syn::parse::Parse for RouteArgs {
27    fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
28        let path: syn::LitStr = input.parse()?;
29        let mut registry = None;
30        if input.peek(syn::Token![,]) {
31            input.parse::<syn::Token![,]>()?;
32            while !input.is_empty() {
33                let ident: syn::Ident = input.parse()?;
34                input.parse::<syn::Token![=]>()?;
35                if ident == "registry" {
36                    registry = Some(input.parse()?);
37                } else {
38                    let _: syn::Expr = input.parse()?;
39                }
40                if input.peek(syn::Token![,]) {
41                    input.parse::<syn::Token![,]>()?;
42                } else {
43                    break;
44                }
45            }
46        }
47        Ok(RouteArgs { path, registry })
48    }
49}
50
51fn route_macro(args: TokenStream, input: TokenStream, method: &str) -> TokenStream {
52    let route_args = parse_macro_input!(args as RouteArgs);
53    let path_lit = &route_args.path;
54    let func = parse_macro_input!(input as ItemFn);
55    let fn_name = &func.sig.ident;
56    let route_fn_name = syn::Ident::new(&format!("{}_route", fn_name), fn_name.span());
57    let method_ident = syn::Ident::new(method, fn_name.span());
58    let vis = &func.vis;
59    let attrs = &func.attrs;
60    let sig = &func.sig;
61    let block = &func.block;
62
63    // Infer state type S from handler arguments
64    let state_ty = func
65        .sig
66        .inputs
67        .iter()
68        .find_map(|input| {
69            if let syn::FnArg::Typed(pat_type) = input {
70                if let Type::Path(type_path) = &*pat_type.ty {
71                    if let Some(segment) = type_path.path.segments.last() {
72                        if segment.ident == "State" {
73                            if let PathArguments::AngleBracketed(ref args) = segment.arguments {
74                                if let Some(GenericArgument::Type(inner_ty)) = args.args.first() {
75                                    return Some(quote! { #inner_ty });
76                                }
77                            }
78                        }
79                    }
80                }
81            }
82            None
83        })
84        .unwrap_or_else(|| quote! { () });
85
86    // Detect request schema from first Json parameter
87    let req_schema = func.sig.inputs.iter().find_map(|input| {
88        if let syn::FnArg::Typed(pat_type) = input {
89            if let Some((name, inner)) = extractor_inner_type(&pat_type.ty) {
90                if name == "Json" || name == "ValidatedJson" {
91                    return Some(
92                        quote! { .with_request_schema(::ukiapi::schema_for::<#inner>()) },
93                    );
94                }
95            }
96        }
97        None
98    });
99
100    let query_schema = func.sig.inputs.iter().find_map(|input| {
101        if let syn::FnArg::Typed(pat_type) = input {
102            if let Some((name, inner)) = extractor_inner_type(&pat_type.ty) {
103                if name == "Query" {
104                    return Some(quote! { .with_query_schema(::ukiapi::schema_for::<#inner>()) });
105                }
106            }
107        }
108        None
109    });
110
111    // Detect response schema from return type
112    let res_schema = match &func.sig.output {
113        syn::ReturnType::Type(_, ret_ty) => {
114            if let Some((name, inner)) = extractor_inner_type(ret_ty.as_ref()) {
115                if name == "Json" {
116                    Some(quote! { .with_response_schema(::ukiapi::schema_for::<#inner>()) })
117                } else {
118                    None
119                }
120            } else {
121                None
122            }
123        }
124        _ => None,
125    };
126
127    let registry_submit = if let Some(reg) = &route_args.registry {
128        if state_ty.to_string() == "()" {
129            quote! { ::ukiapi::submit_route!(#route_fn_name, #reg, stateless); }
130        } else {
131            quote! { ::ukiapi::submit_route!(#route_fn_name, #reg, stateful); }
132        }
133    } else if state_ty.to_string() == "()" {
134        quote! { ::ukiapi::submit_route!(#route_fn_name); }
135    } else {
136        quote! {}
137    };
138
139    let expanded = quote! {
140        #(#attrs)*
141        #vis #sig #block
142
143
144        #[doc(hidden)]
145        pub fn #route_fn_name() -> ::ukiapi::Route<#state_ty> {
146            ::ukiapi::Route::#method_ident(#path_lit, #fn_name)
147                #req_schema
148                #res_schema
149                #query_schema
150        }
151
152        #registry_submit
153    };
154
155    expanded.into()
156}
157
158/// Macro to define a data model. Derives Serialize, Deserialize, JsonSchema, Validate, Clone, and TS.
159#[proc_macro_attribute]
160pub fn model(_args: TokenStream, input: TokenStream) -> TokenStream {
161    let item: proc_macro2::TokenStream = input.into();
162    let expanded = quote! {
163        #[derive(::ukiapi::Serialize, ::ukiapi::Deserialize, ::ukiapi::JsonSchema, ::validator::Validate, Clone, ::ukiapi::ts_rs::TS)]
164        #[ts(export, crate = "::ukiapi::ts_rs")]
165        #item
166    };
167    expanded.into()
168}
169
170/// Define a GET endpoint.
171#[proc_macro_attribute]
172pub fn get(args: TokenStream, input: TokenStream) -> TokenStream {
173    route_macro(args, input, "get")
174}
175
176/// Define a POST endpoint.
177#[proc_macro_attribute]
178pub fn post(args: TokenStream, input: TokenStream) -> TokenStream {
179    route_macro(args, input, "post")
180}
181
182/// Define a PUT endpoint.
183#[proc_macro_attribute]
184pub fn put(args: TokenStream, input: TokenStream) -> TokenStream {
185    route_macro(args, input, "put")
186}
187
188/// Define a DELETE endpoint.
189#[proc_macro_attribute]
190pub fn delete(args: TokenStream, input: TokenStream) -> TokenStream {
191    route_macro(args, input, "delete")
192}
193
194/// Define a PATCH endpoint.
195#[proc_macro_attribute]
196pub fn patch(args: TokenStream, input: TokenStream) -> TokenStream {
197    route_macro(args, input, "patch")
198}
199
200/// Define a WebSocket endpoint.
201#[proc_macro_attribute]
202pub fn websocket(args: TokenStream, input: TokenStream) -> TokenStream {
203    route_macro(args, input, "websocket")
204}
205
206/// Macro to define the main entry point for a UkiApi application.
207///
208/// Sets up the tokio runtime, environment variables, and logger.
209///
210/// # Example
211/// ```rust,ignore
212/// use ukiapi::{get, routes};
213///
214/// #[get("/hello")]
215/// async fn hello() -> &'static str {
216///     "Hello from UkiApi!"
217/// }
218///
219/// #[ukiapi::main]
220/// async fn main() {
221///     routes![(),
222///         hello_route().with_state::<()>()
223///     ]
224///     .serve(())
225///     .await;
226/// }
227/// ```
228#[proc_macro_attribute]
229pub fn main(_args: TokenStream, input: TokenStream) -> TokenStream {
230    let func = parse_macro_input!(input as ItemFn);
231    let block = &func.block;
232    let sig = &func.sig;
233    let attrs = &func.attrs;
234    let vis = &func.vis;
235
236    let expanded = quote! {
237        #[tokio::main]
238        #(#attrs)*
239        #vis #sig {
240            if std::env::var("UKIAPI_HOST").is_err() {
241                std::env::set_var("UKIAPI_HOST", "127.0.0.1");
242            }
243            if std::env::var("UKIAPI_PORT").is_err() {
244                std::env::set_var("UKIAPI_PORT", "3000");
245            }
246            env_logger::init();
247
248            #block
249        }
250    };
251    expanded.into()
252}