use std::collections::BTreeMap;
pub const MAGIC: &[u8; 4] = b"TSOD";
pub const VERSION: u8 = 1;
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum DenseRecordError {
#[error("dense record too short: {0} bytes")]
TooShort(usize),
#[error("bad magic: {0:?}")]
BadMagic([u8; 4]),
#[error("unsupported dense record version {0}")]
UnsupportedVersion(u8),
#[error("checksum mismatch: stored {stored:#010x}, computed {computed:#010x}")]
ChecksumMismatch { stored: u32, computed: u32 },
#[error("malformed dense record body")]
Malformed,
#[error("non-utf8 key in dense record")]
NonUtf8Key,
}
pub fn encode(map: &BTreeMap<String, u64>, cap: u64) -> Vec<u8> {
debug_assert!(map.len() <= u32::MAX as usize, "key_count exceeds u32");
let mut buf = Vec::with_capacity(4 + 1 + 8 + 4 + map.len() * 24 + 4);
buf.extend_from_slice(MAGIC);
buf.push(VERSION);
buf.extend_from_slice(&cap.to_le_bytes());
buf.extend_from_slice(&(map.len() as u32).to_le_bytes());
for (key, counter) in map {
let kb = key.as_bytes();
debug_assert!(kb.len() <= u16::MAX as usize, "key length exceeds u16");
buf.extend_from_slice(&(kb.len() as u16).to_le_bytes());
buf.extend_from_slice(kb);
buf.extend_from_slice(&counter.to_le_bytes());
}
let crc = crc32c::crc32c(&buf);
buf.extend_from_slice(&crc.to_le_bytes());
buf
}
pub fn decode(bytes: &[u8]) -> Result<(BTreeMap<String, u64>, u64), DenseRecordError> {
if bytes.len() < 21 {
return Err(DenseRecordError::TooShort(bytes.len()));
}
if &bytes[0..4] != MAGIC {
return Err(DenseRecordError::BadMagic([
bytes[0], bytes[1], bytes[2], bytes[3],
]));
}
if bytes[4] != VERSION {
return Err(DenseRecordError::UnsupportedVersion(bytes[4]));
}
let body = &bytes[..bytes.len() - 4];
let stored = u32::from_le_bytes(
bytes[bytes.len() - 4..]
.try_into()
.map_err(|_| DenseRecordError::Malformed)?,
);
let computed = crc32c::crc32c(body);
if stored != computed {
return Err(DenseRecordError::ChecksumMismatch { stored, computed });
}
let cap = u64::from_le_bytes(
bytes[5..13]
.try_into()
.map_err(|_| DenseRecordError::Malformed)?,
);
let key_count = u32::from_le_bytes(
bytes[13..17]
.try_into()
.map_err(|_| DenseRecordError::Malformed)?,
);
let mut pos = 17;
let mut map = BTreeMap::new();
let mut prev_key: Option<String> = None;
for _ in 0..key_count {
if pos + 2 > body.len() {
return Err(DenseRecordError::Malformed);
}
let klen = u16::from_le_bytes(
body[pos..pos + 2]
.try_into()
.map_err(|_| DenseRecordError::Malformed)?,
) as usize;
pos += 2;
if pos + klen + 8 > body.len() {
return Err(DenseRecordError::Malformed);
}
let key = std::str::from_utf8(&body[pos..pos + klen])
.map_err(|_| DenseRecordError::NonUtf8Key)?
.to_string();
pos += klen;
let counter = u64::from_le_bytes(
body[pos..pos + 8]
.try_into()
.map_err(|_| DenseRecordError::Malformed)?,
);
pos += 8;
if let Some(prev) = &prev_key {
if key <= *prev {
return Err(DenseRecordError::Malformed);
}
}
prev_key = Some(key.clone());
map.insert(key, counter);
}
if pos != body.len() {
return Err(DenseRecordError::Malformed);
}
Ok((map, cap))
}
#[cfg(test)]
mod tests {
use super::*;
fn sample() -> BTreeMap<String, u64> {
let mut m = BTreeMap::new();
m.insert("orders".to_string(), 42);
m.insert("users".to_string(), 7);
m
}
#[test]
fn round_trips() {
let m = sample();
let bytes = encode(&m, 10_000);
let (decoded, cap) = decode(&bytes).unwrap();
assert_eq!(decoded, m);
assert_eq!(cap, 10_000);
}
#[test]
fn empty_map_round_trips() {
let m = BTreeMap::new();
let bytes = encode(&m, 500);
let (decoded, cap) = decode(&bytes).unwrap();
assert!(decoded.is_empty());
assert_eq!(cap, 500);
}
#[test]
fn detects_corruption() {
let mut bytes = encode(&sample(), 10_000);
let last = bytes.len() - 1;
bytes[last] ^= 0xff; assert!(matches!(
decode(&bytes),
Err(DenseRecordError::ChecksumMismatch { .. })
));
}
#[test]
fn rejects_bad_magic() {
let mut bytes = encode(&sample(), 1);
bytes[0] = b'X';
assert!(matches!(decode(&bytes), Err(DenseRecordError::BadMagic(_))));
}
#[test]
fn rejects_unsupported_version() {
let mut bytes = encode(&sample(), 1);
bytes[4] = 0; assert!(matches!(
decode(&bytes),
Err(DenseRecordError::UnsupportedVersion(0))
));
}
#[test]
fn rejects_non_utf8_key() {
let mut body = Vec::new();
body.extend_from_slice(MAGIC);
body.push(VERSION);
body.extend_from_slice(&1u64.to_le_bytes()); body.extend_from_slice(&1u32.to_le_bytes()); body.extend_from_slice(&2u16.to_le_bytes()); body.extend_from_slice(&[0xff, 0xfe]); body.extend_from_slice(&7u64.to_le_bytes()); let crc = crc32c::crc32c(&body);
body.extend_from_slice(&crc.to_le_bytes());
assert_eq!(decode(&body), Err(DenseRecordError::NonUtf8Key));
}
#[test]
fn rejects_short() {
assert!(matches!(
decode(&[0u8; 3]),
Err(DenseRecordError::TooShort(3))
));
}
fn encode_raw(cap: u64, entries: &[(&str, u64)]) -> Vec<u8> {
let mut body = Vec::new();
body.extend_from_slice(MAGIC);
body.push(VERSION);
body.extend_from_slice(&cap.to_le_bytes());
body.extend_from_slice(&(entries.len() as u32).to_le_bytes());
for (k, c) in entries {
body.extend_from_slice(&(k.len() as u16).to_le_bytes());
body.extend_from_slice(k.as_bytes());
body.extend_from_slice(&c.to_le_bytes());
}
let crc = crc32c::crc32c(&body);
body.extend_from_slice(&crc.to_le_bytes());
body
}
#[test]
fn rejects_duplicate_keys() {
let bytes = encode_raw(10_000, &[("orders", 1), ("orders", 2)]);
assert_eq!(decode(&bytes), Err(DenseRecordError::Malformed));
}
#[test]
fn rejects_out_of_order_keys() {
let bytes = encode_raw(10_000, &[("users", 7), ("orders", 42)]);
assert_eq!(decode(&bytes), Err(DenseRecordError::Malformed));
}
#[test]
fn rejects_truncated_key_entry() {
let mut body = Vec::new();
body.extend_from_slice(MAGIC);
body.push(VERSION);
body.extend_from_slice(&1u64.to_le_bytes()); body.extend_from_slice(&1u32.to_le_bytes()); body.extend_from_slice(&3u16.to_le_bytes()); body.extend_from_slice(b"abc"); let crc = crc32c::crc32c(&body);
body.extend_from_slice(&crc.to_le_bytes());
assert_eq!(decode(&body), Err(DenseRecordError::Malformed));
}
#[test]
fn accepts_canonical_ascending_keys() {
let bytes = encode_raw(10_000, &[("orders", 42), ("users", 7)]);
let (map, cap) = decode(&bytes).unwrap();
assert_eq!(cap, 10_000);
assert_eq!(map.get("orders"), Some(&42));
assert_eq!(map.get("users"), Some(&7));
assert_eq!(encode(&map, cap), bytes);
}
}