Skip to main content

basalt/
storage.rs

1//! On-disk snapshot storage.
2//!
3//! The storage file is a small page container.  A snapshot is encoded into a
4//! sequence of fixed-size pages, each page carrying its own length and CRC.
5//! This keeps the file format inspectable and lets recovery distinguish a
6//! complete snapshot from a torn write without relying on external crates.
7
8use std::ffi::{OsStr, OsString};
9use std::fs::{self, File, OpenOptions};
10use std::io::{self, Read, Write};
11use std::path::Path;
12use std::sync::atomic::{AtomicU64, Ordering};
13
14use crate::crc::crc32;
15use crate::db::{DbError, DbErrorKind, State};
16
17pub const PAGE_SIZE: usize = 4096;
18/// Maximum encoded snapshot size accepted by the file and byte APIs.
19pub const MAX_SNAPSHOT_BYTES: usize = 256 * 1024 * 1024;
20const FILE_MAGIC: &[u8; 8] = b"BASALTDB";
21const FILE_VERSION: u32 = 1;
22const FILE_HEADER: usize = 64;
23const PAGE_HEADER: usize = 24;
24/// Maximum state payload that can fit in a valid snapshot of the configured
25/// maximum size.
26pub const MAX_SNAPSHOT_PAYLOAD_BYTES: usize =
27    ((MAX_SNAPSHOT_BYTES - FILE_HEADER) / PAGE_SIZE) * (PAGE_SIZE - PAGE_HEADER);
28
29static TEMP_COUNTER: AtomicU64 = AtomicU64::new(0);
30
31fn io_error(context: &str, e: io::Error) -> DbError {
32    DbError::new(
33        DbErrorKind::Io(format!("{context}: {e}")),
34        format!("{context}: {e}"),
35    )
36}
37
38/// Write a complete database snapshot atomically.
39pub fn write_snapshot(path: &Path, state: &State, generation: u64) -> Result<(), DbError> {
40    let payload = state.encode();
41    if payload.len() > MAX_SNAPSHOT_PAYLOAD_BYTES {
42        return Err(limit("database state is too large for a snapshot"));
43    }
44    let page_payload = PAGE_SIZE - PAGE_HEADER;
45    let page_count = payload.len().div_ceil(page_payload).max(1);
46    let file_len = FILE_HEADER
47        .checked_add(
48            page_count
49                .checked_mul(PAGE_SIZE)
50                .ok_or_else(|| corrupt("database snapshot is too large"))?,
51        )
52        .ok_or_else(|| corrupt("database snapshot is too large"))?;
53    if file_len > MAX_SNAPSHOT_BYTES {
54        return Err(corrupt("database snapshot is too large"));
55    }
56
57    let mut bytes = vec![0u8; file_len];
58    bytes[..8].copy_from_slice(FILE_MAGIC);
59    bytes[8..12].copy_from_slice(&FILE_VERSION.to_le_bytes());
60    bytes[12..16].copy_from_slice(&(PAGE_SIZE as u32).to_le_bytes());
61    bytes[16..24].copy_from_slice(&generation.to_le_bytes());
62    bytes[24..32].copy_from_slice(&(payload.len() as u64).to_le_bytes());
63    bytes[32..40].copy_from_slice(&(page_count as u64).to_le_bytes());
64    let header_crc = crc32(&bytes[..40]);
65    bytes[40..44].copy_from_slice(&header_crc.to_le_bytes());
66
67    for page in 0..page_count {
68        let source_start = page * page_payload;
69        let source_end = (source_start + page_payload).min(payload.len());
70        let chunk = &payload[source_start..source_end];
71        let offset = FILE_HEADER + page * PAGE_SIZE;
72        bytes[offset..offset + 8].copy_from_slice(&(page as u64).to_le_bytes());
73        bytes[offset + 8..offset + 16].copy_from_slice(&(chunk.len() as u64).to_le_bytes());
74        bytes[offset + 16..offset + 20].copy_from_slice(&crc32(chunk).to_le_bytes());
75        bytes[offset + 20..offset + 24].copy_from_slice(&0u32.to_le_bytes());
76        bytes[offset + PAGE_HEADER..offset + PAGE_HEADER + chunk.len()].copy_from_slice(chunk);
77    }
78
79    ensure_not_symlink(path, "database snapshot")?;
80    let tmp = temporary_path(path);
81    if let Some(parent) = path
82        .parent()
83        .filter(|parent| !parent.as_os_str().is_empty())
84    {
85        fs::create_dir_all(parent).map_err(|e| io_error("create database directory", e))?;
86    }
87    let mut file = OpenOptions::new()
88        .create_new(true)
89        .write(true)
90        .open(&tmp)
91        .map_err(|e| io_error("open snapshot temporary file", e))?;
92    file.write_all(&bytes)
93        .map_err(|e| io_error("write snapshot", e))?;
94    file.sync_all().map_err(|e| io_error("sync snapshot", e))?;
95    drop(file);
96    ensure_not_symlink(path, "database snapshot")?;
97    let install_result = install_snapshot(&tmp, path);
98    if install_result.is_err() {
99        let _ = fs::remove_file(tmp);
100    }
101    install_result?;
102    sync_parent(path)
103}
104
105#[cfg(not(windows))]
106fn install_snapshot(tmp: &Path, path: &Path) -> Result<(), DbError> {
107    fs::rename(tmp, path).map_err(|e| io_error("install snapshot", e))
108}
109
110#[cfg(windows)]
111fn install_snapshot(tmp: &Path, path: &Path) -> Result<(), DbError> {
112    match fs::rename(tmp, path) {
113        Ok(()) => Ok(()),
114        Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {
115            // Windows does not replace an existing file with rename. The
116            // synced WAL remains the recovery source if the process stops
117            // between removing the old snapshot and installing the new one.
118            fs::remove_file(path).map_err(|e| io_error("replace snapshot", e))?;
119            fs::rename(tmp, path).map_err(|e| io_error("install snapshot", e))
120        }
121        Err(error) => Err(io_error("install snapshot", error)),
122    }
123}
124
125/// Read a snapshot.  A missing file is treated as an empty database.
126pub fn read_snapshot(path: &Path) -> Result<(State, u64), DbError> {
127    let metadata = match fs::symlink_metadata(path) {
128        Ok(metadata) => metadata,
129        Err(error) if error.kind() == io::ErrorKind::NotFound => {
130            return Ok((State::empty(), 0));
131        }
132        Err(error) => return Err(io_error("inspect database", error)),
133    };
134    if metadata.file_type().is_symlink() {
135        return Err(path_error("database snapshot cannot be a symbolic link"));
136    }
137    if !metadata.is_file() {
138        return Err(path_error("database snapshot is not a regular file"));
139    }
140    let file_len = metadata.len();
141    if file_len > MAX_SNAPSHOT_BYTES as u64 {
142        return Err(corrupt("database snapshot is too large"));
143    }
144    let file = File::open(path).map_err(|e| io_error("open database", e))?;
145    let mut bytes = Vec::with_capacity(file_len as usize);
146    file.take((MAX_SNAPSHOT_BYTES + 1) as u64)
147        .read_to_end(&mut bytes)
148        .map_err(|e| io_error("read database", e))?;
149    if bytes.len() > MAX_SNAPSHOT_BYTES {
150        return Err(corrupt("database snapshot is too large"));
151    }
152    read_snapshot_bytes(&bytes)
153}
154
155/// Read only the generation from a snapshot header. This lets database
156/// recovery decide whether a WAL frame is newer than a damaged snapshot
157/// without decoding the entire file.
158pub(crate) fn read_snapshot_generation(path: &Path) -> Result<Option<u64>, DbError> {
159    let metadata = match fs::symlink_metadata(path) {
160        Ok(metadata) => metadata,
161        Err(error) if error.kind() == io::ErrorKind::NotFound => return Ok(None),
162        Err(error) => return Err(io_error("inspect database", error)),
163    };
164    if metadata.file_type().is_symlink() {
165        return Err(path_error("database snapshot cannot be a symbolic link"));
166    }
167    if !metadata.is_file() {
168        return Err(path_error("database snapshot is not a regular file"));
169    }
170    if metadata.len() < FILE_HEADER as u64 {
171        return Err(corrupt("database header is truncated"));
172    }
173    let mut header = [0u8; FILE_HEADER];
174    File::open(path)
175        .map_err(|e| io_error("open database", e))?
176        .read_exact(&mut header)
177        .map_err(|e| io_error("read database header", e))?;
178    if &header[..8] != FILE_MAGIC {
179        return Err(corrupt("invalid database magic"));
180    }
181    if u32_at(&header, 8)? != FILE_VERSION {
182        return Err(corrupt("unsupported database version"));
183    }
184    if u32_at(&header, 12)? as usize != PAGE_SIZE {
185        return Err(corrupt("unsupported database page size"));
186    }
187    let header_crc = u32_at(&header, 40)?;
188    if crc32(&header[..40]) != header_crc {
189        return Err(corrupt("database header checksum mismatch"));
190    }
191    Ok(Some(u64_at(&header, 16)?))
192}
193
194/// Validate and decode snapshot bytes without touching the filesystem.
195///
196/// This is useful for embedded callers that already control the bytes and for
197/// exercising the on-disk format boundary without creating a temporary file.
198pub fn read_snapshot_bytes(bytes: &[u8]) -> Result<(State, u64), DbError> {
199    if bytes.len() > MAX_SNAPSHOT_BYTES {
200        return Err(corrupt("database snapshot is too large"));
201    }
202    if bytes.len() < FILE_HEADER {
203        return Err(corrupt("database header is truncated"));
204    }
205    if &bytes[..8] != FILE_MAGIC {
206        return Err(corrupt("invalid database magic"));
207    }
208    if u32_at(bytes, 8)? != FILE_VERSION {
209        return Err(corrupt("unsupported database version"));
210    }
211    if u32_at(bytes, 12)? as usize != PAGE_SIZE {
212        return Err(corrupt("unsupported database page size"));
213    }
214    let header_crc = u32_at(bytes, 40)?;
215    if crc32(&bytes[..40]) != header_crc {
216        return Err(corrupt("database header checksum mismatch"));
217    }
218    let generation = u64_at(bytes, 16)?;
219    let payload_len = usize::try_from(u64_at(bytes, 24)?)
220        .map_err(|_| corrupt("database payload is too large"))?;
221    let page_count = usize::try_from(u64_at(bytes, 32)?)
222        .map_err(|_| corrupt("database page count is too large"))?;
223    if page_count == 0
224        || page_count > (MAX_SNAPSHOT_BYTES - FILE_HEADER) / PAGE_SIZE
225        || payload_len > MAX_SNAPSHOT_PAYLOAD_BYTES
226        || payload_len > page_count.saturating_mul(PAGE_SIZE - PAGE_HEADER)
227    {
228        return Err(corrupt("invalid database payload size"));
229    }
230    let expected = FILE_HEADER
231        .checked_add(
232            page_count
233                .checked_mul(PAGE_SIZE)
234                .ok_or_else(|| corrupt("database is too large"))?,
235        )
236        .ok_or_else(|| corrupt("database is too large"))?;
237    if bytes.len() != expected {
238        return Err(corrupt(
239            "database page area is truncated or has trailing data",
240        ));
241    }
242    let mut payload = Vec::with_capacity(payload_len);
243    for page in 0..page_count {
244        let offset = FILE_HEADER + page * PAGE_SIZE;
245        if u64_at(bytes, offset)? != page as u64 {
246            return Err(corrupt("database page sequence mismatch"));
247        }
248        let len = usize::try_from(u64_at(bytes, offset + 8)?)
249            .map_err(|_| corrupt("database page is too large"))?;
250        let payload_end = payload
251            .len()
252            .checked_add(len)
253            .ok_or_else(|| corrupt("database payload is too large"))?;
254        if len > PAGE_SIZE - PAGE_HEADER || payload_end > payload_len {
255            return Err(corrupt("invalid database page length"));
256        }
257        let checksum = u32_at(bytes, offset + 16)?;
258        let chunk = &bytes[offset + PAGE_HEADER..offset + PAGE_HEADER + len];
259        if crc32(chunk) != checksum {
260            return Err(corrupt("database page checksum mismatch"));
261        }
262        payload.extend_from_slice(chunk);
263    }
264    payload.truncate(payload_len);
265    let state = State::decode(&payload)?;
266    Ok((state, generation))
267}
268
269#[cfg(unix)]
270fn sync_parent(path: &Path) -> Result<(), DbError> {
271    if let Some(parent) = path
272        .parent()
273        .filter(|parent| !parent.as_os_str().is_empty())
274    {
275        let dir = File::open(parent).map_err(|e| io_error("open database directory", e))?;
276        dir.sync_all()
277            .map_err(|e| io_error("sync database directory", e))?;
278    }
279    Ok(())
280}
281
282#[cfg(not(unix))]
283fn sync_parent(_path: &Path) -> Result<(), DbError> {
284    Ok(())
285}
286
287fn ensure_not_symlink(path: &Path, label: &str) -> Result<(), DbError> {
288    match fs::symlink_metadata(path) {
289        Ok(metadata) if metadata.file_type().is_symlink() => {
290            Err(path_error(&format!("{label} cannot be a symbolic link")))
291        }
292        Ok(_) => Ok(()),
293        Err(error) if error.kind() == io::ErrorKind::NotFound => Ok(()),
294        Err(error) => Err(io_error(&format!("inspect {label}"), error)),
295    }
296}
297
298fn temporary_path(path: &Path) -> std::path::PathBuf {
299    let mut name = path
300        .file_name()
301        .map(OsStr::to_os_string)
302        .unwrap_or_else(|| OsString::from("database"));
303    let counter = TEMP_COUNTER.fetch_add(1, Ordering::Relaxed);
304    name.push(format!(
305        ".basalt-snapshot-tmp-{}-{counter}",
306        std::process::id()
307    ));
308    path.with_file_name(name)
309}
310
311fn path_error(message: &str) -> DbError {
312    DbError::new(DbErrorKind::Io(message.to_string()), message)
313}
314
315fn limit(message: &str) -> DbError {
316    DbError::new(DbErrorKind::Limit, message)
317}
318
319fn corrupt(message: &str) -> DbError {
320    DbError::new(
321        DbErrorKind::Io(message.to_string()),
322        format!("corrupt database: {message}"),
323    )
324}
325
326fn u32_at(bytes: &[u8], offset: usize) -> Result<u32, DbError> {
327    let end = offset
328        .checked_add(4)
329        .ok_or_else(|| corrupt("offset overflow"))?;
330    let raw = bytes
331        .get(offset..end)
332        .ok_or_else(|| corrupt("database header is truncated"))?;
333    Ok(u32::from_le_bytes(raw.try_into().unwrap()))
334}
335
336fn u64_at(bytes: &[u8], offset: usize) -> Result<u64, DbError> {
337    let end = offset
338        .checked_add(8)
339        .ok_or_else(|| corrupt("offset overflow"))?;
340    let raw = bytes
341        .get(offset..end)
342        .ok_or_else(|| corrupt("database header is truncated"))?;
343    Ok(u64::from_le_bytes(raw.try_into().unwrap()))
344}
345
346#[cfg(test)]
347mod tests {
348    use super::*;
349    use crate::db::State;
350    use crate::engine;
351    use crate::sql::parser::parse;
352
353    #[test]
354    fn empty_snapshot_round_trips() {
355        let dir = std::env::temp_dir().join(format!("basalt-storage-{}", std::process::id()));
356        let _ = fs::remove_dir_all(&dir);
357        fs::create_dir_all(&dir).unwrap();
358        let path = dir.join("db");
359        write_snapshot(&path, &State::empty(), 7).unwrap();
360        let (loaded, generation) = read_snapshot(&path).unwrap();
361        assert!(loaded.tables.is_empty());
362        assert_eq!(generation, 7);
363        let _ = fs::remove_dir_all(dir);
364    }
365
366    #[test]
367    fn rewrites_an_existing_snapshot() {
368        let dir =
369            std::env::temp_dir().join(format!("basalt-storage-rewrite-{}", std::process::id()));
370        let _ = fs::remove_dir_all(&dir);
371        fs::create_dir_all(&dir).unwrap();
372        let path = dir.join("db");
373        write_snapshot(&path, &State::empty(), 1).unwrap();
374        write_snapshot(&path, &State::empty(), 2).unwrap();
375        let (_, generation) = read_snapshot(&path).unwrap();
376        assert_eq!(generation, 2);
377        let _ = fs::remove_dir_all(dir);
378    }
379
380    #[test]
381    fn table_snapshot_round_trips_tombstones_and_indexes() {
382        let dir = std::env::temp_dir().join(format!("basalt-storage-rows-{}", std::process::id()));
383        let _ = fs::remove_dir_all(&dir);
384        fs::create_dir_all(&dir).unwrap();
385        let path = dir.join("db");
386        let mut state = State::empty();
387        for sql in [
388            "CREATE TABLE t (id INTEGER PRIMARY KEY, value INTEGER)",
389            "INSERT INTO t VALUES (1, 10), (2, 20)",
390            "CREATE INDEX value_idx ON t(value)",
391            "DELETE FROM t WHERE id = 1",
392        ] {
393            let statement = &parse(sql).unwrap()[0];
394            engine::execute(&mut state, statement).unwrap();
395        }
396        write_snapshot(&path, &state, 4).unwrap();
397        let (loaded, generation) = read_snapshot(&path).unwrap();
398        assert_eq!(generation, 4);
399        let table = loaded.table("t").unwrap();
400        assert_eq!(table.row_count(), 1);
401        assert!(table.get_row(0).is_none());
402        assert_eq!(
403            table.get_row(1).unwrap()[0],
404            crate::types::Value::Integer(2)
405        );
406        assert!(table.index(1).is_some());
407        let _ = fs::remove_dir_all(dir);
408    }
409
410    #[test]
411    fn page_checksum_rejects_mutation() {
412        let dir = std::env::temp_dir().join(format!("basalt-storage-crc-{}", std::process::id()));
413        let _ = fs::remove_dir_all(&dir);
414        fs::create_dir_all(&dir).unwrap();
415        let path = dir.join("db");
416        write_snapshot(&path, &State::empty(), 0).unwrap();
417        let mut bytes = fs::read(&path).unwrap();
418        bytes[FILE_HEADER + PAGE_HEADER] ^= 1;
419        fs::write(&path, bytes).unwrap();
420        assert!(read_snapshot(&path).is_err());
421        let _ = fs::remove_dir_all(dir);
422    }
423
424    #[cfg(unix)]
425    #[test]
426    fn refuses_a_symbolic_link_snapshot() {
427        use std::os::unix::fs::symlink;
428
429        let dir =
430            std::env::temp_dir().join(format!("basalt-storage-symlink-{}", std::process::id()));
431        let _ = fs::remove_dir_all(&dir);
432        fs::create_dir_all(&dir).unwrap();
433        let target = dir.join("outside.db");
434        let path = dir.join("db");
435        write_snapshot(&target, &State::empty(), 0).unwrap();
436        symlink(&target, &path).unwrap();
437
438        let read_error = read_snapshot(&path).unwrap_err();
439        let write_error = write_snapshot(&path, &State::empty(), 1).unwrap_err();
440
441        assert!(read_error.message.contains("symbolic link"));
442        assert!(write_error.message.contains("symbolic link"));
443        let _ = fs::remove_dir_all(dir);
444    }
445}