1use std::path::PathBuf;
17use std::sync::Arc;
18
19use rusqlite::Connection;
20use tokio::sync::Mutex;
21use tracing::info;
22
23#[derive(Clone)]
27pub enum DbConfig {
28 Path(PathBuf),
29 InMemory,
30}
31
32impl DbConfig {
33 pub fn new(path: impl Into<PathBuf>) -> Self {
34 Self::Path(path.into())
35 }
36
37 pub fn in_memory() -> Self {
38 Self::InMemory
39 }
40
41 fn is_in_memory(&self) -> bool {
42 matches!(self, Self::InMemory)
43 }
44}
45
46#[derive(Debug, thiserror::Error)]
47pub enum DbError {
48 #[error("open {path}: {source}")]
49 Open { path: String, source: rusqlite::Error },
50 #[error("pragma init: {0}")]
51 Pragma(rusqlite::Error),
52 #[error("migration: {0}")]
53 Migration(rusqlite::Error),
54}
55
56#[derive(Clone)]
62pub struct Db {
63 inner: Arc<Mutex<Connection>>,
64 label: Arc<str>,
65}
66
67impl Db {
68 pub fn open(cfg: &DbConfig) -> Result<Self, DbError> {
73 let (conn, label) = match cfg {
74 DbConfig::Path(path) => {
75 if let Some(parent) = path.parent() {
76 if !parent.as_os_str().is_empty() {
77 let _ = std::fs::create_dir_all(parent);
78 }
79 }
80 let conn = Connection::open(path)
81 .map_err(|e| DbError::Open { path: path.display().to_string(), source: e })?;
82 (conn, path.display().to_string())
83 }
84 DbConfig::InMemory => {
85 let conn = Connection::open_in_memory()
86 .map_err(|e| DbError::Open { path: ":memory:".into(), source: e })?;
87 (conn, ":memory:".to_owned())
88 }
89 };
90
91 if !cfg.is_in_memory() {
92 conn.pragma_update(None, "journal_mode", "WAL").map_err(DbError::Pragma)?;
93 }
94 conn.pragma_update(None, "foreign_keys", "ON").map_err(DbError::Pragma)?;
95
96 info!(path = %label, "db opened");
97 Ok(Self { inner: Arc::new(Mutex::new(conn)), label: Arc::from(label.as_str()) })
98 }
99
100 pub fn label(&self) -> &str {
101 &self.label
102 }
103
104 pub fn run_migrations_blocking(&self, runner: MigrationRunner) -> Result<(), DbError> {
112 let mut guard = self.inner.try_lock().expect(
113 "run_migrations_blocking called with a held lock -- call it before any other Db handle is in use",
114 );
115 runner.run(&mut guard).map_err(DbError::Migration)
116 }
117
118 pub fn read_blocking<F, T>(&self, f: F) -> rusqlite::Result<T>
123 where
124 F: FnOnce(&Connection) -> rusqlite::Result<T>,
125 {
126 let guard = self.inner.blocking_lock();
127 f(&guard)
128 }
129
130 pub fn write_blocking<F, T>(&self, f: F) -> rusqlite::Result<T>
133 where
134 F: FnOnce(&mut Connection) -> rusqlite::Result<T>,
135 {
136 let mut guard = self.inner.blocking_lock();
137 f(&mut guard)
138 }
139}
140
141#[derive(Debug, Clone)]
144pub struct Migration {
145 pub version: u32,
146 pub label: String,
147 pub sql: String,
148}
149
150impl Migration {
151 pub fn new(version: u32, label: impl Into<String>, sql: impl Into<String>) -> Self {
152 Self { version, label: label.into(), sql: sql.into() }
153 }
154}
155
156pub struct MigrationRunner {
157 migrations: Vec<Migration>,
158}
159
160impl MigrationRunner {
161 pub fn new(migrations: Vec<Migration>) -> Self {
165 for w in migrations.windows(2) {
166 assert!(
167 w[0].version < w[1].version,
168 "migrations must be in strictly increasing version order: v{} ({:?}) >= v{} ({:?})",
169 w[0].version,
170 w[0].label,
171 w[1].version,
172 w[1].label,
173 );
174 }
175 Self { migrations }
176 }
177
178 pub fn run(self, conn: &mut Connection) -> rusqlite::Result<()> {
187 let prev_busy_ms: i64 = conn.query_row("PRAGMA busy_timeout;", [], |r| r.get(0)).unwrap_or(0);
188 conn.pragma_update(None, "busy_timeout", 30_000)?;
189
190 let result = (|| -> rusqlite::Result<()> {
191 let tx = conn.transaction_with_behavior(rusqlite::TransactionBehavior::Exclusive)?;
192
193 tx.execute_batch(
194 "CREATE TABLE IF NOT EXISTS schema_migrations (
195 version INTEGER PRIMARY KEY,
196 label TEXT NOT NULL,
197 applied_at TEXT NOT NULL DEFAULT (datetime('now'))
198 );",
199 )?;
200
201 for m in self.migrations {
202 let already: Option<u32> = tx
203 .query_row("SELECT version FROM schema_migrations WHERE version = ?1", [m.version], |r| r.get(0))
204 .ok();
205 if already.is_some() {
206 continue;
207 }
208 tx.execute_batch(&m.sql)?;
209 tx.execute(
210 "INSERT INTO schema_migrations (version, label) VALUES (?1, ?2)",
211 rusqlite::params![m.version, m.label],
212 )?;
213 info!(version = m.version, label = %m.label, "applied migration");
214 }
215
216 tx.commit()?;
217 Ok(())
218 })();
219
220 let _ = conn.pragma_update(None, "busy_timeout", prev_busy_ms);
221 result
222 }
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228
229 #[test]
230 fn in_memory_open_sets_foreign_keys_on() {
231 let db = Db::open(&DbConfig::in_memory()).expect("opens");
232 let fk: u32 =
233 db.read_blocking(|c| c.query_row("PRAGMA foreign_keys;", [], |r| r.get(0))).expect("reads pragma");
234 assert_eq!(fk, 1);
235 assert_eq!(db.label(), ":memory:");
236 }
237
238 #[test]
239 #[should_panic(expected = "strictly increasing")]
240 fn out_of_order_migrations_panic() {
241 let _ = MigrationRunner::new(vec![Migration::new(2, "b", "SELECT 1;"), Migration::new(1, "a", "SELECT 1;")]);
242 }
243
244 #[test]
245 fn migrations_are_idempotent_when_run_twice() {
246 let db = Db::open(&DbConfig::in_memory()).expect("opens");
247 let migrations = || vec![Migration::new(1, "t", "CREATE TABLE t (n INTEGER);")];
248 db.run_migrations_blocking(MigrationRunner::new(migrations())).expect("first run applies the schema");
249 db.run_migrations_blocking(MigrationRunner::new(migrations())).expect("second run is a no-op");
250 db.write_blocking(|c| c.execute("INSERT INTO t VALUES (1)", []).map(|_| ()))
251 .expect("table is still usable after a repeat migration run");
252 }
253
254 #[test]
255 fn appending_a_new_migration_applies_only_the_new_one() {
256 let db = Db::open(&DbConfig::in_memory()).expect("opens");
257 db.run_migrations_blocking(MigrationRunner::new(vec![Migration::new(
258 1,
259 "users",
260 "CREATE TABLE users (id INTEGER PRIMARY KEY);",
261 )]))
262 .expect("v1 applies");
263
264 db.run_migrations_blocking(MigrationRunner::new(vec![
265 Migration::new(1, "users", "SELECT 1;"),
266 Migration::new(2, "add_name", "ALTER TABLE users ADD COLUMN name TEXT;"),
267 ]))
268 .expect("v2 applies on top");
269
270 let count: u32 = db
271 .read_blocking(|c| c.query_row("SELECT count(*) FROM schema_migrations", [], |r| r.get(0)))
272 .expect("reads the migration ledger");
273 assert_eq!(count, 2);
274 }
275}