#[derive(Debug, Clone, PartialEq, Eq)]
pub enum DecodeError {
TooShort,
BadMagic,
WrongFamily,
WrongResultWidth,
NonZeroReserved,
BadGeometry,
LengthMismatch,
BadSeed,
BadChecksum,
}
impl core::fmt::Display for DecodeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(match self {
DecodeError::TooShort => "buffer too short for a valid filter",
DecodeError::BadMagic => "not a pleat filter buffer (bad magic)",
DecodeError::WrongFamily => "filter family does not match",
DecodeError::WrongResultWidth => "result width does not match",
DecodeError::NonZeroReserved => "reserved header bytes must be zero",
DecodeError::BadGeometry => "inconsistent filter geometry",
DecodeError::LengthMismatch => "payload length does not match declared segments",
DecodeError::BadSeed => "seed out of range",
DecodeError::BadChecksum => "checksum mismatch",
})
}
}
impl std::error::Error for DecodeError {}
pub(crate) const MAGIC: [u8; 4] = *b"PLT1";
pub(crate) const HEADER_LEN: usize = 32;
pub(crate) const CHECKSUM_LEN: usize = 8;
pub(crate) const FAMILY_HOMOG: u8 = 0;
pub(crate) const FAMILY_STD: u8 = 1;
pub(crate) fn fnv1a(bytes: &[u8]) -> u64 {
let mut h = 0xcbf2_9ce4_8422_2325u64;
for &b in bytes {
h = (h ^ b as u64).wrapping_mul(0x0000_0100_0000_01b3);
}
h
}
#[derive(Debug)]
pub(crate) struct Header {
pub seed: u64,
pub num_starts: u64,
}
pub(crate) fn write_header(
family: u8,
r: u8,
seed: u64,
num_starts: u64,
segment_count: u64,
) -> Vec<u8> {
let mut out = Vec::with_capacity(HEADER_LEN);
out.extend_from_slice(&MAGIC);
out.push(family);
out.push(r);
out.extend_from_slice(&[0u8, 0u8]); out.extend_from_slice(&seed.to_le_bytes());
out.extend_from_slice(&num_starts.to_le_bytes());
out.extend_from_slice(&segment_count.to_le_bytes());
debug_assert_eq!(out.len(), HEADER_LEN);
out
}
pub(crate) fn finish(mut buf: Vec<u8>) -> Vec<u8> {
let sum = fnv1a(&buf);
buf.extend_from_slice(&sum.to_le_bytes());
buf
}
pub(crate) fn decode(
bytes: &[u8],
expected_family: u8,
expected_r: u8,
elem_size: usize,
w: usize,
) -> Result<(Header, &[u8]), DecodeError> {
if bytes.len() < HEADER_LEN + CHECKSUM_LEN {
return Err(DecodeError::TooShort);
}
if bytes[0..4] != MAGIC {
return Err(DecodeError::BadMagic);
}
let checksum_at = bytes.len() - CHECKSUM_LEN;
let stored = u64::from_le_bytes(bytes[checksum_at..].try_into().unwrap());
if fnv1a(&bytes[..checksum_at]) != stored {
return Err(DecodeError::BadChecksum);
}
let family = bytes[4];
if family != expected_family {
return Err(DecodeError::WrongFamily);
}
let r = bytes[5];
if r != expected_r {
return Err(DecodeError::WrongResultWidth);
}
if bytes[6..8] != [0, 0] {
return Err(DecodeError::NonZeroReserved);
}
let seed = u64::from_le_bytes(bytes[8..16].try_into().unwrap());
let num_starts = u64::from_le_bytes(bytes[16..24].try_into().unwrap());
let segment_count = u64::from_le_bytes(bytes[24..32].try_into().unwrap());
let w = w as u64;
let num_slots = num_starts
.checked_add(w - 1)
.ok_or(DecodeError::BadGeometry)?;
if num_starts == 0 || num_slots % w != 0 || num_slots < 2 * w {
return Err(DecodeError::BadGeometry);
}
usize::try_from(num_slots).map_err(|_| DecodeError::BadGeometry)?;
let num_blocks = num_slots / w;
let expect_segments = num_blocks
.checked_mul(r as u64)
.ok_or(DecodeError::BadGeometry)?;
if segment_count != expect_segments {
return Err(DecodeError::BadGeometry);
}
let segment_count = usize::try_from(segment_count).map_err(|_| DecodeError::BadGeometry)?;
let payload_len = segment_count
.checked_mul(elem_size)
.ok_or(DecodeError::BadGeometry)?;
let total_len = HEADER_LEN
.checked_add(payload_len)
.and_then(|n| n.checked_add(CHECKSUM_LEN))
.ok_or(DecodeError::BadGeometry)?;
if total_len != bytes.len() {
return Err(DecodeError::LengthMismatch);
}
Ok((
Header { seed, num_starts },
&bytes[HEADER_LEN..HEADER_LEN + payload_len],
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_short_bad_magic_and_corrupt() {
assert_eq!(
decode(&[], FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::TooShort
);
let mut buf = write_header(FAMILY_HOMOG, 7, 0, 129, 21);
buf.extend_from_slice(&[0u8; 21 * 8]);
let good = finish(buf);
assert!(decode(&good, FAMILY_HOMOG, 7, 8, 64).is_ok());
assert_eq!(
decode(&good, FAMILY_STD, 7, 8, 64).unwrap_err(),
DecodeError::WrongFamily
);
assert_eq!(
decode(&good, FAMILY_HOMOG, 8, 8, 64).unwrap_err(),
DecodeError::WrongResultWidth
);
let mut reserved = good.clone();
reserved[6] = 1;
let checksum_at = reserved.len() - CHECKSUM_LEN;
let checksum = fnv1a(&reserved[..checksum_at]);
reserved[checksum_at..].copy_from_slice(&checksum.to_le_bytes());
assert_eq!(
decode(&reserved, FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::NonZeroReserved
);
let mut bad = good.clone();
bad[HEADER_LEN] ^= 0xFF;
assert_eq!(
decode(&bad, FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::BadChecksum
);
let mut nm = good.clone();
nm[0] = b'X';
assert_eq!(
decode(&nm, FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::BadMagic
);
}
#[test]
fn rejects_inconsistent_geometry() {
let mut buf = write_header(FAMILY_HOMOG, 7, 0, 129, 8);
buf.extend_from_slice(&[0u8; 8 * 8]);
let b = finish(buf);
assert_eq!(
decode(&b, FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::BadGeometry
);
}
#[cfg(target_pointer_width = "32")]
#[test]
fn rejects_geometry_that_cannot_fit_usize() {
let num_starts = (1u64 << 38) + 1;
let segment_count = 7 * ((1u64 << 32) + 1);
let mut buf = write_header(FAMILY_HOMOG, 7, 0, num_starts, segment_count);
buf.extend_from_slice(&[0u8; 7 * 8]);
let b = finish(buf);
assert_eq!(
decode(&b, FAMILY_HOMOG, 7, 8, 64).unwrap_err(),
DecodeError::BadGeometry
);
}
}