use std::collections::BTreeMap;
use std::io::{BufWriter, Read, Write};
use std::path::{Path, PathBuf};
use std::sync::Mutex;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct WalPos {
pub segment: u64,
}
#[derive(Debug, Clone)]
pub struct WalItem {
pub stream: String,
pub id: String,
pub seq: i64,
pub tombstone: bool,
pub doc: std::sync::Arc<str>,
}
#[derive(Debug, Clone)]
pub struct WalRecord {
pub stream: String,
pub id: String,
pub seq: i64,
pub tombstone: bool,
pub doc: Vec<u8>,
pub pos: WalPos,
}
const V2_FLAG: u16 = 0x8000;
const TOMBSTONE_FLAG: u16 = 0x8000;
struct SegmentState {
outstanding: u64,
sealed: bool,
}
struct WalInner {
current_seq: u64,
current_len: u64,
writer: BufWriter<std::fs::File>,
segments: BTreeMap<u64, SegmentState>,
written: u64,
}
pub struct Wal {
dir: PathBuf,
max_segment_bytes: u64,
inner: Mutex<WalInner>,
synced: std::sync::atomic::AtomicU64,
sync_gate: Mutex<()>,
}
fn segment_path(dir: &Path, seq: u64) -> PathBuf {
dir.join(format!("wal-{seq:016}.log"))
}
impl Wal {
pub fn open(dir: impl Into<PathBuf>, max_segment_bytes: u64) -> std::io::Result<(Self, WalReplay)> {
let dir = dir.into();
std::fs::create_dir_all(&dir)?;
let mut seqs: Vec<u64> = Vec::new();
for entry in std::fs::read_dir(&dir)? {
let name = entry?.file_name().to_string_lossy().to_string();
if let Some(num) = name
.strip_prefix("wal-")
.and_then(|s| s.strip_suffix(".log"))
&& let Ok(seq) = num.parse::<u64>()
{
seqs.push(seq);
}
}
seqs.sort_unstable();
let mut segments = BTreeMap::new();
for &seq in &seqs {
let path = segment_path(&dir, seq);
let mut cursor = SegmentCursor::open(&path, seq)?;
let mut count = 0u64;
while cursor.advance() {
count += 1;
}
segments.insert(
seq,
SegmentState {
outstanding: count,
sealed: true,
},
);
}
let replay = WalReplay {
segments: seqs
.iter()
.map(|&seq| (seq, segment_path(&dir, seq)))
.collect::<Vec<_>>()
.into_iter(),
current: None,
};
let current_seq = seqs.last().map(|s| s + 1).unwrap_or(0);
let file = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(segment_path(&dir, current_seq))?;
segments.insert(
current_seq,
SegmentState {
outstanding: 0,
sealed: false,
},
);
Ok((
Self {
dir,
max_segment_bytes,
inner: Mutex::new(WalInner {
current_seq,
current_len: 0,
writer: BufWriter::new(file),
segments,
written: 0,
}),
synced: std::sync::atomic::AtomicU64::new(0),
sync_gate: Mutex::new(()),
},
replay,
))
}
pub fn append_batch(&self, items: &[WalItem]) -> std::io::Result<Vec<WalPos>> {
use std::sync::atomic::Ordering;
let mut positions = Vec::with_capacity(items.len());
let mut sealed: Vec<std::fs::File> = Vec::new();
let mut victims: Vec<PathBuf> = Vec::new();
let (fd, my_gen) = {
let mut inner = self.inner.lock().unwrap();
for item in items {
let doc = item.doc.as_bytes();
if inner.current_len >= self.max_segment_bytes {
let (old_file, mut gc) = self.rotate(&mut inner)?;
sealed.push(old_file);
victims.append(&mut gc);
}
let stream_bytes = item.stream.as_bytes();
let id_bytes = item.id.as_bytes();
if stream_bytes.len() >= V2_FLAG as usize
|| id_bytes.len() >= TOMBSTONE_FLAG as usize
{
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"stream name or document id too long for the WAL",
));
}
let payload_len = 2 + stream_bytes.len() + 2 + id_bytes.len() + 8 + doc.len();
let len_field = stream_bytes.len() as u16 | V2_FLAG;
let id_len_field =
id_bytes.len() as u16 | if item.tombstone { TOMBSTONE_FLAG } else { 0 };
let seq_bytes = item.seq.to_le_bytes();
let mut hasher = crc32fast::Hasher::new();
hasher.update(&len_field.to_le_bytes());
hasher.update(stream_bytes);
hasher.update(&id_len_field.to_le_bytes());
hasher.update(id_bytes);
hasher.update(&seq_bytes);
hasher.update(doc);
let crc = hasher.finalize();
inner.writer.write_all(&(payload_len as u32).to_le_bytes())?;
inner.writer.write_all(&crc.to_le_bytes())?;
inner.writer.write_all(&len_field.to_le_bytes())?;
inner.writer.write_all(stream_bytes)?;
inner.writer.write_all(&id_len_field.to_le_bytes())?;
inner.writer.write_all(id_bytes)?;
inner.writer.write_all(&seq_bytes)?;
inner.writer.write_all(doc)?;
inner.current_len += 8 + payload_len as u64;
let seq = inner.current_seq;
match inner.segments.get_mut(&seq) {
Some(state) => state.outstanding += 1,
None => {
return Err(std::io::Error::other(
"WAL invariant violated: current segment untracked",
));
}
}
positions.push(WalPos { segment: seq });
}
inner.writer.flush()?;
let fd = inner.writer.get_ref().try_clone()?;
inner.written += items.len() as u64;
(fd, inner.written)
};
for old in &sealed {
old.sync_data()?;
}
for path in victims {
let _ = std::fs::remove_file(path);
}
if self.synced.load(Ordering::Acquire) < my_gen {
let _gate = self.sync_gate.lock().unwrap();
if self.synced.load(Ordering::Acquire) < my_gen {
let covered = self.inner.lock().unwrap().written;
fd.sync_data()?;
self.synced.fetch_max(covered, Ordering::AcqRel);
}
}
Ok(positions)
}
fn rotate(&self, inner: &mut WalInner) -> std::io::Result<(std::fs::File, Vec<PathBuf>)> {
inner.writer.flush()?;
let old_seq = inner.current_seq;
if let Some(state) = inner.segments.get_mut(&old_seq) {
state.sealed = true;
}
let new_seq = old_seq + 1;
let file = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(segment_path(&self.dir, new_seq))?;
let old = std::mem::replace(&mut inner.writer, BufWriter::new(file));
let old_file = old
.into_inner()
.map_err(|e| std::io::Error::other(e.to_string()))?;
inner.current_seq = new_seq;
inner.current_len = 0;
inner.segments.insert(
new_seq,
SegmentState {
outstanding: 0,
sealed: false,
},
);
let victims = self.collect_confirmed(inner);
Ok((old_file, victims))
}
pub fn confirm(&self, positions: &[WalPos]) {
let victims = {
let mut inner = self.inner.lock().unwrap();
for pos in positions {
if let Some(state) = inner.segments.get_mut(&pos.segment) {
state.outstanding = state.outstanding.saturating_sub(1);
}
}
self.collect_confirmed(&mut inner)
};
for path in victims {
let _ = std::fs::remove_file(path);
}
}
fn collect_confirmed(&self, inner: &mut WalInner) -> Vec<PathBuf> {
let seqs: Vec<u64> = inner
.segments
.iter()
.filter(|(_, s)| s.sealed && s.outstanding == 0)
.map(|(&seq, _)| seq)
.collect();
seqs.into_iter()
.map(|seq| {
inner.segments.remove(&seq);
segment_path(&self.dir, seq)
})
.collect()
}
pub fn segment_count(&self) -> usize {
self.inner.lock().unwrap().segments.len()
}
pub fn outstanding(&self) -> u64 {
self.inner
.lock()
.unwrap()
.segments
.values()
.map(|s| s.outstanding)
.sum()
}
}
struct SegmentCursor {
seq: u64,
reader: std::io::BufReader<std::fs::File>,
buf: Vec<u8>,
layout: RecordLayout,
}
#[derive(Default, Clone, Copy)]
struct RecordLayout {
stream: (usize, usize),
id: Option<(usize, usize)>,
seq: i64,
tombstone: bool,
doc_start: usize,
}
impl SegmentCursor {
fn open(path: &Path, seq: u64) -> std::io::Result<Self> {
Ok(Self {
seq,
reader: std::io::BufReader::new(std::fs::File::open(path)?),
buf: Vec::new(),
layout: RecordLayout::default(),
})
}
fn advance(&mut self) -> bool {
let mut header = [0u8; 8];
if self.reader.read_exact(&mut header).is_err() {
return false;
}
let len = u32::from_le_bytes(header[0..4].try_into().unwrap()) as usize;
let crc = u32::from_le_bytes(header[4..8].try_into().unwrap());
if len < 2 {
return false;
}
self.buf.clear();
match (&mut self.reader).take(len as u64).read_to_end(&mut self.buf) {
Ok(n) if n == len => {}
_ => return false, }
if crc32fast::hash(&self.buf) != crc {
return false; }
let len_field = u16::from_le_bytes(self.buf[0..2].try_into().unwrap());
let stream_len = (len_field & !V2_FLAG) as usize;
if 2 + stream_len > self.buf.len() {
return false;
}
let mut layout = RecordLayout {
stream: (2, 2 + stream_len),
..RecordLayout::default()
};
if len_field & V2_FLAG == 0 {
layout.doc_start = 2 + stream_len;
} else {
let mut at = 2 + stream_len;
if at + 2 > self.buf.len() {
return false;
}
let id_len_field = u16::from_le_bytes(self.buf[at..at + 2].try_into().unwrap());
let id_len = (id_len_field & !TOMBSTONE_FLAG) as usize;
layout.tombstone = id_len_field & TOMBSTONE_FLAG != 0;
at += 2;
if at + id_len + 8 > self.buf.len() {
return false;
}
layout.id = Some((at, at + id_len));
at += id_len;
layout.seq = i64::from_le_bytes(self.buf[at..at + 8].try_into().unwrap());
layout.doc_start = at + 8;
}
self.layout = layout;
true
}
fn record(&self) -> WalRecord {
let (s0, s1) = self.layout.stream;
let id = match self.layout.id {
Some((i0, i1)) => String::from_utf8_lossy(&self.buf[i0..i1]).to_string(),
None => uuid::Uuid::new_v4().simple().to_string(),
};
WalRecord {
stream: String::from_utf8_lossy(&self.buf[s0..s1]).to_string(),
id,
seq: self.layout.seq,
tombstone: self.layout.tombstone,
doc: self.buf[self.layout.doc_start..].to_vec(),
pos: WalPos { segment: self.seq },
}
}
}
pub struct WalReplay {
segments: std::vec::IntoIter<(u64, PathBuf)>,
current: Option<SegmentCursor>,
}
impl Iterator for WalReplay {
type Item = std::io::Result<WalRecord>;
fn next(&mut self) -> Option<Self::Item> {
loop {
if self.current.is_none() {
let (seq, path) = self.segments.next()?;
match SegmentCursor::open(&path, seq) {
Ok(cursor) => self.current = Some(cursor),
Err(e) => return Some(Err(e)),
}
}
let cursor = self.current.as_mut().expect("cursor set above");
if cursor.advance() {
return Some(Ok(cursor.record()));
}
self.current = None; }
}
}
#[cfg(test)]
mod tests {
use super::*;
fn items(n: usize, stream: &str) -> Vec<WalItem> {
(0..n)
.map(|i| WalItem {
stream: stream.to_string(),
id: format!("id-{i}"),
seq: 1_000 + i as i64,
tombstone: i % 2 == 1,
doc: std::sync::Arc::from(format!("{{\"n\":{i}}}")),
})
.collect()
}
fn append_legacy(path: &Path, stream: &str, doc: &str) {
let mut payload = Vec::new();
payload.extend_from_slice(&(stream.len() as u16).to_le_bytes());
payload.extend_from_slice(stream.as_bytes());
payload.extend_from_slice(doc.as_bytes());
let mut out = std::fs::OpenOptions::new().append(true).open(path).unwrap();
out.write_all(&(payload.len() as u32).to_le_bytes()).unwrap();
out.write_all(&crc32fast::hash(&payload).to_le_bytes()).unwrap();
out.write_all(&payload).unwrap();
}
#[test]
fn replays_identity_and_legacy_records() {
let dir = tempfile::tempdir().unwrap();
{
let (wal, _) = Wal::open(dir.path(), 1 << 20).unwrap();
wal.append_batch(&items(2, "app")).unwrap();
}
append_legacy(&segment_path(dir.path(), 0), "legacy", "{\"old\":true}");
let (_, replayed) = Wal::open(dir.path(), 1 << 20).unwrap();
let records: Vec<WalRecord> = replayed.collect::<std::io::Result<_>>().unwrap();
assert_eq!(records.len(), 3);
assert_eq!(records[0].id, "id-0");
assert_eq!(records[0].seq, 1_000);
assert!(!records[0].tombstone);
assert_eq!(records[1].doc, b"{\"n\":1}");
assert!(records[1].tombstone);
assert_eq!(records[2].stream, "legacy");
assert_eq!(records[2].seq, 0);
assert!(!records[2].tombstone);
assert_eq!(records[2].id.len(), 32, "legacy records get a generated id");
assert_eq!(records[2].doc, b"{\"old\":true}");
}
#[test]
fn append_confirm_deletes_sealed_segments() {
let dir = tempfile::tempdir().unwrap();
let (wal, replayed) = Wal::open(dir.path(), 128).unwrap();
assert_eq!(replayed.count(), 0);
let positions = wal.append_batch(&items(20, "s")).unwrap();
assert_eq!(wal.outstanding(), 20);
assert!(wal.segment_count() > 1, "should have rotated");
wal.confirm(&positions);
assert_eq!(wal.outstanding(), 0);
assert_eq!(wal.segment_count(), 1);
}
#[test]
fn replay_after_restart_returns_unconfirmed() {
let dir = tempfile::tempdir().unwrap();
{
let (wal, _) = Wal::open(dir.path(), 1 << 20).unwrap();
wal.append_batch(&items(5, "app")).unwrap();
}
let (wal, replayed) = Wal::open(dir.path(), 1 << 20).unwrap();
let replayed: Vec<WalRecord> = replayed.collect::<std::io::Result<_>>().unwrap();
assert_eq!(replayed.len(), 5);
assert_eq!(replayed[0].stream, "app");
assert_eq!(wal.outstanding(), 5);
let positions: Vec<WalPos> = replayed.iter().map(|r| r.pos).collect();
wal.confirm(&positions);
assert_eq!(wal.outstanding(), 0);
}
#[test]
fn concurrent_group_commit_loses_nothing() {
use std::sync::Arc;
let dir = tempfile::tempdir().unwrap();
let (wal, _) = Wal::open(dir.path(), 1 << 20).unwrap();
let wal = Arc::new(wal);
let handles: Vec<_> = (0..8)
.map(|t| {
let wal = wal.clone();
std::thread::spawn(move || {
for _ in 0..100 {
wal.append_batch(&items(1, &format!("s{t}"))).unwrap();
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
assert_eq!(wal.outstanding(), 800);
drop(wal);
let (_, replayed) = Wal::open(dir.path(), 1 << 20).unwrap();
assert_eq!(replayed.map(|r| r.unwrap()).count(), 800);
}
#[test]
fn torn_tail_is_ignored() {
let dir = tempfile::tempdir().unwrap();
{
let (wal, _) = Wal::open(dir.path(), 1 << 20).unwrap();
wal.append_batch(&items(3, "s")).unwrap();
}
let seg = std::fs::read_dir(dir.path())
.unwrap()
.map(|e| e.unwrap().path())
.find(|p| p.to_string_lossy().contains("wal-0000000000000000"))
.unwrap();
let mut data = std::fs::read(&seg).unwrap();
let cut = data.len() - 3;
data.truncate(cut);
data.extend_from_slice(&[0xFF; 2]);
std::fs::write(&seg, data).unwrap();
let (wal, replayed) = Wal::open(dir.path(), 1 << 20).unwrap();
assert_eq!(replayed.map(|r| r.unwrap()).count(), 2);
assert_eq!(wal.outstanding(), 2);
}
#[test]
fn replay_streams_in_order_across_segments() {
let dir = tempfile::tempdir().unwrap();
{
let (wal, _) = Wal::open(dir.path(), 64).unwrap();
wal.append_batch(&items(50, "s")).unwrap();
assert!(wal.segment_count() > 1, "should have rotated");
}
let (wal, replayed) = Wal::open(dir.path(), 64).unwrap();
let records: Vec<WalRecord> = replayed.collect::<std::io::Result<_>>().unwrap();
assert_eq!(records.len(), 50);
assert_eq!(wal.outstanding(), 50);
for (i, record) in records.iter().enumerate() {
assert_eq!(record.doc, format!("{{\"n\":{i}}}").into_bytes());
}
let positions: Vec<WalPos> = records.iter().map(|r| r.pos).collect();
wal.confirm(&positions);
assert_eq!(wal.outstanding(), 0);
assert_eq!(wal.segment_count(), 1);
}
}