use std::ffi::OsString;
use std::fmt;
use std::fs::{File, OpenOptions};
use std::path::PathBuf;
use std::sync::Mutex;
use std::time::Duration;
use fs2::FileExt;
use rusqlite::{
params, Connection, ErrorCode, OptionalExtension, Transaction, TransactionBehavior,
};
const SCHEMA_MARKER_TABLE: &str = "_harn_sqlite_schema_versions";
const CREATE_SCHEMA_MARKER_TABLE: &str =
"CREATE TABLE IF NOT EXISTS main._harn_sqlite_schema_versions (
name TEXT PRIMARY KEY,
version INTEGER NOT NULL CHECK(version > 0)
);";
static TRANSIENT_INITIALIZATION_LOCK: Mutex<()> = Mutex::new(());
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct SchemaVersion {
name: &'static str,
version: i64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum SqliteContention {
Busy,
Locked,
}
#[must_use]
pub fn sqlite_contention(error: &rusqlite::Error) -> Option<SqliteContention> {
match error {
rusqlite::Error::SqliteFailure(failure, _) => match failure.code {
ErrorCode::DatabaseBusy => Some(SqliteContention::Busy),
ErrorCode::DatabaseLocked => Some(SqliteContention::Locked),
_ => None,
},
_ => None,
}
}
impl SchemaVersion {
#[must_use]
pub const fn new(name: &'static str, version: i64) -> Self {
assert!(!name.is_empty(), "SQLite schema name must not be empty");
assert!(version > 0, "SQLite schema version must be positive");
Self { name, version }
}
}
#[derive(Debug)]
#[non_exhaustive]
pub enum InitializationError<E> {
BusyTimeoutTooLarge { milliseconds: u128 },
BusyTimeout(rusqlite::Error),
JournalModeQuery(rusqlite::Error),
DatabasePathUnavailable,
FileBackedTransient { path: PathBuf },
DatabasePath {
path: PathBuf,
source: std::io::Error,
},
InitializationLockOpen {
path: PathBuf,
source: std::io::Error,
},
InitializationLockAcquire {
path: PathBuf,
source: std::io::Error,
},
WalPragma(rusqlite::Error),
WalNotEnabled { mode: String },
WalBusyNotWal { mode: String },
WalBusyQuery {
wal_error: Box<rusqlite::Error>,
query_error: Box<rusqlite::Error>,
},
Synchronous(rusqlite::Error),
SchemaReadiness(rusqlite::Error),
SchemaNotInitialized { name: &'static str, version: i64 },
NewerSchemaVersion {
name: &'static str,
stored: i64,
supported: i64,
},
Transaction(rusqlite::Error),
Initialize(E),
SchemaMarker(rusqlite::Error),
Commit(rusqlite::Error),
}
impl InitializationError<rusqlite::Error> {
#[must_use]
pub fn is_busy_or_locked(&self) -> bool {
if let Self::Initialize(error) = self {
is_sqlite_busy_or_locked(error)
} else {
initialization_stage_is_busy_or_locked(self)
}
}
}
impl<E: fmt::Display> fmt::Display for InitializationError<E> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::BusyTimeoutTooLarge { milliseconds } => write!(
f,
"busy_timeout {milliseconds}ms exceeds SQLite's maximum"
),
Self::BusyTimeout(error) => write!(f, "busy_timeout failed: {error}"),
Self::JournalModeQuery(error) => write!(f, "journal_mode query failed: {error}"),
Self::DatabasePathUnavailable => {
write!(f, "SQLite connection has no file-backed main database path")
}
Self::FileBackedTransient { path } => write!(
f,
"transient SQLite initialization requires a private non-file database, got {}",
path.display()
),
Self::DatabasePath { path, source } => write!(
f,
"could not resolve SQLite database path {}: {source}",
path.display()
),
Self::InitializationLockOpen { path, source } => write!(
f,
"could not open SQLite initialization lock {}: {source}",
path.display()
),
Self::InitializationLockAcquire { path, source } => write!(
f,
"could not acquire SQLite initialization lock {}: {source}",
path.display()
),
Self::WalPragma(error) => write!(f, "WAL journal_mode pragma failed: {error}"),
Self::WalNotEnabled { mode } => {
write!(f, "WAL journal_mode request returned {mode}")
}
Self::WalBusyNotWal { mode } => {
write!(f, "WAL journal_mode request left journal_mode at {mode}")
}
Self::WalBusyQuery {
wal_error,
query_error,
} => write!(
f,
"WAL journal_mode pragma failed: {wal_error}; journal_mode query also failed: {query_error}"
),
Self::Synchronous(error) => write!(f, "synchronous pragma failed: {error}"),
Self::SchemaReadiness(error) => write!(f, "schema readiness query failed: {error}"),
Self::SchemaNotInitialized { name, version } => {
write!(f, "SQLite schema {name} version {version} is not initialized")
}
Self::NewerSchemaVersion {
name,
stored,
supported,
} => write!(
f,
"SQLite schema {name} version {stored} is newer than supported version {supported}"
),
Self::Transaction(error) => write!(f, "schema transaction failed: {error}"),
Self::Initialize(error) => write!(f, "schema initialization failed: {error}"),
Self::SchemaMarker(error) => write!(f, "schema marker update failed: {error}"),
Self::Commit(error) => write!(f, "schema transaction commit failed: {error}"),
}
}
}
impl<E: std::error::Error + 'static> std::error::Error for InitializationError<E> {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match self {
Self::BusyTimeout(error)
| Self::JournalModeQuery(error)
| Self::WalPragma(error)
| Self::Synchronous(error)
| Self::SchemaReadiness(error)
| Self::Transaction(error)
| Self::SchemaMarker(error)
| Self::Commit(error) => Some(error),
Self::DatabasePath { source, .. }
| Self::InitializationLockOpen { source, .. }
| Self::InitializationLockAcquire { source, .. } => Some(source),
Self::WalBusyQuery { wal_error, .. } => Some(wal_error),
Self::Initialize(error) => Some(error),
Self::BusyTimeoutTooLarge { .. }
| Self::DatabasePathUnavailable
| Self::FileBackedTransient { .. }
| Self::SchemaNotInitialized { .. }
| Self::WalNotEnabled { .. }
| Self::WalBusyNotWal { .. }
| Self::NewerSchemaVersion { .. } => None,
}
}
}
pub fn initialize_file<E, F>(
connection: &Connection,
busy_timeout: Duration,
schema: SchemaVersion,
initialize: F,
) -> Result<(), InitializationError<E>>
where
F: FnOnce(&Transaction<'_>) -> Result<(), E>,
{
configure_busy_timeout(connection, busy_timeout)?;
if fast_path_is_ready(connection, schema)? {
return configure_connection(connection);
}
let _initialization_lock = acquire_initialization_lock(connection)?;
ensure_wal_journal_mode(connection)?;
configure_connection(connection)?;
initialize_schema(connection, schema, initialize)
}
fn fast_path_is_ready<E>(
connection: &Connection,
schema: SchemaVersion,
) -> Result<bool, InitializationError<E>> {
match is_wal_journal_mode(connection) {
Ok(true) => {}
Ok(false) => return Ok(false),
Err(error) if initialization_stage_is_busy_or_locked(&error) => return Ok(false),
Err(error) => return Err(error),
}
match schema_is_ready(connection, schema) {
Ok(ready) => Ok(ready),
Err(error) if initialization_stage_is_busy_or_locked(&error) => Ok(false),
Err(error) => Err(error),
}
}
fn initialization_stage_is_busy_or_locked<E>(error: &InitializationError<E>) -> bool {
match error {
InitializationError::BusyTimeout(error)
| InitializationError::JournalModeQuery(error)
| InitializationError::WalPragma(error)
| InitializationError::Synchronous(error)
| InitializationError::SchemaReadiness(error)
| InitializationError::Transaction(error)
| InitializationError::SchemaMarker(error)
| InitializationError::Commit(error) => is_sqlite_busy_or_locked(error),
InitializationError::WalBusyNotWal { .. } => true,
InitializationError::WalBusyQuery {
wal_error,
query_error,
} => is_sqlite_busy_or_locked(wal_error) || is_sqlite_busy_or_locked(query_error),
InitializationError::BusyTimeoutTooLarge { .. }
| InitializationError::DatabasePath { .. }
| InitializationError::DatabasePathUnavailable
| InitializationError::FileBackedTransient { .. }
| InitializationError::InitializationLockOpen { .. }
| InitializationError::InitializationLockAcquire { .. }
| InitializationError::SchemaNotInitialized { .. }
| InitializationError::NewerSchemaVersion { .. }
| InitializationError::WalNotEnabled { .. }
| InitializationError::Initialize(_) => false,
}
}
pub fn require_file_initialized<E>(
connection: &Connection,
busy_timeout: Duration,
schema: SchemaVersion,
) -> Result<(), InitializationError<E>> {
require_file_initialized_impl(connection, busy_timeout, schema, || {})
}
fn require_file_initialized_impl<E>(
connection: &Connection,
busy_timeout: Duration,
schema: SchemaVersion,
on_readiness_contention: impl FnOnce(),
) -> Result<(), InitializationError<E>> {
configure_busy_timeout(connection, busy_timeout)?;
if fast_path_is_ready(connection, schema)? {
return Ok(());
}
let _readiness_lock = acquire_readiness_lock(connection, schema, on_readiness_contention)?;
if is_wal_journal_mode(connection)? && schema_is_ready(connection, schema)? {
return Ok(());
}
Err(InitializationError::SchemaNotInitialized {
name: schema.name,
version: schema.version,
})
}
pub fn initialize_transient<E, F>(
connection: &Connection,
busy_timeout: Duration,
schema: SchemaVersion,
initialize: F,
) -> Result<(), InitializationError<E>>
where
F: FnOnce(&Transaction<'_>) -> Result<(), E>,
{
configure_busy_timeout(connection, busy_timeout)?;
if let Some(path) = main_database_path(connection) {
return Err(InitializationError::FileBackedTransient { path });
}
let _initialization_lock = TRANSIENT_INITIALIZATION_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
configure_connection(connection)?;
if schema_is_ready(connection, schema)? {
return Ok(());
}
initialize_schema(connection, schema, initialize)
}
fn initialize_schema<E, F>(
connection: &Connection,
schema: SchemaVersion,
initialize: F,
) -> Result<(), InitializationError<E>>
where
F: FnOnce(&Transaction<'_>) -> Result<(), E>,
{
let transaction = Transaction::new_unchecked(connection, TransactionBehavior::Immediate)
.map_err(InitializationError::Transaction)?;
transaction
.execute_batch(CREATE_SCHEMA_MARKER_TABLE)
.map_err(InitializationError::SchemaMarker)?;
if schema_marker_is_ready(&transaction, schema)? {
return transaction.commit().map_err(InitializationError::Commit);
}
initialize(&transaction).map_err(InitializationError::Initialize)?;
transaction
.execute(
"INSERT INTO main._harn_sqlite_schema_versions(name, version) VALUES (?1, ?2)
ON CONFLICT(name) DO UPDATE SET version = excluded.version",
params![schema.name, schema.version],
)
.map_err(InitializationError::SchemaMarker)?;
transaction.commit().map_err(InitializationError::Commit)
}
fn configure_busy_timeout<E>(
connection: &Connection,
busy_timeout: Duration,
) -> Result<(), InitializationError<E>> {
let milliseconds = busy_timeout.as_millis();
if milliseconds > i32::MAX as u128 {
return Err(InitializationError::BusyTimeoutTooLarge { milliseconds });
}
connection
.busy_timeout(busy_timeout)
.map_err(InitializationError::BusyTimeout)
}
fn configure_connection<E>(connection: &Connection) -> Result<(), InitializationError<E>> {
connection
.pragma_update(None, "synchronous", "NORMAL")
.map_err(InitializationError::Synchronous)
}
fn acquire_initialization_lock<E>(
connection: &Connection,
) -> Result<SqliteInitializationLock, InitializationError<E>> {
let path = initialization_lock_path(connection)?;
let file = OpenOptions::new()
.create(true)
.truncate(false)
.read(true)
.write(true)
.open(&path)
.map_err(|source| InitializationError::InitializationLockOpen {
path: path.clone(),
source,
})?;
file.lock_exclusive()
.map_err(|source| InitializationError::InitializationLockAcquire {
path: path.clone(),
source,
})?;
Ok(SqliteInitializationLock { file })
}
fn acquire_readiness_lock<E>(
connection: &Connection,
schema: SchemaVersion,
on_contention: impl FnOnce(),
) -> Result<SqliteInitializationLock, InitializationError<E>> {
let path = initialization_lock_path(connection)?;
let file = match OpenOptions::new().read(true).open(&path) {
Ok(file) => file,
Err(source) if source.kind() == std::io::ErrorKind::NotFound => {
return Err(InitializationError::SchemaNotInitialized {
name: schema.name,
version: schema.version,
});
}
Err(source) => {
return Err(InitializationError::InitializationLockOpen { path, source });
}
};
match FileExt::try_lock_shared(&file) {
Ok(()) => {}
Err(source) if source.kind() == std::io::ErrorKind::WouldBlock => {
on_contention();
FileExt::lock_shared(&file).map_err(|source| {
InitializationError::InitializationLockAcquire {
path: path.clone(),
source,
}
})?;
}
Err(source) => {
return Err(InitializationError::InitializationLockAcquire { path, source });
}
}
Ok(SqliteInitializationLock { file })
}
fn initialization_lock_path<E>(connection: &Connection) -> Result<PathBuf, InitializationError<E>> {
let database_path =
main_database_path(connection).ok_or(InitializationError::DatabasePathUnavailable)?;
let canonical = std::fs::canonicalize(&database_path).map_err(|source| {
InitializationError::DatabasePath {
path: database_path.clone(),
source,
}
})?;
let mut path = OsString::from(canonical.as_os_str());
path.push(".harn-init.lock");
Ok(PathBuf::from(path))
}
#[cfg(unix)]
fn main_database_path(connection: &Connection) -> Option<PathBuf> {
use std::ffi::{CStr, OsStr};
use std::os::unix::ffi::OsStrExt;
let filename = unsafe {
let pointer =
rusqlite::ffi::sqlite3_db_filename(connection.handle(), rusqlite::MAIN_DB.as_ptr());
(!pointer.is_null()).then(|| CStr::from_ptr(pointer).to_bytes())
}?;
(!filename.is_empty()).then(|| PathBuf::from(OsStr::from_bytes(filename)))
}
#[cfg(not(unix))]
fn main_database_path(connection: &Connection) -> Option<PathBuf> {
connection
.path()
.filter(|path| !path.is_empty())
.map(PathBuf::from)
}
struct SqliteInitializationLock {
file: File,
}
impl Drop for SqliteInitializationLock {
fn drop(&mut self) {
let _ = FileExt::unlock(&self.file);
}
}
fn schema_is_ready<E>(
connection: &Connection,
schema: SchemaVersion,
) -> Result<bool, InitializationError<E>> {
let marker_exists = connection
.query_row(
"SELECT EXISTS(
SELECT 1 FROM main.sqlite_schema WHERE type = 'table' AND name = ?1
)",
params![SCHEMA_MARKER_TABLE],
|row| row.get::<_, bool>(0),
)
.map_err(InitializationError::SchemaReadiness)?;
if !marker_exists {
return Ok(false);
}
schema_marker_is_ready(connection, schema)
}
fn schema_marker_is_ready<E>(
connection: &Connection,
schema: SchemaVersion,
) -> Result<bool, InitializationError<E>> {
let stored = connection
.query_row(
"SELECT version FROM main._harn_sqlite_schema_versions WHERE name = ?1",
params![schema.name],
|row| row.get::<_, i64>(0),
)
.optional()
.map_err(InitializationError::SchemaReadiness)?;
match stored {
Some(version) if version > schema.version => Err(InitializationError::NewerSchemaVersion {
name: schema.name,
stored: version,
supported: schema.version,
}),
Some(version) => Ok(version == schema.version),
None => Ok(false),
}
}
fn is_wal_journal_mode<E>(connection: &Connection) -> Result<bool, InitializationError<E>> {
current_journal_mode(connection)
.map(|mode| mode.eq_ignore_ascii_case("wal"))
.map_err(InitializationError::JournalModeQuery)
}
fn ensure_wal_journal_mode<E>(connection: &Connection) -> Result<(), InitializationError<E>> {
if is_wal_journal_mode(connection)? {
return Ok(());
}
match connection.query_row("PRAGMA journal_mode = WAL", [], |row| {
row.get::<_, String>(0)
}) {
Ok(mode) if mode.eq_ignore_ascii_case("wal") => Ok(()),
Ok(mode) => Err(InitializationError::WalNotEnabled { mode }),
Err(error) if is_sqlite_busy_or_locked(&error) => match current_journal_mode(connection) {
Ok(mode) if mode.eq_ignore_ascii_case("wal") => Ok(()),
Ok(mode) => Err(InitializationError::WalBusyNotWal { mode }),
Err(query_error) => Err(InitializationError::WalBusyQuery {
wal_error: Box::new(error),
query_error: Box::new(query_error),
}),
},
Err(error) => Err(InitializationError::WalPragma(error)),
}
}
fn current_journal_mode(connection: &Connection) -> Result<String, rusqlite::Error> {
connection.query_row("PRAGMA journal_mode", [], |row| row.get::<_, String>(0))
}
fn is_sqlite_busy_or_locked(error: &rusqlite::Error) -> bool {
sqlite_contention(error).is_some()
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;