use std::fmt;
use std::io::Read;
use std::io::Write;
use std::str::FromStr;
use bytes::Bytes;
#[repr(C)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Compression {
Auto,
Zstd,
Brotli,
Gzip,
Deflate,
}
impl Compression {
pub const CODINGS: &[Self] = &[Self::Zstd, Self::Brotli, Self::Gzip, Self::Deflate];
pub const COUNT: usize = Self::CODINGS.len() + 1;
pub const ACCEPTED: &str = "zstd, br, gzip, deflate";
pub const IDENTITY: &str = "identity";
pub const ZSTD_LEVEL: i32 = 3;
pub const BROTLI_QUALITY: i32 = 5;
pub const BROTLI_WINDOW: i32 = 22;
pub const BUFFER: usize = 8192;
pub fn as_str(&self) -> &'static str {
match self {
Self::Auto => "",
Self::Zstd => "zstd",
Self::Brotli => "br",
Self::Gzip => "gzip",
Self::Deflate => "deflate",
}
}
pub fn parse(token: &str) -> Option<Self> {
let token = token.trim();
Self::CODINGS
.iter()
.copied()
.find(|coding| token.eq_ignore_ascii_case(coding.as_str()))
.or_else(|| token.eq_ignore_ascii_case("x-gzip").then_some(Self::Gzip))
}
pub fn accepted<'a>(values: impl Iterator<Item = &'a str>) -> Option<Self> {
let mut quality = [None; Self::COUNT];
let mut wildcard = None;
for coding in values.flat_map(Coding::list) {
match coding.compression() {
Some(compression) => quality[compression as usize] = Some(coding.quality),
None if coding.wildcard() => wildcard = Some(coding.quality),
None => {}
}
}
let permitted = |coding: &Self| quality[*coding as usize].or(wildcard).unwrap_or(Coding::NONE) > Coding::NONE;
Self::CODINGS.iter().copied().find(permitted)
}
pub fn applied<'a>(values: impl Iterator<Item = &'a str>) -> Option<Self> {
let mut applied = None;
for coding in values.flat_map(Coding::list) {
if coding.token.eq_ignore_ascii_case(Self::IDENTITY) {
continue;
}
if applied.is_some() {
return None;
}
applied = Some(coding.compression()?);
}
applied
}
pub fn encoded<'a>(values: impl Iterator<Item = &'a str>) -> bool {
values.flat_map(Coding::list).any(|coding| !coding.token.eq_ignore_ascii_case(Self::IDENTITY))
}
pub fn drain(reader: impl Read, max: u64, out: &mut Vec<u8>) -> Result<(), Error> {
let start = out.len();
let mut bounded = reader.take(max.saturating_add(1));
match std::io::copy(&mut bounded, out) {
Ok(produced) if produced <= max => Ok(()),
Ok(_) => {
out.truncate(start);
Err(Error::TooLarge(max))
}
Err(err) => {
out.truncate(start);
Err(Error::coding(err))
}
}
}
pub fn encode(&self, input: &[u8]) -> Result<Bytes, Error> {
let mut out = Vec::with_capacity(input.len() / 2 + Self::BUFFER.min(input.len() + 64));
self.encode_into(input, &mut out)?;
Ok(Bytes::from(out))
}
pub fn encode_into(&self, input: &[u8], out: &mut Vec<u8>) -> Result<(), Error> {
match self {
Self::Auto => Err(Error::Settled),
Self::Zstd => zstd::stream::copy_encode(input, out, Self::ZSTD_LEVEL).map_err(Error::coding),
Self::Brotli => {
let params = brotli::enc::BrotliEncoderParams { quality: Self::BROTLI_QUALITY, lgwin: Self::BROTLI_WINDOW, ..Default::default() };
let mut source = input;
brotli::BrotliCompress(&mut source, out, ¶ms).map(drop).map_err(Error::coding)
}
Self::Gzip => {
let mut encoder = flate2::write::GzEncoder::new(out, flate2::Compression::default());
encoder.write_all(input).map_err(Error::coding)?;
encoder.finish().map(drop).map_err(Error::coding)
}
Self::Deflate => {
let mut encoder = flate2::write::ZlibEncoder::new(out, flate2::Compression::default());
encoder.write_all(input).map_err(Error::coding)?;
encoder.finish().map(drop).map_err(Error::coding)
}
}
}
pub fn decode(&self, input: &[u8], max: u64) -> Result<Bytes, Error> {
let mut out = Vec::new();
self.decode_into(input, max, &mut out)?;
Ok(Bytes::from(out))
}
pub fn decode_into(&self, input: &[u8], max: u64, out: &mut Vec<u8>) -> Result<(), Error> {
match self {
Self::Auto => Err(Error::Settled),
Self::Zstd => Self::drain(zstd::stream::read::Decoder::new(input).map_err(Error::coding)?, max, out),
Self::Brotli => Self::drain(brotli::Decompressor::new(input, Self::BUFFER), max, out),
Self::Gzip => Self::drain(flate2::read::GzDecoder::new(input), max, out),
Self::Deflate => match Self::drain(flate2::read::ZlibDecoder::new(input), max, out) {
Err(Error::Coding(_)) => Self::drain(flate2::read::DeflateDecoder::new(input), max, out),
settled => settled,
},
}
}
}
impl fmt::Display for Compression {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for Compression {
type Err = ();
fn from_str(text: &str) -> Result<Self, Self::Err> {
Self::parse(text).ok_or(())
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Coding<'a> {
pub token: &'a str,
pub quality: f32,
}
impl<'a> Coding<'a> {
pub const WILDCARD: &'static str = "*";
pub const FULL: f32 = 1.0;
pub const NONE: f32 = 0.0;
pub fn parse(entry: &'a str) -> Self {
let mut parts = entry.split(';');
let token = parts.next().unwrap_or_default().trim();
let written = parts.filter_map(|parameter| parameter.split_once('=')).find(|(name, _)| name.trim().eq_ignore_ascii_case("q"));
let quality = match written {
Some((_, value)) => value.trim().parse::<f32>().ok().filter(|quality| (Self::NONE..=Self::FULL).contains(quality)).unwrap_or(Self::NONE),
None => Self::FULL,
};
Self { token, quality }
}
pub fn list(value: &'a str) -> impl Iterator<Item = Self> {
value.split(',').map(Self::parse).filter(|coding| !coding.token.is_empty())
}
pub fn compression(&self) -> Option<Compression> {
Compression::parse(self.token)
}
pub fn wildcard(&self) -> bool {
self.token == Self::WILDCARD
}
pub fn accepts(&self) -> bool {
self.quality > Self::NONE
}
}
#[derive(Debug, PartialEq, Eq)]
pub enum Error {
Settled,
TooLarge(u64),
Coding(String),
}
impl Error {
pub fn coding(error: impl fmt::Display) -> Self {
Self::Coding(error.to_string())
}
}
impl fmt::Display for Error {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Settled => write!(f, "the content coding was never settled"),
Self::TooLarge(max) => write!(f, "the decoded body exceeds {max} octets"),
Self::Coding(reason) => write!(f, "the content coding failed: {reason}"),
}
}
}
impl std::error::Error for Error {}