1use std::sync::Arc;
6
7use connectrpc::Router;
8
9use crate::reflector::Reflector;
10
11#[derive(Clone)]
29pub struct ReflectionService {
30 reflector: Arc<Reflector>,
31}
32
33impl ReflectionService {
34 #[must_use]
36 pub fn new(reflector: Reflector) -> Self {
37 Self {
38 reflector: Arc::new(reflector),
39 }
40 }
41
42 #[must_use]
44 pub fn from_arc(reflector: Arc<Reflector>) -> Self {
45 Self { reflector }
46 }
47}
48
49#[must_use]
63pub fn install(router: Router, reflector: Reflector) -> Router {
64 let service = Arc::new(ReflectionService::new(reflector));
65 let router = crate::connect::grpc::reflection::v1::ServerReflectionExt::register(
66 Arc::clone(&service),
67 router,
68 );
69 let router =
70 crate::connect::grpc::reflection::v1alpha::ServerReflectionExt::register(service, router);
71 apply_request_limits(router, request_limits())
72}
73
74pub const MAX_REQUEST_BYTES: usize = 16 * 1024;
84
85#[must_use]
90pub fn request_limits() -> connectrpc::Limits {
91 connectrpc::Limits::default()
92 .with_max_request_body_size(MAX_REQUEST_BYTES + connectrpc::envelope::HEADER_SIZE)
93 .with_max_message_size(MAX_REQUEST_BYTES)
94}
95
96#[must_use]
112pub fn apply_request_limits(mut router: Router, limits: connectrpc::Limits) -> Router {
113 let mut applied = false;
114 for spec in [
115 crate::connect::grpc::reflection::v1::SERVER_REFLECTION_SERVER_REFLECTION_INFO_SPEC,
116 crate::connect::grpc::reflection::v1alpha::SERVER_REFLECTION_SERVER_REFLECTION_INFO_SPEC,
117 ] {
118 if router.has_method(spec.procedure) {
119 router = router.with_route_limits(spec.procedure, limits);
120 applied = true;
121 }
122 }
123 assert!(
124 applied,
125 "connectrpc_reflection::apply_request_limits: no reflection route is registered \
126 on this router — register `ReflectionService` before applying its limits"
127 );
128 router
129}
130
131macro_rules! impl_server_reflection {
137 () => {
138 impl rpc::ServerReflection for crate::ReflectionService {
139 async fn server_reflection_info(
140 &self,
141 _ctx: ::connectrpc::RequestContext,
142 requests: ::connectrpc::ServiceStream<
143 ::connectrpc::StreamMessage<pb::ServerReflectionRequest>,
144 >,
145 ) -> ::connectrpc::ServiceResult<
146 ::connectrpc::ServiceStream<pb::ServerReflectionResponse>,
147 > {
148 use futures::StreamExt;
149 let reflector = ::std::sync::Arc::clone(&self.reflector);
150 let responses = requests.map(move |request| {
151 let request = request?.to_owned_message();
152 respond(&reflector, request)
153 });
154 ::connectrpc::Response::stream_ok(responses)
155 }
156 }
157
158 fn respond(
163 reflector: &$crate::reflector::Reflector,
164 request: pb::ServerReflectionRequest,
165 ) -> Result<pb::ServerReflectionResponse, ::connectrpc::ConnectError> {
166 use pb::server_reflection_request::MessageRequest;
167 use pb::server_reflection_response::MessageResponse;
168 use $crate::reflector::Answer;
169
170 let Some(message_request) = &request.message_request else {
171 return Err(::connectrpc::ConnectError::invalid_argument(
172 "ServerReflectionRequest.message_request is not set",
173 ));
174 };
175
176 let answer = match message_request {
177 MessageRequest::FileByFilename(name) => reflector.file_by_filename(name),
178 MessageRequest::FileContainingSymbol(symbol) => {
179 reflector.file_containing_symbol(symbol)
180 }
181 MessageRequest::FileContainingExtension(ext) => {
182 reflector.file_containing_extension(&ext.containing_type, ext.extension_number)
183 }
184 MessageRequest::AllExtensionNumbersOfType(name) => {
185 reflector.all_extension_numbers_of_type(name)
186 }
187 MessageRequest::ListServices(_) => reflector.list_services(),
188 };
189
190 let message_response = match answer {
191 Answer::Files(file_descriptor_proto) => {
192 MessageResponse::from(pb::FileDescriptorResponse {
193 file_descriptor_proto,
194 ..Default::default()
195 })
196 }
197 Answer::ExtensionNumbers { base_type, numbers } => {
198 MessageResponse::from(pb::ExtensionNumberResponse {
199 base_type_name: base_type,
200 extension_number: numbers,
201 ..Default::default()
202 })
203 }
204 Answer::Services(names) => MessageResponse::from(pb::ListServiceResponse {
205 service: names
206 .into_iter()
207 .map(|name| pb::ServiceResponse {
208 name,
209 ..Default::default()
210 })
211 .collect(),
212 ..Default::default()
213 }),
214 Answer::NotFound(message) => MessageResponse::from(pb::ErrorResponse {
215 error_code: 5,
218 error_message: message,
219 ..Default::default()
220 }),
221 };
222
223 Ok(pb::ServerReflectionResponse {
224 valid_host: request.host.clone(),
225 original_request: ::buffa::MessageField::some(request),
226 message_response: Some(message_response),
227 ..Default::default()
228 })
229 }
230 };
231}
232
233mod v1 {
234 use crate::connect::grpc::reflection::v1 as rpc;
235 use crate::proto::grpc::reflection::v1 as pb;
236
237 impl_server_reflection!();
238}
239
240mod v1alpha {
241 use crate::connect::grpc::reflection::v1alpha as rpc;
242 use crate::proto::grpc::reflection::v1alpha as pb;
243
244 impl_server_reflection!();
245}
246
247#[cfg(test)]
248mod tests {
249 use buffa::Message;
250 use buffa_descriptor::generated::descriptor::{
251 FileDescriptorProto, FileDescriptorSet, ServiceDescriptorProto,
252 };
253 use connectrpc::client::{ClientConfig, HttpClient};
254 use tokio::net::TcpListener;
255
256 use super::*;
257 use crate::ServerReflectionClient;
261 use crate::wire::v1::ServerReflectionRequest;
262 use crate::wire::v1::server_reflection_request::MessageRequest;
263 use crate::wire::v1::server_reflection_response::MessageResponse;
264
265 fn test_set_bytes() -> Vec<u8> {
266 FileDescriptorSet {
267 file: vec![FileDescriptorProto {
268 name: Some("acme/api.proto".into()),
269 package: Some("acme.api".into()),
270 service: vec![ServiceDescriptorProto {
271 name: Some("Search".into()),
272 ..Default::default()
273 }],
274 ..Default::default()
275 }],
276 ..Default::default()
277 }
278 .encode_to_vec()
279 }
280
281 async fn spawn_reflection_server() -> ServerReflectionClient<HttpClient> {
284 let reflector = Reflector::from_descriptor_set_bytes(&test_set_bytes()).unwrap();
285 spawn_router(install(Router::new(), reflector)).await
286 }
287
288 async fn spawn_router(router: Router) -> ServerReflectionClient<HttpClient> {
289 let app = router.into_axum_router();
290 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
291 let addr = listener.local_addr().unwrap();
292 tokio::spawn(async move {
293 axum::serve(listener, app).await.unwrap();
294 });
295 let config = ClientConfig::new(format!("http://{addr}").parse().unwrap());
296 ServerReflectionClient::new(HttpClient::plaintext(), config)
297 }
298
299 fn request(message_request: MessageRequest) -> ServerReflectionRequest {
300 ServerReflectionRequest {
301 host: "test-host".into(),
302 message_request: Some(message_request),
303 ..Default::default()
304 }
305 }
306
307 #[tokio::test]
310 async fn oversized_request_is_refused() {
311 let client = spawn_reflection_server().await;
312 let mut stream = client.server_reflection_info().await.unwrap();
313 stream
314 .send(request(MessageRequest::FileContainingSymbol(
315 "x".repeat(2 * crate::MAX_REQUEST_BYTES),
316 )))
317 .await
318 .unwrap();
319 stream.close_send();
320 let err = stream.message().await.unwrap_err();
321 assert_eq!(err.code, connectrpc::ErrorCode::ResourceExhausted);
322 }
323
324 #[tokio::test]
328 async fn integrator_limits_replace_the_bundled_profile() {
329 let reflector = Reflector::from_descriptor_set_bytes(&test_set_bytes()).unwrap();
330 let router = apply_request_limits(
331 install(Router::new(), reflector),
332 connectrpc::Limits::default().with_max_message_size(1024),
333 );
334 let client = spawn_router(router).await;
335 let mut stream = client.server_reflection_info().await.unwrap();
336 stream
337 .send(request(MessageRequest::FileContainingSymbol(
338 "x".repeat(2048),
339 )))
340 .await
341 .unwrap();
342 stream.close_send();
343 let err = stream.message().await.unwrap_err();
344 assert_eq!(err.code, connectrpc::ErrorCode::ResourceExhausted);
345 }
346
347 #[tokio::test]
348 async fn full_stream_round_trip() {
349 let client = spawn_reflection_server().await;
350 let mut stream = client.server_reflection_info().await.unwrap();
351
352 stream
353 .send(request(MessageRequest::ListServices(String::new())))
354 .await
355 .unwrap();
356 stream
357 .send(request(MessageRequest::FileContainingSymbol(
358 "acme.api.Search".into(),
359 )))
360 .await
361 .unwrap();
362 stream
363 .send(request(MessageRequest::FileByFilename("nope.proto".into())))
364 .await
365 .unwrap();
366 stream.close_send();
367
368 let resp = stream.message().await.unwrap().unwrap().to_owned_message();
370 assert_eq!(resp.valid_host, "test-host");
371 assert!(matches!(
372 resp.original_request
373 .as_option()
374 .and_then(|r| r.message_request.as_ref()),
375 Some(MessageRequest::ListServices(_))
376 ));
377 match resp.message_response.unwrap() {
378 MessageResponse::ListServicesResponse(list) => {
379 let names: Vec<_> = list.service.iter().map(|s| s.name.as_str()).collect();
380 assert_eq!(
381 names,
382 [
383 "acme.api.Search",
384 "grpc.reflection.v1.ServerReflection",
385 "grpc.reflection.v1alpha.ServerReflection",
386 ]
387 );
388 }
389 other => panic!("expected list_services_response, got {other:?}"),
390 }
391
392 let resp = stream.message().await.unwrap().unwrap().to_owned_message();
394 match resp.message_response.unwrap() {
395 MessageResponse::FileDescriptorResponse(fd) => {
396 assert_eq!(fd.file_descriptor_proto.len(), 1);
397 let file =
398 FileDescriptorProto::decode_from_slice(&fd.file_descriptor_proto[0]).unwrap();
399 assert_eq!(file.name.as_deref(), Some("acme/api.proto"));
400 }
401 other => panic!("expected file_descriptor_response, got {other:?}"),
402 }
403
404 let resp = stream.message().await.unwrap().unwrap().to_owned_message();
406 match resp.message_response.unwrap() {
407 MessageResponse::ErrorResponse(err) => {
408 assert_eq!(err.error_code, 5);
409 assert!(err.error_message.contains("nope.proto"));
410 }
411 other => panic!("expected error_response, got {other:?}"),
412 }
413
414 assert!(stream.message().await.unwrap().is_none());
415 }
416
417 #[test]
418 fn crate_descriptor_set_makes_reflection_self_describing() {
419 let reflector = Reflector::from_descriptor_set_bytes(crate::FILE_DESCRIPTOR_SET).unwrap();
420 assert_eq!(
421 reflector.service_names(),
422 [
423 crate::SERVER_REFLECTION_SERVICE_NAME,
424 crate::SERVER_REFLECTION_V1ALPHA_SERVICE_NAME,
425 ]
426 );
427 assert!(matches!(
428 reflector
429 .file_containing_symbol("grpc.reflection.v1.ServerReflection.ServerReflectionInfo"),
430 crate::reflector::Answer::Files(_)
431 ));
432 }
433
434 #[tokio::test]
435 async fn v1alpha_route_is_served() {
436 use crate::connect::grpc::reflection::v1alpha::ServerReflectionClient as AlphaClient;
441 use crate::proto::grpc::reflection::v1alpha::ServerReflectionRequest;
442 use crate::proto::grpc::reflection::v1alpha::server_reflection_request::MessageRequest as AlphaRequest;
443 use crate::proto::grpc::reflection::v1alpha::server_reflection_response::MessageResponse as AlphaResponse;
444
445 let reflector = Reflector::from_descriptor_set_bytes(&test_set_bytes()).unwrap();
446 let router = install(Router::new(), reflector);
447 let app = router.into_axum_router();
448 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
449 let addr = listener.local_addr().unwrap();
450 tokio::spawn(async move {
451 axum::serve(listener, app).await.unwrap();
452 });
453 let config = ClientConfig::new(format!("http://{addr}").parse().unwrap());
454 let client = AlphaClient::new(HttpClient::plaintext(), config);
455
456 let mut stream = client.server_reflection_info().await.unwrap();
457 stream
458 .send(ServerReflectionRequest {
459 message_request: Some(AlphaRequest::ListServices(String::new())),
460 ..Default::default()
461 })
462 .await
463 .unwrap();
464 stream.close_send();
465
466 let resp = stream.message().await.unwrap().unwrap().to_owned_message();
467 match resp.message_response.unwrap() {
468 AlphaResponse::ListServicesResponse(list) => {
469 assert_eq!(list.service.len(), 3);
470 }
471 other => panic!("expected list_services_response, got {other:?}"),
472 }
473 }
474}