use crate::{align_up, get_u8, get_u16, get_u32, get_u64, put_u8, put_u16, put_u32, put_u64};
use yo_common::{Code, Error, Result, crc32c};
pub const HEADER_LEN: usize = 16;
pub const HEADER_LEN_TTL: usize = 24;
pub const TRAILER_LEN: usize = 4;
pub const MAX_KEY_LEN: usize = u16::MAX as usize;
pub mod record_flags {
pub const TIERED: u8 = 1 << 0;
pub const COMPRESSED: u8 = 1 << 1;
pub const HAS_TTL: u8 = 1 << 2;
pub const SHAPE_TAGGED: u8 = 1 << 3;
pub const CHECKSUMMED: u8 = 1 << 4;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum RecordKind {
String = 0,
CollectionChunk = 1,
Document = 2,
Vector = 3,
GraphNode = 4,
GraphAdj = 5,
Checkpoint = 6,
Tombstone = 7,
IndexDelta = 8,
}
impl RecordKind {
pub const ALL: [RecordKind; 9] = [
RecordKind::String,
RecordKind::CollectionChunk,
RecordKind::Document,
RecordKind::Vector,
RecordKind::GraphNode,
RecordKind::GraphAdj,
RecordKind::Checkpoint,
RecordKind::Tombstone,
RecordKind::IndexDelta,
];
#[must_use]
pub const fn from_u8(b: u8) -> Option<RecordKind> {
match b {
0 => Some(RecordKind::String),
1 => Some(RecordKind::CollectionChunk),
2 => Some(RecordKind::Document),
3 => Some(RecordKind::Vector),
4 => Some(RecordKind::GraphNode),
5 => Some(RecordKind::GraphAdj),
6 => Some(RecordKind::Checkpoint),
7 => Some(RecordKind::Tombstone),
8 => Some(RecordKind::IndexDelta),
_ => None,
}
}
#[must_use]
pub const fn as_u8(self) -> u8 {
self as u8
}
#[must_use]
pub const fn carries_a_key(self) -> bool {
!matches!(self, RecordKind::CollectionChunk)
}
}
pub fn total_len(flags: u8, klen: usize, vlen: usize) -> Result<usize> {
if klen > MAX_KEY_LEN {
return Err(
Error::new(Code::Invalid, "the key is longer than 65535 bytes")
.with_detail(format!("klen={klen}")),
);
}
let n = header_len(flags) + klen + vlen + trailer_len(flags);
if n > u32::MAX as usize {
return Err(Error::new(
Code::Invalid,
"the record does not fit in a u32",
));
}
Ok(n)
}
#[inline]
#[must_use]
pub const fn header_len(flags: u8) -> usize {
if flags & record_flags::HAS_TTL != 0 {
HEADER_LEN_TTL
} else {
HEADER_LEN
}
}
#[inline]
#[must_use]
pub const fn trailer_len(flags: u8) -> usize {
if flags & record_flags::CHECKSUMMED != 0 {
TRAILER_LEN
} else {
0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecordHeader {
pub kind: u8,
pub flags: u8,
pub prev: u64,
pub ttl_ms: u64,
}
impl RecordHeader {
#[must_use]
pub const fn new(kind: RecordKind) -> RecordHeader {
RecordHeader {
kind: kind.as_u8(),
flags: record_flags::CHECKSUMMED,
prev: 0,
ttl_ms: 0,
}
}
#[must_use]
pub const fn with_ttl(mut self, unix_ms: u64) -> RecordHeader {
self.flags |= record_flags::HAS_TTL;
self.ttl_ms = unix_ms;
self
}
#[must_use]
pub const fn after(mut self, prev: u64) -> RecordHeader {
self.prev = prev;
self
}
pub fn fill(&self, buf: &mut [u8], key: &[u8], value: &[u8]) -> Result<usize> {
let flags = self.flags | record_flags::CHECKSUMMED;
let n = total_len(flags, key.len(), value.len())?;
if buf.len() < n {
return Err(
Error::new(Code::Full, "the record does not fit in the buffer")
.with_detail(format!("need={n} have={}", buf.len())),
);
}
let h = header_len(flags);
put_u8(buf, 4, self.kind);
put_u8(buf, 5, flags);
put_u16(buf, 6, key.len() as u16);
put_u64(buf, 8, self.prev);
if flags & record_flags::HAS_TTL != 0 {
put_u64(buf, 16, self.ttl_ms);
}
buf[h..h + key.len()].copy_from_slice(key);
let v = h + key.len();
buf[v..v + value.len()].copy_from_slice(value);
let c = crc32c(0, &(n as u32).to_le_bytes());
let c = crc32c(c, &buf[4..n - TRAILER_LEN]);
put_u32(buf, n - TRAILER_LEN, c);
Ok(n)
}
}
pub fn seal_len(buf: &mut [u8], len: usize) {
assert!(len != 0, "zero is the end of log sentinel, not a length");
put_u32(buf, 0, len as u32);
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecordRef<'a> {
pub len: u32,
pub kind: u8,
pub flags: u8,
pub prev: u64,
pub ttl_ms: Option<u64>,
pub key: &'a [u8],
pub value: &'a [u8],
}
impl<'a> RecordRef<'a> {
pub fn parse(bytes: &'a [u8]) -> Result<Option<RecordRef<'a>>> {
if bytes.len() < 4 {
return Ok(None);
}
let len = get_u32(bytes, 0) as usize;
if len == 0 {
return Ok(None);
}
let flags = get_u8(bytes, 5);
if flags & record_flags::CHECKSUMMED == 0 {
return Err(
Error::new(Code::Corrupt, "a record with its checksum flag clear")
.with_detail(format!("flags={flags:#04x}")),
);
}
let h = header_len(flags);
let t = trailer_len(flags);
let klen = get_u16(bytes, 6) as usize;
if len < h + klen + t {
return Err(
Error::new(Code::Corrupt, "the record is shorter than its own header")
.with_detail(format!("len={len} header={h} klen={klen} trailer={t}")),
);
}
if len > bytes.len() {
return Err(
Error::new(Code::Corrupt, "the record runs past the end of the page")
.with_detail(format!("len={len} available={}", bytes.len())),
);
}
if flags & record_flags::CHECKSUMMED != 0 {
let want = get_u32(bytes, len - TRAILER_LEN);
let got = crc32c(0, &bytes[..len - TRAILER_LEN]);
if want != got {
return Err(Error::new(Code::Corrupt, "record checksum mismatch")
.with_detail(format!("stored={want:#010x} computed={got:#010x}")));
}
}
let ttl_ms = if flags & record_flags::HAS_TTL != 0 {
Some(get_u64(bytes, 16))
} else {
None
};
Ok(Some(RecordRef {
len: len as u32,
kind: get_u8(bytes, 4),
flags,
prev: get_u64(bytes, 8),
ttl_ms,
key: &bytes[h..h + klen],
value: &bytes[h + klen..len - t],
}))
}
#[must_use]
pub fn stride(&self) -> usize {
align_up(self.len as usize)
}
#[must_use]
pub fn kind(&self) -> Option<RecordKind> {
RecordKind::from_u8(self.kind)
}
#[must_use]
pub fn is_tombstone(&self) -> bool {
self.kind == RecordKind::Tombstone.as_u8()
}
#[must_use]
pub fn is_tiered(&self) -> bool {
self.flags & record_flags::TIERED != 0
}
}
pub struct RecordIter<'a> {
bytes: &'a [u8],
at: usize,
done: bool,
}
impl<'a> RecordIter<'a> {
#[must_use]
pub const fn new(bytes: &'a [u8]) -> RecordIter<'a> {
RecordIter {
bytes,
at: 0,
done: false,
}
}
#[must_use]
pub const fn offset(&self) -> usize {
self.at
}
}
impl<'a> Iterator for RecordIter<'a> {
type Item = Result<RecordRef<'a>>;
fn next(&mut self) -> Option<Self::Item> {
if self.done || self.at >= self.bytes.len() {
return None;
}
match RecordRef::parse(&self.bytes[self.at..]) {
Ok(Some(r)) => {
self.at += r.stride();
Some(Ok(r))
}
Ok(None) => {
self.done = true;
None
}
Err(e) => {
self.done = true;
Some(Err(e))
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::RECORD_ALIGN;
fn write(h: RecordHeader, key: &[u8], value: &[u8]) -> Vec<u8> {
let mut buf = vec![0u8; 4096];
let n = h.fill(&mut buf, key, value).unwrap();
seal_len(&mut buf, n);
buf.truncate(align_up(n));
buf
}
#[test]
fn a_record_round_trips() {
let h = RecordHeader::new(RecordKind::String).after(4096);
let buf = write(h, b"greeting", b"hello");
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert_eq!(r.kind(), Some(RecordKind::String));
assert_eq!(r.key, b"greeting");
assert_eq!(r.value, b"hello");
assert_eq!(r.prev, 4096);
assert_eq!(r.ttl_ms, None);
assert!(!r.is_tombstone());
}
#[test]
fn every_field_lands_where_the_specification_says() {
let h = RecordHeader::new(RecordKind::Document)
.with_ttl(1_700_000_000_000)
.after(0x1122_3344_5566_7788);
let buf = write(h, b"k", b"v");
assert_eq!(get_u32(&buf, 0) as usize, HEADER_LEN_TTL + 1 + 1 + 4);
assert_eq!(get_u8(&buf, 4), 2, "document is kind 2");
assert_eq!(
get_u8(&buf, 5),
record_flags::CHECKSUMMED | record_flags::HAS_TTL
);
assert_eq!(get_u16(&buf, 6), 1);
assert_eq!(get_u64(&buf, 8), 0x1122_3344_5566_7788);
assert_eq!(get_u64(&buf, 16), 1_700_000_000_000);
assert_eq!(buf[24], b'k');
assert_eq!(buf[25], b'v');
}
#[test]
fn a_ttl_costs_eight_bytes_and_moves_the_key() {
let plain = write(RecordHeader::new(RecordKind::String), b"key", b"value");
let ttl = write(
RecordHeader::new(RecordKind::String).with_ttl(1),
b"key",
b"value",
);
assert_eq!(get_u32(&ttl, 0) - get_u32(&plain, 0), 8);
let r = RecordRef::parse(&ttl).unwrap().unwrap();
assert_eq!(r.ttl_ms, Some(1));
assert_eq!(r.key, b"key");
assert_eq!(r.value, b"value");
}
#[test]
fn a_zero_length_is_the_end_of_the_log_and_not_an_error() {
assert!(RecordRef::parse(&[0u8; 64]).unwrap().is_none());
assert!(RecordRef::parse(&[]).unwrap().is_none());
assert!(RecordRef::parse(&[1, 2, 3]).unwrap().is_none());
}
#[test]
fn len_is_exact_so_a_value_of_any_length_survives() {
for n in 0..=32usize {
let value: Vec<u8> = (0..n).map(|i| i as u8).collect();
let buf = write(RecordHeader::new(RecordKind::String), b"k", &value);
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert_eq!(r.value, &value[..], "value of {n} bytes came back wrong");
assert_eq!(r.value.len(), n);
}
}
#[test]
fn the_stride_is_aligned_even_when_the_length_is_not() {
let buf = write(RecordHeader::new(RecordKind::String), b"k", b"abc");
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert_eq!(r.len as usize, HEADER_LEN + 1 + 3 + TRAILER_LEN);
assert_eq!(r.len % RECORD_ALIGN as u32, 0);
let buf = write(RecordHeader::new(RecordKind::String), b"k", b"ab");
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert_eq!(r.len as usize, 23);
assert_eq!(r.stride(), 24, "the padding is between records, not inside");
}
#[test]
fn walking_a_page_by_stride_stays_in_step() {
let mut page = vec![0u8; 4096];
let mut at = 0usize;
let mut written = Vec::new();
for i in 0..40usize {
let key = format!("key{i}");
let value = vec![b'v'; i];
let h = RecordHeader::new(RecordKind::String).after(at as u64);
let n = h.fill(&mut page[at..], key.as_bytes(), &value).unwrap();
seal_len(&mut page[at..], n);
written.push((key, value));
at += align_up(n);
}
let got: Vec<_> = RecordIter::new(&page).map(|r| r.unwrap()).collect();
assert_eq!(got.len(), 40);
for (r, (key, value)) in got.iter().zip(&written) {
assert_eq!(r.key, key.as_bytes());
assert_eq!(r.value, &value[..]);
}
}
#[test]
fn the_iterator_reports_where_it_stopped() {
let mut page = vec![0u8; 512];
let h = RecordHeader::new(RecordKind::String);
let n = h.fill(&mut page, b"a", b"bb").unwrap();
seal_len(&mut page, n);
let mut it = RecordIter::new(&page);
assert!(it.next().is_some());
assert!(it.next().is_none());
assert_eq!(it.offset(), align_up(n), "the tail is here");
}
#[test]
fn a_flipped_bit_anywhere_in_a_checksummed_record_is_caught() {
let good = write(
RecordHeader::new(RecordKind::String).with_ttl(99),
b"the key",
b"the value, which is long enough to be worth checking",
);
let len = get_u32(&good, 0) as usize;
for i in 0..len {
let mut bad = good.clone();
bad[i] ^= 0x20;
let r = RecordRef::parse(&bad);
match r {
Err(_) => {}
Ok(None) => {}
Ok(Some(rec)) => panic!("byte {i} was not caught, got {rec:?}"),
}
}
}
#[test]
fn a_header_written_without_the_checksum_flag_gets_one_anyway() {
let h = RecordHeader {
kind: RecordKind::String.as_u8(),
flags: 0,
prev: 0,
ttl_ms: 0,
};
let buf = write(h, b"k", b"v");
assert_eq!(get_u32(&buf, 0) as usize, HEADER_LEN + 2 + TRAILER_LEN);
assert_ne!(get_u8(&buf, 5) & record_flags::CHECKSUMMED, 0);
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert_eq!(r.value, b"v");
}
#[test]
fn a_record_with_its_checksum_flag_cleared_is_corruption() {
let mut buf = write(RecordHeader::new(RecordKind::String), b"key", b"value");
buf[5] &= !record_flags::CHECKSUMMED;
let err = RecordRef::parse(&buf).unwrap_err();
assert_eq!(err.code(), Code::Corrupt);
}
#[test]
fn a_length_that_is_shorter_than_the_header_is_corruption() {
let mut buf = write(RecordHeader::new(RecordKind::String), b"key", b"value");
put_u32(&mut buf, 0, 12);
let err = RecordRef::parse(&buf).unwrap_err();
assert_eq!(err.code(), Code::Corrupt);
assert!(err.detail().unwrap().contains("len=12"));
}
#[test]
fn a_record_cut_off_by_a_torn_write_is_corruption() {
let buf = write(RecordHeader::new(RecordKind::String), b"key", b"value");
let err = RecordRef::parse(&buf[..8]).unwrap_err();
assert_eq!(err.code(), Code::Corrupt);
assert!(err.detail().unwrap().contains("available=8"));
}
#[test]
fn an_unknown_kind_is_skipped_rather_than_refused() {
let h = RecordHeader {
kind: 200,
flags: record_flags::CHECKSUMMED,
prev: 0,
ttl_ms: 0,
};
let mut page = vec![0u8; 512];
let n = h.fill(&mut page, b"future", b"stuff").unwrap();
seal_len(&mut page, n);
let after = align_up(n);
let m = RecordHeader::new(RecordKind::String)
.fill(&mut page[after..], b"k", b"v")
.unwrap();
seal_len(&mut page[after..], m);
let got: Vec<_> = RecordIter::new(&page).map(|r| r.unwrap()).collect();
assert_eq!(got.len(), 2);
assert_eq!(got[0].kind(), None, "not a kind this version knows");
assert_eq!(got[0].key, b"future");
assert_eq!(got[1].kind(), Some(RecordKind::String));
}
#[test]
fn a_key_larger_than_a_u16_is_refused_rather_than_truncated() {
let key = vec![b'k'; MAX_KEY_LEN + 1];
let mut buf = vec![0u8; MAX_KEY_LEN + 64];
let err = RecordHeader::new(RecordKind::String)
.fill(&mut buf, &key, b"v")
.unwrap_err();
assert_eq!(err.code(), Code::Invalid);
}
#[test]
fn a_buffer_with_no_room_says_how_much_it_needed() {
let mut buf = [0u8; 8];
let err = RecordHeader::new(RecordKind::String)
.fill(&mut buf, b"key", b"value")
.unwrap_err();
assert_eq!(err.code(), Code::Full);
assert!(err.detail().unwrap().contains("have=8"));
}
#[test]
fn kinds_round_trip_and_chunks_have_no_key() {
for k in RecordKind::ALL {
assert_eq!(RecordKind::from_u8(k.as_u8()), Some(k));
}
assert_eq!(RecordKind::from_u8(9), None);
assert!(!RecordKind::CollectionChunk.carries_a_key());
assert!(RecordKind::String.carries_a_key());
assert!(RecordKind::Tombstone.carries_a_key());
}
#[test]
fn a_tombstone_is_a_record_with_no_value() {
let buf = write(RecordHeader::new(RecordKind::Tombstone), b"gone", b"");
let r = RecordRef::parse(&buf).unwrap().unwrap();
assert!(r.is_tombstone());
assert_eq!(r.value, b"");
assert_eq!(r.key, b"gone");
}
}