use flate2::{Compress, Compression, Decompress, FlushCompress, FlushDecompress};
use hyper::header::{HeaderMap, HeaderValue, SEC_WEBSOCKET_EXTENSIONS};
const DEFLATE_TAIL: [u8; 4] = [0x00, 0x00, 0xff, 0xff];
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub struct DeflateConfig {
pub server_no_context_takeover: bool,
pub client_no_context_takeover: bool,
pub server_max_window_bits: u8,
}
impl Default for DeflateConfig {
fn default() -> Self {
Self {
server_no_context_takeover: false,
client_no_context_takeover: false,
server_max_window_bits: 15,
}
}
}
#[derive(Debug, Default, Clone, Copy)]
pub(super) struct Offer {
server_no_context_takeover: bool,
client_no_context_takeover: bool,
server_max_window_bits_offered: bool,
}
#[derive(Debug, Clone, Copy)]
pub(super) struct Agreement {
pub(super) server_no_context_takeover: bool,
pub(super) client_no_context_takeover: bool,
pub(super) server_max_window_bits: u8,
echo_server_max_window_bits: bool,
}
pub(super) fn parse_offers(headers: &HeaderMap) -> Vec<Offer> {
let mut offers = Vec::new();
for header in headers.get_all(SEC_WEBSOCKET_EXTENSIONS) {
let Ok(text) = header.to_str() else { continue };
for extension in text.split(',') {
let mut parts = extension.split(';').map(str::trim);
let Some(name) = parts.next() else { continue };
if !name.eq_ignore_ascii_case("permessage-deflate") {
continue;
}
if let Some(offer) = parse_params(parts) {
offers.push(offer);
}
}
}
offers
}
fn parse_window_bits(value: Option<&str>) -> Result<(), ()> {
match value.map(str::parse::<u8>) {
None | Some(Ok(9..=15)) => Ok(()),
Some(_) => Err(()),
}
}
fn parse_params<'a>(params: impl Iterator<Item = &'a str>) -> Option<Offer> {
let mut offer = Offer::default();
for param in params {
if param.is_empty() {
continue;
}
let (key, value) = match param.split_once('=') {
Some((k, v)) => (k.trim(), Some(v.trim().trim_matches('"'))),
None => (param, None),
};
match key {
"server_no_context_takeover" if value.is_none() => {
offer.server_no_context_takeover = true;
}
"client_no_context_takeover" if value.is_none() => {
offer.client_no_context_takeover = true;
}
"server_max_window_bits" => {
parse_window_bits(value).ok()?;
offer.server_max_window_bits_offered = true;
}
"client_max_window_bits" => {
parse_window_bits(value).ok()?;
}
_ => return None,
}
}
Some(offer)
}
pub(super) fn negotiate(offers: &[Offer], config: DeflateConfig) -> Option<Agreement> {
let offer = offers.first()?;
let server_max_window_bits = config.server_max_window_bits.clamp(9, 15);
Some(Agreement {
server_no_context_takeover: config.server_no_context_takeover
|| offer.server_no_context_takeover,
client_no_context_takeover: config.client_no_context_takeover
|| offer.client_no_context_takeover,
server_max_window_bits,
echo_server_max_window_bits: offer.server_max_window_bits_offered
&& server_max_window_bits < 15,
})
}
pub(super) fn agreement_header_value(agreement: Agreement) -> HeaderValue {
let mut value = String::from("permessage-deflate");
if agreement.server_no_context_takeover {
value.push_str("; server_no_context_takeover");
}
if agreement.client_no_context_takeover {
value.push_str("; client_no_context_takeover");
}
if agreement.echo_server_max_window_bits {
value.push_str("; server_max_window_bits=");
value.push_str(&agreement.server_max_window_bits.to_string());
}
HeaderValue::from_str(&value).unwrap_or_else(|_| HeaderValue::from_static("permessage-deflate"))
}
pub(super) struct PerMessageDeflate {
compress: Compress,
decompress: Decompress,
server_no_context_takeover: bool,
client_no_context_takeover: bool,
}
impl PerMessageDeflate {
pub(super) fn new(agreement: Agreement) -> Self {
Self {
compress: Compress::new_with_window_bits(
Compression::default(),
false,
agreement.server_max_window_bits,
),
decompress: Decompress::new_with_window_bits(false, 15),
server_no_context_takeover: agreement.server_no_context_takeover,
client_no_context_takeover: agreement.client_no_context_takeover,
}
}
pub(super) fn compress_if_smaller(
&mut self,
data: &[u8],
) -> Result<Option<Vec<u8>>, crate::http::error::Error> {
let compressed = self.compress_raw(data)?;
if compressed.len() < data.len() {
if self.server_no_context_takeover {
self.compress.reset();
}
Ok(Some(compressed))
} else {
self.compress.reset();
Ok(None)
}
}
fn compress_raw(&mut self, data: &[u8]) -> Result<Vec<u8>, crate::http::error::Error> {
let total_in_before = self.compress.total_in();
let mut out = Vec::with_capacity(data.len() + 32);
loop {
grow(&mut out, 1024.max(data.len()));
self.compress
.compress_vec(data, &mut out, FlushCompress::Sync)
.map_err(|e| crate::http::error::Error::Internal(e.to_string()))?;
let consumed =
usize::try_from(self.compress.total_in() - total_in_before).unwrap_or(usize::MAX);
if consumed >= data.len() {
break;
}
}
out.truncate(out.len().saturating_sub(DEFLATE_TAIL.len()));
Ok(out)
}
pub(super) fn decompress(
&mut self,
data: &[u8],
max_size: Option<usize>,
) -> Result<Vec<u8>, crate::http::error::Error> {
let max_size = max_size.unwrap_or(usize::MAX);
let mut input = Vec::with_capacity(data.len() + DEFLATE_TAIL.len());
input.extend_from_slice(data);
input.extend_from_slice(&DEFLATE_TAIL);
let total_in_before = self.decompress.total_in();
let mut out = Vec::with_capacity((data.len() * 3 + 32).min(max_size));
loop {
grow(&mut out, 1024);
let consumed_before =
usize::try_from(self.decompress.total_in() - total_in_before).unwrap_or(usize::MAX);
self.decompress
.decompress_vec(&input[consumed_before..], &mut out, FlushDecompress::Sync)
.map_err(|e| crate::http::error::Error::Internal(e.to_string()))?;
if out.len() > max_size {
return Err(crate::http::error::Error::Internal(
"decompressed message exceeds the configured maximum size".to_string(),
));
}
let consumed =
usize::try_from(self.decompress.total_in() - total_in_before).unwrap_or(usize::MAX);
if consumed >= input.len() {
break;
}
}
if self.client_no_context_takeover {
self.decompress.reset(false);
}
Ok(out)
}
}
fn grow(out: &mut Vec<u8>, min_initial: usize) {
let additional = out.capacity().max(min_initial);
out.reserve(additional);
}