Skip to main content

basalt/
wal.rs

1//! Append-only write-ahead log for committed state snapshots.
2//!
3//! A frame is considered committed only after its complete header, payload,
4//! and checksums have reached the WAL file. Recovery accepts valid complete
5//! frames and repairs and ignores an incomplete final frame, which is the
6//! normal result of a killed process during append.
7
8use std::fs::{self, File, OpenOptions};
9use std::io::{self, Read, Write};
10use std::path::Path;
11
12use crate::crc::crc32;
13use crate::db::{DbError, DbErrorKind};
14use crate::storage;
15
16const MAGIC: &[u8; 4] = b"BSWL";
17const LEGACY_VERSION: u32 = 1;
18const VERSION: u32 = 2;
19const HEADER: usize = 32;
20/// Maximum total WAL size before callers must checkpoint.
21pub const MAX_WAL_BYTES: u64 = (storage::MAX_SNAPSHOT_BYTES as u64) * 4;
22const MAX_PAYLOAD_BYTES: usize = storage::MAX_SNAPSHOT_PAYLOAD_BYTES;
23
24#[derive(Debug, Clone)]
25pub struct Frame {
26    pub generation: u64,
27    pub payload: Vec<u8>,
28}
29
30fn io_error(context: &str, e: io::Error) -> DbError {
31    DbError::new(
32        DbErrorKind::Io(format!("{context}: {e}")),
33        format!("{context}: {e}"),
34    )
35}
36
37/// Append a committed frame.
38///
39/// Callers must provide a generation greater than every complete frame already
40/// in the file. [`latest`] validates that invariant during recovery; the
41/// database commit path supplies generations from its serialized commit lock
42/// without rescanning the full log for every append.
43pub fn append(path: &Path, generation: u64, payload: &[u8]) -> Result<(), DbError> {
44    if payload.len() > MAX_PAYLOAD_BYTES {
45        return Err(limit("database state is too large for the WAL"));
46    }
47    let frame_len = HEADER
48        .checked_add(payload.len())
49        .ok_or_else(|| limit("WAL frame is too large"))?;
50    let existing_len = existing_file_len(path)?.unwrap_or(0);
51    ensure_wal_size(existing_len, frame_len as u64)?;
52    if let Some(parent) = path
53        .parent()
54        .filter(|parent| !parent.as_os_str().is_empty())
55    {
56        fs::create_dir_all(parent).map_err(|e| io_error("create WAL directory", e))?;
57    }
58    let mut header = [0u8; HEADER];
59    header[..4].copy_from_slice(MAGIC);
60    header[4..8].copy_from_slice(&VERSION.to_le_bytes());
61    header[8..16].copy_from_slice(&generation.to_le_bytes());
62    header[16..24].copy_from_slice(&(payload.len() as u64).to_le_bytes());
63    header[24..28].copy_from_slice(&crc32(payload).to_le_bytes());
64    let header_checksum = crc32(&header[..28]);
65    header[28..32].copy_from_slice(&header_checksum.to_le_bytes());
66    let mut file = OpenOptions::new()
67        .create(true)
68        .append(true)
69        .open(path)
70        .map_err(|e| io_error("open WAL", e))?;
71    let actual_len = file
72        .metadata()
73        .map_err(|e| io_error("inspect WAL", e))?
74        .len();
75    ensure_wal_size(actual_len, frame_len as u64)?;
76    file.write_all(&header)
77        .map_err(|e| io_error("write WAL header", e))?;
78    file.write_all(payload)
79        .map_err(|e| io_error("write WAL payload", e))?;
80    file.sync_all().map_err(|e| io_error("sync WAL", e))?;
81    sync_parent(path)
82}
83
84/// Return the highest valid frame.  A partial/corrupt tail is ignored; an
85/// invalid frame before the tail is an error because it would hide later data.
86pub fn latest(path: &Path) -> Result<Option<Frame>, DbError> {
87    let Some(file_len) = existing_file_len(path)? else {
88        return Ok(None);
89    };
90    if file_len > MAX_WAL_BYTES {
91        return Err(limit(
92            "WAL is too large; checkpoint the database before retrying",
93        ));
94    }
95    let mut file = File::open(path).map_err(|e| io_error("open WAL", e))?;
96    let mut offset = 0u64;
97    let mut latest = None;
98    let mut previous_generation = None;
99    loop {
100        let mut header = [0u8; HEADER];
101        let header_len =
102            read_prefix(&mut file, &mut header).map_err(|e| io_error("read WAL header", e))?;
103        if header_len == 0 {
104            break;
105        }
106        if header_len < HEADER {
107            truncate_to(path, offset)?;
108            break;
109        }
110        if &header[..4] != MAGIC {
111            return Err(corrupt("invalid WAL magic"));
112        }
113        let version = u32_at(&header, 4)?;
114        if version == VERSION {
115            let header_checksum = u32_at(&header, 28)?;
116            if crc32(&header[..28]) != header_checksum {
117                return Err(corrupt("WAL header checksum mismatch"));
118            }
119        } else if version != LEGACY_VERSION {
120            return Err(corrupt("unsupported WAL version"));
121        }
122        let generation = u64_at(&header, 8)?;
123        if previous_generation.is_some_and(|previous| generation <= previous) {
124            return Err(corrupt("WAL generations are not strictly increasing"));
125        }
126        previous_generation = Some(generation);
127        let declared_len = u64_at(&header, 16)?;
128        if declared_len > MAX_PAYLOAD_BYTES as u64 {
129            return Err(limit("WAL frame payload is too large"));
130        }
131        let len = match usize::try_from(declared_len) {
132            Ok(len) => len,
133            Err(_) => return Err(limit("WAL frame payload is too large")),
134        };
135        let frame_len = (HEADER as u64)
136            .checked_add(declared_len)
137            .ok_or_else(|| limit("WAL frame is too large"))?;
138        let end = offset
139            .checked_add(frame_len)
140            .ok_or_else(|| limit("WAL offset is too large"))?;
141        if end > file_len {
142            truncate_to(path, offset)?;
143            break;
144        }
145        let mut payload = vec![0u8; len];
146        if let Err(error) = file.read_exact(&mut payload) {
147            if error.kind() == io::ErrorKind::UnexpectedEof {
148                truncate_to(path, offset)?;
149                break;
150            }
151            return Err(io_error("read WAL payload", error));
152        }
153        let checksum = u32_at(&header, 24)?;
154        if crc32(&payload) != checksum {
155            return Err(corrupt("WAL frame checksum mismatch"));
156        }
157        if latest
158            .as_ref()
159            .map(|f: &Frame| generation > f.generation)
160            .unwrap_or(true)
161        {
162            latest = Some(Frame {
163                generation,
164                payload,
165            });
166        }
167        offset = end;
168    }
169    Ok(latest)
170}
171
172pub fn truncate(path: &Path) -> Result<(), DbError> {
173    if existing_file_len(path)?.is_none() {
174        return Ok(());
175    }
176    let file = OpenOptions::new()
177        .write(true)
178        .truncate(true)
179        .open(path)
180        .map_err(|e| io_error("truncate WAL", e))?;
181    file.sync_all()
182        .map_err(|e| io_error("sync truncated WAL", e))?;
183    sync_parent(path)
184}
185
186fn truncate_to(path: &Path, length: u64) -> Result<(), DbError> {
187    let file = OpenOptions::new()
188        .write(true)
189        .open(path)
190        .map_err(|e| io_error("open WAL for tail repair", e))?;
191    file.set_len(length)
192        .map_err(|e| io_error("truncate incomplete WAL frame", e))?;
193    file.sync_all()
194        .map_err(|e| io_error("sync repaired WAL", e))?;
195    sync_parent(path)
196}
197
198fn corrupt(message: &str) -> DbError {
199    DbError::new(
200        DbErrorKind::Io(message.to_string()),
201        format!("corrupt WAL: {message}"),
202    )
203}
204
205fn limit(message: &str) -> DbError {
206    DbError::new(DbErrorKind::Limit, message)
207}
208
209fn existing_file_len(path: &Path) -> Result<Option<u64>, DbError> {
210    let metadata = match fs::symlink_metadata(path) {
211        Ok(metadata) => metadata,
212        Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(None),
213        Err(error) => return Err(io_error("inspect WAL", error)),
214    };
215    if metadata.file_type().is_symlink() {
216        return Err(path_error("WAL cannot be a symbolic link"));
217    }
218    if !metadata.is_file() {
219        return Err(path_error("WAL is not a regular file"));
220    }
221    Ok(Some(metadata.len()))
222}
223
224fn ensure_wal_size(existing_len: u64, additional_len: u64) -> Result<(), DbError> {
225    let total = existing_len
226        .checked_add(additional_len)
227        .ok_or_else(|| limit("WAL is too large; checkpoint the database before retrying"))?;
228    if total > MAX_WAL_BYTES {
229        return Err(limit(
230            "WAL is full; checkpoint the database before retrying the write",
231        ));
232    }
233    Ok(())
234}
235
236fn read_prefix(file: &mut File, bytes: &mut [u8]) -> io::Result<usize> {
237    let mut read = 0;
238    while read < bytes.len() {
239        let count = file.read(&mut bytes[read..])?;
240        if count == 0 {
241            break;
242        }
243        read += count;
244    }
245    Ok(read)
246}
247
248#[cfg(unix)]
249fn sync_parent(path: &Path) -> Result<(), DbError> {
250    if let Some(parent) = path
251        .parent()
252        .filter(|parent| !parent.as_os_str().is_empty())
253    {
254        let dir = File::open(parent).map_err(|e| io_error("open WAL directory", e))?;
255        dir.sync_all()
256            .map_err(|e| io_error("sync WAL directory", e))?;
257    }
258    Ok(())
259}
260
261#[cfg(not(unix))]
262fn sync_parent(_path: &Path) -> Result<(), DbError> {
263    Ok(())
264}
265
266fn path_error(message: &str) -> DbError {
267    DbError::new(DbErrorKind::Io(message.to_string()), message)
268}
269
270fn u32_at(bytes: &[u8], offset: usize) -> Result<u32, DbError> {
271    let raw = bytes
272        .get(offset..offset + 4)
273        .ok_or_else(|| corrupt("WAL header is truncated"))?;
274    Ok(u32::from_le_bytes(raw.try_into().unwrap()))
275}
276
277fn u64_at(bytes: &[u8], offset: usize) -> Result<u64, DbError> {
278    let raw = bytes
279        .get(offset..offset + 8)
280        .ok_or_else(|| corrupt("WAL header is truncated"))?;
281    Ok(u64::from_le_bytes(raw.try_into().unwrap()))
282}
283
284#[cfg(test)]
285mod tests {
286    use super::*;
287
288    #[test]
289    fn ignores_torn_tail() {
290        let dir = std::env::temp_dir().join(format!("basalt-wal-{}", std::process::id()));
291        let _ = fs::remove_dir_all(&dir);
292        fs::create_dir_all(&dir).unwrap();
293        let path = dir.join("db.wal");
294        append(&path, 1, b"one").unwrap();
295        let mut file = OpenOptions::new().append(true).open(&path).unwrap();
296        file.write_all(b"BSWL").unwrap();
297        file.sync_all().unwrap();
298        assert_eq!(latest(&path).unwrap().unwrap().payload, b"one");
299        truncate(&path).unwrap();
300        assert!(latest(&path).unwrap().is_none());
301        let _ = fs::remove_dir_all(dir);
302    }
303
304    #[test]
305    fn rejects_a_complete_corrupt_frame() {
306        let dir = std::env::temp_dir().join(format!("basalt-wal-corrupt-{}", std::process::id()));
307        let _ = fs::remove_dir_all(&dir);
308        fs::create_dir_all(&dir).unwrap();
309        let path = dir.join("db.wal");
310        append(&path, 1, b"one").unwrap();
311        let mut bytes = fs::read(&path).unwrap();
312        *bytes.last_mut().unwrap() ^= 1;
313        fs::write(&path, bytes).unwrap();
314        assert!(latest(&path).is_err());
315        let _ = fs::remove_dir_all(dir);
316    }
317
318    #[test]
319    fn repairs_a_torn_tail_before_a_later_commit() {
320        let dir = std::env::temp_dir().join(format!("basalt-wal-tail-{}", std::process::id()));
321        let _ = fs::remove_dir_all(&dir);
322        fs::create_dir_all(&dir).unwrap();
323        let path = dir.join("db.wal");
324        append(&path, 1, b"one").unwrap();
325        OpenOptions::new()
326            .append(true)
327            .open(&path)
328            .unwrap()
329            .write_all(b"BSWL")
330            .unwrap();
331
332        assert_eq!(latest(&path).unwrap().unwrap().generation, 1);
333        append(&path, 2, b"two").unwrap();
334        let frame = latest(&path).unwrap().unwrap();
335        assert_eq!(frame.generation, 2);
336        assert_eq!(frame.payload, b"two");
337        let _ = fs::remove_dir_all(dir);
338    }
339
340    #[test]
341    fn rejects_an_oversized_frame_before_allocating_its_payload() {
342        let dir = std::env::temp_dir().join(format!("basalt-wal-limit-{}", std::process::id()));
343        let _ = fs::remove_dir_all(&dir);
344        fs::create_dir_all(&dir).unwrap();
345        let path = dir.join("db.wal");
346        let mut header = [0u8; HEADER];
347        header[..4].copy_from_slice(MAGIC);
348        header[4..8].copy_from_slice(&VERSION.to_le_bytes());
349        header[8..16].copy_from_slice(&1u64.to_le_bytes());
350        header[16..24].copy_from_slice(&(MAX_PAYLOAD_BYTES as u64 + 1).to_le_bytes());
351        header[24..28].copy_from_slice(&0u32.to_le_bytes());
352        let header_checksum = crc32(&header[..28]);
353        header[28..32].copy_from_slice(&header_checksum.to_le_bytes());
354        fs::write(&path, header).unwrap();
355
356        let error = latest(&path).unwrap_err();
357
358        assert_eq!(error.kind, DbErrorKind::Limit);
359        assert!(error.message.contains("payload is too large"));
360        let _ = fs::remove_dir_all(dir);
361    }
362
363    #[test]
364    fn rejects_a_wal_file_above_the_total_limit() {
365        let dir =
366            std::env::temp_dir().join(format!("basalt-wal-total-limit-{}", std::process::id()));
367        let _ = fs::remove_dir_all(&dir);
368        fs::create_dir_all(&dir).unwrap();
369        let path = dir.join("db.wal");
370        let file = OpenOptions::new()
371            .create(true)
372            .truncate(false)
373            .write(true)
374            .open(&path)
375            .unwrap();
376        file.set_len(MAX_WAL_BYTES + 1).unwrap();
377        drop(file);
378
379        let error = latest(&path).unwrap_err();
380
381        assert_eq!(error.kind, DbErrorKind::Limit);
382        assert!(error.message.contains("WAL is too large"));
383        let _ = fs::remove_dir_all(dir);
384    }
385
386    #[test]
387    fn rejects_a_changed_v2_header_even_when_the_payload_is_intact() {
388        let dir =
389            std::env::temp_dir().join(format!("basalt-wal-header-corrupt-{}", std::process::id()));
390        let _ = fs::remove_dir_all(&dir);
391        fs::create_dir_all(&dir).unwrap();
392        let path = dir.join("db.wal");
393        append(&path, 1, b"one").unwrap();
394        let mut bytes = fs::read(&path).unwrap();
395        bytes[8] ^= 1;
396        fs::write(&path, bytes).unwrap();
397
398        let error = latest(&path).unwrap_err();
399
400        assert!(error.message.contains("header checksum mismatch"));
401        let _ = fs::remove_dir_all(dir);
402    }
403
404    #[test]
405    fn rejects_non_monotonic_wal_generations_during_recovery() {
406        let dir = std::env::temp_dir().join(format!(
407            "basalt-wal-generation-order-{}",
408            std::process::id()
409        ));
410        let _ = fs::remove_dir_all(&dir);
411        fs::create_dir_all(&dir).unwrap();
412        let path = dir.join("db.wal");
413        append(&path, 2, b"two").unwrap();
414        let payload = b"one";
415        let mut header = [0u8; HEADER];
416        header[..4].copy_from_slice(MAGIC);
417        header[4..8].copy_from_slice(&VERSION.to_le_bytes());
418        header[8..16].copy_from_slice(&1u64.to_le_bytes());
419        header[16..24].copy_from_slice(&(payload.len() as u64).to_le_bytes());
420        header[24..28].copy_from_slice(&crc32(payload).to_le_bytes());
421        let header_checksum = crc32(&header[..28]);
422        header[28..32].copy_from_slice(&header_checksum.to_le_bytes());
423        let mut file = OpenOptions::new().append(true).open(&path).unwrap();
424        file.write_all(&header).unwrap();
425        file.write_all(payload).unwrap();
426        file.sync_all().unwrap();
427
428        let error = latest(&path).unwrap_err();
429
430        assert!(error.message.contains("not strictly increasing"));
431        let _ = fs::remove_dir_all(dir);
432    }
433
434    #[test]
435    fn reads_legacy_v1_frames_during_upgrade() {
436        let dir = std::env::temp_dir().join(format!("basalt-wal-legacy-{}", std::process::id()));
437        let _ = fs::remove_dir_all(&dir);
438        fs::create_dir_all(&dir).unwrap();
439        let path = dir.join("db.wal");
440        let payload = b"legacy";
441        let mut header = [0u8; HEADER];
442        header[..4].copy_from_slice(MAGIC);
443        header[4..8].copy_from_slice(&LEGACY_VERSION.to_le_bytes());
444        header[8..16].copy_from_slice(&1u64.to_le_bytes());
445        header[16..24].copy_from_slice(&(payload.len() as u64).to_le_bytes());
446        header[24..28].copy_from_slice(&crc32(payload).to_le_bytes());
447        fs::write(&path, [header.as_slice(), payload].concat()).unwrap();
448
449        let frame = latest(&path).unwrap().unwrap();
450
451        assert_eq!(frame.generation, 1);
452        assert_eq!(frame.payload, payload);
453        let _ = fs::remove_dir_all(dir);
454    }
455
456    #[cfg(unix)]
457    #[test]
458    fn refuses_a_symbolic_link_wal() {
459        use std::os::unix::fs::symlink;
460
461        let dir = std::env::temp_dir().join(format!("basalt-wal-symlink-{}", std::process::id()));
462        let _ = fs::remove_dir_all(&dir);
463        fs::create_dir_all(&dir).unwrap();
464        let target = dir.join("outside.wal");
465        let path = dir.join("db.wal");
466        fs::write(&target, b"").unwrap();
467        symlink(&target, &path).unwrap();
468
469        let error = latest(&path).unwrap_err();
470
471        assert_eq!(
472            error.kind,
473            DbErrorKind::Io("WAL cannot be a symbolic link".into())
474        );
475        let _ = fs::remove_dir_all(dir);
476    }
477}