1use std::{io, pin::Pin};
18
19use async_compression::tokio::bufread::{
20 BrotliDecoder, BrotliEncoder, GzipDecoder, GzipEncoder, ZlibDecoder, ZlibEncoder, ZstdDecoder,
21 ZstdEncoder,
22};
23use bytes::Bytes;
24use futures::{Stream, TryStreamExt};
25use http::header::{CONTENT_ENCODING, CONTENT_LENGTH, HeaderMap};
26use tokio::io::AsyncReadExt;
27use tokio_util::io::{ReaderStream, StreamReader};
28
29pub type ByteStream = dyn Stream<Item = Result<Bytes, String>> + Send + Sync;
34
35pub const DEFAULT_ACCEPT_ENCODING: &str = "zstd,gzip,deflate,br";
40
41#[derive(Clone, Copy, Debug, PartialEq, Eq)]
46pub enum Coding {
47 Gzip,
48 Deflate,
49 Brotli,
50 Zstd,
51}
52
53impl Coding {
54 pub fn from_option(value: &str) -> Option<Self> {
62 match value {
63 "gzip" => Some(Self::Gzip),
64 "deflate" => Some(Self::Deflate),
65 "br" => Some(Self::Brotli),
66 "zstd" => Some(Self::Zstd),
67 _ => None,
68 }
69 }
70
71 pub fn token(self) -> &'static str {
73 match self {
74 Self::Gzip => "gzip",
75 Self::Deflate => "deflate",
76 Self::Brotli => "br",
77 Self::Zstd => "zstd",
78 }
79 }
80
81 pub fn from_token(token: &str) -> Option<Self> {
84 let token = token.trim();
85 if token.eq_ignore_ascii_case("gzip") || token.eq_ignore_ascii_case("x-gzip") {
86 Some(Self::Gzip)
87 } else if token.eq_ignore_ascii_case("deflate") {
88 Some(Self::Deflate)
89 } else if token.eq_ignore_ascii_case("br") {
90 Some(Self::Brotli)
91 } else if token.eq_ignore_ascii_case("zstd") {
92 Some(Self::Zstd)
93 } else {
94 None
95 }
96 }
97}
98
99pub fn decision(headers: &HeaderMap, accept: &AcceptEncoding) -> Option<Coding> {
106 let mut codings = Vec::new();
110 for value in headers.get_all(CONTENT_ENCODING) {
111 let value = value.to_str().ok()?;
114 codings.extend(value.split(',').map(str::trim).filter(|c| !c.is_empty()));
115 }
116
117 let [single] = codings[..] else {
118 return None;
119 };
120 let coding = Coding::from_token(single)?;
121 accept.accepts(coding).then_some(coding)
122}
123
124pub fn strip_decoded_headers(headers: &mut HeaderMap) {
126 headers.remove(CONTENT_ENCODING);
127 headers.remove(CONTENT_LENGTH);
128}
129
130#[derive(Clone, Copy, Debug, Default)]
135pub struct AcceptEncoding {
136 gzip: Option<u16>,
137 deflate: Option<u16>,
138 brotli: Option<u16>,
139 zstd: Option<u16>,
140 star: Option<u16>,
141}
142
143impl AcceptEncoding {
144 pub fn parse(value: &str) -> Self {
145 let mut accept = Self::default();
146 for element in value.split(',') {
147 let mut parts = element.split(';');
148 let Some(token) = parts.next().map(str::trim) else {
149 continue;
150 };
151 if token.is_empty() {
152 continue;
153 }
154
155 let mut quality = 1000;
156 for param in parts {
157 let param = param.trim();
158 if let Some(rest) = param
159 .strip_prefix("q=")
160 .or_else(|| param.strip_prefix("Q="))
161 {
162 quality = parse_quality(rest).unwrap_or(0);
163 }
164 }
165
166 let slot = if token == "*" {
167 &mut accept.star
168 } else {
169 match Coding::from_token(token) {
170 Some(Coding::Gzip) => &mut accept.gzip,
171 Some(Coding::Deflate) => &mut accept.deflate,
172 Some(Coding::Brotli) => &mut accept.brotli,
173 Some(Coding::Zstd) => &mut accept.zstd,
174 None => continue,
175 }
176 };
177 *slot = Some(quality);
178 }
179 accept
180 }
181
182 fn accepts(&self, coding: Coding) -> bool {
186 let named = match coding {
187 Coding::Gzip => self.gzip,
188 Coding::Deflate => self.deflate,
189 Coding::Brotli => self.brotli,
190 Coding::Zstd => self.zstd,
191 };
192 match named {
193 Some(quality) => quality > 0,
194 None => matches!(self.star, Some(quality) if quality > 0),
195 }
196 }
197}
198
199fn parse_quality(value: &str) -> Option<u16> {
201 let value = value.trim();
202 let mut chars = value.chars();
203 let mut quality: u16 = match chars.next()? {
204 '0' => 0,
205 '1' => 1000,
206 _ => return None,
207 };
208 if let Some(dot) = chars.next() {
209 if dot != '.' {
210 return None;
211 }
212 let mut scale = 100;
213 for digit in chars {
214 quality += digit.to_digit(10)? as u16 * scale;
215 if scale == 1 {
216 break;
217 }
218 scale /= 10;
219 }
220 }
221 Some(quality.min(1000))
222}
223
224pub fn decode_stream(input: Pin<Box<ByteStream>>, coding: Coding) -> Pin<Box<ByteStream>> {
229 let reader = StreamReader::new(input.map_err(io::Error::other));
230 match coding {
231 Coding::Gzip => reader_stream(GzipDecoder::new(reader)),
232 Coding::Deflate => reader_stream(ZlibDecoder::new(reader)),
233 Coding::Brotli => reader_stream(BrotliDecoder::new(reader)),
234 Coding::Zstd => {
235 let mut decoder = ZstdDecoder::new(reader);
237 decoder.multiple_members(true);
238 reader_stream(decoder)
239 }
240 }
241}
242
243fn reader_stream<R>(reader: R) -> Pin<Box<ByteStream>>
244where
245 R: tokio::io::AsyncRead + Send + Sync + 'static,
246{
247 Box::pin(ReaderStream::new(reader).map_err(|err| err.to_string()))
248}
249
250pub type RequestStream = Pin<Box<dyn Stream<Item = io::Result<Bytes>> + Send>>;
252
253pub async fn compress_buffer(input: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
259 let mut output = Vec::new();
260 match coding {
261 Coding::Gzip => GzipEncoder::new(input).read_to_end(&mut output).await?,
262 Coding::Deflate => ZlibEncoder::new(input).read_to_end(&mut output).await?,
263 Coding::Brotli => BrotliEncoder::new(input).read_to_end(&mut output).await?,
264 Coding::Zstd => ZstdEncoder::new(input).read_to_end(&mut output).await?,
265 };
266 Ok(output)
267}
268
269pub fn compress_stream<S>(input: S, coding: Coding) -> RequestStream
276where
277 S: Stream<Item = io::Result<Bytes>> + Send + 'static,
278{
279 let reader = StreamReader::new(input);
280 match coding {
281 Coding::Gzip => encoder_stream(GzipEncoder::new(reader)),
282 Coding::Deflate => encoder_stream(ZlibEncoder::new(reader)),
283 Coding::Brotli => encoder_stream(BrotliEncoder::new(reader)),
284 Coding::Zstd => encoder_stream(ZstdEncoder::new(reader)),
285 }
286}
287
288fn encoder_stream<R>(reader: R) -> RequestStream
289where
290 R: tokio::io::AsyncRead + Send + 'static,
291{
292 Box::pin(ReaderStream::new(reader))
293}
294
295pub fn layer_content_encoding(declared: Option<&str>, applied: Coding) -> String {
301 match declared.map(str::trim).filter(|value| !value.is_empty()) {
302 Some(declared) => format!("{declared}, {}", applied.token()),
303 None => applied.token().to_owned(),
304 }
305}
306
307#[cfg(test)]
308mod tests {
309 use http::header::{CONTENT_ENCODING, HeaderMap, HeaderValue};
310
311 use super::*;
312
313 fn decide(content_encoding: &str, accept: &str) -> Option<Coding> {
314 let mut headers = HeaderMap::new();
315 headers.insert(
316 CONTENT_ENCODING,
317 HeaderValue::from_str(content_encoding).unwrap(),
318 );
319 decision(&headers, &AcceptEncoding::parse(accept))
320 }
321
322 #[test]
323 fn decodes_a_negotiated_coding() {
324 assert_eq!(decide("gzip", DEFAULT_ACCEPT_ENCODING), Some(Coding::Gzip));
325 assert_eq!(decide("br", DEFAULT_ACCEPT_ENCODING), Some(Coding::Brotli));
326 assert_eq!(decide("zstd", DEFAULT_ACCEPT_ENCODING), Some(Coding::Zstd));
327 assert_eq!(
328 decide("deflate", DEFAULT_ACCEPT_ENCODING),
329 Some(Coding::Deflate)
330 );
331 }
332
333 #[test]
334 fn a_coding_named_alone_decodes_only_itself() {
335 assert_eq!(decide("gzip", "gzip"), Some(Coding::Gzip));
336 assert_eq!(decide("br", "gzip"), None);
337 }
338
339 #[test]
340 fn identity_leaves_a_compressed_body_alone() {
341 assert_eq!(decide("gzip", "identity"), None);
342 }
343
344 #[test]
345 fn a_zero_quality_value_refuses() {
346 assert_eq!(decide("gzip", "gzip;q=0"), None);
347 assert_eq!(decide("gzip", "gzip;q=0.000"), None);
348 }
349
350 #[test]
351 fn a_named_coding_settles_the_question_over_star() {
352 assert_eq!(decide("gzip", "gzip;q=0, *"), None);
354 assert_eq!(decide("br", "gzip;q=0, *"), Some(Coding::Brotli));
355 assert_eq!(decide("zstd", "gzip;q=0, *"), Some(Coding::Zstd));
356 }
357
358 #[test]
359 fn star_covers_what_is_not_named() {
360 assert_eq!(decide("gzip", "*"), Some(Coding::Gzip));
361 assert_eq!(decide("gzip", "br, *"), Some(Coding::Gzip));
362 }
363
364 #[test]
365 fn a_star_with_zero_quality_accepts_nothing_unnamed() {
366 assert_eq!(decide("gzip", "*;q=0"), None);
367 assert_eq!(decide("gzip", "gzip, *;q=0"), Some(Coding::Gzip));
368 }
369
370 #[test]
371 fn more_than_one_coding_is_delivered_as_received() {
372 assert_eq!(decide("gzip, br", DEFAULT_ACCEPT_ENCODING), None);
373 assert_eq!(decide("br, gzip", DEFAULT_ACCEPT_ENCODING), None);
374 assert_eq!(decide("identity, gzip", DEFAULT_ACCEPT_ENCODING), None);
375 }
376
377 #[test]
378 fn codings_split_across_header_lines_count_together() {
379 let mut headers = HeaderMap::new();
381 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
382 headers.append(CONTENT_ENCODING, HeaderValue::from_static("br"));
383 assert_eq!(
384 decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
385 None
386 );
387 }
388
389 #[test]
390 fn one_coding_split_across_lines_with_an_empty_line_still_decodes() {
391 let mut headers = HeaderMap::new();
393 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
394 headers.append(CONTENT_ENCODING, HeaderValue::from_static(""));
395 assert_eq!(
396 decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
397 Some(Coding::Gzip)
398 );
399 }
400
401 #[test]
402 fn a_non_ascii_line_is_delivered_as_received() {
403 let mut headers = HeaderMap::new();
404 headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
405 headers.append(CONTENT_ENCODING, HeaderValue::from_bytes(b"\xff").unwrap());
406 assert_eq!(
407 decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
408 None
409 );
410 }
411
412 #[test]
413 fn a_coding_faith_cannot_decode_is_delivered_as_received() {
414 assert_eq!(decide("compress", DEFAULT_ACCEPT_ENCODING), None);
415 }
416
417 #[test]
418 fn no_content_encoding_means_nothing_to_decode() {
419 let headers = HeaderMap::new();
420 assert_eq!(
421 decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
422 None
423 );
424 }
425
426 #[test]
427 fn quality_values_parse_to_thousandths() {
428 assert_eq!(parse_quality("0"), Some(0));
429 assert_eq!(parse_quality("1"), Some(1000));
430 assert_eq!(parse_quality("0.5"), Some(500));
431 assert_eq!(parse_quality("0.001"), Some(1));
432 assert_eq!(parse_quality("1.0"), Some(1000));
433 }
434
435 #[test]
436 fn the_compress_option_names_a_coding_by_its_wire_token() {
437 assert_eq!(Coding::from_option("gzip"), Some(Coding::Gzip));
438 assert_eq!(Coding::from_option("deflate"), Some(Coding::Deflate));
439 assert_eq!(Coding::from_option("br"), Some(Coding::Brotli));
440 assert_eq!(Coding::from_option("zstd"), Some(Coding::Zstd));
441 }
442
443 #[test]
444 fn the_compress_option_matches_its_tokens_exactly() {
445 assert_eq!(Coding::from_token("x-gzip"), Some(Coding::Gzip));
448 assert_eq!(Coding::from_option("x-gzip"), None);
449 assert_eq!(Coding::from_token("GZIP"), Some(Coding::Gzip));
450 assert_eq!(Coding::from_option("GZIP"), None);
451 assert_eq!(Coding::from_option(" gzip"), None);
452 assert_eq!(Coding::from_option("brotli"), None);
453 assert_eq!(Coding::from_option("identity"), None);
454 assert_eq!(Coding::from_option(""), None);
455 }
456
457 #[tokio::test]
458 async fn a_compressed_body_decodes_back_to_what_went_in() {
459 let input = b"the quick brown fox jumps over the lazy dog".repeat(20);
462 for coding in [Coding::Gzip, Coding::Deflate, Coding::Brotli, Coding::Zstd] {
463 let compressed = compress_buffer(&input, coding).await.unwrap();
464 assert!(
465 compressed.len() < input.len(),
466 "{coding:?} did not compress repetitive input"
467 );
468
469 let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
470 let decoded: Vec<u8> = decode_stream(Box::pin(source), coding)
471 .try_fold(Vec::new(), |mut acc, chunk| async move {
472 acc.extend_from_slice(&chunk);
473 Ok(acc)
474 })
475 .await
476 .unwrap();
477 assert_eq!(decoded, input, "{coding:?} round trip");
478 }
479 }
480
481 #[tokio::test]
482 async fn a_streaming_body_compresses_across_its_chunks() {
483 let chunks = ["first chunk, ", "second chunk, ", "third chunk"];
484 let source = futures::stream::iter(
485 chunks
486 .into_iter()
487 .map(|chunk| Ok(Bytes::from_static(chunk.as_bytes()))),
488 );
489
490 let compressed: Vec<u8> = compress_stream(source, Coding::Zstd)
491 .try_fold(Vec::new(), |mut acc, chunk| async move {
492 acc.extend_from_slice(&chunk);
493 Ok(acc)
494 })
495 .await
496 .unwrap();
497
498 let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
499 let decoded: Vec<u8> = decode_stream(Box::pin(source), Coding::Zstd)
500 .try_fold(Vec::new(), |mut acc, chunk| async move {
501 acc.extend_from_slice(&chunk);
502 Ok(acc)
503 })
504 .await
505 .unwrap();
506 assert_eq!(decoded, chunks.concat().as_bytes());
507 }
508
509 #[test]
510 fn faiths_coding_is_named_after_the_codings_the_caller_declared() {
511 assert_eq!(
512 layer_content_encoding(Some("gzip"), Coding::Zstd),
513 "gzip, zstd"
514 );
515 assert_eq!(
516 layer_content_encoding(Some("gzip, br"), Coding::Deflate),
517 "gzip, br, deflate"
518 );
519 }
520
521 #[test]
522 fn a_request_declaring_nothing_names_only_the_coding_faith_applied() {
523 assert_eq!(layer_content_encoding(None, Coding::Brotli), "br");
524 assert_eq!(layer_content_encoding(Some(""), Coding::Gzip), "gzip");
525 assert_eq!(layer_content_encoding(Some(" "), Coding::Gzip), "gzip");
526 }
527}