use crate::common::{crc32, Error, Result, WalSyncPolicy};
use std::fs::{File, OpenOptions};
use std::io::{BufReader, BufWriter, Read, Write};
use std::path::{Path, PathBuf};
const WAL_MAGIC: [u8; 4] = [0x57, 0x41, 0x4C, 0x31]; const OP_PUT: u8 = 1;
const OP_DELETE: u8 = 2;
#[derive(Debug, Clone)]
pub struct WalEntry {
pub sequence: u64,
pub op: WalOp,
}
#[derive(Debug, Clone)]
pub enum WalOp {
Put { key: String, value: Vec<u8> },
Delete { key: String },
}
pub struct Wal {
path: PathBuf,
writer: BufWriter<File>,
next_sequence: u64,
sync_policy: WalSyncPolicy,
}
impl Wal {
pub fn open(path: impl AsRef<Path>, sync_policy: WalSyncPolicy) -> Result<Self> {
let path = path.as_ref().to_path_buf();
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
let file = OpenOptions::new()
.create(true)
.append(true)
.read(true)
.open(&path)?;
let next_sequence = Self::find_last_sequence(&path)?;
Ok(Self {
path,
writer: BufWriter::new(file),
next_sequence,
sync_policy,
})
}
fn find_last_sequence(path: &Path) -> Result<u64> {
let file = match File::open(path) {
Ok(f) => f,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(0),
Err(e) => return Err(e.into()),
};
let mut reader = BufReader::new(file);
let mut max_seq = None;
loop {
match Self::read_entry_internal(&mut reader) {
Ok(Some(entry)) => {
max_seq = Some(max_seq.unwrap_or(0).max(entry.sequence));
}
Ok(None) => break,
Err(_) => break, }
}
Ok(max_seq.map(|s| s + 1).unwrap_or(0))
}
pub fn append_put(&mut self, key: &str, value: &[u8]) -> Result<u64> {
let sequence = self.next_sequence;
self.next_sequence += 1;
self.write_entry(sequence, OP_PUT, key, Some(value))?;
self.maybe_sync()?;
Ok(sequence)
}
pub fn append_delete(&mut self, key: &str) -> Result<u64> {
let sequence = self.next_sequence;
self.next_sequence += 1;
self.write_entry(sequence, OP_DELETE, key, None)?;
self.maybe_sync()?;
Ok(sequence)
}
fn write_entry(
&mut self,
sequence: u64,
op: u8,
key: &str,
value: Option<&[u8]>,
) -> Result<()> {
let key_bytes = key.as_bytes();
let val_bytes = value.unwrap_or(&[]);
self.writer.write_all(&WAL_MAGIC)?;
self.writer.write_all(&sequence.to_le_bytes())?;
self.writer.write_all(&[op])?;
self.writer
.write_all(&(key_bytes.len() as u32).to_le_bytes())?;
self.writer
.write_all(&(val_bytes.len() as u32).to_le_bytes())?;
self.writer.write_all(key_bytes)?;
if op == OP_PUT {
self.writer.write_all(val_bytes)?;
}
let mut checksum_data = Vec::new();
checksum_data.extend_from_slice(&sequence.to_le_bytes());
checksum_data.push(op);
checksum_data.extend_from_slice(&(key_bytes.len() as u32).to_le_bytes());
checksum_data.extend_from_slice(&(val_bytes.len() as u32).to_le_bytes());
checksum_data.extend_from_slice(key_bytes);
if op == OP_PUT {
checksum_data.extend_from_slice(val_bytes);
}
let checksum = crc32(&checksum_data);
self.writer.write_all(&checksum.to_le_bytes())?;
Ok(())
}
fn maybe_sync(&mut self) -> Result<()> {
match self.sync_policy {
WalSyncPolicy::Always => {
self.writer.flush()?;
self.writer.get_ref().sync_all()?;
}
WalSyncPolicy::Interval => {
self.writer.flush()?;
}
WalSyncPolicy::Never => {}
}
Ok(())
}
pub fn replay<F>(path: impl AsRef<Path>, mut callback: F) -> Result<()>
where
F: FnMut(WalEntry) -> Result<()>,
{
let file = match File::open(path.as_ref()) {
Ok(f) => f,
Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(()),
Err(e) => return Err(e.into()),
};
let mut reader = BufReader::new(file);
loop {
match Self::read_entry_internal(&mut reader) {
Ok(Some(entry)) => callback(entry)?,
Ok(None) => break,
Err(e) => {
tracing::warn!("WAL replay stopped at corrupted entry: {}", e);
break;
}
}
}
Ok(())
}
fn read_entry_internal<R: Read>(reader: &mut R) -> Result<Option<WalEntry>> {
let mut magic = [0u8; 4];
match reader.read_exact(&mut magic) {
Ok(_) => {}
Err(e) if e.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(e) => return Err(e.into()),
}
if magic != WAL_MAGIC {
return Err(Error::Wal("Invalid WAL magic".into()));
}
let mut seq_bytes = [0u8; 8];
reader.read_exact(&mut seq_bytes)?;
let sequence = u64::from_le_bytes(seq_bytes);
let mut op = [0u8; 1];
reader.read_exact(&mut op)?;
let mut key_len_bytes = [0u8; 4];
reader.read_exact(&mut key_len_bytes)?;
let key_len = u32::from_le_bytes(key_len_bytes) as usize;
let mut val_len_bytes = [0u8; 4];
reader.read_exact(&mut val_len_bytes)?;
let val_len = u32::from_le_bytes(val_len_bytes) as usize;
let mut key_bytes = vec![0u8; key_len];
reader.read_exact(&mut key_bytes)?;
let key =
String::from_utf8(key_bytes).map_err(|_| Error::Wal("Invalid UTF-8 in key".into()))?;
let value = if op[0] == OP_PUT {
let mut val = vec![0u8; val_len];
reader.read_exact(&mut val)?;
Some(val)
} else {
None
};
let mut checksum_bytes = [0u8; 4];
reader.read_exact(&mut checksum_bytes)?;
let stored_checksum = u32::from_le_bytes(checksum_bytes);
let mut checksum_data = Vec::new();
checksum_data.extend_from_slice(&seq_bytes);
checksum_data.push(op[0]);
checksum_data.extend_from_slice(&key_len_bytes);
checksum_data.extend_from_slice(&val_len_bytes);
checksum_data.extend_from_slice(key.as_bytes());
if let Some(ref v) = value {
checksum_data.extend_from_slice(v);
}
let computed_checksum = crc32(&checksum_data);
if computed_checksum != stored_checksum {
return Err(Error::Wal("Checksum mismatch".into()));
}
let wal_op = match op[0] {
OP_PUT => WalOp::Put {
key,
value: value.unwrap(),
},
OP_DELETE => WalOp::Delete { key },
_ => return Err(Error::Wal(format!("Unknown op code: {}", op[0]))),
};
Ok(Some(WalEntry {
sequence,
op: wal_op,
}))
}
pub fn truncate(&mut self) -> Result<()> {
self.writer.flush()?;
drop(std::mem::replace(
&mut self.writer,
BufWriter::new(File::open(&self.path)?),
));
let file = OpenOptions::new()
.write(true)
.truncate(true)
.open(&self.path)?;
self.writer = BufWriter::new(file);
self.next_sequence = 0;
Ok(())
}
pub fn sync(&mut self) -> Result<()> {
self.writer.flush()?;
self.writer.get_ref().sync_all()?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[test]
fn test_wal_basic() {
let dir = tempdir().unwrap();
let wal_path = dir.path().join("test.wal");
{
let mut wal = Wal::open(&wal_path, WalSyncPolicy::Always).unwrap();
let seq1 = wal.append_put("key1", b"value1").unwrap();
let seq2 = wal.append_put("key2", b"value2").unwrap();
let seq3 = wal.append_delete("key1").unwrap();
assert_eq!(seq1, 0);
assert_eq!(seq2, 1);
assert_eq!(seq3, 2);
wal.sync().unwrap();
}
let mut entries = Vec::new();
Wal::replay(&wal_path, |entry| {
entries.push(entry);
Ok(())
})
.unwrap();
assert_eq!(entries.len(), 3);
assert_eq!(entries[0].sequence, 0);
assert_eq!(entries[1].sequence, 1);
assert_eq!(entries[2].sequence, 2);
match &entries[0].op {
WalOp::Put { key, value } => {
assert_eq!(key, "key1");
assert_eq!(value, b"value1");
}
_ => panic!("Expected Put"),
}
}
#[test]
fn test_wal_reopen() {
let dir = tempdir().unwrap();
let wal_path = dir.path().join("reopen.wal");
{
let mut wal = Wal::open(&wal_path, WalSyncPolicy::Always).unwrap();
wal.append_put("key1", b"value1").unwrap();
wal.append_put("key2", b"value2").unwrap();
wal.sync().unwrap();
}
{
let mut wal = Wal::open(&wal_path, WalSyncPolicy::Always).unwrap();
assert_eq!(wal.next_sequence, 2);
let seq = wal.append_put("key3", b"value3").unwrap();
assert_eq!(seq, 2);
wal.sync().unwrap();
}
let mut count = 0;
Wal::replay(&wal_path, |_| {
count += 1;
Ok(())
})
.unwrap();
assert_eq!(count, 3);
}
}