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