1 use super::*; 2 use tonic::codec::CompressionEncoding; 3 4 util::parametrized_tests! { 5 client_enabled_server_enabled, 6 zstd: CompressionEncoding::Zstd, 7 gzip: CompressionEncoding::Gzip, 8 } 9 10 #[allow(dead_code)] 11 async fn client_enabled_server_enabled(encoding: CompressionEncoding) { 12 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 13 14 #[derive(Clone, Copy)] 15 struct AssertCorrectAcceptEncoding<S> { 16 service: S, 17 encoding: CompressionEncoding, 18 } 19 20 impl<S, B> Service<http::Request<B>> for AssertCorrectAcceptEncoding<S> 21 where 22 S: Service<http::Request<B>>, 23 { 24 type Response = S::Response; 25 type Error = S::Error; 26 type Future = S::Future; 27 28 fn poll_ready( 29 &mut self, 30 cx: &mut std::task::Context<'_>, 31 ) -> std::task::Poll<Result<(), Self::Error>> { 32 self.service.poll_ready(cx) 33 } 34 35 fn call(&mut self, req: http::Request<B>) -> Self::Future { 36 let expected = match self.encoding { 37 CompressionEncoding::Gzip => "gzip", 38 CompressionEncoding::Zstd => "zstd", 39 _ => panic!("unexpected encoding {:?}", self.encoding), 40 }; 41 assert_eq!( 42 req.headers() 43 .get("grpc-accept-encoding") 44 .unwrap() 45 .to_str() 46 .unwrap(), 47 format!("{},identity", expected) 48 ); 49 self.service.call(req) 50 } 51 } 52 53 let svc = test_server::TestServer::new(Svc::default()).send_compressed(encoding); 54 55 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 56 57 tokio::spawn({ 58 let response_bytes_counter = response_bytes_counter.clone(); 59 async move { 60 Server::builder() 61 .layer( 62 ServiceBuilder::new() 63 .layer(layer_fn(|service| AssertCorrectAcceptEncoding { 64 service, 65 encoding, 66 })) 67 .layer(MapResponseBodyLayer::new(move |body| { 68 util::CountBytesBody { 69 inner: body, 70 counter: response_bytes_counter.clone(), 71 } 72 })) 73 .into_inner(), 74 ) 75 .add_service(svc) 76 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 77 .await 78 .unwrap(); 79 } 80 }); 81 82 let mut client = 83 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 84 85 let expected = match encoding { 86 CompressionEncoding::Gzip => "gzip", 87 CompressionEncoding::Zstd => "zstd", 88 _ => panic!("unexpected encoding {:?}", encoding), 89 }; 90 91 for _ in 0..3 { 92 let res = client.compress_output_unary(()).await.unwrap(); 93 assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected); 94 let bytes_sent = response_bytes_counter.load(SeqCst); 95 assert!(bytes_sent < UNCOMPRESSED_MIN_BODY_SIZE); 96 } 97 } 98 99 util::parametrized_tests! { 100 client_enabled_server_disabled, 101 zstd: CompressionEncoding::Zstd, 102 gzip: CompressionEncoding::Gzip, 103 } 104 105 #[allow(dead_code)] 106 async fn client_enabled_server_disabled(encoding: CompressionEncoding) { 107 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 108 109 let svc = test_server::TestServer::new(Svc::default()); 110 111 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 112 113 tokio::spawn({ 114 let response_bytes_counter = response_bytes_counter.clone(); 115 async move { 116 Server::builder() 117 // no compression enable on the server so responses should not be compressed 118 .layer( 119 ServiceBuilder::new() 120 .layer(MapResponseBodyLayer::new(move |body| { 121 util::CountBytesBody { 122 inner: body, 123 counter: response_bytes_counter.clone(), 124 } 125 })) 126 .into_inner(), 127 ) 128 .add_service(svc) 129 .serve_with_incoming(tokio_stream::iter(vec![Ok::<_, std::io::Error>(server)])) 130 .await 131 .unwrap(); 132 } 133 }); 134 135 let mut client = 136 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 137 138 let res = client.compress_output_unary(()).await.unwrap(); 139 140 assert!(res.metadata().get("grpc-encoding").is_none()); 141 142 let bytes_sent = response_bytes_counter.load(SeqCst); 143 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 144 } 145 146 #[tokio::test(flavor = "multi_thread")] 147 async fn client_enabled_server_disabled_multi_encoding() { 148 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 149 150 let svc = test_server::TestServer::new(Svc::default()); 151 152 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 153 154 tokio::spawn({ 155 let response_bytes_counter = response_bytes_counter.clone(); 156 async move { 157 Server::builder() 158 // no compression enable on the server so responses should not be compressed 159 .layer( 160 ServiceBuilder::new() 161 .layer(MapResponseBodyLayer::new(move |body| { 162 util::CountBytesBody { 163 inner: body, 164 counter: response_bytes_counter.clone(), 165 } 166 })) 167 .into_inner(), 168 ) 169 .add_service(svc) 170 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 171 .await 172 .unwrap(); 173 } 174 }); 175 176 let mut client = test_client::TestClient::new(mock_io_channel(client).await) 177 .accept_compressed(CompressionEncoding::Gzip) 178 .accept_compressed(CompressionEncoding::Zstd); 179 180 let res = client.compress_output_unary(()).await.unwrap(); 181 182 assert!(res.metadata().get("grpc-encoding").is_none()); 183 184 let bytes_sent = response_bytes_counter.load(SeqCst); 185 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 186 } 187 188 util::parametrized_tests! { 189 client_disabled, 190 zstd: CompressionEncoding::Zstd, 191 gzip: CompressionEncoding::Gzip, 192 } 193 194 #[allow(dead_code)] 195 async fn client_disabled(encoding: CompressionEncoding) { 196 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 197 198 #[derive(Clone, Copy)] 199 struct AssertCorrectAcceptEncoding<S>(S); 200 201 impl<S, B> Service<http::Request<B>> for AssertCorrectAcceptEncoding<S> 202 where 203 S: Service<http::Request<B>>, 204 { 205 type Response = S::Response; 206 type Error = S::Error; 207 type Future = S::Future; 208 209 fn poll_ready( 210 &mut self, 211 cx: &mut std::task::Context<'_>, 212 ) -> std::task::Poll<Result<(), Self::Error>> { 213 self.0.poll_ready(cx) 214 } 215 216 fn call(&mut self, req: http::Request<B>) -> Self::Future { 217 assert!(req.headers().get("grpc-accept-encoding").is_none()); 218 self.0.call(req) 219 } 220 } 221 222 let svc = test_server::TestServer::new(Svc::default()).send_compressed(encoding); 223 224 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 225 226 tokio::spawn({ 227 let response_bytes_counter = response_bytes_counter.clone(); 228 async move { 229 Server::builder() 230 .layer( 231 ServiceBuilder::new() 232 .layer(layer_fn(AssertCorrectAcceptEncoding)) 233 .layer(MapResponseBodyLayer::new(move |body| { 234 util::CountBytesBody { 235 inner: body, 236 counter: response_bytes_counter.clone(), 237 } 238 })) 239 .into_inner(), 240 ) 241 .add_service(svc) 242 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 243 .await 244 .unwrap(); 245 } 246 }); 247 248 let mut client = test_client::TestClient::new(mock_io_channel(client).await); 249 250 let res = client.compress_output_unary(()).await.unwrap(); 251 252 assert!(res.metadata().get("grpc-encoding").is_none()); 253 254 let bytes_sent = response_bytes_counter.load(SeqCst); 255 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 256 } 257 258 util::parametrized_tests! { 259 server_replying_with_unsupported_encoding, 260 zstd: CompressionEncoding::Zstd, 261 gzip: CompressionEncoding::Gzip, 262 } 263 264 #[allow(dead_code)] 265 async fn server_replying_with_unsupported_encoding(encoding: CompressionEncoding) { 266 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 267 268 let svc = test_server::TestServer::new(Svc::default()).send_compressed(encoding); 269 270 fn add_weird_content_encoding<B>(mut response: http::Response<B>) -> http::Response<B> { 271 response 272 .headers_mut() 273 .insert("grpc-encoding", "br".parse().unwrap()); 274 response 275 } 276 277 tokio::spawn(async move { 278 Server::builder() 279 .layer( 280 ServiceBuilder::new() 281 .map_response(add_weird_content_encoding) 282 .into_inner(), 283 ) 284 .add_service(svc) 285 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 286 .await 287 .unwrap(); 288 }); 289 290 let mut client = 291 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 292 let status: Status = client.compress_output_unary(()).await.unwrap_err(); 293 294 assert_eq!(status.code(), tonic::Code::Unimplemented); 295 assert_eq!( 296 status.message(), 297 "Content is compressed with `br` which isn't supported" 298 ); 299 } 300 301 util::parametrized_tests! { 302 disabling_compression_on_single_response, 303 zstd: CompressionEncoding::Zstd, 304 gzip: CompressionEncoding::Gzip, 305 } 306 307 #[allow(dead_code)] 308 async fn disabling_compression_on_single_response(encoding: CompressionEncoding) { 309 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 310 311 let svc = test_server::TestServer::new(Svc { 312 disable_compressing_on_response: true, 313 }) 314 .send_compressed(encoding); 315 316 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 317 318 tokio::spawn({ 319 let response_bytes_counter = response_bytes_counter.clone(); 320 async move { 321 Server::builder() 322 .layer( 323 ServiceBuilder::new() 324 .layer(MapResponseBodyLayer::new(move |body| { 325 util::CountBytesBody { 326 inner: body, 327 counter: response_bytes_counter.clone(), 328 } 329 })) 330 .into_inner(), 331 ) 332 .add_service(svc) 333 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 334 .await 335 .unwrap(); 336 } 337 }); 338 339 let mut client = 340 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 341 342 let res = client.compress_output_unary(()).await.unwrap(); 343 344 let expected = match encoding { 345 CompressionEncoding::Gzip => "gzip", 346 CompressionEncoding::Zstd => "zstd", 347 _ => panic!("unexpected encoding {:?}", encoding), 348 }; 349 assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected); 350 351 let bytes_sent = response_bytes_counter.load(SeqCst); 352 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 353 } 354 355 util::parametrized_tests! { 356 disabling_compression_on_response_but_keeping_compression_on_stream, 357 zstd: CompressionEncoding::Zstd, 358 gzip: CompressionEncoding::Gzip, 359 } 360 361 #[allow(dead_code)] 362 async fn disabling_compression_on_response_but_keeping_compression_on_stream( 363 encoding: CompressionEncoding, 364 ) { 365 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 366 367 let svc = test_server::TestServer::new(Svc { 368 disable_compressing_on_response: true, 369 }) 370 .send_compressed(encoding); 371 372 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 373 374 tokio::spawn({ 375 let response_bytes_counter = response_bytes_counter.clone(); 376 async move { 377 Server::builder() 378 .layer( 379 ServiceBuilder::new() 380 .layer(MapResponseBodyLayer::new(move |body| { 381 util::CountBytesBody { 382 inner: body, 383 counter: response_bytes_counter.clone(), 384 } 385 })) 386 .into_inner(), 387 ) 388 .add_service(svc) 389 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 390 .await 391 .unwrap(); 392 } 393 }); 394 395 let mut client = 396 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 397 398 let res = client.compress_output_server_stream(()).await.unwrap(); 399 400 let expected = match encoding { 401 CompressionEncoding::Gzip => "gzip", 402 CompressionEncoding::Zstd => "zstd", 403 _ => panic!("unexpected encoding {:?}", encoding), 404 }; 405 assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected); 406 407 let mut stream: Streaming<SomeData> = res.into_inner(); 408 409 stream 410 .next() 411 .await 412 .expect("stream empty") 413 .expect("item was error"); 414 assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); 415 416 stream 417 .next() 418 .await 419 .expect("stream empty") 420 .expect("item was error"); 421 assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE); 422 } 423 424 util::parametrized_tests! { 425 disabling_compression_on_response_from_client_stream, 426 zstd: CompressionEncoding::Zstd, 427 gzip: CompressionEncoding::Gzip, 428 } 429 430 #[allow(dead_code)] 431 async fn disabling_compression_on_response_from_client_stream(encoding: CompressionEncoding) { 432 let (client, server) = tokio::io::duplex(UNCOMPRESSED_MIN_BODY_SIZE * 10); 433 434 let svc = test_server::TestServer::new(Svc { 435 disable_compressing_on_response: true, 436 }) 437 .send_compressed(encoding); 438 439 let response_bytes_counter = Arc::new(AtomicUsize::new(0)); 440 441 tokio::spawn({ 442 let response_bytes_counter = response_bytes_counter.clone(); 443 async move { 444 Server::builder() 445 .layer( 446 ServiceBuilder::new() 447 .layer(MapResponseBodyLayer::new(move |body| { 448 util::CountBytesBody { 449 inner: body, 450 counter: response_bytes_counter.clone(), 451 } 452 })) 453 .into_inner(), 454 ) 455 .add_service(svc) 456 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server))) 457 .await 458 .unwrap(); 459 } 460 }); 461 462 let mut client = 463 test_client::TestClient::new(mock_io_channel(client).await).accept_compressed(encoding); 464 465 let req = Request::new(Box::pin(tokio_stream::empty())); 466 467 let res = client.compress_output_client_stream(req).await.unwrap(); 468 469 let expected = match encoding { 470 CompressionEncoding::Gzip => "gzip", 471 CompressionEncoding::Zstd => "zstd", 472 _ => panic!("unexpected encoding {:?}", encoding), 473 }; 474 assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected); 475 let bytes_sent = response_bytes_counter.load(SeqCst); 476 assert!(bytes_sent > UNCOMPRESSED_MIN_BODY_SIZE); 477 } 478