1use std::{io, pin::Pin};
4
5use async_compression::tokio::bufread::{BrotliDecoder, GzipDecoder, ZlibDecoder, ZstdDecoder};
6use bytes::Bytes;
7use futures::{Stream, TryStreamExt};
8use http::header::{ACCEPT_ENCODING, HeaderMap};
9use tokio_util::io::{ReaderStream, StreamReader};
10
11use crate::{Coding, ContentEncoding};
12
13pub type ByteStream = dyn Stream<Item = Result<Bytes, String>> + Send + Sync;
15
16pub const DEFAULT_ACCEPT_ENCODING: &str = "zstd,gzip,deflate,br";
18
19#[derive(Clone, Debug)]
27pub struct AcceptEncoding {
28 codings: Vec<(Coding, u16)>,
30 star: Option<u16>,
32}
33
34impl AcceptEncoding {
35 fn parse(value: &str) -> Self {
36 let mut accept = Self {
37 codings: Vec::new(),
38 star: None,
39 };
40 accept.merge(value);
41 accept
42 }
43
44 fn merge(&mut self, value: &str) {
46 let accept = self;
47 for element in value.split(',') {
48 let mut parts = element.split(';');
49 let Some(token) = parts.next().map(str::trim) else {
50 continue;
51 };
52 if token.is_empty() {
53 continue;
54 }
55
56 let mut quality = 1000;
57 for param in parts {
58 let param = param.trim();
59 if let Some(rest) = param
60 .strip_prefix("q=")
61 .or_else(|| param.strip_prefix("Q="))
62 {
63 quality = parse_quality(rest).unwrap_or(0);
64 }
65 }
66
67 if token == "*" {
68 accept.star = Some(quality);
69 continue;
70 }
71
72 let coding = Coding::from_token(token);
73 match accept
76 .codings
77 .iter_mut()
78 .find(|(named, _)| *named == coding)
79 {
80 Some((_, existing)) => *existing = quality,
81 None => accept.codings.push((coding, quality)),
82 }
83 }
84 }
85
86 pub fn iter(&self) -> impl Iterator<Item = (&Coding, u16)> {
91 self.codings
92 .iter()
93 .map(|(coding, quality)| (coding, *quality))
94 }
95
96 pub fn quality(&self, coding: &Coding) -> Option<u16> {
100 self.codings
101 .iter()
102 .find(|(named, _)| named == coding)
103 .map(|(_, quality)| *quality)
104 }
105
106 pub fn star(&self) -> Option<u16> {
108 self.star
109 }
110
111 pub fn accepts(&self, coding: &Coding) -> bool {
117 match self.quality(coding) {
118 Some(quality) => quality > 0,
119 None => matches!(self.star, Some(quality) if quality > 0),
120 }
121 }
122}
123
124impl From<&HeaderMap> for AcceptEncoding {
125 fn from(headers: &HeaderMap) -> Self {
130 let mut this = Self {
131 codings: Vec::new(),
132 star: None,
133 };
134 for value in headers.get_all(ACCEPT_ENCODING) {
135 let Ok(value) = value.to_str() else { continue };
136 this.merge(value);
137 }
138 this
139 }
140}
141
142impl From<&str> for AcceptEncoding {
143 fn from(value: &str) -> Self {
144 Self::parse(value)
145 }
146}
147
148impl Default for AcceptEncoding {
149 fn default() -> Self {
150 Self::parse(DEFAULT_ACCEPT_ENCODING)
151 }
152}
153
154fn parse_quality(value: &str) -> Option<u16> {
156 let value = value.trim();
157 let mut chars = value.chars();
158 let mut quality: u16 = match chars.next()? {
159 '0' => 0,
160 '1' => 1000,
161 _ => return None,
162 };
163 if let Some(dot) = chars.next() {
164 if dot != '.' {
165 return None;
166 }
167 let mut scale = 100;
168 for digit in chars {
169 quality += digit.to_digit(10)? as u16 * scale;
170 if scale == 1 {
171 break;
172 }
173 scale /= 10;
174 }
175 }
176 Some(quality.min(1000))
177}
178
179pub fn decode(
188 headers: &mut HeaderMap,
189 body: Pin<Box<ByteStream>>,
190 accept: &AcceptEncoding,
191) -> Pin<Box<ByteStream>> {
192 match ContentEncoding::peel_one_header(headers, accept) {
193 Some(coding) => decode_stream(body, coding),
194 None => body,
195 }
196}
197
198pub fn decode_stream(input: Pin<Box<ByteStream>>, coding: Coding) -> Pin<Box<ByteStream>> {
204 let reader = StreamReader::new(input.map_err(io::Error::other));
205 match coding {
206 Coding::Gzip => reader_stream(GzipDecoder::new(reader)),
207 Coding::Deflate => reader_stream(ZlibDecoder::new(reader)),
208 Coding::Brotli => reader_stream(BrotliDecoder::new(reader)),
209 Coding::Zstd => {
210 let mut decoder = ZstdDecoder::new(reader);
212 decoder.multiple_members(true);
213 reader_stream(decoder)
214 }
215 _ => reader_stream(reader),
216 }
217}
218
219fn reader_stream<R>(reader: R) -> Pin<Box<ByteStream>>
220where
221 R: tokio::io::AsyncRead + Send + Sync + 'static,
222{
223 Box::pin(ReaderStream::new(reader).map_err(|err| err.to_string()))
224}
225
226#[cfg(test)]
227mod tests {
228 use http::header::{CONTENT_ENCODING, CONTENT_LENGTH, HeaderMap, HeaderValue};
229
230 use super::*;
231 use crate::request::{compress_buffer, compress_stream};
232
233 fn decide(content_encoding: &str, accept: &str) -> Option<Coding> {
234 let mut headers = HeaderMap::new();
235 headers.insert(
236 CONTENT_ENCODING,
237 HeaderValue::from_str(content_encoding).unwrap(),
238 );
239 ContentEncoding::from(&headers).can_decode_as(&AcceptEncoding::from(accept))
240 }
241
242 #[test]
243 fn decodes_a_negotiated_coding() {
244 assert_eq!(decide("gzip", DEFAULT_ACCEPT_ENCODING), Some(Coding::Gzip));
245 assert_eq!(decide("br", DEFAULT_ACCEPT_ENCODING), Some(Coding::Brotli));
246 assert_eq!(decide("zstd", DEFAULT_ACCEPT_ENCODING), Some(Coding::Zstd));
247 assert_eq!(
248 decide("deflate", DEFAULT_ACCEPT_ENCODING),
249 Some(Coding::Deflate)
250 );
251 }
252
253 #[test]
254 fn a_coding_named_alone_decodes_only_itself() {
255 assert_eq!(decide("gzip", "gzip"), Some(Coding::Gzip));
256 assert_eq!(decide("br", "gzip"), None);
257 }
258
259 #[test]
260 fn identity_leaves_a_compressed_body_alone() {
261 assert_eq!(decide("gzip", "identity"), None);
262 }
263
264 #[test]
265 fn a_zero_quality_value_refuses() {
266 assert_eq!(decide("gzip", "gzip;q=0"), None);
267 assert_eq!(decide("gzip", "gzip;q=0.000"), None);
268 }
269
270 #[test]
271 fn a_named_coding_settles_the_question_over_star() {
272 assert_eq!(decide("gzip", "gzip;q=0, *"), None);
274 assert_eq!(decide("br", "gzip;q=0, *"), Some(Coding::Brotli));
275 assert_eq!(decide("zstd", "gzip;q=0, *"), Some(Coding::Zstd));
276 }
277
278 #[test]
279 fn star_covers_what_is_not_named() {
280 assert_eq!(decide("gzip", "*"), Some(Coding::Gzip));
281 assert_eq!(decide("gzip", "br, *"), Some(Coding::Gzip));
282 }
283
284 #[test]
285 fn a_star_with_zero_quality_accepts_nothing_unnamed() {
286 assert_eq!(decide("gzip", "*;q=0"), None);
287 assert_eq!(decide("gzip", "gzip, *;q=0"), Some(Coding::Gzip));
288 }
289
290 #[test]
291 fn the_outermost_of_several_codings_is_the_one_decoded() {
292 assert_eq!(
294 decide("gzip, br", DEFAULT_ACCEPT_ENCODING),
295 Some(Coding::Brotli)
296 );
297 assert_eq!(
298 decide("br, gzip", DEFAULT_ACCEPT_ENCODING),
299 Some(Coding::Gzip)
300 );
301 assert_eq!(
302 decide("identity, gzip", DEFAULT_ACCEPT_ENCODING),
303 Some(Coding::Gzip)
304 );
305
306 assert_eq!(decide("gzip, identity", DEFAULT_ACCEPT_ENCODING), None);
309 }
310
311 #[test]
312 fn peeling_a_layer_leaves_the_headers_describing_the_rest() {
313 let mut headers = HeaderMap::new();
314 headers.insert(CONTENT_ENCODING, HeaderValue::from_static("gzip, br"));
315 headers.insert(CONTENT_LENGTH, HeaderValue::from_static("42"));
316
317 let accept = AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING);
318 assert_eq!(
319 ContentEncoding::peel_one_header(&mut headers, &accept),
320 Some(Coding::Brotli)
321 );
322
323 assert_eq!(headers[CONTENT_ENCODING], "gzip");
325 assert!(!headers.contains_key(CONTENT_LENGTH));
327
328 assert_eq!(
330 ContentEncoding::peel_one_header(&mut headers, &accept),
331 Some(Coding::Gzip)
332 );
333 assert!(!headers.contains_key(CONTENT_ENCODING));
334 }
335
336 #[test]
337 fn a_coding_this_crate_cannot_decode_survives_under_one_it_can() {
338 let accept = AcceptEncoding::from("gzip, custom-thing");
340 let mut headers = HeaderMap::new();
341 headers.insert(
342 CONTENT_ENCODING,
343 HeaderValue::from_static("custom-thing, gzip"),
344 );
345 headers.insert(CONTENT_LENGTH, HeaderValue::from_static("42"));
346
347 assert_eq!(
349 ContentEncoding::peel_one_header(&mut headers, &accept),
350 Some(Coding::Gzip)
351 );
352 assert_eq!(headers[CONTENT_ENCODING], "custom-thing");
354 assert!(!headers.contains_key(CONTENT_LENGTH));
355
356 assert_eq!(
358 ContentEncoding::peel_one_header(&mut headers, &accept),
359 None
360 );
361 assert_eq!(headers[CONTENT_ENCODING], "custom-thing");
362 }
363
364 #[test]
365 fn a_coding_this_crate_cannot_decode_hides_one_it_can() {
366 let accept = AcceptEncoding::from("gzip, custom-thing");
369 let mut headers = HeaderMap::new();
370 headers.insert(
371 CONTENT_ENCODING,
372 HeaderValue::from_static("gzip, custom-thing"),
373 );
374
375 assert_eq!(
376 ContentEncoding::peel_one_header(&mut headers, &accept),
377 None
378 );
379 assert_eq!(headers[CONTENT_ENCODING], "gzip, custom-thing");
380 }
381
382 #[test]
383 fn nothing_to_peel_leaves_the_headers_alone() {
384 let mut headers = HeaderMap::new();
385 headers.insert(CONTENT_ENCODING, HeaderValue::from_static("identity"));
386 headers.insert(CONTENT_LENGTH, HeaderValue::from_static("42"));
387
388 let accept = AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING);
389 assert_eq!(
390 ContentEncoding::peel_one_header(&mut headers, &accept),
391 None
392 );
393
394 assert_eq!(headers[CONTENT_ENCODING], "identity");
395 assert_eq!(headers[CONTENT_LENGTH], "42");
396 }
397
398 #[test]
399 fn codings_split_across_header_lines_count_together() {
400 let mut headers = HeaderMap::new();
402 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
403 headers.append(CONTENT_ENCODING, HeaderValue::from_static("br"));
404 assert_eq!(
405 ContentEncoding::from(&headers)
406 .can_decode_as(&AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING)),
407 Some(Coding::Brotli)
408 );
409 }
410
411 #[test]
412 fn one_coding_split_across_lines_with_an_empty_line_still_decodes() {
413 let mut headers = HeaderMap::new();
415 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
416 headers.append(CONTENT_ENCODING, HeaderValue::from_static(""));
417 assert_eq!(
418 ContentEncoding::from(&headers)
419 .can_decode_as(&AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING)),
420 Some(Coding::Gzip)
421 );
422 }
423
424 #[test]
425 fn a_non_ascii_line_is_delivered_as_received() {
426 let mut headers = HeaderMap::new();
427 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
428 headers.append(CONTENT_ENCODING, HeaderValue::from_bytes(b"\xff").unwrap());
429 assert_eq!(
430 ContentEncoding::from(&headers)
431 .can_decode_as(&AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING)),
432 None
433 );
434 }
435
436 #[test]
437 fn a_coding_faith_cannot_decode_is_delivered_as_received() {
438 assert_eq!(decide("compress", DEFAULT_ACCEPT_ENCODING), None);
439 }
440
441 #[test]
442 fn no_content_encoding_means_nothing_to_decode() {
443 let headers = HeaderMap::new();
444 assert_eq!(
445 ContentEncoding::from(&headers)
446 .can_decode_as(&AcceptEncoding::from(DEFAULT_ACCEPT_ENCODING)),
447 None
448 );
449 }
450
451 #[test]
452 fn the_codings_a_response_declared_are_readable() {
453 let mut headers = HeaderMap::new();
454 headers.append(CONTENT_ENCODING, HeaderValue::from_static("br, gzip"));
455 headers.append(CONTENT_ENCODING, HeaderValue::from_static("identity"));
456
457 assert_eq!(
459 ContentEncoding::from(&headers).codings(),
460 [
461 Coding::Brotli,
462 Coding::Gzip,
463 Coding::Other("identity".into())
464 ]
465 );
466
467 assert_eq!(
469 ContentEncoding::from(&headers).can_decode_as(&AcceptEncoding::default()),
470 None
471 );
472
473 assert!(ContentEncoding::default().codings().is_empty());
474 }
475
476 #[test]
477 fn quality_values_parse_to_thousandths() {
478 assert_eq!(parse_quality("0"), Some(0));
479 assert_eq!(parse_quality("1"), Some(1000));
480 assert_eq!(parse_quality("0.5"), Some(500));
481 assert_eq!(parse_quality("0.001"), Some(1));
482 assert_eq!(parse_quality("1.0"), Some(1000));
483 }
484
485 #[tokio::test]
486 async fn a_compressed_body_decodes_back_to_what_went_in() {
487 let input = b"the quick brown fox jumps over the lazy dog".repeat(20);
490 for coding in [Coding::Gzip, Coding::Deflate, Coding::Brotli, Coding::Zstd] {
491 let compressed = compress_buffer(&input, coding.clone()).await.unwrap();
492 assert!(
493 compressed.len() < input.len(),
494 "{coding:?} did not compress repetitive input"
495 );
496
497 let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
498 let decoded: Vec<u8> = decode_stream(Box::pin(source), coding.clone())
499 .try_fold(Vec::new(), |mut acc, chunk| async move {
500 acc.extend_from_slice(&chunk);
501 Ok(acc)
502 })
503 .await
504 .unwrap();
505 assert_eq!(decoded, input, "{coding:?} round trip");
506 }
507 }
508
509 #[tokio::test]
510 async fn a_streaming_body_compresses_across_its_chunks() {
511 let chunks = ["first chunk, ", "second chunk, ", "third chunk"];
512 let source = futures::stream::iter(
513 chunks
514 .into_iter()
515 .map(|chunk| Ok(Bytes::from_static(chunk.as_bytes()))),
516 );
517
518 let compressed: Vec<u8> = compress_stream(source, Coding::Zstd)
519 .expect("zstd compresses")
520 .try_fold(Vec::new(), |mut acc, chunk| async move {
521 acc.extend_from_slice(&chunk);
522 Ok(acc)
523 })
524 .await
525 .unwrap();
526
527 let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
528 let decoded: Vec<u8> = decode_stream(Box::pin(source), Coding::Zstd)
529 .try_fold(Vec::new(), |mut acc, chunk| async move {
530 acc.extend_from_slice(&chunk);
531 Ok(acc)
532 })
533 .await
534 .unwrap();
535 assert_eq!(decoded, chunks.concat().as_bytes());
536 }
537}