xref: /tonic/tonic-build/src/server.rs (revision 63f2e95b)
1 use std::collections::HashSet;
2 
3 use super::{Attributes, Method, Service};
4 use crate::{format_method_name, generate_doc_comment, generate_doc_comments, naive_snake_case};
5 use proc_macro2::{Span, TokenStream};
6 use quote::quote;
7 use syn::{Ident, Lit, LitStr};
8 
9 /// Generate service for Server.
10 ///
11 /// This takes some `Service` and will generate a `TokenStream` that contains
12 /// a public module containing the server service and handler trait.
13 #[deprecated(since = "0.8.3", note = "Use CodeGenBuilder::generate_server")]
14 pub fn generate<T: Service>(
15     service: &T,
16     emit_package: bool,
17     proto_path: &str,
18     compile_well_known_types: bool,
19     attributes: &Attributes,
20 ) -> TokenStream {
21     generate_internal(
22         service,
23         emit_package,
24         proto_path,
25         compile_well_known_types,
26         attributes,
27         &HashSet::default(),
28     )
29 }
30 
31 pub(crate) fn generate_internal<T: Service>(
32     service: &T,
33     emit_package: bool,
34     proto_path: &str,
35     compile_well_known_types: bool,
36     attributes: &Attributes,
37     disable_comments: &HashSet<String>,
38 ) -> TokenStream {
39     let methods = generate_methods(service, proto_path, compile_well_known_types);
40 
41     let server_service = quote::format_ident!("{}Server", service.name());
42     let server_trait = quote::format_ident!("{}", service.name());
43     let server_mod = quote::format_ident!("{}_server", naive_snake_case(service.name()));
44     let generated_trait = generate_trait(
45         service,
46         emit_package,
47         proto_path,
48         compile_well_known_types,
49         server_trait.clone(),
50         disable_comments,
51     );
52     let package = if emit_package { service.package() } else { "" };
53     // Transport based implementations
54     let path = format!(
55         "{}{}{}",
56         package,
57         if package.is_empty() { "" } else { "." },
58         service.identifier()
59     );
60 
61     let service_doc = if disable_comments.contains(&path) {
62         TokenStream::new()
63     } else {
64         generate_doc_comments(service.comment())
65     };
66 
67     let named = generate_named(&server_service, &server_trait, &path);
68     let mod_attributes = attributes.for_mod(package);
69     let struct_attributes = attributes.for_struct(&path);
70 
71     let configure_compression_methods = quote! {
72         /// Enable decompressing requests with the given encoding.
73         #[must_use]
74         pub fn accept_compressed(mut self, encoding: CompressionEncoding) -> Self {
75             self.accept_compression_encodings.enable(encoding);
76             self
77         }
78 
79         /// Compress responses with the given encoding, if the client supports it.
80         #[must_use]
81         pub fn send_compressed(mut self, encoding: CompressionEncoding) -> Self {
82             self.send_compression_encodings.enable(encoding);
83             self
84         }
85     };
86 
87     quote! {
88         /// Generated server implementations.
89         #(#mod_attributes)*
90         pub mod #server_mod {
91             #![allow(
92                 unused_variables,
93                 dead_code,
94                 missing_docs,
95                 // will trigger if compression is disabled
96                 clippy::let_unit_value,
97             )]
98             use tonic::codegen::*;
99 
100             #generated_trait
101 
102             #service_doc
103             #(#struct_attributes)*
104             #[derive(Debug)]
105             pub struct #server_service<T: #server_trait> {
106                 inner: _Inner<T>,
107                 accept_compression_encodings: EnabledCompressionEncodings,
108                 send_compression_encodings: EnabledCompressionEncodings,
109             }
110 
111             struct _Inner<T>(Arc<T>);
112 
113             impl<T: #server_trait> #server_service<T> {
114                 pub fn new(inner: T) -> Self {
115                     Self::from_arc(Arc::new(inner))
116                 }
117 
118                 pub fn from_arc(inner: Arc<T>) -> Self {
119                     let inner = _Inner(inner);
120                     Self {
121                         inner,
122                         accept_compression_encodings: Default::default(),
123                         send_compression_encodings: Default::default(),
124                     }
125                 }
126 
127                 pub fn with_interceptor<F>(inner: T, interceptor: F) -> InterceptedService<Self, F>
128                 where
129                     F: tonic::service::Interceptor,
130                 {
131                     InterceptedService::new(Self::new(inner), interceptor)
132                 }
133 
134                 #configure_compression_methods
135             }
136 
137             impl<T, B> tonic::codegen::Service<http::Request<B>> for #server_service<T>
138                 where
139                     T: #server_trait,
140                     B: Body + Send + 'static,
141                     B::Error: Into<StdError> + Send + 'static,
142             {
143                 type Response = http::Response<tonic::body::BoxBody>;
144                 type Error = std::convert::Infallible;
145                 type Future = BoxFuture<Self::Response, Self::Error>;
146 
147                 fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<std::result::Result<(), Self::Error>> {
148                     Poll::Ready(Ok(()))
149                 }
150 
151                 fn call(&mut self, req: http::Request<B>) -> Self::Future {
152                     let inner = self.inner.clone();
153 
154                     match req.uri().path() {
155                         #methods
156 
157                         _ => Box::pin(async move {
158                             Ok(http::Response::builder()
159                                .status(200)
160                                .header("grpc-status", "12")
161                                .header("content-type", "application/grpc")
162                                .body(empty_body())
163                                .unwrap())
164                         }),
165                     }
166                 }
167             }
168 
169             impl<T: #server_trait> Clone for #server_service<T> {
170                 fn clone(&self) -> Self {
171                     let inner = self.inner.clone();
172                     Self {
173                         inner,
174                         accept_compression_encodings: self.accept_compression_encodings,
175                         send_compression_encodings: self.send_compression_encodings,
176                     }
177                 }
178             }
179 
180             impl<T: #server_trait> Clone for _Inner<T> {
181                 fn clone(&self) -> Self {
182                     Self(Arc::clone(&self.0))
183                 }
184             }
185 
186             impl<T: std::fmt::Debug> std::fmt::Debug for _Inner<T> {
187                 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
188                    write!(f, "{:?}", self.0)
189                 }
190             }
191 
192             #named
193         }
194     }
195 }
196 
197 fn generate_trait<T: Service>(
198     service: &T,
199     emit_package: bool,
200     proto_path: &str,
201     compile_well_known_types: bool,
202     server_trait: Ident,
203     disable_comments: &HashSet<String>,
204 ) -> TokenStream {
205     let methods = generate_trait_methods(
206         service,
207         emit_package,
208         proto_path,
209         compile_well_known_types,
210         disable_comments,
211     );
212     let trait_doc = generate_doc_comment(format!(
213         " Generated trait containing gRPC methods that should be implemented for use with {}Server.",
214         service.name()
215     ));
216 
217     quote! {
218         #trait_doc
219         #[async_trait]
220         pub trait #server_trait : Send + Sync + 'static {
221             #methods
222         }
223     }
224 }
225 
226 fn generate_trait_methods<T: Service>(
227     service: &T,
228     emit_package: bool,
229     proto_path: &str,
230     compile_well_known_types: bool,
231     disable_comments: &HashSet<String>,
232 ) -> TokenStream {
233     let mut stream = TokenStream::new();
234 
235     let package = if emit_package { service.package() } else { "" };
236     for method in service.methods() {
237         let name = quote::format_ident!("{}", method.name());
238 
239         let (req_message, res_message) =
240             method.request_response_name(proto_path, compile_well_known_types);
241 
242         let method_doc = if disable_comments.contains(&format_method_name(package, service, method))
243         {
244             TokenStream::new()
245         } else {
246             generate_doc_comments(method.comment())
247         };
248 
249         let method = match (method.client_streaming(), method.server_streaming()) {
250             (false, false) => {
251                 quote! {
252                     #method_doc
253                     async fn #name(&self, request: tonic::Request<#req_message>)
254                         -> std::result::Result<tonic::Response<#res_message>, tonic::Status>;
255                 }
256             }
257             (true, false) => {
258                 quote! {
259                     #method_doc
260                     async fn #name(&self, request: tonic::Request<tonic::Streaming<#req_message>>)
261                         -> std::result::Result<tonic::Response<#res_message>, tonic::Status>;
262                 }
263             }
264             (false, true) => {
265                 let stream = quote::format_ident!("{}Stream", method.identifier());
266                 let stream_doc = generate_doc_comment(format!(
267                     " Server streaming response type for the {} method.",
268                     method.identifier()
269                 ));
270 
271                 quote! {
272                     #stream_doc
273                     type #stream: futures_core::Stream<Item = std::result::Result<#res_message, tonic::Status>> + Send + 'static;
274 
275                     #method_doc
276                     async fn #name(&self, request: tonic::Request<#req_message>)
277                         -> std::result::Result<tonic::Response<Self::#stream>, tonic::Status>;
278                 }
279             }
280             (true, true) => {
281                 let stream = quote::format_ident!("{}Stream", method.identifier());
282                 let stream_doc = generate_doc_comment(format!(
283                     " Server streaming response type for the {} method.",
284                     method.identifier()
285                 ));
286 
287                 quote! {
288                     #stream_doc
289                     type #stream: futures_core::Stream<Item = std::result::Result<#res_message, tonic::Status>> + Send + 'static;
290 
291                     #method_doc
292                     async fn #name(&self, request: tonic::Request<tonic::Streaming<#req_message>>)
293                         -> std::result::Result<tonic::Response<Self::#stream>, tonic::Status>;
294                 }
295             }
296         };
297 
298         stream.extend(method);
299     }
300 
301     stream
302 }
303 
304 fn generate_named(
305     server_service: &syn::Ident,
306     server_trait: &syn::Ident,
307     service_name: &str,
308 ) -> TokenStream {
309     let service_name = syn::LitStr::new(service_name, proc_macro2::Span::call_site());
310 
311     quote! {
312         impl<T: #server_trait> tonic::server::NamedService for #server_service<T> {
313             const NAME: &'static str = #service_name;
314         }
315     }
316 }
317 
318 fn generate_methods<T: Service>(
319     service: &T,
320     proto_path: &str,
321     compile_well_known_types: bool,
322 ) -> TokenStream {
323     let mut stream = TokenStream::new();
324 
325     for method in service.methods() {
326         let path = format!(
327             "/{}{}{}/{}",
328             service.package(),
329             if service.package().is_empty() {
330                 ""
331             } else {
332                 "."
333             },
334             service.identifier(),
335             method.identifier()
336         );
337         let method_path = Lit::Str(LitStr::new(&path, Span::call_site()));
338         let ident = quote::format_ident!("{}", method.name());
339         let server_trait = quote::format_ident!("{}", service.name());
340 
341         let method_stream = match (method.client_streaming(), method.server_streaming()) {
342             (false, false) => generate_unary(
343                 method,
344                 proto_path,
345                 compile_well_known_types,
346                 ident,
347                 server_trait,
348             ),
349 
350             (false, true) => generate_server_streaming(
351                 method,
352                 proto_path,
353                 compile_well_known_types,
354                 ident.clone(),
355                 server_trait,
356             ),
357             (true, false) => generate_client_streaming(
358                 method,
359                 proto_path,
360                 compile_well_known_types,
361                 ident.clone(),
362                 server_trait,
363             ),
364 
365             (true, true) => generate_streaming(
366                 method,
367                 proto_path,
368                 compile_well_known_types,
369                 ident.clone(),
370                 server_trait,
371             ),
372         };
373 
374         let method = quote! {
375             #method_path => {
376                 #method_stream
377             }
378         };
379         stream.extend(method);
380     }
381 
382     stream
383 }
384 
385 fn generate_unary<T: Method>(
386     method: &T,
387     proto_path: &str,
388     compile_well_known_types: bool,
389     method_ident: Ident,
390     server_trait: Ident,
391 ) -> TokenStream {
392     let codec_name = syn::parse_str::<syn::Path>(method.codec_path()).unwrap();
393 
394     let service_ident = quote::format_ident!("{}Svc", method.identifier());
395 
396     let (request, response) = method.request_response_name(proto_path, compile_well_known_types);
397 
398     quote! {
399         #[allow(non_camel_case_types)]
400         struct #service_ident<T: #server_trait >(pub Arc<T>);
401 
402         impl<T: #server_trait> tonic::server::UnaryService<#request> for #service_ident<T> {
403             type Response = #response;
404             type Future = BoxFuture<tonic::Response<Self::Response>, tonic::Status>;
405 
406             fn call(&mut self, request: tonic::Request<#request>) -> Self::Future {
407                 let inner = Arc::clone(&self.0);
408                 let fut = async move {
409                     (*inner).#method_ident(request).await
410                 };
411                 Box::pin(fut)
412             }
413         }
414 
415         let accept_compression_encodings = self.accept_compression_encodings;
416         let send_compression_encodings = self.send_compression_encodings;
417         let inner = self.inner.clone();
418         let fut = async move {
419             let inner = inner.0;
420             let method = #service_ident(inner);
421             let codec = #codec_name::default();
422 
423             let mut grpc = tonic::server::Grpc::new(codec)
424                 .apply_compression_config(accept_compression_encodings, send_compression_encodings);
425 
426             let res = grpc.unary(method, req).await;
427             Ok(res)
428         };
429 
430         Box::pin(fut)
431     }
432 }
433 
434 fn generate_server_streaming<T: Method>(
435     method: &T,
436     proto_path: &str,
437     compile_well_known_types: bool,
438     method_ident: Ident,
439     server_trait: Ident,
440 ) -> TokenStream {
441     let codec_name = syn::parse_str::<syn::Path>(method.codec_path()).unwrap();
442 
443     let service_ident = quote::format_ident!("{}Svc", method.identifier());
444 
445     let (request, response) = method.request_response_name(proto_path, compile_well_known_types);
446 
447     let response_stream = quote::format_ident!("{}Stream", method.identifier());
448 
449     quote! {
450         #[allow(non_camel_case_types)]
451         struct #service_ident<T: #server_trait >(pub Arc<T>);
452 
453         impl<T: #server_trait> tonic::server::ServerStreamingService<#request> for #service_ident<T> {
454             type Response = #response;
455             type ResponseStream = T::#response_stream;
456             type Future = BoxFuture<tonic::Response<Self::ResponseStream>, tonic::Status>;
457 
458             fn call(&mut self, request: tonic::Request<#request>) -> Self::Future {
459                 let inner = Arc::clone(&self.0);
460                 let fut = async move {
461                     (*inner).#method_ident(request).await
462                 };
463                 Box::pin(fut)
464             }
465         }
466 
467         let accept_compression_encodings = self.accept_compression_encodings;
468         let send_compression_encodings = self.send_compression_encodings;
469         let inner = self.inner.clone();
470         let fut = async move {
471             let inner = inner.0;
472             let method = #service_ident(inner);
473             let codec = #codec_name::default();
474 
475             let mut grpc = tonic::server::Grpc::new(codec)
476                 .apply_compression_config(accept_compression_encodings, send_compression_encodings);
477 
478             let res = grpc.server_streaming(method, req).await;
479             Ok(res)
480         };
481 
482         Box::pin(fut)
483     }
484 }
485 
486 fn generate_client_streaming<T: Method>(
487     method: &T,
488     proto_path: &str,
489     compile_well_known_types: bool,
490     method_ident: Ident,
491     server_trait: Ident,
492 ) -> TokenStream {
493     let service_ident = quote::format_ident!("{}Svc", method.identifier());
494 
495     let (request, response) = method.request_response_name(proto_path, compile_well_known_types);
496     let codec_name = syn::parse_str::<syn::Path>(method.codec_path()).unwrap();
497 
498     quote! {
499         #[allow(non_camel_case_types)]
500         struct #service_ident<T: #server_trait >(pub Arc<T>);
501 
502         impl<T: #server_trait> tonic::server::ClientStreamingService<#request> for #service_ident<T>
503         {
504             type Response = #response;
505             type Future = BoxFuture<tonic::Response<Self::Response>, tonic::Status>;
506 
507             fn call(&mut self, request: tonic::Request<tonic::Streaming<#request>>) -> Self::Future {
508                 let inner = Arc::clone(&self.0);
509                 let fut = async move {
510                     (*inner).#method_ident(request).await
511 
512                 };
513                 Box::pin(fut)
514             }
515         }
516 
517         let accept_compression_encodings = self.accept_compression_encodings;
518         let send_compression_encodings = self.send_compression_encodings;
519         let inner = self.inner.clone();
520         let fut = async move {
521             let inner = inner.0;
522             let method = #service_ident(inner);
523             let codec = #codec_name::default();
524 
525             let mut grpc = tonic::server::Grpc::new(codec)
526                 .apply_compression_config(accept_compression_encodings, send_compression_encodings);
527 
528             let res = grpc.client_streaming(method, req).await;
529             Ok(res)
530         };
531 
532         Box::pin(fut)
533     }
534 }
535 
536 fn generate_streaming<T: Method>(
537     method: &T,
538     proto_path: &str,
539     compile_well_known_types: bool,
540     method_ident: Ident,
541     server_trait: Ident,
542 ) -> TokenStream {
543     let codec_name = syn::parse_str::<syn::Path>(method.codec_path()).unwrap();
544 
545     let service_ident = quote::format_ident!("{}Svc", method.identifier());
546 
547     let (request, response) = method.request_response_name(proto_path, compile_well_known_types);
548 
549     let response_stream = quote::format_ident!("{}Stream", method.identifier());
550 
551     quote! {
552         #[allow(non_camel_case_types)]
553         struct #service_ident<T: #server_trait>(pub Arc<T>);
554 
555         impl<T: #server_trait> tonic::server::StreamingService<#request> for #service_ident<T>
556         {
557             type Response = #response;
558             type ResponseStream = T::#response_stream;
559             type Future = BoxFuture<tonic::Response<Self::ResponseStream>, tonic::Status>;
560 
561             fn call(&mut self, request: tonic::Request<tonic::Streaming<#request>>) -> Self::Future {
562                 let inner = Arc::clone(&self.0);
563                 let fut = async move {
564                     (*inner).#method_ident(request).await
565                 };
566                 Box::pin(fut)
567             }
568         }
569 
570         let accept_compression_encodings = self.accept_compression_encodings;
571         let send_compression_encodings = self.send_compression_encodings;
572         let inner = self.inner.clone();
573         let fut = async move {
574             let inner = inner.0;
575             let method = #service_ident(inner);
576             let codec = #codec_name::default();
577 
578             let mut grpc = tonic::server::Grpc::new(codec)
579                 .apply_compression_config(accept_compression_encodings, send_compression_encodings);
580 
581             let res = grpc.streaming(method, req).await;
582             Ok(res)
583         };
584 
585         Box::pin(fut)
586     }
587 }
588