1use rusqlite::{
9 params, Connection, ErrorCode, OptionalExtension, Transaction, TransactionBehavior,
10};
11use std::ffi::OsString;
12use std::fmt;
13use std::fs::{File, OpenOptions, TryLockError};
14use std::path::PathBuf;
15use std::sync::Mutex;
16use std::time::Duration;
17
18const SCHEMA_MARKER_TABLE: &str = "_harn_sqlite_schema_versions";
19const CREATE_SCHEMA_MARKER_TABLE: &str =
20 "CREATE TABLE IF NOT EXISTS main._harn_sqlite_schema_versions (
21 name TEXT PRIMARY KEY,
22 version INTEGER NOT NULL CHECK(version > 0)
23);";
24static TRANSIENT_INITIALIZATION_LOCK: Mutex<()> = Mutex::new(());
25
26#[derive(Clone, Copy, Debug, PartialEq, Eq)]
28pub struct SchemaVersion {
29 name: &'static str,
30 version: i64,
31}
32
33#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum SqliteContention {
36 Busy,
38 Locked,
40}
41
42#[must_use]
44pub fn sqlite_contention(error: &rusqlite::Error) -> Option<SqliteContention> {
45 match error {
46 rusqlite::Error::SqliteFailure(failure, _) => match failure.code {
47 ErrorCode::DatabaseBusy => Some(SqliteContention::Busy),
48 ErrorCode::DatabaseLocked => Some(SqliteContention::Locked),
49 _ => None,
50 },
51 _ => None,
52 }
53}
54
55impl SchemaVersion {
56 #[must_use]
60 pub const fn new(name: &'static str, version: i64) -> Self {
61 assert!(!name.is_empty(), "SQLite schema name must not be empty");
62 assert!(version > 0, "SQLite schema version must be positive");
63 Self { name, version }
64 }
65}
66
67#[derive(Debug)]
69#[non_exhaustive]
70pub enum InitializationError<E> {
71 BusyTimeoutTooLarge { milliseconds: u128 },
73 BusyTimeout(rusqlite::Error),
75 JournalModeQuery(rusqlite::Error),
77 DatabasePathUnavailable,
79 FileBackedTransient { path: PathBuf },
81 DatabasePath {
83 path: PathBuf,
84 source: std::io::Error,
85 },
86 InitializationLockOpen {
88 path: PathBuf,
89 source: std::io::Error,
90 },
91 InitializationLockAcquire {
93 path: PathBuf,
94 source: std::io::Error,
95 },
96 WalPragma(rusqlite::Error),
98 WalNotEnabled { mode: String },
100 WalBusyNotWal { mode: String },
102 WalBusyQuery {
104 wal_error: Box<rusqlite::Error>,
105 query_error: Box<rusqlite::Error>,
106 },
107 Synchronous(rusqlite::Error),
109 SchemaReadiness(rusqlite::Error),
111 SchemaNotInitialized { name: &'static str, version: i64 },
114 NewerSchemaVersion {
116 name: &'static str,
117 stored: i64,
118 supported: i64,
119 },
120 Transaction(rusqlite::Error),
122 Initialize(E),
124 SchemaMarker(rusqlite::Error),
126 Commit(rusqlite::Error),
128}
129
130impl InitializationError<rusqlite::Error> {
131 #[must_use]
133 pub fn is_busy_or_locked(&self) -> bool {
134 if let Self::Initialize(error) = self {
135 is_sqlite_busy_or_locked(error)
136 } else {
137 initialization_stage_is_busy_or_locked(self)
138 }
139 }
140}
141
142impl<E: fmt::Display> fmt::Display for InitializationError<E> {
143 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
144 match self {
145 Self::BusyTimeoutTooLarge { milliseconds } => write!(
146 f,
147 "busy_timeout {milliseconds}ms exceeds SQLite's maximum"
148 ),
149 Self::BusyTimeout(error) => write!(f, "busy_timeout failed: {error}"),
150 Self::JournalModeQuery(error) => write!(f, "journal_mode query failed: {error}"),
151 Self::DatabasePathUnavailable => {
152 write!(f, "SQLite connection has no file-backed main database path")
153 }
154 Self::FileBackedTransient { path } => write!(
155 f,
156 "transient SQLite initialization requires a private non-file database, got {}",
157 path.display()
158 ),
159 Self::DatabasePath { path, source } => write!(
160 f,
161 "could not resolve SQLite database path {}: {source}",
162 path.display()
163 ),
164 Self::InitializationLockOpen { path, source } => write!(
165 f,
166 "could not open SQLite initialization lock {}: {source}",
167 path.display()
168 ),
169 Self::InitializationLockAcquire { path, source } => write!(
170 f,
171 "could not acquire SQLite initialization lock {}: {source}",
172 path.display()
173 ),
174 Self::WalPragma(error) => write!(f, "WAL journal_mode pragma failed: {error}"),
175 Self::WalNotEnabled { mode } => {
176 write!(f, "WAL journal_mode request returned {mode}")
177 }
178 Self::WalBusyNotWal { mode } => {
179 write!(f, "WAL journal_mode request left journal_mode at {mode}")
180 }
181 Self::WalBusyQuery {
182 wal_error,
183 query_error,
184 } => write!(
185 f,
186 "WAL journal_mode pragma failed: {wal_error}; journal_mode query also failed: {query_error}"
187 ),
188 Self::Synchronous(error) => write!(f, "synchronous pragma failed: {error}"),
189 Self::SchemaReadiness(error) => write!(f, "schema readiness query failed: {error}"),
190 Self::SchemaNotInitialized { name, version } => {
191 write!(f, "SQLite schema {name} version {version} is not initialized")
192 }
193 Self::NewerSchemaVersion {
194 name,
195 stored,
196 supported,
197 } => write!(
198 f,
199 "SQLite schema {name} version {stored} is newer than supported version {supported}"
200 ),
201 Self::Transaction(error) => write!(f, "schema transaction failed: {error}"),
202 Self::Initialize(error) => write!(f, "schema initialization failed: {error}"),
203 Self::SchemaMarker(error) => write!(f, "schema marker update failed: {error}"),
204 Self::Commit(error) => write!(f, "schema transaction commit failed: {error}"),
205 }
206 }
207}
208
209impl<E: std::error::Error + 'static> std::error::Error for InitializationError<E> {
210 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
211 match self {
212 Self::BusyTimeout(error)
213 | Self::JournalModeQuery(error)
214 | Self::WalPragma(error)
215 | Self::Synchronous(error)
216 | Self::SchemaReadiness(error)
217 | Self::Transaction(error)
218 | Self::SchemaMarker(error)
219 | Self::Commit(error) => Some(error),
220 Self::DatabasePath { source, .. }
221 | Self::InitializationLockOpen { source, .. }
222 | Self::InitializationLockAcquire { source, .. } => Some(source),
223 Self::WalBusyQuery { wal_error, .. } => Some(wal_error),
224 Self::Initialize(error) => Some(error),
225 Self::BusyTimeoutTooLarge { .. }
226 | Self::DatabasePathUnavailable
227 | Self::FileBackedTransient { .. }
228 | Self::SchemaNotInitialized { .. }
229 | Self::WalNotEnabled { .. }
230 | Self::WalBusyNotWal { .. }
231 | Self::NewerSchemaVersion { .. } => None,
232 }
233 }
234}
235
236pub fn initialize_file<E, F>(
247 connection: &Connection,
248 busy_timeout: Duration,
249 schema: SchemaVersion,
250 initialize: F,
251) -> Result<(), InitializationError<E>>
252where
253 F: FnOnce(&Transaction<'_>) -> Result<(), E>,
254{
255 configure_busy_timeout(connection, busy_timeout)?;
256 if fast_path_is_ready(connection, schema)? {
257 return configure_connection(connection);
258 }
259
260 let _initialization_lock = acquire_initialization_lock(connection)?;
261 ensure_wal_journal_mode(connection)?;
262 configure_connection(connection)?;
263 initialize_schema(connection, schema, initialize)
264}
265
266fn fast_path_is_ready<E>(
267 connection: &Connection,
268 schema: SchemaVersion,
269) -> Result<bool, InitializationError<E>> {
270 match is_wal_journal_mode(connection) {
271 Ok(true) => {}
272 Ok(false) => return Ok(false),
273 Err(error) if initialization_stage_is_busy_or_locked(&error) => return Ok(false),
274 Err(error) => return Err(error),
275 }
276 match schema_is_ready(connection, schema) {
277 Ok(ready) => Ok(ready),
278 Err(error) if initialization_stage_is_busy_or_locked(&error) => Ok(false),
279 Err(error) => Err(error),
280 }
281}
282
283fn initialization_stage_is_busy_or_locked<E>(error: &InitializationError<E>) -> bool {
284 match error {
285 InitializationError::BusyTimeout(error)
286 | InitializationError::JournalModeQuery(error)
287 | InitializationError::WalPragma(error)
288 | InitializationError::Synchronous(error)
289 | InitializationError::SchemaReadiness(error)
290 | InitializationError::Transaction(error)
291 | InitializationError::SchemaMarker(error)
292 | InitializationError::Commit(error) => is_sqlite_busy_or_locked(error),
293 InitializationError::WalBusyNotWal { .. } => true,
294 InitializationError::WalBusyQuery {
295 wal_error,
296 query_error,
297 } => is_sqlite_busy_or_locked(wal_error) || is_sqlite_busy_or_locked(query_error),
298 InitializationError::BusyTimeoutTooLarge { .. }
299 | InitializationError::DatabasePath { .. }
300 | InitializationError::DatabasePathUnavailable
301 | InitializationError::FileBackedTransient { .. }
302 | InitializationError::InitializationLockOpen { .. }
303 | InitializationError::InitializationLockAcquire { .. }
304 | InitializationError::SchemaNotInitialized { .. }
305 | InitializationError::NewerSchemaVersion { .. }
306 | InitializationError::WalNotEnabled { .. }
307 | InitializationError::Initialize(_) => false,
308 }
309}
310
311pub fn require_file_initialized<E>(
320 connection: &Connection,
321 busy_timeout: Duration,
322 schema: SchemaVersion,
323) -> Result<(), InitializationError<E>> {
324 require_file_initialized_impl(connection, busy_timeout, schema, || {})
325}
326
327fn require_file_initialized_impl<E>(
328 connection: &Connection,
329 busy_timeout: Duration,
330 schema: SchemaVersion,
331 on_readiness_contention: impl FnOnce(),
332) -> Result<(), InitializationError<E>> {
333 configure_busy_timeout(connection, busy_timeout)?;
334 if fast_path_is_ready(connection, schema)? {
335 return Ok(());
336 }
337
338 let _readiness_lock = acquire_readiness_lock(connection, schema, on_readiness_contention)?;
339 if is_wal_journal_mode(connection)? && schema_is_ready(connection, schema)? {
340 return Ok(());
341 }
342 Err(InitializationError::SchemaNotInitialized {
343 name: schema.name,
344 version: schema.version,
345 })
346}
347
348pub fn initialize_transient<E, F>(
357 connection: &Connection,
358 busy_timeout: Duration,
359 schema: SchemaVersion,
360 initialize: F,
361) -> Result<(), InitializationError<E>>
362where
363 F: FnOnce(&Transaction<'_>) -> Result<(), E>,
364{
365 configure_busy_timeout(connection, busy_timeout)?;
366 if let Some(path) = main_database_path(connection) {
367 return Err(InitializationError::FileBackedTransient { path });
368 }
369 let _initialization_lock = TRANSIENT_INITIALIZATION_LOCK
370 .lock()
371 .unwrap_or_else(std::sync::PoisonError::into_inner);
372 configure_connection(connection)?;
373 if schema_is_ready(connection, schema)? {
374 return Ok(());
375 }
376 initialize_schema(connection, schema, initialize)
377}
378
379fn initialize_schema<E, F>(
380 connection: &Connection,
381 schema: SchemaVersion,
382 initialize: F,
383) -> Result<(), InitializationError<E>>
384where
385 F: FnOnce(&Transaction<'_>) -> Result<(), E>,
386{
387 let transaction = Transaction::new_unchecked(connection, TransactionBehavior::Immediate)
388 .map_err(InitializationError::Transaction)?;
389 transaction
390 .execute_batch(CREATE_SCHEMA_MARKER_TABLE)
391 .map_err(InitializationError::SchemaMarker)?;
392 if schema_marker_is_ready(&transaction, schema)? {
393 return transaction.commit().map_err(InitializationError::Commit);
394 }
395 initialize(&transaction).map_err(InitializationError::Initialize)?;
396 transaction
397 .execute(
398 "INSERT INTO main._harn_sqlite_schema_versions(name, version) VALUES (?1, ?2)
399 ON CONFLICT(name) DO UPDATE SET version = excluded.version",
400 params![schema.name, schema.version],
401 )
402 .map_err(InitializationError::SchemaMarker)?;
403 transaction.commit().map_err(InitializationError::Commit)
404}
405
406fn configure_busy_timeout<E>(
407 connection: &Connection,
408 busy_timeout: Duration,
409) -> Result<(), InitializationError<E>> {
410 let milliseconds = busy_timeout.as_millis();
411 if milliseconds > i32::MAX as u128 {
412 return Err(InitializationError::BusyTimeoutTooLarge { milliseconds });
413 }
414 connection
415 .busy_timeout(busy_timeout)
416 .map_err(InitializationError::BusyTimeout)
417}
418
419fn configure_connection<E>(connection: &Connection) -> Result<(), InitializationError<E>> {
420 connection
421 .pragma_update(None, "synchronous", "NORMAL")
422 .map_err(InitializationError::Synchronous)
423}
424
425fn acquire_initialization_lock<E>(
426 connection: &Connection,
427) -> Result<SqliteInitializationLock, InitializationError<E>> {
428 let path = initialization_lock_path(connection)?;
429 let file = OpenOptions::new()
430 .create(true)
431 .truncate(false)
432 .read(true)
433 .write(true)
434 .open(&path)
435 .map_err(|source| InitializationError::InitializationLockOpen {
436 path: path.clone(),
437 source,
438 })?;
439 file.lock()
440 .map_err(|source| InitializationError::InitializationLockAcquire {
441 path: path.clone(),
442 source,
443 })?;
444 Ok(SqliteInitializationLock { file })
445}
446
447fn acquire_readiness_lock<E>(
448 connection: &Connection,
449 schema: SchemaVersion,
450 on_contention: impl FnOnce(),
451) -> Result<SqliteInitializationLock, InitializationError<E>> {
452 let path = initialization_lock_path(connection)?;
453 let file = match OpenOptions::new().read(true).open(&path) {
454 Ok(file) => file,
455 Err(source) if source.kind() == std::io::ErrorKind::NotFound => {
456 return Err(InitializationError::SchemaNotInitialized {
457 name: schema.name,
458 version: schema.version,
459 });
460 }
461 Err(source) => {
462 return Err(InitializationError::InitializationLockOpen { path, source });
463 }
464 };
465 match file.try_lock_shared() {
466 Ok(()) => {}
467 Err(TryLockError::WouldBlock) => {
468 on_contention();
469 file.lock_shared().map_err(|source| {
470 InitializationError::InitializationLockAcquire {
471 path: path.clone(),
472 source,
473 }
474 })?;
475 }
476 Err(TryLockError::Error(source)) => {
477 return Err(InitializationError::InitializationLockAcquire { path, source });
478 }
479 }
480 Ok(SqliteInitializationLock { file })
481}
482
483fn initialization_lock_path<E>(connection: &Connection) -> Result<PathBuf, InitializationError<E>> {
484 let database_path =
485 main_database_path(connection).ok_or(InitializationError::DatabasePathUnavailable)?;
486 let canonical = std::fs::canonicalize(&database_path).map_err(|source| {
487 InitializationError::DatabasePath {
488 path: database_path.clone(),
489 source,
490 }
491 })?;
492 let mut path = OsString::from(canonical.as_os_str());
493 path.push(".harn-init.lock");
494 Ok(PathBuf::from(path))
495}
496
497#[cfg(unix)]
498fn main_database_path(connection: &Connection) -> Option<PathBuf> {
499 use std::ffi::{CStr, OsStr};
500 use std::os::unix::ffi::OsStrExt;
501
502 let filename = unsafe {
505 let pointer =
506 rusqlite::ffi::sqlite3_db_filename(connection.handle(), rusqlite::MAIN_DB.as_ptr());
507 (!pointer.is_null()).then(|| CStr::from_ptr(pointer).to_bytes())
508 }?;
509 (!filename.is_empty()).then(|| PathBuf::from(OsStr::from_bytes(filename)))
510}
511
512#[cfg(not(unix))]
513fn main_database_path(connection: &Connection) -> Option<PathBuf> {
514 connection
515 .path()
516 .filter(|path| !path.is_empty())
517 .map(PathBuf::from)
518}
519
520struct SqliteInitializationLock {
521 file: File,
522}
523
524impl Drop for SqliteInitializationLock {
525 fn drop(&mut self) {
526 let _ = self.file.unlock();
529 }
530}
531
532fn schema_is_ready<E>(
533 connection: &Connection,
534 schema: SchemaVersion,
535) -> Result<bool, InitializationError<E>> {
536 let marker_exists = connection
537 .query_row(
538 "SELECT EXISTS(
539 SELECT 1 FROM main.sqlite_schema WHERE type = 'table' AND name = ?1
540 )",
541 params![SCHEMA_MARKER_TABLE],
542 |row| row.get::<_, bool>(0),
543 )
544 .map_err(InitializationError::SchemaReadiness)?;
545 if !marker_exists {
546 return Ok(false);
547 }
548 schema_marker_is_ready(connection, schema)
549}
550
551fn schema_marker_is_ready<E>(
552 connection: &Connection,
553 schema: SchemaVersion,
554) -> Result<bool, InitializationError<E>> {
555 let stored = connection
556 .query_row(
557 "SELECT version FROM main._harn_sqlite_schema_versions WHERE name = ?1",
558 params![schema.name],
559 |row| row.get::<_, i64>(0),
560 )
561 .optional()
562 .map_err(InitializationError::SchemaReadiness)?;
563 match stored {
564 Some(version) if version > schema.version => Err(InitializationError::NewerSchemaVersion {
565 name: schema.name,
566 stored: version,
567 supported: schema.version,
568 }),
569 Some(version) => Ok(version == schema.version),
570 None => Ok(false),
571 }
572}
573
574fn is_wal_journal_mode<E>(connection: &Connection) -> Result<bool, InitializationError<E>> {
575 current_journal_mode(connection)
576 .map(|mode| mode.eq_ignore_ascii_case("wal"))
577 .map_err(InitializationError::JournalModeQuery)
578}
579
580fn ensure_wal_journal_mode<E>(connection: &Connection) -> Result<(), InitializationError<E>> {
581 if is_wal_journal_mode(connection)? {
582 return Ok(());
583 }
584 match connection.query_row("PRAGMA journal_mode = WAL", [], |row| {
585 row.get::<_, String>(0)
586 }) {
587 Ok(mode) if mode.eq_ignore_ascii_case("wal") => Ok(()),
588 Ok(mode) => Err(InitializationError::WalNotEnabled { mode }),
589 Err(error) if is_sqlite_busy_or_locked(&error) => match current_journal_mode(connection) {
590 Ok(mode) if mode.eq_ignore_ascii_case("wal") => Ok(()),
591 Ok(mode) => Err(InitializationError::WalBusyNotWal { mode }),
592 Err(query_error) => Err(InitializationError::WalBusyQuery {
593 wal_error: Box::new(error),
594 query_error: Box::new(query_error),
595 }),
596 },
597 Err(error) => Err(InitializationError::WalPragma(error)),
598 }
599}
600
601fn current_journal_mode(connection: &Connection) -> Result<String, rusqlite::Error> {
602 connection.query_row("PRAGMA journal_mode", [], |row| row.get::<_, String>(0))
603}
604
605fn is_sqlite_busy_or_locked(error: &rusqlite::Error) -> bool {
606 sqlite_contention(error).is_some()
607}
608
609#[cfg(test)]
610#[path = "tests.rs"]
611mod tests;