1use 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;
18pub 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;
24pub 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
38pub 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 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
125pub 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
155pub(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
194pub 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}