xref: /tonic/tonic-reflection/tests/server.rs (revision a562a3ce)
1 use futures::stream;
2 use futures_util::FutureExt;
3 use prost::Message;
4 use std::net::SocketAddr;
5 use tokio::sync::oneshot;
6 use tokio_stream::{wrappers::TcpListenerStream, StreamExt};
7 use tonic::{transport::Server, Request};
8 use tonic_reflection::{
9     proto::{
10         server_reflection_client::ServerReflectionClient,
11         server_reflection_request::MessageRequest, server_reflection_response::MessageResponse,
12         ServerReflectionRequest, ServiceResponse, FILE_DESCRIPTOR_SET,
13     },
14     server::Builder,
15 };
16 
17 pub(crate) fn get_encoded_reflection_service_fd() -> Vec<u8> {
18     let mut expected = Vec::new();
19     prost_types::FileDescriptorSet::decode(FILE_DESCRIPTOR_SET)
20         .expect("decode reflection service file descriptor set")
21         .file[0]
22         .encode(&mut expected)
23         .expect("encode reflection service file descriptor");
24     expected
25 }
26 
27 #[tokio::test]
28 async fn test_list_services() {
29     let response = make_test_reflection_request(ServerReflectionRequest {
30         host: "".to_string(),
31         message_request: Some(MessageRequest::ListServices(String::new())),
32     })
33     .await;
34 
35     if let MessageResponse::ListServicesResponse(services) = response {
36         assert_eq!(
37             services.service,
38             vec![ServiceResponse {
39                 name: String::from("grpc.reflection.v1alpha.ServerReflection")
40             }]
41         );
42     } else {
43         panic!("Expected a ListServicesResponse variant");
44     }
45 }
46 
47 #[tokio::test]
48 async fn test_file_by_filename() {
49     let response = make_test_reflection_request(ServerReflectionRequest {
50         host: "".to_string(),
51         message_request: Some(MessageRequest::FileByFilename(String::from(
52             "reflection.proto",
53         ))),
54     })
55     .await;
56 
57     if let MessageResponse::FileDescriptorResponse(descriptor) = response {
58         let file_descriptor_proto = descriptor
59             .file_descriptor_proto
60             .first()
61             .expect("descriptor");
62         assert_eq!(
63             file_descriptor_proto.as_ref(),
64             get_encoded_reflection_service_fd()
65         );
66     } else {
67         panic!("Expected a FileDescriptorResponse variant");
68     }
69 }
70 
71 #[tokio::test]
72 async fn test_file_containing_symbol() {
73     let response = make_test_reflection_request(ServerReflectionRequest {
74         host: "".to_string(),
75         message_request: Some(MessageRequest::FileContainingSymbol(String::from(
76             "grpc.reflection.v1alpha.ServerReflection",
77         ))),
78     })
79     .await;
80 
81     if let MessageResponse::FileDescriptorResponse(descriptor) = response {
82         let file_descriptor_proto = descriptor
83             .file_descriptor_proto
84             .first()
85             .expect("descriptor");
86         assert_eq!(
87             file_descriptor_proto.as_ref(),
88             get_encoded_reflection_service_fd()
89         );
90     } else {
91         panic!("Expected a FileDescriptorResponse variant");
92     }
93 }
94 
95 async fn make_test_reflection_request(request: ServerReflectionRequest) -> MessageResponse {
96     // Run a test server
97     let (shutdown_tx, shutdown_rx) = oneshot::channel();
98 
99     let addr: SocketAddr = "127.0.0.1:0".parse().expect("SocketAddr parse");
100     let listener = tokio::net::TcpListener::bind(addr).await.expect("bind");
101     let local_addr = format!("http://{}", listener.local_addr().expect("local address"));
102     let jh = tokio::spawn(async move {
103         let service = Builder::configure()
104             .register_encoded_file_descriptor_set(FILE_DESCRIPTOR_SET)
105             .build()
106             .unwrap();
107 
108         Server::builder()
109             .add_service(service)
110             .serve_with_incoming_shutdown(TcpListenerStream::new(listener), shutdown_rx.map(drop))
111             .await
112             .unwrap();
113     });
114 
115     // Give the test server a few ms to become available
116     tokio::time::sleep(std::time::Duration::from_millis(100)).await;
117 
118     // Construct client and send request, extract response
119     let conn = tonic::transport::Endpoint::new(local_addr)
120         .unwrap()
121         .connect()
122         .await
123         .unwrap();
124     let mut client = ServerReflectionClient::new(conn);
125 
126     let request = Request::new(stream::iter(vec![request]));
127     let mut inbound = client
128         .server_reflection_info(request)
129         .await
130         .expect("request")
131         .into_inner();
132 
133     let response = inbound
134         .next()
135         .await
136         .expect("steamed response")
137         .expect("successful response")
138         .message_response
139         .expect("some MessageResponse");
140 
141     // We only expect one response per request
142     assert!(inbound.next().await.is_none());
143 
144     // Shut down test server
145     shutdown_tx.send(()).expect("send shutdown");
146     jh.await.expect("server shutdown");
147 
148     response
149 }
150