xref: /tonic/tonic-reflection/tests/server.rs (revision 7173808a)
1 use prost::Message;
2 use std::net::SocketAddr;
3 use tokio::sync::oneshot;
4 use tokio_stream::{wrappers::TcpListenerStream, StreamExt};
5 use tonic::{transport::Server, Request};
6 use tonic_reflection::{
7     pb::v1::{
8         server_reflection_client::ServerReflectionClient,
9         server_reflection_request::MessageRequest, server_reflection_response::MessageResponse,
10         ServerReflectionRequest, ServiceResponse, FILE_DESCRIPTOR_SET,
11     },
12     server::Builder,
13 };
14 
15 pub(crate) fn get_encoded_reflection_service_fd() -> Vec<u8> {
16     let mut expected = Vec::new();
17     prost_types::FileDescriptorSet::decode(FILE_DESCRIPTOR_SET)
18         .expect("decode reflection service file descriptor set")
19         .file[0]
20         .encode(&mut expected)
21         .expect("encode reflection service file descriptor");
22     expected
23 }
24 
25 #[tokio::test]
26 async fn test_list_services() {
27     let response = make_test_reflection_request(ServerReflectionRequest {
28         host: "".to_string(),
29         message_request: Some(MessageRequest::ListServices(String::new())),
30     })
31     .await;
32 
33     if let MessageResponse::ListServicesResponse(services) = response {
34         assert_eq!(
35             services.service,
36             vec![ServiceResponse {
37                 name: String::from("grpc.reflection.v1.ServerReflection")
38             }]
39         );
40     } else {
41         panic!("Expected a ListServicesResponse variant");
42     }
43 }
44 
45 #[tokio::test]
46 async fn test_file_by_filename() {
47     let response = make_test_reflection_request(ServerReflectionRequest {
48         host: "".to_string(),
49         message_request: Some(MessageRequest::FileByFilename(String::from(
50             "reflection_v1.proto",
51         ))),
52     })
53     .await;
54 
55     if let MessageResponse::FileDescriptorResponse(descriptor) = response {
56         let file_descriptor_proto = descriptor
57             .file_descriptor_proto
58             .first()
59             .expect("descriptor");
60         assert_eq!(
61             file_descriptor_proto.as_ref(),
62             get_encoded_reflection_service_fd()
63         );
64     } else {
65         panic!("Expected a FileDescriptorResponse variant");
66     }
67 }
68 
69 #[tokio::test]
70 async fn test_file_containing_symbol() {
71     let response = make_test_reflection_request(ServerReflectionRequest {
72         host: "".to_string(),
73         message_request: Some(MessageRequest::FileContainingSymbol(String::from(
74             "grpc.reflection.v1.ServerReflection",
75         ))),
76     })
77     .await;
78 
79     if let MessageResponse::FileDescriptorResponse(descriptor) = response {
80         let file_descriptor_proto = descriptor
81             .file_descriptor_proto
82             .first()
83             .expect("descriptor");
84         assert_eq!(
85             file_descriptor_proto.as_ref(),
86             get_encoded_reflection_service_fd()
87         );
88     } else {
89         panic!("Expected a FileDescriptorResponse variant");
90     }
91 }
92 
93 async fn make_test_reflection_request(request: ServerReflectionRequest) -> MessageResponse {
94     // Run a test server
95     let (shutdown_tx, shutdown_rx) = oneshot::channel();
96 
97     let addr: SocketAddr = "127.0.0.1:0".parse().expect("SocketAddr parse");
98     let listener = tokio::net::TcpListener::bind(addr).await.expect("bind");
99     let local_addr = format!("http://{}", listener.local_addr().expect("local address"));
100     let jh = tokio::spawn(async move {
101         let service = Builder::configure()
102             .register_encoded_file_descriptor_set(FILE_DESCRIPTOR_SET)
103             .build_v1()
104             .unwrap();
105 
106         Server::builder()
107             .add_service(service)
108             .serve_with_incoming_shutdown(TcpListenerStream::new(listener), async {
109                 drop(shutdown_rx.await)
110             })
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(tokio_stream::once(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