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 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 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 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#[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#[proc_macro_attribute]
172pub fn get(args: TokenStream, input: TokenStream) -> TokenStream {
173 route_macro(args, input, "get")
174}
175
176#[proc_macro_attribute]
178pub fn post(args: TokenStream, input: TokenStream) -> TokenStream {
179 route_macro(args, input, "post")
180}
181
182#[proc_macro_attribute]
184pub fn put(args: TokenStream, input: TokenStream) -> TokenStream {
185 route_macro(args, input, "put")
186}
187
188#[proc_macro_attribute]
190pub fn delete(args: TokenStream, input: TokenStream) -> TokenStream {
191 route_macro(args, input, "delete")
192}
193
194#[proc_macro_attribute]
196pub fn patch(args: TokenStream, input: TokenStream) -> TokenStream {
197 route_macro(args, input, "patch")
198}
199
200#[proc_macro_attribute]
202pub fn websocket(args: TokenStream, input: TokenStream) -> TokenStream {
203 route_macro(args, input, "websocket")
204}
205
206#[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}