use crate::error::{corrupt, DbResult};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum WalByteOrder {
Big,
Little,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq, Default)]
pub struct WalChecksum {
pub s0: u32,
pub s1: u32,
}
impl WalChecksum {
pub fn new(s0: u32, s1: u32) -> WalChecksum {
WalChecksum { s0, s1 }
}
pub fn update(&mut self, data: &[u8], order: WalByteOrder) -> DbResult<()> {
if !data.len().is_multiple_of(8) {
return Err(corrupt(
"WAL checksum input is not a whole number of 8-byte blocks",
));
}
let (blocks, _) = data.as_chunks::<8>();
for block in blocks {
let (first, second) = split_block(block, order);
self.s0 = self.s0.wrapping_add(first).wrapping_add(self.s1);
self.s1 = self.s1.wrapping_add(second).wrapping_add(self.s0);
}
Ok(())
}
pub fn extended(self, data: &[u8], order: WalByteOrder) -> DbResult<WalChecksum> {
let mut next = self;
next.update(data, order)?;
Ok(next)
}
}
fn split_block(block: &[u8; 8], order: WalByteOrder) -> (u32, u32) {
let word = |bytes: &[u8]| -> u32 {
let mut value = [0u8; 4];
for (slot, byte) in value.iter_mut().zip(bytes.iter()) {
*slot = *byte;
}
match order {
WalByteOrder::Big => u32::from_be_bytes(value),
WalByteOrder::Little => u32::from_le_bytes(value),
}
};
(
word(block.get(..4).unwrap_or(&[])),
word(block.get(4..).unwrap_or(&[])),
)
}
const CRC32_TABLE: [u32; 256] = build_crc32_table();
#[allow(clippy::arithmetic_side_effects, clippy::indexing_slicing)]
const fn build_crc32_table() -> [u32; 256] {
let mut table = [0u32; 256];
let mut index = 0usize;
while index < 256 {
let mut value = index as u32;
let mut bit = 0;
while bit < 8 {
value = if value & 1 == 1 {
0xedb8_8320 ^ (value >> 1)
} else {
value >> 1
};
bit += 1;
}
table[index] = value;
index += 1;
}
table
}
const CRC32_SLICES: [[u32; 256]; 7] = build_crc32_slices();
#[allow(clippy::arithmetic_side_effects, clippy::indexing_slicing)]
const fn build_crc32_slices() -> [[u32; 256]; 7] {
let mut slices = [[0u32; 256]; 7];
let mut index = 0usize;
while index < 256 {
let mut previous = CRC32_TABLE[index];
let mut level = 0usize;
while level < 7 {
previous = (previous >> 8) ^ CRC32_TABLE[(previous & 0xff) as usize];
slices[level][index] = previous;
level += 1;
}
index += 1;
}
slices
}
fn slice(level: usize, index: u32) -> u32 {
CRC32_SLICES
.get(level)
.and_then(|table| table.get((index & 0xff) as usize))
.copied()
.unwrap_or(0)
}
fn byte_entry(index: u32) -> u32 {
CRC32_TABLE
.get((index & 0xff) as usize)
.copied()
.unwrap_or(0)
}
pub fn crc32(data: &[u8]) -> u32 {
crc32_continue(0, data)
}
pub fn crc32_continue(previous: u32, data: &[u8]) -> u32 {
let mut crc = !previous;
let mut chunks = data.chunks_exact(8);
for chunk in &mut chunks {
let (low, high) = chunk.split_at(4);
let mut first = [0u8; 4];
let mut second = [0u8; 4];
first.copy_from_slice(low);
second.copy_from_slice(high);
let one = u32::from_le_bytes(first) ^ crc;
let two = u32::from_le_bytes(second);
crc = slice(6, one)
^ slice(5, one >> 8)
^ slice(4, one >> 16)
^ slice(3, one >> 24)
^ slice(2, two)
^ slice(1, two >> 8)
^ slice(0, two >> 16)
^ byte_entry(two >> 24);
}
for byte in chunks.remainder() {
crc = byte_entry(crc ^ u32::from(*byte)) ^ (crc >> 8);
}
!crc
}
#[cfg(test)]
fn crc32_one_byte_at_a_time(previous: u32, data: &[u8]) -> u32 {
let mut crc = !previous;
for byte in data {
crc = byte_entry(crc ^ u32::from(*byte)) ^ (crc >> 8);
}
!crc
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::PrimaryCode;
use crate::rng::Rng;
#[test]
fn crc32_matches_the_published_check_value() {
assert_eq!(crc32(b"123456789"), 0xcbf4_3926);
assert_eq!(crc32(b""), 0);
}
#[test]
fn crc32_is_the_same_whether_it_arrives_whole_or_in_pieces() {
let mut rng = Rng::new(0x1782_0003);
let mut data = vec![0u8; 4096];
rng.fill(&mut data);
let whole = crc32(&data);
let mut running = 0;
for chunk in data.chunks(97) {
running = crc32_continue(running, chunk);
}
assert_eq!(running, whole);
}
#[test]
fn crc32_matches_the_byte_at_a_time_form() {
let mut rng = Rng::new(0x1833_0001);
let mut data = vec![0u8; 8_192];
rng.fill(&mut data);
for length in 0..=16usize {
let piece = data.get(..length).unwrap_or(&[]);
assert_eq!(
crc32(piece),
crc32_one_byte_at_a_time(0, piece),
"length {length}"
);
}
assert_eq!(crc32(&data), crc32_one_byte_at_a_time(0, &data));
let mut running = 0;
let mut slow = 0;
for chunk in data.chunks(101) {
running = crc32_continue(running, chunk);
slow = crc32_one_byte_at_a_time(slow, chunk);
}
assert_eq!(running, slow);
}
#[test]
fn crc32_detects_single_bit_flips() {
let mut rng = Rng::new(0x1782_0004);
let mut data = vec![0u8; 512];
rng.fill(&mut data);
let baseline = crc32(&data);
let stride = if cfg!(miri) { 16 } else { 1 };
for index in (0..data.len()).step_by(stride) {
for bit in 0..8 {
data[index] ^= 1 << bit;
assert_ne!(
crc32(&data),
baseline,
"flip at {index}:{bit} was invisible"
);
data[index] ^= 1 << bit;
}
}
}
#[test]
fn wal_checksum_depends_on_the_declared_byte_order() {
let data = [1u8, 2, 3, 4, 5, 6, 7, 8];
let big = WalChecksum::default()
.extended(&data, WalByteOrder::Big)
.unwrap();
let little = WalChecksum::default()
.extended(&data, WalByteOrder::Little)
.unwrap();
assert_ne!(big, little);
assert_eq!(big, WalChecksum::new(0x0102_0304, 0x0608_0a0c));
}
#[test]
fn wal_checksum_refuses_a_partial_block() {
let error = WalChecksum::default()
.update(&[0u8; 5], WalByteOrder::Big)
.expect_err("five bytes is not a whole block");
assert_eq!(error.code(), PrimaryCode::Corrupt);
}
#[test]
fn wal_checksum_chains_across_calls() {
let mut rng = Rng::new(0x1782_0005);
let mut data = vec![0u8; 1024];
rng.fill(&mut data);
let whole = WalChecksum::default()
.extended(&data, WalByteOrder::Big)
.unwrap();
let mut running = WalChecksum::default();
running.update(&data[..512], WalByteOrder::Big).unwrap();
running.update(&data[512..], WalByteOrder::Big).unwrap();
assert_eq!(running, whole);
}
}