use std::borrow::Cow;
use celeriant_wal::compression_type::CompressionType;
use celeriant_wire::codec::compression::{compress_with_dict, decompress_with_dict};
use celeriant_wire::network::wire_error::WireError;
use celeriant_wire::network::wire_header::{
WireHeader, checked_u32_len, wire_header_write_variable_size_raw,
};
use futures_lite::{AsyncReadExt, AsyncWriteExt};
const COMPRESSION_LEVEL: i32 = 3;
pub struct OutFrame {
type_id: u32,
compression: CompressionType,
uncompressed_size: u32,
body: Vec<u8>,
}
pub fn build_frame(
type_id: u32,
uncompressed: Vec<u8>,
dict: Option<&[u8]>,
compress: bool,
max_size_bytes: u64,
) -> Result<OutFrame, WireError> {
let uncompressed_size = checked_u32_len(uncompressed.len())?;
if uncompressed_size as u64 > max_size_bytes {
return Err(WireError::MessageTooLarge {
message_length: uncompressed_size as u64,
max_size_bytes,
});
}
let (compression, body) = match (compress, dict) {
(true, Some(dict)) => (
CompressionType::ZstdDict,
compress_with_dict(&uncompressed, COMPRESSION_LEVEL, dict)?,
),
_ => (CompressionType::None, uncompressed),
};
Ok(OutFrame { type_id, compression, uncompressed_size, body })
}
pub async fn write_frame<W>(
writer: &mut W,
frame: &OutFrame,
max_size_bytes: u64,
version: u32,
) -> Result<(), WireError>
where
W: AsyncWriteExt + Unpin,
{
wire_header_write_variable_size_raw(
writer,
&frame.body,
frame.type_id,
frame.compression,
frame.uncompressed_size,
max_size_bytes,
version,
)
.await
}
pub struct InFrame {
pub message_type: u32,
compression: CompressionType,
uncompressed_size: u32,
body: Vec<u8>,
}
pub async fn read_frame<R>(reader: &mut R, max_size_bytes: u64) -> Result<InFrame, WireError>
where
R: AsyncReadExt + Unpin,
{
let header = WireHeader::from_reader(reader, max_size_bytes).await?;
let body = header.read_variable_body_raw(reader).await?;
Ok(InFrame {
message_type: header.message_type,
compression: header.compression_type,
uncompressed_size: header.uncompressed_length,
body,
})
}
pub fn decompress<'a>(frame: &'a InFrame, dict: Option<&[u8]>) -> Result<Cow<'a, [u8]>, WireError> {
decompress_body(frame.compression, frame.uncompressed_size, &frame.body, dict)
}
pub fn decompress_body<'a>(
compression: CompressionType,
uncompressed_size: u32,
body: &'a [u8],
dict: Option<&[u8]>,
) -> Result<Cow<'a, [u8]>, WireError> {
match (compression, dict) {
(CompressionType::None, _) => Ok(Cow::Borrowed(body)),
(CompressionType::ZstdDict, Some(dict)) => {
Ok(Cow::Owned(decompress_with_dict(body, uncompressed_size as usize, dict)?))
}
(CompressionType::ZstdDict, None) => Err(WireError::MalformedFrame(
"ZstdDict frame but the connection has no compression dictionary".into(),
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use celeriant_wire::network::wire_header::PROTOCOL_VERSION_V2;
use futures_lite::future::block_on;
use futures_lite::io::Cursor;
const MAX: u64 = 1 << 20;
fn dict() -> Vec<u8> {
let mut d = vec![0u8; 14 * 1024];
for (i, b) in d.iter_mut().enumerate() {
*b = (i % 251) as u8;
}
d
}
struct Conn {
dict: Vec<u8>,
}
impl Conn {
async fn round_trip(&mut self, type_id: u32, body: Vec<u8>, compress: bool) -> Result<Vec<u8>, WireError> {
let frame = build_frame(type_id, body, Some(&self.dict), compress, MAX)?;
let mut wire = Vec::new();
write_frame(&mut wire, &frame, MAX, PROTOCOL_VERSION_V2).await?;
let mut reader = Cursor::new(wire);
let inframe = read_frame(&mut reader, MAX).await?;
assert_eq!(inframe.message_type, type_id);
Ok(decompress(&inframe, Some(&self.dict))?.into_owned())
}
}
#[test]
fn request_future_is_send_and_round_trips_compressed() {
fn assert_send<T: Send>(_: &T) {}
let mut conn = Conn { dict: dict() };
let payload = vec![7u8; 4096];
let fut = conn.round_trip(3, payload.clone(), true);
assert_send(&fut);
assert_eq!(block_on(fut).unwrap(), payload);
}
#[test]
fn an_oversized_body_is_refused_before_it_is_compressed() {
const TINY_MAX: u64 = 4096;
let body = vec![0u8; 1_048_576];
match build_frame(3, body, Some(&dict()), true, TINY_MAX) {
Ok(_) => panic!("a body over the cap must not reach the compressor"),
Err(WireError::MessageTooLarge { message_length, max_size_bytes }) => {
assert_eq!(message_length, 1_048_576, "the uncompressed length is what was judged");
assert_eq!(max_size_bytes, TINY_MAX);
}
Err(other) => panic!("expected MessageTooLarge, got {other:?}"),
}
}
#[test]
fn uncompressed_round_trips() {
let mut conn = Conn { dict: dict() };
let payload = b"small control verb".to_vec();
assert_eq!(block_on(conn.round_trip(8, payload.clone(), false)).unwrap(), payload);
}
}