Skip to main content

silent_openapi_macros/
lib.rs

1use convert_case::Casing;
2use proc_macro::TokenStream;
3use quote::{format_ident, quote};
4use syn::Token;
5use syn::punctuated::Punctuated;
6use syn::{
7    Expr, ExprLit, FnArg, ItemFn, Lit, Meta, Result as SynResult, parse::Parse, parse::ParseStream,
8};
9
10fn endpoint_impl(
11    attr: proc_macro2::TokenStream,
12    item: proc_macro2::TokenStream,
13) -> proc_macro2::TokenStream {
14    struct MetaArgs(Punctuated<Meta, Token![,]>);
15    impl Parse for MetaArgs {
16        fn parse(input: ParseStream) -> SynResult<Self> {
17            Ok(MetaArgs(Punctuated::parse_terminated(input)?))
18        }
19    }
20    let MetaArgs(args) = syn::parse2::<MetaArgs>(attr).expect("parse attr");
21    let mut summary_arg: Option<String> = None;
22    let mut description_arg: Option<String> = None;
23    let mut deprecated_flag = false;
24    let mut tags_arg: Vec<String> = Vec::new();
25    let mut extra_responses: Vec<(u16, String)> = Vec::new();
26
27    for meta in args {
28        match &meta {
29            // deprecated(无值标志)
30            Meta::Path(path) if path.is_ident("deprecated") => {
31                deprecated_flag = true;
32            }
33            // summary = "...", description = "...", tags = "..."
34            Meta::NameValue(nv) => {
35                if nv.path.is_ident("summary")
36                    && let Expr::Lit(ExprLit {
37                        lit: Lit::Str(s), ..
38                    }) = &nv.value
39                {
40                    summary_arg = Some(s.value());
41                } else if nv.path.is_ident("description")
42                    && let Expr::Lit(ExprLit {
43                        lit: Lit::Str(s), ..
44                    }) = &nv.value
45                {
46                    description_arg = Some(s.value());
47                } else if nv.path.is_ident("tags")
48                    && let Expr::Lit(ExprLit {
49                        lit: Lit::Str(s), ..
50                    }) = &nv.value
51                {
52                    for tag in s.value().split(',') {
53                        let t = tag.trim().to_string();
54                        if !t.is_empty() {
55                            tags_arg.push(t);
56                        }
57                    }
58                }
59            }
60            // response(status = 400, description = "...")
61            Meta::List(list) if list.path.is_ident("response") => {
62                let mut status: Option<u16> = None;
63                let mut desc: Option<String> = None;
64                let _ = list.parse_nested_meta(|nested| {
65                    if nested.path.is_ident("status") {
66                        let value = nested.value()?;
67                        let lit: syn::LitInt = value.parse()?;
68                        status = Some(lit.base10_parse()?);
69                    } else if nested.path.is_ident("description") {
70                        let value = nested.value()?;
71                        let lit: syn::LitStr = value.parse()?;
72                        desc = Some(lit.value());
73                    }
74                    Ok(())
75                });
76                if let (Some(st), Some(d)) = (status, desc) {
77                    extra_responses.push((st, d));
78                }
79            }
80            _ => {}
81        }
82    }
83
84    let input: ItemFn = syn::parse2(item).expect("parse item fn");
85    let vis = &input.vis;
86    let sig = input.sig.clone();
87    let attrs = &input.attrs;
88    let block = &input.block;
89    let name = &sig.ident;
90
91    // 收集文档注释作为默认 summary/description
92    let mut doc_lines: Vec<String> = Vec::new();
93    for a in attrs.iter() {
94        if a.path().is_ident("doc") {
95            let _ = a.parse_nested_meta(|meta| {
96                let lit: syn::LitStr = meta.value()?.parse()?;
97                let v = lit.value();
98                doc_lines.push(v.trim().to_string());
99                Ok(())
100            });
101        }
102    }
103    let (def_summary, def_description) = if !doc_lines.is_empty() {
104        let mut it = doc_lines.into_iter().filter(|s| !s.is_empty());
105        if let Some(first) = it.next() {
106            let rest = it.collect::<Vec<_>>().join("\n");
107            (Some(first), if rest.is_empty() { None } else { Some(rest) })
108        } else {
109            (None, None)
110        }
111    } else {
112        (None, None)
113    };
114
115    let summary = summary_arg.or(def_summary);
116    let description = description_arg.or(def_description);
117
118    // 真实处理函数改名
119    let impl_name = format_ident!("{}_impl", name);
120    // 生成实现函数签名(重命名)
121    let mut impl_sig = sig.clone();
122    impl_sig.ident = impl_name.clone();
123
124    // 端点类型 + 常量(实现与原 `.get(get_xxx)` 风格兼容)
125    let ep_ty = format_ident!(
126        "{}Endpoint",
127        name.to_string().to_case(convert_case::Case::UpperCamel)
128    );
129    let sum_tokens = if let Some(s) = &summary {
130        let lit = syn::LitStr::new(s, proc_macro2::Span::call_site());
131        quote!(Some(#lit))
132    } else {
133        quote!(None)
134    };
135    let desc_tokens = if let Some(s) = &description {
136        let lit = syn::LitStr::new(s, proc_macro2::Span::call_site());
137        quote!(Some(#lit))
138    } else {
139        quote!(None)
140    };
141
142    // deprecated / tags / extra responses 的 token
143    let deprecated_tokens = deprecated_flag;
144    let tags_tokens = {
145        let tag_lits: Vec<_> = tags_arg
146            .iter()
147            .map(|t| syn::LitStr::new(t, proc_macro2::Span::call_site()))
148            .collect();
149        quote!(&[#(#tag_lits),*])
150    };
151    let extra_response_tokens = {
152        let stmts: Vec<_> = extra_responses
153            .iter()
154            .map(|(status, desc)| {
155                let st = *status;
156                let d = syn::LitStr::new(desc, proc_macro2::Span::call_site());
157                quote! {
158                    ::silent_openapi::doc::register_extra_response_by_ptr(ptr, #st, #d);
159                }
160            })
161            .collect();
162        quote!(#(#stmts)*)
163    };
164
165    // 解析返回类型 Ok(T) -> ResponseMeta
166    let ret_meta = {
167        match &sig.output {
168            syn::ReturnType::Type(_, ty) => {
169                if let syn::Type::Path(tp) = ty.as_ref() {
170                    if let Some(seg) = tp.path.segments.last() {
171                        if seg.ident == "Result" || seg.ident == "SilentResult" {
172                            if let syn::PathArguments::AngleBracketed(args) = &seg.arguments {
173                                if let Some(syn::GenericArgument::Type(ok_ty)) = args.args.first() {
174                                    match ok_ty {
175                                        syn::Type::Path(tpath) => {
176                                            if let Some(id) = tpath.path.segments.last() {
177                                                if id.ident == "Response" {
178                                                    quote!(None)
179                                                } else if id.ident == "String" {
180                                                    quote!(Some(::silent_openapi::doc::ResponseMeta::TextPlain))
181                                                } else {
182                                                    let tn = id.ident.to_string();
183                                                    quote!(Some(::silent_openapi::doc::ResponseMeta::Json { type_name: #tn }))
184                                                }
185                                            } else {
186                                                quote!(None)
187                                            }
188                                        }
189                                        syn::Type::Reference(r) => {
190                                            if let syn::Type::Path(tp2) = r.elem.as_ref() {
191                                                if let Some(id) = tp2.path.segments.last() {
192                                                    if id.ident == "str" {
193                                                        quote!(Some(::silent_openapi::doc::ResponseMeta::TextPlain))
194                                                    } else {
195                                                        let tn = id.ident.to_string();
196                                                        quote!(Some(::silent_openapi::doc::ResponseMeta::Json { type_name: #tn }))
197                                                    }
198                                                } else {
199                                                    quote!(None)
200                                                }
201                                            } else {
202                                                quote!(None)
203                                            }
204                                        }
205                                        _ => quote!(None),
206                                    }
207                                } else {
208                                    quote!(None)
209                                }
210                            } else {
211                                quote!(None)
212                            }
213                        } else {
214                            quote!(None)
215                        }
216                    } else {
217                        quote!(None)
218                    }
219                } else {
220                    quote!(None)
221                }
222            }
223            _ => quote!(None),
224        }
225    };
226
227    // 为自定义 Ok(T) 注册 ToSchema 完整 schema
228    let ret_schema_register = {
229        match &sig.output {
230            syn::ReturnType::Type(_, ty) => {
231                if let syn::Type::Path(tp) = ty.as_ref() {
232                    if let Some(seg) = tp.path.segments.last() {
233                        if seg.ident == "Result" || seg.ident == "SilentResult" {
234                            if let syn::PathArguments::AngleBracketed(args) = &seg.arguments {
235                                if let Some(syn::GenericArgument::Type(ok_ty)) = args.args.first() {
236                                    match ok_ty {
237                                        syn::Type::Path(tpath) => {
238                                            if let Some(id) = tpath.path.segments.last() {
239                                                if id.ident == "Response" || id.ident == "String" {
240                                                    quote!()
241                                                } else {
242                                                    let ty = ok_ty.clone();
243                                                    quote!(::silent_openapi::doc::register_schema_for::<#ty>();)
244                                                }
245                                            } else {
246                                                quote!()
247                                            }
248                                        }
249                                        syn::Type::Reference(r) => {
250                                            if let syn::Type::Path(tp2) = r.elem.as_ref() {
251                                                if let Some(id) = tp2.path.segments.last() {
252                                                    if id.ident == "str" {
253                                                        quote!()
254                                                    } else {
255                                                        let inner = tp2.clone();
256                                                        quote!(::silent_openapi::doc::register_schema_for::<#inner>();)
257                                                    }
258                                                } else {
259                                                    quote!()
260                                                }
261                                            } else {
262                                                quote!()
263                                            }
264                                        }
265                                        _ => quote!(),
266                                    }
267                                } else {
268                                    quote!()
269                                }
270                            } else {
271                                quote!()
272                            }
273                        } else {
274                            quote!()
275                        }
276                    } else {
277                        quote!()
278                    }
279                } else {
280                    quote!()
281                }
282            }
283            _ => quote!(),
284        }
285    };
286
287    // 从提取器类型中生成请求元信息注册代码
288    fn gen_request_meta_register(ty: &syn::Type) -> proc_macro2::TokenStream {
289        if let syn::Type::Path(tp) = ty {
290            if let Some(seg) = tp.path.segments.last() {
291                let ident = seg.ident.to_string();
292                if let syn::PathArguments::AngleBracketed(args) = &seg.arguments {
293                    if let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() {
294                        // 获取内部类型名称
295                        let inner_name = if let syn::Type::Path(inner_tp) = inner_ty {
296                            inner_tp
297                                .path
298                                .segments
299                                .last()
300                                .map(|s| s.ident.to_string())
301                                .unwrap_or_default()
302                        } else {
303                            String::new()
304                        };
305
306                        if !inner_name.is_empty() {
307                            match ident.as_str() {
308                                "Json" => {
309                                    let inner = inner_ty.clone();
310                                    return quote! {
311                                        ::silent_openapi::doc::register_request_by_ptr(
312                                            ptr,
313                                            ::silent_openapi::doc::RequestMeta::JsonBody { type_name: #inner_name },
314                                        );
315                                        ::silent_openapi::doc::register_schema_for::<#inner>();
316                                    };
317                                }
318                                "Form" => {
319                                    let inner = inner_ty.clone();
320                                    return quote! {
321                                        ::silent_openapi::doc::register_request_by_ptr(
322                                            ptr,
323                                            ::silent_openapi::doc::RequestMeta::FormBody { type_name: #inner_name },
324                                        );
325                                        ::silent_openapi::doc::register_schema_for::<#inner>();
326                                    };
327                                }
328                                "Query" => {
329                                    let inner = inner_ty.clone();
330                                    return quote! {
331                                        ::silent_openapi::doc::register_request_by_ptr(
332                                            ptr,
333                                            ::silent_openapi::doc::RequestMeta::QueryParams { type_name: #inner_name },
334                                        );
335                                        ::silent_openapi::doc::register_schema_for::<#inner>();
336                                    };
337                                }
338                                _ => {}
339                            }
340                        }
341                    }
342                }
343            }
344        }
345        quote!()
346    }
347
348    // 根据函数参数形态生成 IntoRouteHandler 实现
349    let inputs = sig.inputs.clone().into_iter().collect::<Vec<_>>();
350    let impls = if inputs.len() == 1 {
351        match &inputs[0] {
352            FnArg::Typed(pat_ty) => {
353                let ty = &pat_ty.ty;
354                // 简单规则:类型标识名为 Request 则认为是 Request 形态
355                let is_request = matches!(
356                    &**ty,
357                    syn::Type::Path(tp) if tp.path.segments.last().map(|s| s.ident == "Request").unwrap_or(false)
358                );
359                if is_request {
360                    quote! {
361                        impl ::silent::prelude::IntoRouteHandler<::silent::Request> for #ep_ty {
362                            fn into_handler(self) -> std::sync::Arc<dyn ::silent::Handler> {
363                                let handler = std::sync::Arc::new(::silent::HandlerWrapper::new(#impl_name));
364                                let ptr = std::sync::Arc::as_ptr(&handler) as *const () as usize;
365                                ::silent_openapi::doc::register_doc_by_ptr_ext(
366                                    ptr,
367                                    #sum_tokens,
368                                    #desc_tokens,
369                                    #deprecated_tokens,
370                                    #tags_tokens,
371                                );
372                                #ret_schema_register
373                                if let Some(meta) = #ret_meta { ::silent_openapi::doc::register_response_by_ptr(ptr, meta); }
374                                #extra_response_tokens
375                                handler
376                            }
377                        }
378                    }
379                } else {
380                    // 单萃取器参数
381                    let req_meta_register = gen_request_meta_register(ty);
382                    quote! {
383                        impl ::silent::prelude::IntoRouteHandler<#ty> for #ep_ty {
384                            fn into_handler(self) -> std::sync::Arc<dyn ::silent::Handler> {
385                                let adapted = ::silent::extractor::handler_from_extractor::<#ty, _, _, _>(#impl_name);
386                                let handler = std::sync::Arc::new(::silent::HandlerWrapper::new(adapted));
387                                let ptr = std::sync::Arc::as_ptr(&handler) as *const () as usize;
388                                ::silent_openapi::doc::register_doc_by_ptr_ext(
389                                    ptr,
390                                    #sum_tokens,
391                                    #desc_tokens,
392                                    #deprecated_tokens,
393                                    #tags_tokens,
394                                );
395                                #ret_schema_register
396                                if let Some(meta) = #ret_meta { ::silent_openapi::doc::register_response_by_ptr(ptr, meta); }
397                                #extra_response_tokens
398                                #req_meta_register
399                                handler
400                            }
401                        }
402                    }
403                }
404            }
405            _ => quote! {},
406        }
407    } else if inputs.len() == 2 {
408        match (&inputs[0], &inputs[1]) {
409            (FnArg::Typed(first), FnArg::Typed(second)) => {
410                let ty1 = &first.ty;
411                let ty2 = &second.ty;
412                // 期望形态: (Request, Args)
413                let is_request_first = matches!(
414                    &**ty1,
415                    syn::Type::Path(tp) if tp.path.segments.last().map(|s| s.ident == "Request").unwrap_or(false)
416                );
417                if is_request_first {
418                    let req_meta_register = gen_request_meta_register(ty2);
419                    quote! {
420                        impl ::silent::prelude::IntoRouteHandler<(::silent::Request, #ty2)> for #ep_ty {
421                            fn into_handler(self) -> std::sync::Arc<dyn ::silent::Handler> {
422                                let adapted = ::silent::extractor::handler_from_extractor_with_request::<#ty2, _, _, _>(#impl_name);
423                                let handler = std::sync::Arc::new(::silent::HandlerWrapper::new(adapted));
424                                let ptr = std::sync::Arc::as_ptr(&handler) as *const () as usize;
425                                ::silent_openapi::doc::register_doc_by_ptr_ext(
426                                    ptr,
427                                    #sum_tokens,
428                                    #desc_tokens,
429                                    #deprecated_tokens,
430                                    #tags_tokens,
431                                );
432                                #ret_schema_register
433                                if let Some(meta) = #ret_meta { ::silent_openapi::doc::register_response_by_ptr(ptr, meta); }
434                                #extra_response_tokens
435                                #req_meta_register
436                                handler
437                            }
438                        }
439                    }
440                } else {
441                    quote! {}
442                }
443            }
444            _ => quote! {},
445        }
446    } else {
447        quote! {}
448    };
449
450    let code = quote! {
451        // 原函数体改名为实现函数
452        #(#attrs)*
453        #impl_sig #block
454
455        // 端点类型(零尺寸) + 常量,同名以保留 `.get(get_xxx)` 调用方式
456        pub struct #ep_ty;
457        #[allow(non_upper_case_globals)]
458        #vis const #name: #ep_ty = #ep_ty;
459
460        #impls
461    };
462
463    code
464}
465
466#[proc_macro_attribute]
467pub fn endpoint(attr: TokenStream, item: TokenStream) -> TokenStream {
468    endpoint_impl(attr.into(), item.into()).into()
469}
470
471#[cfg(test)]
472mod tests {
473    use quote::quote;
474
475    fn render(ts: proc_macro2::TokenStream) -> String {
476        ts.to_string()
477    }
478
479    #[test]
480    fn generates_endpoint_type_and_const_for_request_sig() {
481        let attr = quote!(summary = "hello", description = "world");
482        let item = quote!(
483            async fn get_hello(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
484                unimplemented!()
485            }
486        );
487        let out = super::endpoint_impl(attr, item);
488        let s = render(out);
489        assert!(s.contains("struct GetHelloEndpoint"));
490        assert!(s.contains("const get_hello"));
491    }
492
493    #[test]
494    fn generates_into_route_handler_for_extractor_sig() {
495        let attr = quote!();
496        let item = quote!(
497            async fn get_user(_id: Path<u64>) -> ::silent::Result<::silent::Response> {
498                unimplemented!()
499            }
500        );
501        let out = super::endpoint_impl(attr, item);
502        let s = render(out);
503        // 生成的端点常量与 IntoRouteHandler 实现
504        assert!(s.contains("struct GetUserEndpoint"));
505        assert!(s.contains("const get_user"));
506        assert!(s.contains("IntoRouteHandler"));
507        assert!(s.contains("GetUserEndpoint"));
508    }
509
510    #[test]
511    fn registers_request_meta_for_json_extractor() {
512        let attr = quote!();
513        let item = quote!(
514            async fn create_user(body: Json<UserInput>) -> ::silent::Result<::silent::Response> {
515                unimplemented!()
516            }
517        );
518        let out = super::endpoint_impl(attr, item);
519        let s = render(out);
520        assert!(s.contains("RequestMeta :: JsonBody"));
521        assert!(s.contains("register_request_by_ptr"));
522        assert!(s.contains("register_schema_for"));
523    }
524
525    #[test]
526    fn registers_request_meta_for_query_extractor() {
527        let attr = quote!();
528        let item = quote!(
529            async fn list_users(params: Query<ListParams>) -> ::silent::Result<::silent::Response> {
530                unimplemented!()
531            }
532        );
533        let out = super::endpoint_impl(attr, item);
534        let s = render(out);
535        assert!(s.contains("RequestMeta :: QueryParams"));
536        assert!(s.contains("register_request_by_ptr"));
537    }
538
539    #[test]
540    fn registers_request_meta_for_form_extractor() {
541        let attr = quote!();
542        let item = quote!(
543            async fn submit_form(data: Form<FormData>) -> ::silent::Result<::silent::Response> {
544                unimplemented!()
545            }
546        );
547        let out = super::endpoint_impl(attr, item);
548        let s = render(out);
549        assert!(s.contains("RequestMeta :: FormBody"));
550        assert!(s.contains("register_request_by_ptr"));
551    }
552
553    #[test]
554    fn registers_request_meta_for_request_with_extractor() {
555        let attr = quote!();
556        let item = quote!(
557            async fn update_user(
558                _req: ::silent::Request,
559                body: Json<UserInput>,
560            ) -> ::silent::Result<::silent::Response> {
561                unimplemented!()
562            }
563        );
564        let out = super::endpoint_impl(attr, item);
565        let s = render(out);
566        assert!(s.contains("RequestMeta :: JsonBody"));
567        assert!(s.contains("register_request_by_ptr"));
568    }
569
570    #[test]
571    fn no_request_meta_for_plain_request() {
572        let attr = quote!();
573        let item = quote!(
574            async fn health(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
575                unimplemented!()
576            }
577        );
578        let out = super::endpoint_impl(attr, item);
579        let s = render(out);
580        assert!(!s.contains("register_request_by_ptr"));
581    }
582
583    #[test]
584    fn registers_schema_for_enum_return_type() {
585        let attr = quote!();
586        let item = quote!(
587            async fn get_status(_req: ::silent::Request) -> ::silent::Result<ApiResponse> {
588                unimplemented!()
589            }
590        );
591        let out = super::endpoint_impl(attr, item);
592        let s = render(out);
593        // 枚举返回类型应生成 Json 响应元信息和 schema 注册
594        assert!(s.contains("ResponseMeta :: Json"));
595        assert!(s.contains("register_schema_for"));
596        assert!(s.contains("ApiResponse"));
597    }
598
599    #[test]
600    fn registers_schema_for_enum_request_body() {
601        let attr = quote!();
602        let item = quote!(
603            async fn create_item(body: Json<CreateAction>) -> ::silent::Result<::silent::Response> {
604                unimplemented!()
605            }
606        );
607        let out = super::endpoint_impl(attr, item);
608        let s = render(out);
609        // 枚举请求体类型同样应注册 schema
610        assert!(s.contains("RequestMeta :: JsonBody"));
611        assert!(s.contains("register_schema_for"));
612        assert!(s.contains("CreateAction"));
613    }
614
615    #[test]
616    fn doc_comment_as_summary_and_description() {
617        let attr = quote!();
618        let item = quote!(
619            /// 获取用户信息
620            ///
621            /// 根据用户 ID 查询完整的用户资料
622            async fn get_user(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
623                unimplemented!()
624            }
625        );
626        let out = super::endpoint_impl(attr, item);
627        let s = render(out);
628        assert!(s.contains("获取用户信息"));
629        assert!(s.contains("根据用户 ID 查询完整的用户资料"));
630    }
631
632    #[test]
633    fn registers_response_meta_for_string() {
634        let attr = quote!();
635        let item = quote!(
636            async fn ping(_req: ::silent::Request) -> ::silent::Result<String> {
637                unimplemented!()
638            }
639        );
640        let out = super::endpoint_impl(attr, item);
641        let s = render(out);
642        // 生成文本响应的注册调用
643        assert!(s.contains("ResponseMeta :: TextPlain"));
644    }
645
646    #[test]
647    fn deprecated_flag_generates_ext_call() {
648        let attr = quote!(deprecated);
649        let item = quote!(
650            async fn old_api(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
651                unimplemented!()
652            }
653        );
654        let out = super::endpoint_impl(attr, item);
655        let s = render(out);
656        assert!(s.contains("register_doc_by_ptr_ext"));
657        assert!(s.contains("true")); // deprecated = true
658    }
659
660    #[test]
661    fn tags_generates_ext_call_with_tags() {
662        let attr = quote!(tags = "users,admin");
663        let item = quote!(
664            async fn list_users(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
665                unimplemented!()
666            }
667        );
668        let out = super::endpoint_impl(attr, item);
669        let s = render(out);
670        assert!(s.contains("register_doc_by_ptr_ext"));
671        assert!(s.contains("\"users\""));
672        assert!(s.contains("\"admin\""));
673    }
674
675    #[test]
676    fn response_generates_extra_response_registration() {
677        let attr = quote!(
678            response(status = 400, description = "Bad request"),
679            response(status = 401, description = "Unauthorized")
680        );
681        let item = quote!(
682            async fn create(_req: ::silent::Request) -> ::silent::Result<::silent::Response> {
683                unimplemented!()
684            }
685        );
686        let out = super::endpoint_impl(attr, item);
687        let s = render(out);
688        assert!(s.contains("register_extra_response_by_ptr"));
689        assert!(s.contains("400"));
690        assert!(s.contains("401"));
691        assert!(s.contains("Bad request"));
692        assert!(s.contains("Unauthorized"));
693    }
694}