1 use super::*; 2 use http_body::Body; 3 use tonic::codec::CompressionEncoding; 4 5 util::parametrized_tests! { 6 client_enabled_server_enabled, 7 zstd: CompressionEncoding::Zstd, 8 gzip: CompressionEncoding::Gzip, 9 } 10 11 #[allow(dead_code)] 12 async fn client_enabled_server_enabled(encoding: CompressionEncoding) { 13 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 14 15 let svc = test_server::TestServer::new(Svc::default()).accept_compressed(encoding); 16 17 let request_bytes_counter = Arc::new(AtomicUsize::new(0)); 18 19 #[derive(Clone)] 20 pub struct AssertRightEncoding { 21 encoding: CompressionEncoding, 22 } 23 24 #[allow(dead_code)] 25 impl AssertRightEncoding { 26 pub fn new(encoding: CompressionEncoding) -> Self { 27 Self { encoding } 28 } 29 30 pub fn call<B: Body>(self, req: http::Request<B>) -> http::Request<B> { 31 let expected = match self.encoding { 32 CompressionEncoding::Gzip => "gzip", 33 CompressionEncoding::Zstd => "zstd", 34 _ => panic!("unexpected encoding {:?}", self.encoding), 35 }; 36 assert_eq!(req.headers().get("grpc-encoding").unwrap(), expected); 37 38 req 39 } 40 } 41 42 tokio::spawn({ 43 let request_bytes_counter = request_bytes_counter.clone(); 44 async move { 45 Server::builder() 46 .layer( 47 ServiceBuilder::new() 48 .map_request(move |req| { 49 AssertRightEncoding::new(encoding).clone().call(req) 50 }) 51 .layer(measure_request_body_size_layer( 52 request_bytes_counter.clone(), 53 )) 54 .into_inner(), 55 ) 56 .add_service(svc) 57 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 58 .await 59 .unwrap(); 60 } 61 }); 62 63 let mut client = 64 test_client::TestClient::new(mock_io_channel(client).await).send_compressed(encoding); 65 66 let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); 67 let stream = tokio_stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); 68 let req = Request::new(Box::pin(stream)); 69 70 client.compress_input_client_stream(req).await.unwrap(); 71 72 let bytes_sent = request_bytes_counter.load(SeqCst); 73 assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); 74 } 75 76 util::parametrized_tests! { 77 client_disabled_server_enabled, 78 zstd: CompressionEncoding::Zstd, 79 gzip: CompressionEncoding::Gzip, 80 } 81 82 #[allow(dead_code)] 83 async fn client_disabled_server_enabled(encoding: CompressionEncoding) { 84 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 85 86 let svc = test_server::TestServer::new(Svc::default()).accept_compressed(encoding); 87 88 let request_bytes_counter = Arc::new(AtomicUsize::new(0)); 89 90 fn assert_right_encoding<B>(req: http::Request<B>) -> http::Request<B> { 91 assert!(req.headers().get("grpc-encoding").is_none()); 92 req 93 } 94 95 tokio::spawn({ 96 let request_bytes_counter = request_bytes_counter.clone(); 97 async move { 98 Server::builder() 99 .layer( 100 ServiceBuilder::new() 101 .map_request(assert_right_encoding) 102 .layer(measure_request_body_size_layer( 103 request_bytes_counter.clone(), 104 )) 105 .into_inner(), 106 ) 107 .add_service(svc) 108 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 109 .await 110 .unwrap(); 111 } 112 }); 113 114 let mut client = test_client::TestClient::new(mock_io_channel(client).await); 115 116 let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); 117 let stream = tokio_stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); 118 let req = Request::new(Box::pin(stream)); 119 120 client.compress_input_client_stream(req).await.unwrap(); 121 122 let bytes_sent = request_bytes_counter.load(SeqCst); 123 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 124 } 125 126 util::parametrized_tests! { 127 client_enabled_server_disabled, 128 zstd: CompressionEncoding::Zstd, 129 gzip: CompressionEncoding::Gzip, 130 } 131 132 #[allow(dead_code)] 133 async fn client_enabled_server_disabled(encoding: CompressionEncoding) { 134 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 135 136 let svc = test_server::TestServer::new(Svc::default()); 137 138 tokio::spawn(async move { 139 Server::builder() 140 .add_service(svc) 141 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 142 .await 143 .unwrap(); 144 }); 145 146 let mut client = 147 test_client::TestClient::new(mock_io_channel(client).await).send_compressed(encoding); 148 149 let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec(); 150 let stream = tokio_stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]); 151 let req = Request::new(Box::pin(stream)); 152 153 let status = client.compress_input_client_stream(req).await.unwrap_err(); 154 155 assert_eq!(status.code(), tonic::Code::Unimplemented); 156 let expected = match encoding { 157 CompressionEncoding::Gzip => "gzip", 158 CompressionEncoding::Zstd => "zstd", 159 _ => panic!("unexpected encoding {:?}", encoding), 160 }; 161 assert_eq!( 162 status.message(), 163 format!( 164 "Content is compressed with `{}` which isn't supported", 165 expected 166 ) 167 ); 168 } 169 170 util::parametrized_tests! { 171 compressing_response_from_client_stream, 172 zstd: CompressionEncoding::Zstd, 173 gzip: CompressionEncoding::Gzip, 174 } 175 176 #[allow(dead_code)] 177 async fn compressing_response_from_client_stream(encoding: CompressionEncoding) { 178 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 179 180 let svc = test_server::TestServer::new(Svc::default()).send_compressed(encoding); 181 182 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 183 184 tokio::spawn({ 185 let response_bytes_counter = response_bytes_counter.clone(); 186 async move { 187 Server::builder() 188 .layer( 189 ServiceBuilder::new() 190 .layer(MapResponseBodyLayer::new(move |body| { 191 util::CountBytesBody { 192 inner: body, 193 counter: response_bytes_counter.clone(), 194 } 195 })) 196 .into_inner(), 197 ) 198 .add_service(svc) 199 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 200 .await 201 .unwrap(); 202 } 203 }); 204 205 let mut client = 206 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 207 208 let req = Request::new(Box::pin(tokio_stream::empty())); 209 210 let res = client.compress_output_client_stream(req).await.unwrap(); 211 let expected = match encoding { 212 CompressionEncoding::Gzip => "gzip", 213 CompressionEncoding::Zstd => "zstd", 214 _ => panic!("unexpected encoding {:?}", encoding), 215 }; 216 assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected); 217 let bytes_sent = response_bytes_counter.load(SeqCst); 218 assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); 219 } 220