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