use std::{io, pin::Pin};
use async_compression::tokio::bufread::{
BrotliDecoder, BrotliEncoder, GzipDecoder, GzipEncoder, ZlibDecoder, ZlibEncoder, ZstdDecoder,
ZstdEncoder,
};
use bytes::Bytes;
use futures::{Stream, TryStreamExt};
use http::header::{CONTENT_ENCODING, CONTENT_LENGTH, HeaderMap};
use tokio::io::AsyncReadExt;
use tokio_util::io::{ReaderStream, StreamReader};
pub type ByteStream = dyn Stream<Item = Result<Bytes, String>> + Send + Sync;
pub const DEFAULT_ACCEPT_ENCODING: &str = "zstd,gzip,deflate,br";
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Coding {
Gzip,
Deflate,
Brotli,
Zstd,
}
impl Coding {
pub fn from_option(value: &str) -> Option<Self> {
match value {
"gzip" => Some(Self::Gzip),
"deflate" => Some(Self::Deflate),
"br" => Some(Self::Brotli),
"zstd" => Some(Self::Zstd),
_ => None,
}
}
pub fn token(self) -> &'static str {
match self {
Self::Gzip => "gzip",
Self::Deflate => "deflate",
Self::Brotli => "br",
Self::Zstd => "zstd",
}
}
pub fn from_token(token: &str) -> Option<Self> {
let token = token.trim();
if token.eq_ignore_ascii_case("gzip") || token.eq_ignore_ascii_case("x-gzip") {
Some(Self::Gzip)
} else if token.eq_ignore_ascii_case("deflate") {
Some(Self::Deflate)
} else if token.eq_ignore_ascii_case("br") {
Some(Self::Brotli)
} else if token.eq_ignore_ascii_case("zstd") {
Some(Self::Zstd)
} else {
None
}
}
}
pub fn decision(headers: &HeaderMap, accept: &AcceptEncoding) -> Option<Coding> {
let mut codings = Vec::new();
for value in headers.get_all(CONTENT_ENCODING) {
let value = value.to_str().ok()?;
codings.extend(value.split(',').map(str::trim).filter(|c| !c.is_empty()));
}
let [single] = codings[..] else {
return None;
};
let coding = Coding::from_token(single)?;
accept.accepts(coding).then_some(coding)
}
pub fn strip_decoded_headers(headers: &mut HeaderMap) {
headers.remove(CONTENT_ENCODING);
headers.remove(CONTENT_LENGTH);
}
#[derive(Clone, Copy, Debug, Default)]
pub struct AcceptEncoding {
gzip: Option<u16>,
deflate: Option<u16>,
brotli: Option<u16>,
zstd: Option<u16>,
star: Option<u16>,
}
impl AcceptEncoding {
pub fn parse(value: &str) -> Self {
let mut accept = Self::default();
for element in value.split(',') {
let mut parts = element.split(';');
let Some(token) = parts.next().map(str::trim) else {
continue;
};
if token.is_empty() {
continue;
}
let mut quality = 1000;
for param in parts {
let param = param.trim();
if let Some(rest) = param
.strip_prefix("q=")
.or_else(|| param.strip_prefix("Q="))
{
quality = parse_quality(rest).unwrap_or(0);
}
}
let slot = if token == "*" {
&mut accept.star
} else {
match Coding::from_token(token) {
Some(Coding::Gzip) => &mut accept.gzip,
Some(Coding::Deflate) => &mut accept.deflate,
Some(Coding::Brotli) => &mut accept.brotli,
Some(Coding::Zstd) => &mut accept.zstd,
None => continue,
}
};
*slot = Some(quality);
}
accept
}
fn accepts(&self, coding: Coding) -> bool {
let named = match coding {
Coding::Gzip => self.gzip,
Coding::Deflate => self.deflate,
Coding::Brotli => self.brotli,
Coding::Zstd => self.zstd,
};
match named {
Some(quality) => quality > 0,
None => matches!(self.star, Some(quality) if quality > 0),
}
}
}
fn parse_quality(value: &str) -> Option<u16> {
let value = value.trim();
let mut chars = value.chars();
let mut quality: u16 = match chars.next()? {
'0' => 0,
'1' => 1000,
_ => return None,
};
if let Some(dot) = chars.next() {
if dot != '.' {
return None;
}
let mut scale = 100;
for digit in chars {
quality += digit.to_digit(10)? as u16 * scale;
if scale == 1 {
break;
}
scale /= 10;
}
}
Some(quality.min(1000))
}
pub fn decode_stream(input: Pin<Box<ByteStream>>, coding: Coding) -> Pin<Box<ByteStream>> {
let reader = StreamReader::new(input.map_err(io::Error::other));
match coding {
Coding::Gzip => reader_stream(GzipDecoder::new(reader)),
Coding::Deflate => reader_stream(ZlibDecoder::new(reader)),
Coding::Brotli => reader_stream(BrotliDecoder::new(reader)),
Coding::Zstd => {
let mut decoder = ZstdDecoder::new(reader);
decoder.multiple_members(true);
reader_stream(decoder)
}
}
}
fn reader_stream<R>(reader: R) -> Pin<Box<ByteStream>>
where
R: tokio::io::AsyncRead + Send + Sync + 'static,
{
Box::pin(ReaderStream::new(reader).map_err(|err| err.to_string()))
}
pub type RequestStream = Pin<Box<dyn Stream<Item = io::Result<Bytes>> + Send>>;
pub async fn compress_buffer(input: &[u8], coding: Coding) -> io::Result<Vec<u8>> {
let mut output = Vec::new();
match coding {
Coding::Gzip => GzipEncoder::new(input).read_to_end(&mut output).await?,
Coding::Deflate => ZlibEncoder::new(input).read_to_end(&mut output).await?,
Coding::Brotli => BrotliEncoder::new(input).read_to_end(&mut output).await?,
Coding::Zstd => ZstdEncoder::new(input).read_to_end(&mut output).await?,
};
Ok(output)
}
pub fn compress_stream<S>(input: S, coding: Coding) -> RequestStream
where
S: Stream<Item = io::Result<Bytes>> + Send + 'static,
{
let reader = StreamReader::new(input);
match coding {
Coding::Gzip => encoder_stream(GzipEncoder::new(reader)),
Coding::Deflate => encoder_stream(ZlibEncoder::new(reader)),
Coding::Brotli => encoder_stream(BrotliEncoder::new(reader)),
Coding::Zstd => encoder_stream(ZstdEncoder::new(reader)),
}
}
fn encoder_stream<R>(reader: R) -> RequestStream
where
R: tokio::io::AsyncRead + Send + 'static,
{
Box::pin(ReaderStream::new(reader))
}
pub fn layer_content_encoding(declared: Option<&str>, applied: Coding) -> String {
match declared.map(str::trim).filter(|value| !value.is_empty()) {
Some(declared) => format!("{declared}, {}", applied.token()),
None => applied.token().to_owned(),
}
}
#[cfg(test)]
mod tests {
use http::header::{CONTENT_ENCODING, HeaderMap, HeaderValue};
use super::*;
fn decide(content_encoding: &str, accept: &str) -> Option<Coding> {
let mut headers = HeaderMap::new();
headers.insert(
CONTENT_ENCODING,
HeaderValue::from_str(content_encoding).unwrap(),
);
decision(&headers, &AcceptEncoding::parse(accept))
}
#[test]
fn decodes_a_negotiated_coding() {
assert_eq!(decide("gzip", DEFAULT_ACCEPT_ENCODING), Some(Coding::Gzip));
assert_eq!(decide("br", DEFAULT_ACCEPT_ENCODING), Some(Coding::Brotli));
assert_eq!(decide("zstd", DEFAULT_ACCEPT_ENCODING), Some(Coding::Zstd));
assert_eq!(
decide("deflate", DEFAULT_ACCEPT_ENCODING),
Some(Coding::Deflate)
);
}
#[test]
fn a_coding_named_alone_decodes_only_itself() {
assert_eq!(decide("gzip", "gzip"), Some(Coding::Gzip));
assert_eq!(decide("br", "gzip"), None);
}
#[test]
fn identity_leaves_a_compressed_body_alone() {
assert_eq!(decide("gzip", "identity"), None);
}
#[test]
fn a_zero_quality_value_refuses() {
assert_eq!(decide("gzip", "gzip;q=0"), None);
assert_eq!(decide("gzip", "gzip;q=0.000"), None);
}
#[test]
fn a_named_coding_settles_the_question_over_star() {
assert_eq!(decide("gzip", "gzip;q=0, *"), None);
assert_eq!(decide("br", "gzip;q=0, *"), Some(Coding::Brotli));
assert_eq!(decide("zstd", "gzip;q=0, *"), Some(Coding::Zstd));
}
#[test]
fn star_covers_what_is_not_named() {
assert_eq!(decide("gzip", "*"), Some(Coding::Gzip));
assert_eq!(decide("gzip", "br, *"), Some(Coding::Gzip));
}
#[test]
fn a_star_with_zero_quality_accepts_nothing_unnamed() {
assert_eq!(decide("gzip", "*;q=0"), None);
assert_eq!(decide("gzip", "gzip, *;q=0"), Some(Coding::Gzip));
}
#[test]
fn more_than_one_coding_is_delivered_as_received() {
assert_eq!(decide("gzip, br", DEFAULT_ACCEPT_ENCODING), None);
assert_eq!(decide("br, gzip", DEFAULT_ACCEPT_ENCODING), None);
assert_eq!(decide("identity, gzip", DEFAULT_ACCEPT_ENCODING), None);
}
#[test]
fn codings_split_across_header_lines_count_together() {
let mut headers = HeaderMap::new();
headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.append(CONTENT_ENCODING, HeaderValue::from_static("br"));
assert_eq!(
decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
None
);
}
#[test]
fn one_coding_split_across_lines_with_an_empty_line_still_decodes() {
let mut headers = HeaderMap::new();
headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.append(CONTENT_ENCODING, HeaderValue::from_static(""));
assert_eq!(
decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
Some(Coding::Gzip)
);
}
#[test]
fn a_non_ascii_line_is_delivered_as_received() {
let mut headers = HeaderMap::new();
headers.append(CONTENT_ENCODING, HeaderValue::from_static("gzip"));
headers.append(CONTENT_ENCODING, HeaderValue::from_bytes(b"\xff").unwrap());
assert_eq!(
decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
None
);
}
#[test]
fn a_coding_faith_cannot_decode_is_delivered_as_received() {
assert_eq!(decide("compress", DEFAULT_ACCEPT_ENCODING), None);
}
#[test]
fn no_content_encoding_means_nothing_to_decode() {
let headers = HeaderMap::new();
assert_eq!(
decision(&headers, &AcceptEncoding::parse(DEFAULT_ACCEPT_ENCODING)),
None
);
}
#[test]
fn quality_values_parse_to_thousandths() {
assert_eq!(parse_quality("0"), Some(0));
assert_eq!(parse_quality("1"), Some(1000));
assert_eq!(parse_quality("0.5"), Some(500));
assert_eq!(parse_quality("0.001"), Some(1));
assert_eq!(parse_quality("1.0"), Some(1000));
}
#[test]
fn the_compress_option_names_a_coding_by_its_wire_token() {
assert_eq!(Coding::from_option("gzip"), Some(Coding::Gzip));
assert_eq!(Coding::from_option("deflate"), Some(Coding::Deflate));
assert_eq!(Coding::from_option("br"), Some(Coding::Brotli));
assert_eq!(Coding::from_option("zstd"), Some(Coding::Zstd));
}
#[test]
fn the_compress_option_matches_its_tokens_exactly() {
assert_eq!(Coding::from_token("x-gzip"), Some(Coding::Gzip));
assert_eq!(Coding::from_option("x-gzip"), None);
assert_eq!(Coding::from_token("GZIP"), Some(Coding::Gzip));
assert_eq!(Coding::from_option("GZIP"), None);
assert_eq!(Coding::from_option(" gzip"), None);
assert_eq!(Coding::from_option("brotli"), None);
assert_eq!(Coding::from_option("identity"), None);
assert_eq!(Coding::from_option(""), None);
}
#[tokio::test]
async fn a_compressed_body_decodes_back_to_what_went_in() {
let input = b"the quick brown fox jumps over the lazy dog".repeat(20);
for coding in [Coding::Gzip, Coding::Deflate, Coding::Brotli, Coding::Zstd] {
let compressed = compress_buffer(&input, coding).await.unwrap();
assert!(
compressed.len() < input.len(),
"{coding:?} did not compress repetitive input"
);
let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
let decoded: Vec<u8> = decode_stream(Box::pin(source), coding)
.try_fold(Vec::new(), |mut acc, chunk| async move {
acc.extend_from_slice(&chunk);
Ok(acc)
})
.await
.unwrap();
assert_eq!(decoded, input, "{coding:?} round trip");
}
}
#[tokio::test]
async fn a_streaming_body_compresses_across_its_chunks() {
let chunks = ["first chunk, ", "second chunk, ", "third chunk"];
let source = futures::stream::iter(
chunks
.into_iter()
.map(|chunk| Ok(Bytes::from_static(chunk.as_bytes()))),
);
let compressed: Vec<u8> = compress_stream(source, Coding::Zstd)
.try_fold(Vec::new(), |mut acc, chunk| async move {
acc.extend_from_slice(&chunk);
Ok(acc)
})
.await
.unwrap();
let source = futures::stream::once(async move { Ok(Bytes::from(compressed)) });
let decoded: Vec<u8> = decode_stream(Box::pin(source), Coding::Zstd)
.try_fold(Vec::new(), |mut acc, chunk| async move {
acc.extend_from_slice(&chunk);
Ok(acc)
})
.await
.unwrap();
assert_eq!(decoded, chunks.concat().as_bytes());
}
#[test]
fn faiths_coding_is_named_after_the_codings_the_caller_declared() {
assert_eq!(
layer_content_encoding(Some("gzip"), Coding::Zstd),
"gzip, zstd"
);
assert_eq!(
layer_content_encoding(Some("gzip, br"), Coding::Deflate),
"gzip, br, deflate"
);
}
#[test]
fn a_request_declaring_nothing_names_only_the_coding_faith_applied() {
assert_eq!(layer_content_encoding(None, Coding::Brotli), "br");
assert_eq!(layer_content_encoding(Some(""), Coding::Gzip), "gzip");
assert_eq!(layer_content_encoding(Some(" "), Coding::Gzip), "gzip");
}
}