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