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())
16         .accept_compressed(encoding)
17         .send_compressed(encoding);
18 
19     let request_bytes_counter = Arc::new(AtomicUsize::new(0));
20     let response_bytes_counter = Arc::new(AtomicUsize::new(0));
21 
22     #[derive(Clone)]
23     pub struct AssertRightEncoding {
24         encoding: CompressionEncoding,
25     }
26 
27     #[allow(dead_code)]
28     impl AssertRightEncoding {
29         pub fn new(encoding: CompressionEncoding) -> Self {
30             Self { encoding }
31         }
32 
33         pub fn call<B: Body>(self, req: http::Request<B>) -> http::Request<B> {
34             let expected = match self.encoding {
35                 CompressionEncoding::Gzip => "gzip",
36                 CompressionEncoding::Zstd => "zstd",
37                 _ => panic!("unexpected encoding {:?}", self.encoding),
38             };
39             assert_eq!(req.headers().get("grpc-encoding").unwrap(), expected);
40 
41             req
42         }
43     }
44 
45     tokio::spawn({
46         let request_bytes_counter = request_bytes_counter.clone();
47         let response_bytes_counter = response_bytes_counter.clone();
48         async move {
49             Server::builder()
50                 .layer(
51                     ServiceBuilder::new()
52                         .map_request(move |req| {
53                             AssertRightEncoding::new(encoding).clone().call(req)
54                         })
55                         .layer(measure_request_body_size_layer(
56                             request_bytes_counter.clone(),
57                         ))
58                         .layer(MapResponseBodyLayer::new(move |body| {
59                             util::CountBytesBody {
60                                 inner: body,
61                                 counter: response_bytes_counter.clone(),
62                             }
63                         }))
64                         .into_inner(),
65                 )
66                 .add_service(svc)
67                 .serve_with_incoming(tokio_stream::once(Ok::<_, std::io::Error>(server)))
68                 .await
69                 .unwrap();
70         }
71     });
72 
73     let mut client = test_client::TestClient::new(mock_io_channel(client).await)
74         .send_compressed(encoding)
75         .accept_compressed(encoding);
76 
77     let data = [0_u8; UNCOMPRESSED_MIN_BODY_SIZE].to_vec();
78     let stream = tokio_stream::iter(vec![SomeData { data: data.clone() }, SomeData { data }]);
79     let req = Request::new(stream);
80 
81     let res = client
82         .compress_input_output_bidirectional_stream(req)
83         .await
84         .unwrap();
85 
86     let expected = match encoding {
87         CompressionEncoding::Gzip => "gzip",
88         CompressionEncoding::Zstd => "zstd",
89         _ => panic!("unexpected encoding {:?}", encoding),
90     };
91     assert_eq!(res.metadata().get("grpc-encoding").unwrap(), expected);
92 
93     let mut stream: Streaming<SomeData> = res.into_inner();
94 
95     stream
96         .next()
97         .await
98         .expect("stream empty")
99         .expect("item was error");
100 
101     stream
102         .next()
103         .await
104         .expect("stream empty")
105         .expect("item was error");
106 
107     assert!(request_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE);
108     assert!(response_bytes_counter.load(SeqCst) < UNCOMPRESSED_MIN_BODY_SIZE);
109 }
110