use std::{
fs::{File, OpenOptions},
io::Write,
path::PathBuf,
sync::atomic::{AtomicU64, Ordering},
};
use sparrowdb_common::{Lsn, Result, TxnId};
use super::codec::{
WalPayload, WalRecord, WalRecordKind, WAL_FORMAT_VERSION, WAL_FORMAT_VERSION_LEGACY,
};
use crate::encryption::EncryptionContext;
pub const SEGMENT_SIZE: u64 = 64 * 1024 * 1024;
pub fn segment_path(wal_dir: &std::path::Path, seg_no: u64) -> PathBuf {
wal_dir.join(format!("segment-{:020}.wal", seg_no))
}
pub struct WalWriter {
wal_dir: PathBuf,
file: File,
seg_no: u64,
seg_offset: u64,
next_lsn: AtomicU64,
enc: EncryptionContext,
}
impl WalWriter {
pub fn open(wal_dir: &std::path::Path) -> Result<Self> {
Self::open_inner(wal_dir, EncryptionContext::none())
}
pub fn open_encrypted(wal_dir: &std::path::Path, key: [u8; 32]) -> Result<Self> {
Self::open_inner(wal_dir, EncryptionContext::with_key(key))
}
fn open_inner(wal_dir: &std::path::Path, enc: EncryptionContext) -> Result<Self> {
std::fs::create_dir_all(wal_dir)?;
let (seg_no, seg_offset, next_lsn) = Self::scan_wal_state(wal_dir)?;
let path = segment_path(wal_dir, seg_no);
if path.exists() {
let f = OpenOptions::new().write(true).open(&path)?;
f.set_len(seg_offset)?;
}
let mut file = OpenOptions::new()
.create(true)
.truncate(false)
.read(true)
.write(true)
.open(&path)?;
let file_len = file.metadata()?.len();
let seg_offset = if file_len == 0 {
file.write_all(&[WAL_FORMAT_VERSION])?;
1
} else {
use std::io::Seek;
file.seek(std::io::SeekFrom::Start(seg_offset))?;
seg_offset
};
Ok(Self {
wal_dir: wal_dir.to_path_buf(),
file,
seg_no,
seg_offset,
next_lsn: AtomicU64::new(next_lsn),
enc,
})
}
fn scan_wal_state(wal_dir: &std::path::Path) -> Result<(u64, u64, u64)> {
let mut segments: Vec<u64> = Vec::new();
if let Ok(entries) = std::fs::read_dir(wal_dir) {
for entry in entries.flatten() {
let name = entry.file_name();
let name = name.to_string_lossy();
if name.starts_with("segment-") && name.ends_with(".wal") {
let num_str = &name["segment-".len()..name.len() - ".wal".len()];
if let Ok(n) = num_str.parse::<u64>() {
segments.push(n);
}
}
}
}
segments.sort();
if segments.is_empty() {
return Ok((0, 0, 1));
}
let last_seg = *segments.last().unwrap();
let path = segment_path(wal_dir, last_seg);
let data = std::fs::read(&path)?;
if data.is_empty() {
return Err(sparrowdb_common::Error::Corruption(
"WAL segment is empty and missing the required version header".to_string(),
));
}
let version = data[0];
if version != WAL_FORMAT_VERSION && version != WAL_FORMAT_VERSION_LEGACY {
return Err(sparrowdb_common::Error::Corruption(format!(
"WAL segment has unrecognised version byte {version}. \
Supported versions: {WAL_FORMAT_VERSION} (current CRC32C), \
{WAL_FORMAT_VERSION_LEGACY} (legacy 0.1.2 CRC32). \
The database cannot be opened."
)));
}
let mut offset = 1usize;
let mut max_lsn = 0u64;
while offset < data.len() {
match WalRecord::decode_with_version(&data[offset..], version) {
Ok((rec, consumed)) => {
if rec.lsn.0 > max_lsn {
max_lsn = rec.lsn.0;
}
offset += consumed;
}
Err(_) => break, }
}
Ok((last_seg, offset as u64, max_lsn + 1))
}
fn alloc_lsn(&self) -> Lsn {
Lsn(self.next_lsn.fetch_add(1, Ordering::Relaxed))
}
pub fn append(
&mut self,
kind: WalRecordKind,
txn_id: TxnId,
payload: WalPayload,
) -> Result<Lsn> {
let lsn = self.alloc_lsn();
let final_payload = if self.enc.is_encrypted() {
let raw_payload_bytes = payload.encode();
if raw_payload_bytes.is_empty() {
payload
} else {
let encrypted = self.enc.encrypt_wal_payload(lsn.0, &raw_payload_bytes)?;
WalPayload::Raw(encrypted)
}
} else {
payload
};
let record = WalRecord {
lsn,
txn_id,
kind,
payload: final_payload,
};
let encoded = record.encode();
let record_len = encoded.len() as u64;
if self.seg_offset + record_len > SEGMENT_SIZE {
self.rotate()?;
}
self.file.write_all(&encoded)?;
self.seg_offset += record_len;
Ok(lsn)
}
fn rotate(&mut self) -> Result<()> {
self.file.flush()?;
self.seg_no += 1;
self.seg_offset = 1;
let path = segment_path(&self.wal_dir, self.seg_no);
let mut new_file = OpenOptions::new()
.create(true)
.truncate(false)
.read(true)
.write(true)
.open(&path)?;
new_file.write_all(&[WAL_FORMAT_VERSION])?;
self.file = new_file;
Ok(())
}
pub fn fsync(&self) -> Result<()> {
self.file.sync_all()?;
Ok(())
}
pub fn wal_dir(&self) -> &std::path::Path {
&self.wal_dir
}
pub fn last_lsn(&self) -> Lsn {
Lsn(self.next_lsn.load(Ordering::Relaxed).saturating_sub(1))
}
pub fn commit_transaction(
&mut self,
txn_id: TxnId,
dirty_pages: &[(u64, Vec<u8>)],
) -> Result<Lsn> {
self.append(WalRecordKind::Begin, txn_id, WalPayload::Empty)?;
for (page_id, image) in dirty_pages {
self.append(
WalRecordKind::Write,
txn_id,
WalPayload::Write {
page_id: *page_id,
image: image.clone(),
},
)?;
}
let commit_lsn = self.append(WalRecordKind::Commit, txn_id, WalPayload::Empty)?;
self.fsync()?;
Ok(commit_lsn)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::TempDir;
fn open_writer(dir: &TempDir) -> WalWriter {
WalWriter::open(dir.path()).unwrap()
}
#[test]
fn test_writer_creates_segment_file() {
let dir = TempDir::new().unwrap();
let _writer = open_writer(&dir);
let seg = segment_path(dir.path(), 0);
assert!(seg.exists());
}
#[test]
fn test_writer_append_returns_monotonic_lsns() {
let dir = TempDir::new().unwrap();
let mut writer = open_writer(&dir);
let txn = TxnId(1);
let l1 = writer
.append(WalRecordKind::Begin, txn, WalPayload::Empty)
.unwrap();
let l2 = writer
.append(WalRecordKind::Commit, txn, WalPayload::Empty)
.unwrap();
assert!(l1 < l2);
}
#[test]
fn test_writer_commit_transaction() {
let dir = TempDir::new().unwrap();
let mut writer = open_writer(&dir);
let dirty = vec![(0u64, vec![0xAAu8; 64])];
let commit_lsn = writer.commit_transaction(TxnId(42), &dirty).unwrap();
assert!(commit_lsn.0 >= 1);
}
#[test]
fn test_writer_segment_rotation() {
let dir = TempDir::new().unwrap();
let mut writer = open_writer(&dir);
let initial_seg = writer.seg_no;
writer.seg_offset = SEGMENT_SIZE; writer
.append(WalRecordKind::Begin, TxnId(1), WalPayload::Empty)
.unwrap();
assert_eq!(writer.seg_no, initial_seg + 1);
}
#[test]
fn test_writer_fsync_does_not_panic() {
let dir = TempDir::new().unwrap();
let mut writer = open_writer(&dir);
writer
.append(WalRecordKind::Begin, TxnId(1), WalPayload::Empty)
.unwrap();
writer.fsync().unwrap();
}
#[test]
fn test_writer_reopen_continues_lsn() {
let dir = TempDir::new().unwrap();
let last_lsn = {
let mut writer = open_writer(&dir);
writer
.append(WalRecordKind::Begin, TxnId(1), WalPayload::Empty)
.unwrap();
writer
.append(WalRecordKind::Commit, TxnId(1), WalPayload::Empty)
.unwrap()
};
let mut writer2 = open_writer(&dir);
let new_lsn = writer2
.append(WalRecordKind::Begin, TxnId(2), WalPayload::Empty)
.unwrap();
assert!(
new_lsn > last_lsn,
"reopened writer must continue from last LSN"
);
}
}