xref: /tonic/tonic-reflection/src/server/v1.rs (revision c8754f3a)
1 use crate::pb::v1::server_reflection_server::{ServerReflection, ServerReflectionServer};
2 
3 use crate::pb::v1::server_reflection_request::MessageRequest;
4 use crate::pb::v1::server_reflection_response::MessageResponse;
5 use crate::pb::v1::{
6     ExtensionNumberResponse, FileDescriptorResponse, ListServiceResponse, ServerReflectionRequest,
7     ServerReflectionResponse, ServiceResponse,
8 };
9 use prost::Message;
10 use prost_types::{
11     DescriptorProto, EnumDescriptorProto, FieldDescriptorProto, FileDescriptorProto,
12     FileDescriptorSet,
13 };
14 use std::collections::HashMap;
15 use std::sync::Arc;
16 use tokio::sync::mpsc;
17 use tokio_stream::{wrappers::ReceiverStream, StreamExt};
18 use tonic::{Request, Response, Status, Streaming};
19 
20 use crate::server::Error;
21 
22 /// A builder used to construct a gRPC Reflection Service.
23 #[derive(Debug)]
24 pub struct Builder<'b> {
25     file_descriptor_sets: Vec<FileDescriptorSet>,
26     encoded_file_descriptor_sets: Vec<&'b [u8]>,
27     include_reflection_service: bool,
28 
29     service_names: Vec<String>,
30     use_all_service_names: bool,
31     symbols: HashMap<String, Arc<FileDescriptorProto>>,
32 }
33 
34 impl<'b> Builder<'b> {
35     /// Create a new builder that can configure a gRPC Reflection Service.
36     pub fn configure() -> Self {
37         Builder {
38             file_descriptor_sets: Vec::new(),
39             encoded_file_descriptor_sets: Vec::new(),
40             include_reflection_service: true,
41 
42             service_names: Vec::new(),
43             use_all_service_names: true,
44             symbols: HashMap::new(),
45         }
46     }
47 
48     /// Registers an instance of `prost_types::FileDescriptorSet` with the gRPC Reflection
49     /// Service builder.
50     pub fn register_file_descriptor_set(mut self, file_descriptor_set: FileDescriptorSet) -> Self {
51         self.file_descriptor_sets.push(file_descriptor_set);
52         self
53     }
54 
55     /// Registers a byte slice containing an encoded `prost_types::FileDescriptorSet` with
56     /// the gRPC Reflection Service builder.
57     pub fn register_encoded_file_descriptor_set(
58         mut self,
59         encoded_file_descriptor_set: &'b [u8],
60     ) -> Self {
61         self.encoded_file_descriptor_sets
62             .push(encoded_file_descriptor_set);
63         self
64     }
65 
66     /// Serve the gRPC Reflection Service descriptor via the Reflection Service. This is enabled
67     /// by default - set `include` to false to disable.
68     pub fn include_reflection_service(mut self, include: bool) -> Self {
69         self.include_reflection_service = include;
70         self
71     }
72 
73     /// Advertise a fully-qualified gRPC service name.
74     ///
75     /// If not called, then all services present in the registered file descriptor sets
76     /// will be advertised.
77     pub fn with_service_name(mut self, name: impl Into<String>) -> Self {
78         self.use_all_service_names = false;
79         self.service_names.push(name.into());
80         self
81     }
82 
83     /// Build a gRPC Reflection Service to be served via Tonic.
84     pub fn build(mut self) -> Result<ServerReflectionServer<impl ServerReflection>, Error> {
85         if self.include_reflection_service {
86             self = self.register_encoded_file_descriptor_set(crate::pb::v1::FILE_DESCRIPTOR_SET);
87         }
88 
89         for encoded in &self.encoded_file_descriptor_sets {
90             let decoded = FileDescriptorSet::decode(*encoded)?;
91             self.file_descriptor_sets.push(decoded);
92         }
93 
94         let all_fds = self.file_descriptor_sets.clone();
95         let mut files: HashMap<String, Arc<FileDescriptorProto>> = HashMap::new();
96 
97         for fds in all_fds {
98             for fd in fds.file {
99                 let name = match fd.name.clone() {
100                     None => {
101                         return Err(Error::InvalidFileDescriptorSet("missing name".to_string()));
102                     }
103                     Some(n) => n,
104                 };
105 
106                 if files.contains_key(&name) {
107                     continue;
108                 }
109 
110                 let fd = Arc::new(fd);
111                 files.insert(name, fd.clone());
112 
113                 self.process_file(fd)?;
114             }
115         }
116 
117         let service_names = self
118             .service_names
119             .iter()
120             .map(|name| ServiceResponse { name: name.clone() })
121             .collect();
122 
123         Ok(ServerReflectionServer::new(ReflectionService {
124             state: Arc::new(ReflectionServiceState {
125                 service_names,
126                 files,
127                 symbols: self.symbols,
128             }),
129         }))
130     }
131 
132     fn process_file(&mut self, fd: Arc<FileDescriptorProto>) -> Result<(), Error> {
133         let prefix = &fd.package.clone().unwrap_or_default();
134 
135         for msg in &fd.message_type {
136             self.process_message(fd.clone(), prefix, msg)?;
137         }
138 
139         for en in &fd.enum_type {
140             self.process_enum(fd.clone(), prefix, en)?;
141         }
142 
143         for service in &fd.service {
144             let service_name = extract_name(prefix, "service", service.name.as_ref())?;
145             if self.use_all_service_names {
146                 self.service_names.push(service_name.clone());
147             }
148             self.symbols.insert(service_name.clone(), fd.clone());
149 
150             for method in &service.method {
151                 let method_name = extract_name(&service_name, "method", method.name.as_ref())?;
152                 self.symbols.insert(method_name, fd.clone());
153             }
154         }
155 
156         Ok(())
157     }
158 
159     fn process_message(
160         &mut self,
161         fd: Arc<FileDescriptorProto>,
162         prefix: &str,
163         msg: &DescriptorProto,
164     ) -> Result<(), Error> {
165         let message_name = extract_name(prefix, "message", msg.name.as_ref())?;
166         self.symbols.insert(message_name.clone(), fd.clone());
167 
168         for nested in &msg.nested_type {
169             self.process_message(fd.clone(), &message_name, nested)?;
170         }
171 
172         for en in &msg.enum_type {
173             self.process_enum(fd.clone(), &message_name, en)?;
174         }
175 
176         for field in &msg.field {
177             self.process_field(fd.clone(), &message_name, field)?;
178         }
179 
180         for oneof in &msg.oneof_decl {
181             let oneof_name = extract_name(&message_name, "oneof", oneof.name.as_ref())?;
182             self.symbols.insert(oneof_name, fd.clone());
183         }
184 
185         Ok(())
186     }
187 
188     fn process_enum(
189         &mut self,
190         fd: Arc<FileDescriptorProto>,
191         prefix: &str,
192         en: &EnumDescriptorProto,
193     ) -> Result<(), Error> {
194         let enum_name = extract_name(prefix, "enum", en.name.as_ref())?;
195         self.symbols.insert(enum_name.clone(), fd.clone());
196 
197         for value in &en.value {
198             let value_name = extract_name(&enum_name, "enum value", value.name.as_ref())?;
199             self.symbols.insert(value_name, fd.clone());
200         }
201 
202         Ok(())
203     }
204 
205     fn process_field(
206         &mut self,
207         fd: Arc<FileDescriptorProto>,
208         prefix: &str,
209         field: &FieldDescriptorProto,
210     ) -> Result<(), Error> {
211         let field_name = extract_name(prefix, "field", field.name.as_ref())?;
212         self.symbols.insert(field_name, fd);
213         Ok(())
214     }
215 }
216 
217 fn extract_name(
218     prefix: &str,
219     name_type: &str,
220     maybe_name: Option<&String>,
221 ) -> Result<String, Error> {
222     match maybe_name {
223         None => Err(Error::InvalidFileDescriptorSet(format!(
224             "missing {} name",
225             name_type
226         ))),
227         Some(name) => {
228             if prefix.is_empty() {
229                 Ok(name.to_string())
230             } else {
231                 Ok(format!("{}.{}", prefix, name))
232             }
233         }
234     }
235 }
236 
237 #[derive(Debug)]
238 struct ReflectionServiceState {
239     service_names: Vec<ServiceResponse>,
240     files: HashMap<String, Arc<FileDescriptorProto>>,
241     symbols: HashMap<String, Arc<FileDescriptorProto>>,
242 }
243 
244 impl ReflectionServiceState {
245     fn list_services(&self) -> MessageResponse {
246         MessageResponse::ListServicesResponse(ListServiceResponse {
247             service: self.service_names.clone(),
248         })
249     }
250 
251     fn symbol_by_name(&self, symbol: &str) -> Result<MessageResponse, Status> {
252         match self.symbols.get(symbol) {
253             None => Err(Status::not_found(format!("symbol '{}' not found", symbol))),
254             Some(fd) => {
255                 let mut encoded_fd = Vec::new();
256                 if fd.clone().encode(&mut encoded_fd).is_err() {
257                     return Err(Status::internal("encoding error"));
258                 };
259 
260                 Ok(MessageResponse::FileDescriptorResponse(
261                     FileDescriptorResponse {
262                         file_descriptor_proto: vec![encoded_fd],
263                     },
264                 ))
265             }
266         }
267     }
268 
269     fn file_by_filename(&self, filename: &str) -> Result<MessageResponse, Status> {
270         match self.files.get(filename) {
271             None => Err(Status::not_found(format!("file '{}' not found", filename))),
272             Some(fd) => {
273                 let mut encoded_fd = Vec::new();
274                 if fd.clone().encode(&mut encoded_fd).is_err() {
275                     return Err(Status::internal("encoding error"));
276                 }
277 
278                 Ok(MessageResponse::FileDescriptorResponse(
279                     FileDescriptorResponse {
280                         file_descriptor_proto: vec![encoded_fd],
281                     },
282                 ))
283             }
284         }
285     }
286 }
287 
288 #[derive(Debug)]
289 struct ReflectionService {
290     state: Arc<ReflectionServiceState>,
291 }
292 
293 #[tonic::async_trait]
294 impl ServerReflection for ReflectionService {
295     type ServerReflectionInfoStream = ReceiverStream<Result<ServerReflectionResponse, Status>>;
296 
297     async fn server_reflection_info(
298         &self,
299         req: Request<Streaming<ServerReflectionRequest>>,
300     ) -> Result<Response<Self::ServerReflectionInfoStream>, Status> {
301         let mut req_rx = req.into_inner();
302         let (resp_tx, resp_rx) = mpsc::channel::<Result<ServerReflectionResponse, Status>>(1);
303 
304         let state = self.state.clone();
305 
306         tokio::spawn(async move {
307             while let Some(req) = req_rx.next().await {
308                 let Ok(req) = req else {
309                     return;
310                 };
311 
312                 let resp_msg = match req.message_request.clone() {
313                     None => Err(Status::invalid_argument("invalid MessageRequest")),
314                     Some(msg) => match msg {
315                         MessageRequest::FileByFilename(s) => state.file_by_filename(&s),
316                         MessageRequest::FileContainingSymbol(s) => state.symbol_by_name(&s),
317                         MessageRequest::FileContainingExtension(_) => {
318                             Err(Status::not_found("extensions are not supported"))
319                         }
320                         MessageRequest::AllExtensionNumbersOfType(_) => {
321                             // NOTE: Workaround. Some grpc clients (e.g. grpcurl) expect this method not to fail.
322                             // https://github.com/hyperium/tonic/issues/1077
323                             Ok(MessageResponse::AllExtensionNumbersResponse(
324                                 ExtensionNumberResponse::default(),
325                             ))
326                         }
327                         MessageRequest::ListServices(_) => Ok(state.list_services()),
328                     },
329                 };
330 
331                 match resp_msg {
332                     Ok(resp_msg) => {
333                         let resp = ServerReflectionResponse {
334                             valid_host: req.host.clone(),
335                             original_request: Some(req.clone()),
336                             message_response: Some(resp_msg),
337                         };
338                         resp_tx.send(Ok(resp)).await.expect("send");
339                     }
340                     Err(status) => {
341                         resp_tx.send(Err(status)).await.expect("send");
342                         return;
343                     }
344                 }
345             }
346         });
347 
348         Ok(Response::new(ReceiverStream::new(resp_rx)))
349     }
350 }
351