use std::collections::HashMap;
use crate::error::GpkgError;
const WAL_MAGIC_BE: u32 = 0x377f_0682;
const WAL_MAGIC_LE: u32 = 0x377f_0683;
const WAL_HEADER_SIZE: usize = 32;
const FRAME_HEADER_SIZE: usize = 24;
fn wal_checksum(data: &[u8], big_endian: bool, s0_init: u32, s1_init: u32) -> (u32, u32) {
debug_assert!(
data.len() % 8 == 0,
"checksum input must be a multiple of 8 bytes"
);
let mut s0 = s0_init;
let mut s1 = s1_init;
let mut i = 0;
while i + 8 <= data.len() {
let word_a = if big_endian {
u32::from_be_bytes([data[i], data[i + 1], data[i + 2], data[i + 3]])
} else {
u32::from_le_bytes([data[i], data[i + 1], data[i + 2], data[i + 3]])
};
let word_b = if big_endian {
u32::from_be_bytes([data[i + 4], data[i + 5], data[i + 6], data[i + 7]])
} else {
u32::from_le_bytes([data[i + 4], data[i + 5], data[i + 6], data[i + 7]])
};
s0 = s0.wrapping_add(word_a).wrapping_add(s1);
s1 = s1.wrapping_add(word_b).wrapping_add(s0);
i += 8;
}
(s0, s1)
}
#[derive(Debug)]
pub struct WalReader {
page_size: u32,
committed_pages: HashMap<u32, Vec<u8>>,
}
impl WalReader {
pub fn from_bytes(data: &[u8]) -> Result<Self, GpkgError> {
if data.is_empty() {
return Ok(Self {
page_size: 0,
committed_pages: HashMap::new(),
});
}
if data.len() < WAL_HEADER_SIZE {
return Err(GpkgError::InvalidFormat(format!(
"WAL data too short: need {WAL_HEADER_SIZE} bytes for header, got {}",
data.len()
)));
}
let magic = u32::from_be_bytes([data[0], data[1], data[2], data[3]]);
let big_endian = match magic {
WAL_MAGIC_BE => true,
WAL_MAGIC_LE => false,
other => return Err(GpkgError::InvalidWalMagic(other)),
};
let page_size = u32::from_be_bytes([data[8], data[9], data[10], data[11]]);
let salt1 = u32::from_be_bytes([data[16], data[17], data[18], data[19]]);
let salt2 = u32::from_be_bytes([data[20], data[21], data[22], data[23]]);
let (hdr_s0, hdr_s1) = wal_checksum(&data[0..24], big_endian, 0, 0);
let stored_hdr_s0 = u32::from_be_bytes([data[24], data[25], data[26], data[27]]);
let stored_hdr_s1 = u32::from_be_bytes([data[28], data[29], data[30], data[31]]);
if hdr_s0 != stored_hdr_s0 || hdr_s1 != stored_hdr_s1 {
return Err(GpkgError::InvalidFormat(
"WAL header checksum mismatch".into(),
));
}
let frame_size = FRAME_HEADER_SIZE + page_size as usize;
let mut committed_pages: HashMap<u32, Vec<u8>> = HashMap::new();
let mut pending: HashMap<u32, Vec<u8>> = HashMap::new();
let mut cumulative_s0: u32 = hdr_s0;
let mut cumulative_s1: u32 = hdr_s1;
let mut offset = WAL_HEADER_SIZE;
while offset + frame_size <= data.len() {
let frame_hdr = &data[offset..offset + FRAME_HEADER_SIZE];
let page_data = &data[offset + FRAME_HEADER_SIZE..offset + frame_size];
let frame_page_no =
u32::from_be_bytes([frame_hdr[0], frame_hdr[1], frame_hdr[2], frame_hdr[3]]);
let db_size_after =
u32::from_be_bytes([frame_hdr[4], frame_hdr[5], frame_hdr[6], frame_hdr[7]]);
let frame_salt1 =
u32::from_be_bytes([frame_hdr[8], frame_hdr[9], frame_hdr[10], frame_hdr[11]]);
let frame_salt2 =
u32::from_be_bytes([frame_hdr[12], frame_hdr[13], frame_hdr[14], frame_hdr[15]]);
let stored_s0 =
u32::from_be_bytes([frame_hdr[16], frame_hdr[17], frame_hdr[18], frame_hdr[19]]);
let stored_s1 =
u32::from_be_bytes([frame_hdr[20], frame_hdr[21], frame_hdr[22], frame_hdr[23]]);
if frame_salt1 != salt1 || frame_salt2 != salt2 {
break;
}
let (after_hdr_s0, after_hdr_s1) =
wal_checksum(&frame_hdr[0..8], big_endian, cumulative_s0, cumulative_s1);
let (frame_s0, frame_s1) =
wal_checksum(page_data, big_endian, after_hdr_s0, after_hdr_s1);
if frame_s0 != stored_s0 || frame_s1 != stored_s1 {
break;
}
cumulative_s0 = frame_s0;
cumulative_s1 = frame_s1;
pending.insert(frame_page_no, page_data.to_vec());
if db_size_after != 0 {
for (pg_no, pg_data) in pending.drain() {
committed_pages.insert(pg_no, pg_data);
}
}
offset += frame_size;
}
Ok(Self {
page_size,
committed_pages,
})
}
pub fn apply_to_main(&self, main: &[u8]) -> Result<Vec<u8>, GpkgError> {
if self.committed_pages.is_empty() {
return Ok(main.to_vec());
}
if self.page_size == 0 {
return Err(GpkgError::InvalidFormat(
"WAL page size is zero; cannot apply overlay".into(),
));
}
let page_size = self.page_size as usize;
let max_wal_page = self.committed_pages.keys().copied().max().unwrap_or(0);
let required_len = (max_wal_page as usize) * page_size;
let output_len = main.len().max(required_len);
let mut output = Vec::with_capacity(output_len);
output.extend_from_slice(main);
if output.len() < required_len {
output.resize(required_len, 0u8);
}
for (&page_no, page_bytes) in &self.committed_pages {
debug_assert!(page_no >= 1, "page numbers are 1-indexed");
let start = (page_no as usize - 1) * page_size;
let end = start + page_size;
if output.len() < end {
output.resize(end, 0u8);
}
output[start..end].copy_from_slice(page_bytes);
}
Ok(output)
}
pub fn page_size(&self) -> u32 {
self.page_size
}
pub fn committed_page_count(&self) -> usize {
self.committed_pages.len()
}
}
pub fn overlay_wal(main: &[u8], wal: &[u8]) -> Result<Vec<u8>, GpkgError> {
let reader = WalReader::from_bytes(wal)?;
reader.apply_to_main(main)
}