use alloc::vec::Vec;
use crate::zerocopy::{Cursor, OutOfSpace, Reader};
pub type Lsn = u64;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum RecordKind {
PageWrite = 1,
Commit = 2,
Checkpoint = 3,
}
impl RecordKind {
fn from_u8(v: u8) -> Option<Self> {
match v {
1 => Some(RecordKind::PageWrite),
2 => Some(RecordKind::Commit),
3 => Some(RecordKind::Checkpoint),
_ => None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct WalRecord {
pub lsn: Lsn,
pub kind: RecordKind,
pub block_id: u64,
pub payload: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WalError {
OutOfSpace,
Malformed,
BadChecksum,
UnknownKind(u8),
}
impl From<OutOfSpace> for WalError {
fn from(_: OutOfSpace) -> Self {
WalError::OutOfSpace
}
}
const HEADER_LEN: usize = 8 + 1 + 8 + 4;
fn crc32(data: &[u8]) -> u32 {
let mut crc: u32 = 0xFFFF_FFFF;
for &b in data {
crc ^= b as u32;
for _ in 0..8 {
let mask = (crc & 1).wrapping_neg();
crc = (crc >> 1) ^ (0xEDB8_8320 & mask);
}
}
!crc
}
impl WalRecord {
pub fn encoded_len(&self) -> usize {
4 + HEADER_LEN + self.payload.len() + 4
}
pub fn encode(&self, out: &mut [u8]) -> Result<usize, WalError> {
let body_len = HEADER_LEN + self.payload.len() + 4; let end_before_crc;
{
let mut c = Cursor::new(out);
c.put_u32(body_len as u32)?;
c.put_u64(self.lsn)?;
c.put_u8(self.kind as u8)?;
c.put_u64(self.block_id)?;
c.put_u32(self.payload.len() as u32)?;
c.put_bytes(&self.payload)?;
end_before_crc = c.position();
}
let crc = crc32(&out[4..end_before_crc]);
let mut c = Cursor::new(&mut out[end_before_crc..]);
c.put_u32(crc)?;
Ok(end_before_crc + 4)
}
pub fn decode(data: &[u8]) -> Result<(WalRecord, usize), WalError> {
let mut r = Reader::new(data);
let body_len = r.read_u32().map_err(|_| WalError::Malformed)? as usize;
if r.remaining() < body_len {
return Err(WalError::Malformed);
}
let body_start = r.position();
let lsn = r.read_u64().map_err(|_| WalError::Malformed)?;
let kind_tag = r.read_u8().map_err(|_| WalError::Malformed)?;
let kind = RecordKind::from_u8(kind_tag).ok_or(WalError::UnknownKind(kind_tag))?;
let block_id = r.read_u64().map_err(|_| WalError::Malformed)?;
let payload_len = r.read_u32().map_err(|_| WalError::Malformed)? as usize;
let payload = r
.read_bytes(payload_len)
.map_err(|_| WalError::Malformed)?
.to_vec();
let stored_crc = r.read_u32().map_err(|_| WalError::Malformed)?;
let body_end = r.position();
let crc_region_end = body_end - 4;
if body_end - body_start != body_len {
return Err(WalError::Malformed);
}
let computed = crc32(&data[body_start..crc_region_end]);
if computed != stored_crc {
return Err(WalError::BadChecksum);
}
Ok((
WalRecord {
lsn,
kind,
block_id,
payload,
},
4 + body_len,
))
}
}
#[derive(Debug, Default, Clone)]
pub struct Wal {
bytes: Vec<u8>,
next_lsn: Lsn,
}
impl Wal {
pub fn new() -> Self {
Self {
bytes: Vec::new(),
next_lsn: 0,
}
}
pub fn next_lsn(&self) -> Lsn {
self.next_lsn
}
pub fn as_bytes(&self) -> &[u8] {
&self.bytes
}
pub fn append(&mut self, kind: RecordKind, block_id: u64, payload: &[u8]) -> Lsn {
let lsn = self.next_lsn;
let rec = WalRecord {
lsn,
kind,
block_id,
payload: payload.to_vec(),
};
let mut frame = alloc::vec![0u8; rec.encoded_len()];
let n = rec.encode(&mut frame).expect("sized buffer");
self.bytes.extend_from_slice(&frame[..n]);
self.next_lsn += 1;
lsn
}
pub fn replay<F: FnMut(&WalRecord)>(&self, mut f: F) -> usize {
let mut offset = 0;
let mut count = 0;
while offset < self.bytes.len() {
match WalRecord::decode(&self.bytes[offset..]) {
Ok((rec, consumed)) => {
f(&rec);
offset += consumed;
count += 1;
}
Err(_) => break, }
}
count
}
pub fn from_bytes(bytes: &[u8]) -> Self {
let mut offset = 0;
let mut last_lsn: Option<Lsn> = None;
while offset < bytes.len() {
match WalRecord::decode(&bytes[offset..]) {
Ok((rec, consumed)) => {
last_lsn = Some(rec.lsn);
offset += consumed;
}
Err(_) => break,
}
}
Self {
bytes: bytes[..offset].to_vec(),
next_lsn: last_lsn.map(|l| l + 1).unwrap_or(0),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn record_round_trips() {
let rec = WalRecord {
lsn: 7,
kind: RecordKind::PageWrite,
block_id: 42,
payload: alloc::vec![1, 2, 3, 4, 5],
};
let mut buf = alloc::vec![0u8; rec.encoded_len()];
let n = rec.encode(&mut buf).unwrap();
let (decoded, consumed) = WalRecord::decode(&buf[..n]).unwrap();
assert_eq!(decoded, rec);
assert_eq!(consumed, n);
}
#[test]
fn append_assigns_monotonic_lsns_and_replays_in_order() {
let mut wal = Wal::new();
assert_eq!(wal.append(RecordKind::PageWrite, 1, b"a"), 0);
assert_eq!(wal.append(RecordKind::PageWrite, 2, b"bb"), 1);
assert_eq!(wal.append(RecordKind::Commit, 0, b""), 2);
let mut seen = Vec::new();
let n = wal.replay(|r| seen.push((r.lsn, r.kind, r.block_id)));
assert_eq!(n, 3);
assert_eq!(seen[0], (0, RecordKind::PageWrite, 1));
assert_eq!(seen[1], (1, RecordKind::PageWrite, 2));
assert_eq!(seen[2], (2, RecordKind::Commit, 0));
}
#[test]
fn torn_tail_write_is_truncated_on_replay() {
let mut wal = Wal::new();
wal.append(RecordKind::PageWrite, 1, b"good");
wal.append(RecordKind::PageWrite, 2, b"also-good");
wal.append(RecordKind::PageWrite, 3, b"torn");
let mut raw = wal.as_bytes().to_vec();
let len = raw.len();
for b in raw.iter_mut().skip(len - 3) {
*b ^= 0xFF;
}
let recovered = Wal::from_bytes(&raw);
let mut lsns = Vec::new();
let n = recovered.replay(|r| lsns.push(r.lsn));
assert_eq!(n, 2, "only the two fully-durable records survive");
assert_eq!(lsns, alloc::vec![0, 1]);
assert_eq!(recovered.next_lsn(), 2);
}
#[test]
fn bad_checksum_is_detected() {
let rec = WalRecord {
lsn: 0,
kind: RecordKind::PageWrite,
block_id: 1,
payload: alloc::vec![9, 9, 9],
};
let mut buf = alloc::vec![0u8; rec.encoded_len()];
rec.encode(&mut buf).unwrap();
buf[HEADER_LEN + 4 + 1] ^= 0x01;
assert_eq!(WalRecord::decode(&buf), Err(WalError::BadChecksum));
}
#[test]
fn mismatched_body_len_is_malformed() {
let rec = WalRecord {
lsn: 0,
kind: RecordKind::PageWrite,
block_id: 1,
payload: alloc::vec![9, 9, 9],
};
let mut buf = alloc::vec![0u8; rec.encoded_len()];
rec.encode(&mut buf).unwrap();
buf[0..4].copy_from_slice(&999u32.to_le_bytes());
assert_eq!(WalRecord::decode(&buf), Err(WalError::Malformed));
}
#[test]
fn empty_log_replays_nothing() {
let wal = Wal::new();
let n = wal.replay(|_| panic!("no records expected"));
assert_eq!(n, 0);
assert_eq!(wal.next_lsn(), 0);
}
}