use crate::{
ProtocolError,
protocol::{IncompleteMessageType, PerMessageDeflateConfig, Role},
};
use flate2::{
Compress, CompressError, Compression, Decompress, DecompressError, FlushCompress,
FlushDecompress, Status,
};
use rama_core::error::{BoxError, ErrorContext};
use std::slice;
#[derive(Debug)]
pub(super) struct PerMessageDeflateState {
pub(super) decompress_incomplete_msg: IncompleteCompressedMessage,
pub(super) encoder: DeflateEncoder,
pub(super) decoder: DeflateDecoder,
}
impl PerMessageDeflateState {
pub(super) fn new(role: Role, cfg: PerMessageDeflateConfig) -> Self {
match role {
Role::Server => Self {
decompress_incomplete_msg: Default::default(),
encoder: DeflateEncoder::new(
Compression::default(),
cfg.server_max_window_bits.unwrap_or(15),
cfg.server_no_context_takeover,
),
decoder: DeflateDecoder::new(
cfg.client_max_window_bits.unwrap_or(15),
cfg.client_no_context_takeover,
),
},
Role::Client => Self {
decompress_incomplete_msg: Default::default(),
encoder: DeflateEncoder::new(
Compression::default(),
cfg.client_max_window_bits.unwrap_or(15),
cfg.client_no_context_takeover,
),
decoder: DeflateDecoder::new(
cfg.server_max_window_bits.unwrap_or(15),
cfg.server_no_context_takeover,
),
},
}
}
}
const DEFLATE_TRAILER: [u8; 4] = [0, 0, 255, 255];
#[derive(Debug)]
pub(super) struct DeflateEncoder {
compress: Compress,
compress_reset: bool,
}
#[derive(Debug, Default)]
pub(super) struct IncompleteCompressedMessage {
pub(super) buffer: Vec<u8>,
pub(super) msg_type: Option<IncompleteMessageType>,
}
impl IncompleteCompressedMessage {
pub(super) fn reset(&mut self, r#type: IncompleteMessageType) {
self.buffer.clear();
self.msg_type = Some(r#type);
}
pub(super) fn fin_buffer(
&mut self,
tail: impl AsRef<[u8]>,
size_limit: Option<usize>,
) -> Result<(&[u8], IncompleteMessageType), ProtocolError> {
match self.msg_type.take() {
Some(t) => {
self.extend(tail, size_limit)?;
Ok((&self.buffer, t))
}
None => Err(ProtocolError::UnexpectedContinueFrame),
}
}
pub(super) fn extend<T: AsRef<[u8]>>(
&mut self,
tail: T,
size_limit: Option<usize>,
) -> Result<(), ProtocolError> {
let max_size = size_limit.unwrap_or_else(usize::max_value);
let my_size = self.buffer.len();
let portion_size = tail.as_ref().len();
if my_size > max_size || portion_size > max_size - my_size {
return Err(ProtocolError::MessageTooLong {
size: my_size + portion_size,
max_size,
});
}
self.buffer.extend(tail.as_ref());
Ok(())
}
}
impl DeflateEncoder {
pub(super) fn new(compression: Compression, mut window_size: u8, compress_reset: bool) -> Self {
if window_size == 8 {
window_size = 9;
}
Self {
compress: Compress::new_with_window_bits(compression, false, window_size),
compress_reset,
}
}
pub(super) fn encode(&mut self, input_data: &[u8]) -> Result<Vec<u8>, BoxError> {
if input_data.is_empty() {
return Ok(vec![0x00]);
}
let mut buf = Vec::with_capacity(input_data.len() * 2);
let before_in = self.compress.total_in();
while self.compress.total_in() - before_in < input_data.as_ref().len() as u64 {
let i = self.compress.total_in() as usize - before_in as usize;
match self
.compress
.buf_compress(&input_data[i..], &mut buf, FlushCompress::Sync)
.context("deflate encode next chunk")?
{
Status::BufError => buf.reserve((buf.len() as f64 * 1.5) as usize),
Status::Ok => (),
Status::StreamEnd => break,
}
}
while !buf.ends_with(&[0, 0, 0xFF, 0xFF]) {
buf.reserve(5);
match self
.compress
.buf_compress(&[], &mut buf, FlushCompress::Sync)
.context("enforce buf to finish")?
{
Status::Ok | Status::BufError => (),
Status::StreamEnd => break,
}
}
buf.truncate(buf.len() - DEFLATE_TRAILER.len());
if self.compress_reset {
self.compress.reset();
}
Ok(buf)
}
}
#[derive(Debug)]
pub(super) struct DeflateDecoder {
decompress: Decompress,
decompress_reset: bool,
}
impl DeflateDecoder {
pub(super) fn new(mut window_size: u8, decompress_reset: bool) -> Self {
if window_size == 8 {
window_size = 9;
}
Self {
decompress: Decompress::new_with_window_bits(false, window_size),
decompress_reset,
}
}
pub(super) fn decode(
&mut self,
compressed_data: &[u8],
size_limit: Option<usize>,
) -> Result<Vec<u8>, ProtocolError> {
let max_size = size_limit.unwrap_or_else(usize::max_value);
let initial_capacity = (compressed_data.len() + DEFLATE_TRAILER.len())
.saturating_mul(2)
.min(max_size);
let mut buf = Vec::with_capacity(initial_capacity);
for payload in [compressed_data, &DEFLATE_TRAILER] {
let before_in = self.decompress.total_in();
while self.decompress.total_in() - before_in < payload.as_ref().len() as u64 {
let i = self.decompress.total_in() as usize - before_in as usize;
match self
.decompress
.buf_decompress(&payload[i..], &mut buf, FlushDecompress::Sync, max_size)
.context("flate2 decode next chunk")
.map_err(ProtocolError::DeflateError)?
{
Status::BufError => grow_inflate_buffer(&mut buf, max_size)?,
Status::Ok => (),
Status::StreamEnd => break,
}
check_decode_size(buf.len(), max_size)?;
}
}
if self.decompress_reset {
self.decompress.reset(false);
}
Ok(buf)
}
}
fn grow_inflate_buffer(buf: &mut Vec<u8>, max_size: usize) -> Result<(), ProtocolError> {
if buf.capacity() >= max_size {
return Err(ProtocolError::MessageTooLong {
size: max_size.saturating_add(1),
max_size,
});
}
let next_capacity = buf
.capacity()
.saturating_add((buf.capacity() / 2).max(1))
.min(max_size);
buf.reserve(next_capacity - buf.capacity());
Ok(())
}
fn check_decode_size(size: usize, max_size: usize) -> Result<(), ProtocolError> {
if size > max_size {
return Err(ProtocolError::MessageTooLong { size, max_size });
}
Ok(())
}
trait BufCompress {
fn buf_compress(
&mut self,
input: &[u8],
output: &mut Vec<u8>,
flush: FlushCompress,
) -> Result<Status, CompressError>;
}
trait BufDecompress {
fn buf_decompress(
&mut self,
input: &[u8],
output: &mut Vec<u8>,
flush: FlushDecompress,
max_size: usize,
) -> Result<Status, DecompressError>;
}
impl BufCompress for Compress {
fn buf_compress(
&mut self,
input: &[u8],
output: &mut Vec<u8>,
flush: FlushCompress,
) -> Result<Status, CompressError> {
op_buf(input, output, self.total_out(), |input, out| {
let ret = self.compress(input, out, flush);
(ret, self.total_out())
})
}
}
impl BufDecompress for Decompress {
fn buf_decompress(
&mut self,
input: &[u8],
output: &mut Vec<u8>,
flush: FlushDecompress,
max_size: usize,
) -> Result<Status, DecompressError> {
op_buf_limited(input, output, self.total_out(), max_size, |input, out| {
let ret = self.decompress(input, out, flush);
(ret, self.total_out())
})
}
}
fn op_buf<Fn, E>(input: &[u8], output: &mut Vec<u8>, before: u64, op: Fn) -> Result<Status, E>
where
Fn: FnOnce(&[u8], &mut [u8]) -> (Result<Status, E>, u64),
{
let cap = output.capacity();
let len = output.len();
#[expect(
clippy::multiple_unsafe_ops_per_block,
reason = "single uninitialized-tail write sequence: ptr.add → from_raw_parts_mut → set_len; splitting them harms readability without changing the safety contract"
)]
unsafe {
let ptr = output.as_mut_ptr().add(len);
let out = slice::from_raw_parts_mut(ptr, cap - len);
let (ret, total_out) = op(input, out);
output.set_len((total_out - before) as usize + len);
ret
}
}
fn op_buf_limited<Fn, E>(
input: &[u8],
output: &mut Vec<u8>,
before: u64,
max_size: usize,
op: Fn,
) -> Result<Status, E>
where
Fn: FnOnce(&[u8], &mut [u8]) -> (Result<Status, E>, u64),
{
let cap = output.capacity().min(max_size);
let len = output.len();
#[expect(
clippy::multiple_unsafe_ops_per_block,
reason = "single uninitialized-tail write sequence: ptr.add → from_raw_parts_mut → set_len; splitting them harms readability without changing the safety contract"
)]
unsafe {
let ptr = output.as_mut_ptr().add(len);
let out = slice::from_raw_parts_mut(ptr, cap - len);
let (ret, total_out) = op(input, out);
output.set_len((total_out - before) as usize + len);
ret
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deflate_decoder_rejects_output_past_size_limit() {
let payload = vec![b'a'; 1024];
let mut encoder = DeflateEncoder::new(Compression::default(), 15, false);
let compressed = encoder.encode(&payload).unwrap();
assert!(compressed.len() < payload.len());
let mut decoder = DeflateDecoder::new(15, false);
assert!(matches!(
decoder.decode(&compressed, Some(128)),
Err(ProtocolError::MessageTooLong {
size: 129,
max_size: 128
})
));
}
#[test]
fn deflate_decoder_allows_output_at_size_limit() {
let payload = vec![b'a'; 128];
let mut encoder = DeflateEncoder::new(Compression::default(), 15, false);
let compressed = encoder.encode(&payload).unwrap();
let mut decoder = DeflateDecoder::new(15, false);
let decoded = decoder.decode(&compressed, Some(payload.len())).unwrap();
assert_eq!(decoded, payload);
}
#[test]
fn deflate_decoder_unbounded_limit_inflates_correctly() {
for len in [0usize, 1, 200, 5000, 100_000] {
let payload = vec![b'a'; len];
let mut encoder = DeflateEncoder::new(Compression::default(), 15, false);
let compressed = encoder.encode(&payload).unwrap();
let mut decoder = DeflateDecoder::new(15, false);
let decoded = decoder.decode(&compressed, None).unwrap();
assert_eq!(decoded, payload, "unbounded decode mismatch for len={len}");
}
}
#[test]
fn deflate_decoder_boundary_accepts_n_rejects_n_minus_one() {
let payload = vec![b'a'; 4096];
let mut encoder = DeflateEncoder::new(Compression::default(), 15, false);
let compressed = encoder.encode(&payload).unwrap();
let mut decoder = DeflateDecoder::new(15, false);
assert_eq!(
decoder.decode(&compressed, Some(payload.len())).unwrap(),
payload,
);
let mut decoder = DeflateDecoder::new(15, false);
assert!(matches!(
decoder.decode(&compressed, Some(payload.len() - 1)),
Err(ProtocolError::MessageTooLong { max_size, .. }) if max_size == payload.len() - 1,
));
}
}