1#![recursion_limit = "512"]
8
9extern crate proc_macro;
10extern crate proc_macro2;
11extern crate quote;
12extern crate syn;
13
14use proc_macro::TokenStream;
15use proc_macro2::TokenStream as TokenStream2;
16use quote::{ToTokens, format_ident, quote};
17use syn::{
18 AttrStyle, Attribute, Expr, FnArg, Ident, Lit, LitBool, MetaNameValue, Pat, PatType, Path,
19 ReturnType, Token, Type, Visibility, braced,
20 ext::IdentExt,
21 parenthesized,
22 parse::{Parse, ParseStream},
23 parse_macro_input, parse_quote,
24 spanned::Spanned,
25 token::Comma,
26};
27
28macro_rules! extend_errors {
32 ($errors: ident, $e: expr) => {
33 match $errors {
34 Ok(_) => $errors = Err($e),
35 Err(ref mut errors) => errors.extend($e),
36 }
37 };
38}
39
40struct Service {
41 attrs: Vec<Attribute>,
42 vis: Visibility,
43 ident: Ident,
44 rpcs: Vec<RpcMethod>,
45}
46
47struct RpcMethod {
48 attrs: Vec<Attribute>,
49 ident: Ident,
50 args: Vec<PatType>,
51 output: ReturnType,
52}
53
54impl Parse for Service {
55 fn parse(input: ParseStream) -> syn::Result<Self> {
56 let attrs = input.call(Attribute::parse_outer)?;
57 let vis = input.parse()?;
58 input.parse::<Token![trait]>()?;
59 let ident: Ident = input.parse()?;
60 let content;
61 braced!(content in input);
62 let mut rpcs = Vec::<RpcMethod>::new();
63 while !content.is_empty() {
64 rpcs.push(content.parse()?);
65 }
66 let mut ident_errors = Ok(());
67 for rpc in &rpcs {
68 if rpc.ident == "new" {
69 extend_errors!(
70 ident_errors,
71 syn::Error::new(
72 rpc.ident.span(),
73 format!(
74 "method name conflicts with generated fn `{}Client::new`",
75 ident.unraw()
76 )
77 )
78 );
79 }
80 if rpc.ident == "serve" {
81 extend_errors!(
82 ident_errors,
83 syn::Error::new(
84 rpc.ident.span(),
85 format!("method name conflicts with generated fn `{ident}::serve`")
86 )
87 );
88 }
89 }
90 ident_errors?;
91
92 Ok(Self {
93 attrs,
94 vis,
95 ident,
96 rpcs,
97 })
98 }
99}
100
101impl Parse for RpcMethod {
102 fn parse(input: ParseStream) -> syn::Result<Self> {
103 let attrs = input.call(Attribute::parse_outer)?;
104 input.parse::<Token![async]>()?;
105 input.parse::<Token![fn]>()?;
106 let ident = input.parse()?;
107 let content;
108 parenthesized!(content in input);
109 let mut args = Vec::new();
110 let mut errors = Ok(());
111 for arg in content.parse_terminated(FnArg::parse, Comma)? {
112 match arg {
113 FnArg::Typed(captured) if matches!(&*captured.pat, Pat::Ident(_)) => {
114 args.push(captured);
115 }
116 FnArg::Typed(captured) => {
117 extend_errors!(
118 errors,
119 syn::Error::new(captured.pat.span(), "patterns aren't allowed in RPC args")
120 );
121 }
122 FnArg::Receiver(_) => {
123 extend_errors!(
124 errors,
125 syn::Error::new(arg.span(), "method args cannot start with self")
126 );
127 }
128 }
129 }
130 errors?;
131 let output = input.parse()?;
132 input.parse::<Token![;]>()?;
133
134 Ok(Self {
135 attrs,
136 ident,
137 args,
138 output,
139 })
140 }
141}
142
143#[derive(Default)]
144struct DeriveMeta {
145 derive: Option<Derive>,
146 warnings: Vec<TokenStream2>,
147}
148
149impl DeriveMeta {
150 fn with_derives(mut self, new: Vec<Path>) -> Self {
151 match self.derive.as_mut() {
152 Some(Derive::Explicit(old)) => old.extend(new),
153 _ => self.derive = Some(Derive::Explicit(new)),
154 }
155
156 self
157 }
158}
159
160enum Derive {
161 Explicit(Vec<Path>),
162 Serde(bool),
163}
164
165impl Parse for DeriveMeta {
166 fn parse(input: ParseStream) -> syn::Result<Self> {
167 let mut result = Ok(DeriveMeta::default());
168
169 let mut derives = Vec::new();
170 let mut derive_serde = Vec::new();
171 let mut has_derive_serde = false;
172 let mut has_explicit_derives = false;
173
174 let meta_items = input.parse_terminated(MetaNameValue::parse, Comma)?;
175 for meta in meta_items {
176 if meta.path.segments.len() != 1 {
177 extend_errors!(
178 result,
179 syn::Error::new(
180 meta.span(),
181 "tarpc::service does not support this meta item"
182 )
183 );
184 continue;
185 }
186 let segment = meta.path.segments.first().unwrap();
187 if segment.ident == "derive" {
188 has_explicit_derives = true;
189 let Expr::Array(ref array) = meta.value else {
190 extend_errors!(
191 result,
192 syn::Error::new(
193 meta.span(),
194 "tarpc::service does not support this meta item"
195 )
196 );
197 continue;
198 };
199
200 let paths = array
201 .elems
202 .iter()
203 .filter_map(|e| {
204 if let Expr::Path(path) = e {
205 Some(path.path.clone())
206 } else {
207 extend_errors!(
208 result,
209 syn::Error::new(e.span(), "Expected Path or Type")
210 );
211 None
212 }
213 })
214 .collect::<Vec<_>>();
215
216 result = result.map(|d| d.with_derives(paths));
217 derives.push(meta);
218 } else if segment.ident == "derive_serde" {
219 has_derive_serde = true;
220 let Expr::Lit(expr_lit) = &meta.value else {
221 extend_errors!(
222 result,
223 syn::Error::new(meta.value.span(), "expected literal")
224 );
225 continue;
226 };
227 match expr_lit.lit {
228 Lit::Bool(LitBool { value: true, .. }) if cfg!(feature = "serde1") => {
229 result = result.map(|d| DeriveMeta {
230 derive: Some(Derive::Serde(true)),
231 ..d
232 })
233 }
234 Lit::Bool(LitBool { value: true, .. }) => {
235 extend_errors!(
236 result,
237 syn::Error::new(
238 meta.span(),
239 "To enable serde, first enable the `serde1` feature of tarpc"
240 )
241 );
242 }
243 Lit::Bool(LitBool { value: false, .. }) => {
244 result = result.map(|d| DeriveMeta {
245 derive: Some(Derive::Serde(false)),
246 ..d
247 })
248 }
249 _ => extend_errors!(
250 result,
251 syn::Error::new(
252 expr_lit.lit.span(),
253 "`derive_serde` expects a value of type `bool`"
254 )
255 ),
256 }
257 derive_serde.push(meta);
258 } else {
259 extend_errors!(
260 result,
261 syn::Error::new(
262 meta.span(),
263 "tarpc::service does not support this meta item"
264 )
265 );
266 continue;
267 }
268 }
269
270 if has_derive_serde {
271 let deprecation_hack = quote! {
272 const _: () = {
273 #[deprecated(
274 note = "\nThe form `tarpc::service(derive_serde = true)` is deprecated.\
275 \nUse `tarpc::service(derive = [Serialize, Deserialize])`."
276 )]
277 const DEPRECATED_SYNTAX: () = ();
278 let _ = DEPRECATED_SYNTAX;
279 };
280 };
281
282 result = result.map(|mut d| {
283 d.warnings.push(deprecation_hack.to_token_stream());
284 d
285 });
286 }
287
288 if has_explicit_derives & has_derive_serde {
289 extend_errors!(
290 result,
291 syn::Error::new(
292 input.span(),
293 "tarpc does not support `derive_serde` and `derive` at the same time"
294 )
295 );
296 }
297
298 if derive_serde.len() > 1 {
299 for (i, derive_serde) in derive_serde.iter().enumerate() {
300 extend_errors!(
301 result,
302 syn::Error::new(
303 derive_serde.span(),
304 format!(
305 "`derive_serde` appears more than once (occurrence #{})",
306 i + 1
307 )
308 )
309 );
310 }
311 }
312
313 if derives.len() > 1 {
314 for (i, derive) in derives.iter().enumerate() {
315 extend_errors!(
316 result,
317 syn::Error::new(
318 derive.span(),
319 format!("`derive` appears more than once (occurrence #{})", i + 1)
320 )
321 );
322 }
323 }
324
325 result
326 }
327}
328
329#[proc_macro_attribute]
339#[cfg(feature = "serde1")]
340pub fn derive_serde(_attr: TokenStream, item: TokenStream) -> TokenStream {
341 let mut derives: proc_macro2::TokenStream = quote! {
342 #[derive(::tarpc::serde::Serialize, ::tarpc::serde::Deserialize)]
343 #[serde(crate = "::tarpc::serde")]
344 };
345 derives.extend(proc_macro2::TokenStream::from(item));
346 proc_macro::TokenStream::from(derives)
347}
348
349fn collect_cfg_attrs(rpcs: &[RpcMethod]) -> Vec<Vec<&Attribute>> {
350 rpcs.iter()
351 .map(|rpc| {
352 rpc.attrs
353 .iter()
354 .filter(|att| {
355 att.style == AttrStyle::Outer
356 && match &att.meta {
357 syn::Meta::List(syn::MetaList { path, .. }) => {
358 path.get_ident() == Some(&Ident::new("cfg", rpc.ident.span()))
359 }
360 _ => false,
361 }
362 })
363 .collect::<Vec<_>>()
364 })
365 .collect::<Vec<_>>()
366}
367
368#[proc_macro_attribute]
414pub fn service(attr: TokenStream, input: TokenStream) -> TokenStream {
415 let derive_meta = parse_macro_input!(attr as DeriveMeta);
416 let unit_type: &Type = &parse_quote!(());
417 let Service {
418 ref attrs,
419 ref vis,
420 ref ident,
421 ref rpcs,
422 } = parse_macro_input!(input as Service);
423
424 let camel_case_fn_names: &Vec<_> = &rpcs
425 .iter()
426 .map(|rpc| snake_to_camel(&rpc.ident.unraw().to_string()))
427 .collect();
428 let args: &[&[PatType]] = &rpcs.iter().map(|rpc| &*rpc.args).collect::<Vec<_>>();
429
430 let derives = match derive_meta.derive.as_ref() {
431 Some(Derive::Explicit(paths)) => {
432 if !paths.is_empty() {
433 Some(quote! {
434 #[derive(
435 #(
436 #paths
437 ),*
438 )]
439 })
440 } else {
441 None
442 }
443 }
444 Some(Derive::Serde(serde)) => {
445 if *serde {
446 Some(quote! {
447 #[derive(::tarpc::serde::Serialize, ::tarpc::serde::Deserialize)]
448 #[serde(crate = "::tarpc::serde")]
449 })
450 } else {
451 None
452 }
453 }
454 None => {
455 if cfg!(feature = "serde1") {
456 Some(quote! {
457 #[derive(::tarpc::serde::Serialize, ::tarpc::serde::Deserialize)]
458 #[serde(crate = "::tarpc::serde")]
459 })
460 } else {
461 None
462 }
463 }
464 };
465
466 let methods = rpcs.iter().map(|rpc| &rpc.ident).collect::<Vec<_>>();
467 let request_names = methods
468 .iter()
469 .map(|m| format!("{ident}.{m}"))
470 .collect::<Vec<_>>();
471
472 ServiceGenerator {
473 service_ident: ident,
474 client_stub_ident: &format_ident!("{}Stub", ident),
475 server_ident: &format_ident!("Serve{}", ident),
476 client_ident: &format_ident!("{}Client", ident),
477 request_ident: &format_ident!("{}Request", ident),
478 response_ident: &format_ident!("{}Response", ident),
479 vis,
480 args,
481 method_attrs: &rpcs.iter().map(|rpc| &*rpc.attrs).collect::<Vec<_>>(),
482 method_cfgs: &collect_cfg_attrs(rpcs),
483 method_idents: &methods,
484 request_names: &request_names,
485 attrs,
486 rpcs,
487 return_types: &rpcs
488 .iter()
489 .map(|rpc| match rpc.output {
490 ReturnType::Type(_, ref ty) => ty.as_ref(),
491 ReturnType::Default => unit_type,
492 })
493 .collect::<Vec<_>>(),
494 arg_pats: &args
495 .iter()
496 .map(|args| args.iter().map(|arg| &*arg.pat).collect())
497 .collect::<Vec<_>>(),
498 camel_case_idents: &rpcs
499 .iter()
500 .zip(camel_case_fn_names.iter())
501 .map(|(rpc, name)| Ident::new(name, rpc.ident.span()))
502 .collect::<Vec<_>>(),
503 derives: derives.as_ref(),
504 warnings: &derive_meta.warnings,
505 }
506 .into_token_stream()
507 .into()
508}
509
510struct ServiceGenerator<'a> {
513 service_ident: &'a Ident,
514 client_stub_ident: &'a Ident,
515 server_ident: &'a Ident,
516 client_ident: &'a Ident,
517 request_ident: &'a Ident,
518 response_ident: &'a Ident,
519 vis: &'a Visibility,
520 attrs: &'a [Attribute],
521 rpcs: &'a [RpcMethod],
522 camel_case_idents: &'a [Ident],
523 method_idents: &'a [&'a Ident],
524 request_names: &'a [String],
525 method_attrs: &'a [&'a [Attribute]],
526 method_cfgs: &'a [Vec<&'a Attribute>],
527 args: &'a [&'a [PatType]],
528 return_types: &'a [&'a Type],
529 arg_pats: &'a [Vec<&'a Pat>],
530 derives: Option<&'a TokenStream2>,
531 warnings: &'a [TokenStream2],
532}
533
534impl ServiceGenerator<'_> {
535 fn trait_service(&self) -> TokenStream2 {
536 let &Self {
537 attrs,
538 rpcs,
539 vis,
540 return_types,
541 service_ident,
542 client_stub_ident,
543 request_ident,
544 response_ident,
545 server_ident,
546 ..
547 } = self;
548
549 let rpc_fns = rpcs
550 .iter()
551 .zip(return_types.iter())
552 .map(
553 |(
554 RpcMethod {
555 attrs, ident, args, ..
556 },
557 output,
558 )| {
559 quote! {
560 #( #attrs )*
561 async fn #ident(self, context: ::tarpc::context::Context, #( #args ),*) -> #output;
562 }
563 },
564 );
565
566 let stub_doc = format!("The stub trait for service [`{service_ident}`].");
567 quote! {
568 #( #attrs )*
569 #vis trait #service_ident: ::core::marker::Sized {
570 #( #rpc_fns )*
571
572 fn serve(self) -> #server_ident<Self> {
575 #server_ident { service: self }
576 }
577 }
578
579 #[doc = #stub_doc]
580 #vis trait #client_stub_ident: ::tarpc::client::stub::Stub<Req = #request_ident, Resp = #response_ident> {
581 }
582
583 impl<S> #client_stub_ident for S
584 where S: ::tarpc::client::stub::Stub<Req = #request_ident, Resp = #response_ident>
585 {
586 }
587 }
588 }
589
590 fn struct_server(&self) -> TokenStream2 {
591 let &Self {
592 vis, server_ident, ..
593 } = self;
594
595 quote! {
596 #[derive(Clone)]
598 #vis struct #server_ident<S> {
599 service: S,
600 }
601 }
602 }
603
604 fn impl_serve_for_server(&self) -> TokenStream2 {
605 let &Self {
606 request_ident,
607 server_ident,
608 service_ident,
609 response_ident,
610 camel_case_idents,
611 arg_pats,
612 method_idents,
613 method_cfgs,
614 ..
615 } = self;
616
617 quote! {
618 impl<S> ::tarpc::server::Serve for #server_ident<S>
619 where S: #service_ident
620 {
621 type Req = #request_ident;
622 type Resp = #response_ident;
623
624
625 async fn serve(self, ctx: ::tarpc::context::Context, req: #request_ident)
626 -> ::core::result::Result<#response_ident, ::tarpc::ServerError> {
627 match req {
628 #(
629 #( #method_cfgs )*
630 #request_ident::#camel_case_idents{ #( #arg_pats ),* } => {
631 ::core::result::Result::Ok(#response_ident::#camel_case_idents(
632 #service_ident::#method_idents(
633 self.service, ctx, #( #arg_pats ),*
634 ).await
635 ))
636 }
637 )*
638 }
639 }
640 }
641 }
642 }
643
644 fn enum_request(&self) -> TokenStream2 {
645 let &Self {
646 derives,
647 vis,
648 request_ident,
649 camel_case_idents,
650 args,
651 request_names,
652 method_cfgs,
653 ..
654 } = self;
655
656 quote! {
657 #[allow(missing_docs)]
659 #[derive(Debug)]
660 #derives
661 #vis enum #request_ident {
662 #(
663 #( #method_cfgs )*
664 #camel_case_idents{ #( #args ),* }
665 ),*
666 }
667 impl ::tarpc::RequestName for #request_ident {
668 fn name(&self) -> &str {
669 match self {
670 #(
671 #( #method_cfgs )*
672 #request_ident::#camel_case_idents{..} => {
673 #request_names
674 }
675 )*
676 }
677 }
678 }
679 }
680 }
681
682 fn enum_response(&self) -> TokenStream2 {
683 let &Self {
684 derives,
685 vis,
686 response_ident,
687 camel_case_idents,
688 return_types,
689 ..
690 } = self;
691
692 quote! {
693 #[allow(missing_docs)]
695 #[derive(Debug)]
696 #derives
697 #vis enum #response_ident {
698 #( #camel_case_idents(#return_types) ),*
699 }
700 }
701 }
702
703 fn struct_client(&self) -> TokenStream2 {
704 let &Self {
705 vis,
706 client_ident,
707 request_ident,
708 response_ident,
709 ..
710 } = self;
711
712 quote! {
713 #[allow(unused)]
714 #[derive(Clone, Debug)]
715 #vis struct #client_ident<
718 Stub = ::tarpc::client::Channel<#request_ident, #response_ident>
719 >(Stub);
720 }
721 }
722
723 fn impl_client_new(&self) -> TokenStream2 {
724 let &Self {
725 client_ident,
726 vis,
727 request_ident,
728 response_ident,
729 ..
730 } = self;
731
732 quote! {
733 impl #client_ident {
734 #vis fn new<T>(config: ::tarpc::client::Config, transport: T)
736 -> ::tarpc::client::NewClient<
737 Self,
738 ::tarpc::client::RequestDispatch<#request_ident, #response_ident, T>
739 >
740 where
741 T: ::tarpc::Transport<::tarpc::ClientMessage<#request_ident>, ::tarpc::Response<#response_ident>>
742 {
743 let new_client = ::tarpc::client::new(config, transport);
744 ::tarpc::client::NewClient {
745 client: #client_ident(new_client.client),
746 dispatch: new_client.dispatch,
747 }
748 }
749 }
750
751 impl<Stub> ::core::convert::From<Stub> for #client_ident<Stub>
752 where Stub: ::tarpc::client::stub::Stub<
753 Req = #request_ident,
754 Resp = #response_ident>
755 {
756 fn from(stub: Stub) -> Self {
758 #client_ident(stub)
759 }
760
761 }
762 }
763 }
764
765 fn impl_client_rpc_methods(&self) -> TokenStream2 {
766 let &Self {
767 client_ident,
768 request_ident,
769 response_ident,
770 method_attrs,
771 vis,
772 method_idents,
773 args,
774 return_types,
775 arg_pats,
776 camel_case_idents,
777 ..
778 } = self;
779
780 quote! {
781 impl<Stub> #client_ident<Stub>
782 where Stub: ::tarpc::client::stub::Stub<
783 Req = #request_ident,
784 Resp = #response_ident>
785 {
786 #(
787 #[allow(unused)]
788 #( #method_attrs )*
789 #vis fn #method_idents(&self, ctx: ::tarpc::context::Context, #( #args ),*)
790 -> impl ::core::future::Future<Output = ::core::result::Result<#return_types, ::tarpc::client::RpcError>> + '_ {
791 let request = #request_ident::#camel_case_idents { #( #arg_pats ),* };
792 let resp = self.0.call(ctx, request);
793 async move {
794 match resp.await? {
795 #response_ident::#camel_case_idents(msg) => ::core::result::Result::Ok(msg),
796 _ => ::core::unreachable!(),
797 }
798 }
799 }
800 )*
801 }
802 }
803 }
804
805 fn emit_warnings(&self) -> TokenStream2 {
806 self.warnings.iter().map(|w| w.to_token_stream()).collect()
807 }
808}
809
810impl ToTokens for ServiceGenerator<'_> {
811 fn to_tokens(&self, output: &mut TokenStream2) {
812 output.extend(vec![
813 self.trait_service(),
814 self.struct_server(),
815 self.impl_serve_for_server(),
816 self.enum_request(),
817 self.enum_response(),
818 self.struct_client(),
819 self.impl_client_new(),
820 self.impl_client_rpc_methods(),
821 self.emit_warnings(),
822 ]);
823 }
824}
825
826fn snake_to_camel(ident_str: &str) -> String {
827 let mut camel_ty = String::with_capacity(ident_str.len());
828
829 let mut last_char_was_underscore = true;
830 for c in ident_str.chars() {
831 match c {
832 '_' => last_char_was_underscore = true,
833 c if last_char_was_underscore => {
834 camel_ty.extend(c.to_uppercase());
835 last_char_was_underscore = false;
836 }
837 c => camel_ty.extend(c.to_lowercase()),
838 }
839 }
840
841 camel_ty.shrink_to_fit();
842 camel_ty
843}
844
845#[test]
846fn snake_to_camel_basic() {
847 assert_eq!(snake_to_camel("abc_def"), "AbcDef");
848}
849
850#[test]
851fn snake_to_camel_underscore_suffix() {
852 assert_eq!(snake_to_camel("abc_def_"), "AbcDef");
853}
854
855#[test]
856fn snake_to_camel_underscore_prefix() {
857 assert_eq!(snake_to_camel("_abc_def"), "AbcDef");
858}
859
860#[test]
861fn snake_to_camel_underscore_consecutive() {
862 assert_eq!(snake_to_camel("abc__def"), "AbcDef");
863}
864
865#[test]
866fn snake_to_camel_capital_in_middle() {
867 assert_eq!(snake_to_camel("aBc_dEf"), "AbcDef");
868}