use std::borrow::Cow;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
pub enum PayloadCodec {
Raw = 0,
Zstd = 1,
}
impl PayloadCodec {
pub fn from_u32(value: u32) -> Result<Self, CodecError> {
match value {
0 => Ok(Self::Raw),
1 => Ok(Self::Zstd),
other => Err(CodecError::Unsupported { codec: other }),
}
}
#[must_use]
pub const fn as_u32(self) -> u32 {
self as u32
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum CodecError {
#[error("unsupported payload codec {codec}")]
Unsupported {
codec: u32,
},
#[error("raw codec: raw_byte_len {raw} != stored_byte_len {stored}")]
RawStoredMismatch {
raw: u64,
stored: u64,
},
#[error("decoded {decoded} bytes, index says raw_byte_len {raw}")]
DecodedLengthMismatch {
decoded: usize,
raw: u64,
},
#[error("zstd: {0}")]
Zstd(String),
}
pub fn encode(codec: PayloadCodec, raw: &[u8]) -> Result<Vec<u8>, CodecError> {
match codec {
PayloadCodec::Raw => Ok(raw.to_vec()),
PayloadCodec::Zstd => zstd::encode_all(raw, 0).map_err(|e| CodecError::Zstd(e.to_string())),
}
}
pub fn decode<'a>(
codec: PayloadCodec,
stored: &'a [u8],
raw_byte_len: u64,
) -> Result<Cow<'a, [u8]>, CodecError> {
match codec {
PayloadCodec::Raw => {
let stored_len = stored.len() as u64;
if stored_len != raw_byte_len {
return Err(CodecError::RawStoredMismatch {
raw: raw_byte_len,
stored: stored_len,
});
}
Ok(Cow::Borrowed(stored))
}
PayloadCodec::Zstd => {
let dec = zstd::decode_all(stored).map_err(|e| CodecError::Zstd(e.to_string()))?;
if dec.len() as u64 != raw_byte_len {
return Err(CodecError::DecodedLengthMismatch {
decoded: dec.len(),
raw: raw_byte_len,
});
}
Ok(Cow::Owned(dec))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn raw_round_trip_borrows() {
let raw = b"hello chunk";
let stored = encode(PayloadCodec::Raw, raw).unwrap();
assert_eq!(stored.as_slice(), raw);
let out = decode(PayloadCodec::Raw, &stored, raw.len() as u64).unwrap();
assert!(matches!(out, Cow::Borrowed(_)));
assert_eq!(&*out, raw);
}
#[test]
fn zstd_round_trip() {
let raw = b"compress me please........";
let stored = encode(PayloadCodec::Zstd, raw).unwrap();
assert_ne!(stored.as_slice(), raw);
let out = decode(PayloadCodec::Zstd, &stored, raw.len() as u64).unwrap();
assert_eq!(&*out, raw);
}
#[test]
fn raw_rejects_length_mismatch() {
let err = decode(PayloadCodec::Raw, b"ab", 3).unwrap_err();
assert_eq!(err, CodecError::RawStoredMismatch { raw: 3, stored: 2 });
}
#[test]
fn from_u32_rejects_unknown() {
assert!(matches!(
PayloadCodec::from_u32(99),
Err(CodecError::Unsupported { codec: 99 })
));
}
}