use crate::error::{StoreError, StoreResult};
pub const FORMAT_RAW: u8 = 0x00;
pub const FORMAT_ZSTD: u8 = 0x01;
pub const DEFAULT_COMPRESSION_LEVEL: i32 = 3;
pub const DEFAULT_MIN_LEN: usize = 96;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Compression {
On {
level: i32,
min_len: usize,
},
Off,
}
impl Default for Compression {
fn default() -> Self {
Compression::On {
level: DEFAULT_COMPRESSION_LEVEL,
min_len: DEFAULT_MIN_LEN,
}
}
}
pub(crate) fn pack_value(plain: &[u8], policy: Compression) -> Vec<u8> {
if let Compression::On { level, min_len } = policy {
if plain.len() >= min_len {
if let Ok(z) = rusty_zstd::compress(plain, level) {
if z.len() < plain.len() {
let mut out = Vec::with_capacity(1 + z.len());
out.push(FORMAT_ZSTD);
out.extend_from_slice(&z);
return out;
}
}
}
}
let mut out = Vec::with_capacity(1 + plain.len());
out.push(FORMAT_RAW);
out.extend_from_slice(plain);
out
}
pub(crate) fn unpack_value(packed: &[u8]) -> StoreResult<Vec<u8>> {
match packed.split_first() {
Some((&FORMAT_RAW, rest)) => Ok(rest.to_vec()),
Some((&FORMAT_ZSTD, rest)) => rusty_zstd::decompress(rest)
.map_err(|e| StoreError::Compression(format!("zstd frame: {e:?}"))),
Some((&byte, _)) => Err(StoreError::Compression(format!(
"unknown row format byte {byte:#04x}"
))),
None => Err(StoreError::Compression("empty sealed payload".into())),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn raw_round_trips() {
let plain = b"hello world";
let packed = pack_value(plain, Compression::Off);
assert_eq!(packed[0], FORMAT_RAW);
assert_eq!(unpack_value(&packed).unwrap(), plain);
}
#[test]
fn compressible_payload_shrinks_and_round_trips() {
let plain = vec![b'a'; 4096];
let packed = pack_value(&plain, Compression::default());
assert_eq!(packed[0], FORMAT_ZSTD);
assert!(packed.len() < plain.len());
assert_eq!(unpack_value(&packed).unwrap(), plain);
}
#[test]
fn incompressible_payload_stays_raw() {
let mut state = 0x9E3779B97F4A7C15u64;
let plain: Vec<u8> = (0..4096)
.map(|_| {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
(state >> 56) as u8
})
.collect();
let packed = pack_value(&plain, Compression::default());
assert_eq!(packed[0], FORMAT_RAW);
assert_eq!(packed.len(), plain.len() + 1);
assert_eq!(unpack_value(&packed).unwrap(), plain);
}
#[test]
fn below_the_floor_is_not_attempted() {
let plain = vec![b'a'; 8];
let packed = pack_value(
&plain,
Compression::On {
level: DEFAULT_COMPRESSION_LEVEL,
min_len: 64,
},
);
assert_eq!(packed[0], FORMAT_RAW);
}
#[test]
fn unknown_format_byte_fails_loudly() {
assert!(unpack_value(&[0x7F, 1, 2, 3]).is_err());
assert!(unpack_value(&[]).is_err());
}
}