use std::path::PathBuf;
use std::sync::Arc;
use rusqlite::Connection;
use tokio::sync::Mutex;
use tracing::info;
#[derive(Clone)]
pub enum DbConfig {
Path(PathBuf),
InMemory,
}
impl DbConfig {
pub fn new(path: impl Into<PathBuf>) -> Self {
Self::Path(path.into())
}
pub fn in_memory() -> Self {
Self::InMemory
}
fn is_in_memory(&self) -> bool {
matches!(self, Self::InMemory)
}
}
#[derive(Debug, thiserror::Error)]
pub enum DbError {
#[error("open {path}: {source}")]
Open { path: String, source: rusqlite::Error },
#[error("pragma init: {0}")]
Pragma(rusqlite::Error),
#[error("migration: {0}")]
Migration(rusqlite::Error),
}
#[derive(Clone)]
pub struct Db {
inner: Arc<Mutex<Connection>>,
label: Arc<str>,
}
impl Db {
pub fn open(cfg: &DbConfig) -> Result<Self, DbError> {
let (conn, label) = match cfg {
DbConfig::Path(path) => {
if let Some(parent) = path.parent() {
if !parent.as_os_str().is_empty() {
let _ = std::fs::create_dir_all(parent);
}
}
let conn = Connection::open(path)
.map_err(|e| DbError::Open { path: path.display().to_string(), source: e })?;
(conn, path.display().to_string())
}
DbConfig::InMemory => {
let conn = Connection::open_in_memory()
.map_err(|e| DbError::Open { path: ":memory:".into(), source: e })?;
(conn, ":memory:".to_owned())
}
};
if !cfg.is_in_memory() {
conn.pragma_update(None, "journal_mode", "WAL").map_err(DbError::Pragma)?;
}
conn.pragma_update(None, "foreign_keys", "ON").map_err(DbError::Pragma)?;
info!(path = %label, "db opened");
Ok(Self { inner: Arc::new(Mutex::new(conn)), label: Arc::from(label.as_str()) })
}
pub fn label(&self) -> &str {
&self.label
}
pub fn run_migrations_blocking(&self, runner: MigrationRunner) -> Result<(), DbError> {
let mut guard = self.inner.try_lock().expect(
"run_migrations_blocking called with a held lock -- call it before any other Db handle is in use",
);
runner.run(&mut guard).map_err(DbError::Migration)
}
pub fn read_blocking<F, T>(&self, f: F) -> rusqlite::Result<T>
where
F: FnOnce(&Connection) -> rusqlite::Result<T>,
{
let guard = self.inner.blocking_lock();
f(&guard)
}
pub fn write_blocking<F, T>(&self, f: F) -> rusqlite::Result<T>
where
F: FnOnce(&mut Connection) -> rusqlite::Result<T>,
{
let mut guard = self.inner.blocking_lock();
f(&mut guard)
}
}
#[derive(Debug, Clone)]
pub struct Migration {
pub version: u32,
pub label: String,
pub sql: String,
}
impl Migration {
pub fn new(version: u32, label: impl Into<String>, sql: impl Into<String>) -> Self {
Self { version, label: label.into(), sql: sql.into() }
}
}
pub struct MigrationRunner {
migrations: Vec<Migration>,
}
impl MigrationRunner {
pub fn new(migrations: Vec<Migration>) -> Self {
for w in migrations.windows(2) {
assert!(
w[0].version < w[1].version,
"migrations must be in strictly increasing version order: v{} ({:?}) >= v{} ({:?})",
w[0].version,
w[0].label,
w[1].version,
w[1].label,
);
}
Self { migrations }
}
pub fn run(self, conn: &mut Connection) -> rusqlite::Result<()> {
let prev_busy_ms: i64 = conn.query_row("PRAGMA busy_timeout;", [], |r| r.get(0)).unwrap_or(0);
conn.pragma_update(None, "busy_timeout", 30_000)?;
let result = (|| -> rusqlite::Result<()> {
let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Exclusive)?;
tx.execute_batch(
"CREATE TABLE IF NOT EXISTS schema_migrations (
version INTEGER PRIMARY KEY,
label TEXT NOT NULL,
applied_at TEXT NOT NULL DEFAULT (datetime('now'))
);",
)?;
for m in self.migrations {
let already: Option<u32> = tx
.query_row("SELECT version FROM schema_migrations WHERE version = ?1", [m.version], |r| r.get(0))
.ok();
if already.is_some() {
continue;
}
tx.execute_batch(&m.sql)?;
tx.execute(
"INSERT INTO schema_migrations (version, label) VALUES (?1, ?2)",
rusqlite::params![m.version, m.label],
)?;
info!(version = m.version, label = %m.label, "applied migration");
}
tx.commit()?;
Ok(())
})();
let _ = conn.pragma_update(None, "busy_timeout", prev_busy_ms);
result
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn in_memory_open_sets_foreign_keys_on() {
let db = Db::open(&DbConfig::in_memory()).expect("opens");
let fk: u32 =
db.read_blocking(|c| c.query_row("PRAGMA foreign_keys;", [], |r| r.get(0))).expect("reads pragma");
assert_eq!(fk, 1);
assert_eq!(db.label(), ":memory:");
}
#[test]
#[should_panic(expected = "strictly increasing")]
fn out_of_order_migrations_panic() {
let _ = MigrationRunner::new(vec![Migration::new(2, "b", "SELECT 1;"), Migration::new(1, "a", "SELECT 1;")]);
}
#[test]
fn migrations_are_idempotent_when_run_twice() {
let db = Db::open(&DbConfig::in_memory()).expect("opens");
let migrations = || vec![Migration::new(1, "t", "CREATE TABLE t (n INTEGER);")];
db.run_migrations_blocking(MigrationRunner::new(migrations())).expect("first run applies the schema");
db.run_migrations_blocking(MigrationRunner::new(migrations())).expect("second run is a no-op");
db.write_blocking(|c| c.execute("INSERT INTO t VALUES (1)", []).map(|_| ()))
.expect("table is still usable after a repeat migration run");
}
#[test]
fn appending_a_new_migration_applies_only_the_new_one() {
let db = Db::open(&DbConfig::in_memory()).expect("opens");
db.run_migrations_blocking(MigrationRunner::new(vec![Migration::new(
1,
"users",
"CREATE TABLE users (id INTEGER PRIMARY KEY);",
)]))
.expect("v1 applies");
db.run_migrations_blocking(MigrationRunner::new(vec![
Migration::new(1, "users", "SELECT 1;"),
Migration::new(2, "add_name", "ALTER TABLE users ADD COLUMN name TEXT;"),
]))
.expect("v2 applies on top");
let count: u32 = db
.read_blocking(|c| c.query_row("SELECT count(*) FROM schema_migrations", [], |r| r.get(0)))
.expect("reads the migration ledger");
assert_eq!(count, 2);
}
}