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