1use 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;
20pub 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
37pub 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
84pub 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}