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 Meta::Path(path) if path.is_ident("deprecated") => {
31 deprecated_flag = true;
32 }
33 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 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 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 let impl_name = format_ident!("{}_impl", name);
120 let mut impl_sig = sig.clone();
122 impl_sig.ident = impl_name.clone();
123
124 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 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 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 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 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 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 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 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 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 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 #(#attrs)*
453 #impl_sig #block
454
455 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 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 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 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 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 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")); }
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}