#![forbid(unsafe_code)]
use jerrycan_core::{App, Error, Extension, Result};
use sea_orm::{ConnectionTrait, Database, DatabaseConnection, Statement, TransactionTrait};
pub const MIGRATION_ADVISORY_KEY: i64 = 0x6A_43_6D_69_67_00_00_01;
pub use sea_orm;
pub use sea_query;
pub use sea_query_binder;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Backend {
Sqlite,
Postgres,
}
#[derive(Clone)]
pub struct Db {
conn: DatabaseConnection,
backend: Backend,
url: String,
}
impl Db {
pub async fn connect(url: &str) -> Result<Self> {
let backend = if url.starts_with("postgres") {
Backend::Postgres
} else if url.starts_with("sqlite") {
Backend::Sqlite
} else {
return Err(Error::internal(format!(
"unsupported database url scheme: `{url}` (sqlite:// or postgres:// in v0)"
)));
};
let max = match backend {
Backend::Sqlite => 1,
Backend::Postgres => 5,
};
let mut opts = sea_orm::ConnectOptions::new(url.to_string());
opts.max_connections(max);
if backend == Backend::Sqlite {
opts.map_sqlx_sqlite_opts(|o| o.foreign_keys(true));
}
let conn = Database::connect(opts).await.map_err(db_error)?;
Ok(Self {
conn,
backend,
url: url.to_string(),
})
}
pub async fn from_env() -> Result<Self> {
let url = std::env::var("JERRYCAN_DATABASE_URL")
.unwrap_or_else(|_| "sqlite::memory:".to_string());
Self::connect(&url).await
}
pub fn conn(&self) -> &DatabaseConnection {
&self.conn
}
pub fn backend(&self) -> Backend {
self.backend
}
pub fn url(&self) -> &str {
&self.url
}
pub fn sql(&self, query: &str) -> String {
translate_placeholders(query, self.backend)
}
pub fn query_builder(&self) -> &'static dyn sea_query::QueryBuilder {
match self.backend {
Backend::Sqlite => &sea_query::SqliteQueryBuilder,
Backend::Postgres => &sea_query::PostgresQueryBuilder,
}
}
fn backend_db(&self) -> sea_orm::DatabaseBackend {
match self.backend {
Backend::Sqlite => sea_orm::DatabaseBackend::Sqlite,
Backend::Postgres => sea_orm::DatabaseBackend::Postgres,
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct Migration {
pub name: &'static str,
pub sqlite: &'static str,
pub postgres: &'static str,
}
#[derive(Debug, Clone)]
pub struct OwnedMigration {
pub name: String,
pub sqlite: String,
pub postgres: String,
}
impl Db {
pub async fn migrate(&self, migrations: &[Migration]) -> Result<Vec<String>> {
self.migrate_iter(migrations.iter().map(|m| (m.name, m.sqlite, m.postgres)))
.await
}
pub async fn migrate_owned(&self, migrations: &[OwnedMigration]) -> Result<Vec<String>> {
self.migrate_iter(
migrations
.iter()
.map(|m| (m.name.as_str(), m.sqlite.as_str(), m.postgres.as_str())),
)
.await
}
async fn migrate_iter<'a>(
&self,
items: impl Iterator<Item = (&'a str, &'a str, &'a str)>,
) -> Result<Vec<String>> {
let txn = self.conn.begin().await.map_err(db_error)?;
if self.backend == Backend::Postgres {
txn.execute(Statement::from_string(
sea_orm::DatabaseBackend::Postgres,
format!("SELECT pg_advisory_xact_lock({MIGRATION_ADVISORY_KEY})"),
))
.await
.map_err(db_error)?;
}
txn.execute_unprepared(
"CREATE TABLE IF NOT EXISTS _jerrycan_migrations (name TEXT PRIMARY KEY, applied_at TEXT NOT NULL)",
)
.await
.map_err(db_error)?;
let mut applied = Vec::new();
for (name, sqlite, postgres) in items {
let seen = txn
.query_one(Statement::from_sql_and_values(
self.backend_db(),
self.sql("SELECT name FROM _jerrycan_migrations WHERE name = ?"),
[name.into()],
))
.await
.map_err(db_error)?;
if seen.is_some() {
continue;
}
let statement = match self.backend {
Backend::Sqlite => sqlite,
Backend::Postgres => postgres,
};
txn.execute_unprepared(statement).await.map_err(|e| {
eprintln!("jerrycan-db: migration `{name}` failed");
db_error(e)
})?;
txn.execute(Statement::from_sql_and_values(
self.backend_db(),
self.sql("INSERT INTO _jerrycan_migrations (name, applied_at) VALUES (?, ?)"),
[name.into(), chrono_free_timestamp().into()],
))
.await
.map_err(db_error)?;
applied.push(name.to_string());
}
txn.commit().await.map_err(db_error)?;
Ok(applied)
}
}
fn chrono_free_timestamp() -> String {
let secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
format!("unix:{secs}")
}
pub fn translate_placeholders(query: &str, backend: Backend) -> String {
match backend {
Backend::Sqlite => query.to_string(),
Backend::Postgres => {
let mut out = String::with_capacity(query.len() + 8);
let mut n = 0;
for ch in query.chars() {
if ch == '?' {
n += 1;
out.push('$');
out.push_str(&n.to_string());
} else {
out.push(ch);
}
}
out
}
}
}
pub fn db_error(e: sea_orm::DbErr) -> Error {
eprintln!("jerrycan-db: {e}");
if matches!(
e.sql_err(),
Some(sea_orm::SqlErr::UniqueConstraintViolation(_))
) {
return Error::conflict("conflict: a row with this key already exists");
}
Error::new(
jerrycan_core::http::StatusCode::INTERNAL_SERVER_ERROR,
"JC0510",
"database error",
)
}
impl Extension for Db {
fn register(self, app: App) -> App {
app.provide(self)
}
}
pub use sqlx;
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn db_exposes_its_connection_url() {
let db = Db::connect("sqlite::memory:").await.unwrap();
assert_eq!(db.url(), "sqlite::memory:");
}
#[tokio::test]
async fn connects_and_executes_via_sea_orm() {
let db = Db::connect("sqlite::memory:").await.unwrap();
assert_eq!(db.backend(), Backend::Sqlite);
db.conn()
.execute_unprepared("CREATE TABLE t (id INTEGER PRIMARY KEY)")
.await
.unwrap();
}
#[test]
fn placeholder_translation_is_backend_aware() {
assert_eq!(
translate_placeholders("INSERT INTO t (a, b) VALUES (?, ?)", Backend::Postgres),
"INSERT INTO t (a, b) VALUES ($1, $2)"
);
assert_eq!(
translate_placeholders("INSERT INTO t (a, b) VALUES (?, ?)", Backend::Sqlite),
"INSERT INTO t (a, b) VALUES (?, ?)"
);
}
#[tokio::test]
async fn from_env_defaults_to_sqlite_memory() {
let db = Db::from_env().await.unwrap();
assert_eq!(db.backend(), Backend::Sqlite);
}
#[test]
fn db_errors_are_jc0510_and_leak_nothing() {
let e = db_error(sea_orm::DbErr::Custom("boom".into()));
assert_eq!(e.code(), "JC0510");
assert_eq!(e.message(), "database error");
}
#[tokio::test]
async fn sea_query_builds_and_executes_via_the_connection() {
use sea_query::{Alias, Expr, Query};
let db = Db::connect("sqlite::memory:").await.unwrap();
db.conn()
.execute_unprepared("CREATE TABLE sq (id INTEGER PRIMARY KEY, title TEXT NOT NULL)")
.await
.unwrap();
let (sql, values) = Query::insert()
.into_table(Alias::new("sq"))
.columns([Alias::new("id"), Alias::new("title")])
.values_panic([7.into(), "hello".into()])
.returning(Query::returning().columns([Alias::new("id")]))
.build_any(db.query_builder());
let row = db
.conn()
.query_one(Statement::from_sql_and_values(db.backend_db(), sql, values))
.await
.unwrap()
.expect("RETURNING id row");
assert_eq!(
row.try_get::<i64>("", "id").unwrap(),
7,
"RETURNING id round-trips"
);
let (sql, values) = Query::select()
.columns([Alias::new("id"), Alias::new("title")])
.from(Alias::new("sq"))
.and_where(Expr::col(Alias::new("id")).eq(7))
.build_any(db.query_builder());
let row = db
.conn()
.query_one(Statement::from_sql_and_values(db.backend_db(), sql, values))
.await
.unwrap()
.expect("select row");
assert_eq!(row.try_get::<String>("", "title").unwrap(), "hello");
}
#[tokio::test]
async fn unique_violations_map_to_409_conflict() {
let db = Db::connect("sqlite::memory:").await.unwrap();
db.conn()
.execute_unprepared("CREATE TABLE u (id INTEGER PRIMARY KEY, t TEXT)")
.await
.unwrap();
db.conn()
.execute_unprepared("INSERT INTO u VALUES (1, 'a')")
.await
.unwrap();
let dup = db
.conn()
.execute_unprepared("INSERT INTO u VALUES (1, 'b')")
.await
.expect_err("duplicate pk must fail");
let e = db_error(dup);
assert_eq!(e.code(), "JC0409");
assert_eq!(e.status().as_u16(), 409);
assert!(!e.message().contains("sqlite"), "{}", e.message());
}
#[tokio::test]
async fn sqlite_foreign_keys_are_enforced_through_the_pool() {
let db = Db::connect("sqlite::memory:").await.unwrap();
let row = db
.conn()
.query_one(Statement::from_string(
sea_orm::DatabaseBackend::Sqlite,
"PRAGMA foreign_keys",
))
.await
.unwrap()
.expect("PRAGMA foreign_keys returns a row");
let on: i64 = row
.try_get::<i64>("", "foreign_keys")
.or_else(|_| row.try_get::<i32>("", "foreign_keys").map(i64::from))
.unwrap();
assert_eq!(on, 1, "foreign_keys must be ON through the pool");
db.conn()
.execute_unprepared("CREATE TABLE parents (id INTEGER PRIMARY KEY)")
.await
.unwrap();
db.conn()
.execute_unprepared(
"CREATE TABLE children (id INTEGER PRIMARY KEY, \
parent_id INTEGER NOT NULL REFERENCES parents(id) ON DELETE CASCADE)",
)
.await
.unwrap();
db.conn()
.execute_unprepared("INSERT INTO parents (id) VALUES (1)")
.await
.unwrap();
let orphan = db
.conn()
.execute_unprepared("INSERT INTO children (id, parent_id) VALUES (10, 999)")
.await
.expect_err("orphan insert must violate the FK");
assert!(
matches!(
orphan.sql_err(),
Some(sea_orm::SqlErr::ForeignKeyConstraintViolation(_))
),
"must be an FK violation, got: {orphan}"
);
db.conn()
.execute_unprepared("INSERT INTO children (id, parent_id) VALUES (11, 1)")
.await
.unwrap();
db.conn()
.execute_unprepared("DELETE FROM parents WHERE id = 1")
.await
.unwrap();
let row = db
.conn()
.query_one(Statement::from_string(
sea_orm::DatabaseBackend::Sqlite,
"SELECT COUNT(*) AS n FROM children",
))
.await
.unwrap()
.unwrap();
let n: i64 = row
.try_get::<i64>("", "n")
.or_else(|_| row.try_get::<i32>("", "n").map(i64::from))
.unwrap();
assert_eq!(n, 0, "ON DELETE CASCADE must remove the child rows");
}
fn demo_migrations() -> Vec<Migration> {
vec![
Migration {
name: "0001_create_todos",
sqlite: "CREATE TABLE todos (id INTEGER PRIMARY KEY AUTOINCREMENT, title TEXT NOT NULL)",
postgres: "CREATE TABLE todos (id BIGSERIAL PRIMARY KEY, title TEXT NOT NULL)",
},
Migration {
name: "0002_add_done",
sqlite: "ALTER TABLE todos ADD COLUMN done BOOLEAN NOT NULL DEFAULT 0",
postgres: "ALTER TABLE todos ADD COLUMN done BOOLEAN NOT NULL DEFAULT FALSE",
},
]
}
#[tokio::test]
async fn migrations_apply_in_order_and_only_once() {
let db = Db::connect("sqlite::memory:").await.unwrap();
let applied = db.migrate(&demo_migrations()).await.unwrap();
assert_eq!(applied, vec!["0001_create_todos", "0002_add_done"]);
let applied = db.migrate(&demo_migrations()).await.unwrap();
assert!(applied.is_empty());
db.conn()
.execute_unprepared("INSERT INTO todos (title, done) VALUES ('x', 1)")
.await
.unwrap();
}
#[tokio::test]
async fn owned_migrations_apply_in_order_and_only_once() {
let db = Db::connect("sqlite::memory:").await.unwrap();
let owned = vec![
OwnedMigration {
name: "0001_create_todos".into(),
sqlite:
"CREATE TABLE todos (id INTEGER PRIMARY KEY AUTOINCREMENT, title TEXT NOT NULL)"
.into(),
postgres: "CREATE TABLE todos (id BIGSERIAL PRIMARY KEY, title TEXT NOT NULL)"
.into(),
},
OwnedMigration {
name: "0002_add_done".into(),
sqlite: "ALTER TABLE todos ADD COLUMN done BOOLEAN NOT NULL DEFAULT 0".into(),
postgres: "ALTER TABLE todos ADD COLUMN done BOOLEAN NOT NULL DEFAULT FALSE".into(),
},
];
let applied = db.migrate_owned(&owned).await.unwrap();
assert_eq!(applied, vec!["0001_create_todos", "0002_add_done"]);
let applied = db.migrate_owned(&owned).await.unwrap();
assert!(applied.is_empty());
}
#[tokio::test]
async fn transactions_roll_back_on_error() {
use sea_orm::TransactionTrait;
let db = Db::connect("sqlite::memory:").await.unwrap();
db.conn()
.execute_unprepared("CREATE TABLE t (id INTEGER PRIMARY KEY)")
.await
.unwrap();
let r = db
.conn()
.transaction::<_, (), sea_orm::DbErr>(|txn| {
Box::pin(async move {
txn.execute_unprepared("INSERT INTO t VALUES (1)").await?;
Err(sea_orm::DbErr::Custom("boom".into()))
})
})
.await;
assert!(r.is_err());
let rows = db
.conn()
.query_all(sea_orm::Statement::from_string(
sea_orm::DatabaseBackend::Sqlite,
"SELECT id FROM t",
))
.await
.unwrap();
assert!(rows.is_empty(), "rollback must leave no rows");
}
#[tokio::test]
async fn a_failing_migration_surfaces_jc0510_and_is_not_recorded() {
let db = Db::connect("sqlite::memory:").await.unwrap();
let bad = vec![Migration {
name: "0001_broken",
sqlite: "CREATE GARBAGE",
postgres: "CREATE GARBAGE",
}];
let err = db.migrate(&bad).await.unwrap_err();
assert_eq!(err.code(), "JC0510");
let good = vec![Migration {
name: "0001_broken",
sqlite: "CREATE TABLE ok (x BIGINT)",
postgres: "CREATE TABLE ok (x BIGINT)",
}];
let applied = db.migrate(&good).await.unwrap();
assert_eq!(applied, vec!["0001_broken"]);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
#[ignore = "needs a local postgres (set JERRYCAN_TEST_PG_URL)"]
async fn concurrent_migrators_do_not_race() {
let Ok(url) = std::env::var("JERRYCAN_TEST_PG_URL") else {
eprintln!("SKIP: JERRYCAN_TEST_PG_URL not set");
return;
};
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let table = format!("mig_race_{nanos}");
let name = format!("{table}_0001");
let migrations = vec![Migration {
name: Box::leak(name.clone().into_boxed_str()),
sqlite: "",
postgres: Box::leak(
format!("CREATE TABLE {table} (id BIGSERIAL PRIMARY KEY, v TEXT NOT NULL)")
.into_boxed_str(),
),
}];
let migrations = std::sync::Arc::new(migrations);
let mut handles = Vec::new();
for _ in 0..8 {
let url = url.clone();
let migrations = migrations.clone();
handles.push(tokio::spawn(async move {
let db = Db::connect(&url).await.expect("connect");
db.migrate(&migrations).await
}));
}
let mut total_applied = 0usize;
for h in handles {
let applied = h.await.expect("task").expect("migrate must not error");
total_applied += applied.len();
}
assert_eq!(
total_applied, 1,
"exactly one migrator applies the migration; the rest see it recorded"
);
let db = Db::connect(&url).await.unwrap();
db.conn()
.execute_unprepared(&format!("INSERT INTO {table} (v) VALUES ('ok')"))
.await
.unwrap();
db.conn()
.execute_unprepared(&format!("DROP TABLE {table}"))
.await
.unwrap();
}
async fn atomic_reserve_one(db: &Db, id: i64) -> Result<bool> {
let stmt = Statement::from_sql_and_values(
db.backend_db(),
db.sql("UPDATE seats SET used = used + 1 WHERE id = ? AND used + 1 <= capacity"),
[id.into()],
);
let res = db.conn().execute(stmt).await.map_err(db_error)?;
Ok(res.rows_affected() == 1)
}
async fn seats_used(db: &Db, id: i64) -> i64 {
let row = db
.conn()
.query_one(Statement::from_sql_and_values(
db.backend_db(),
db.sql("SELECT used FROM seats WHERE id = ?"),
[id.into()],
))
.await
.unwrap()
.expect("seats row");
row.try_get::<i64>("", "used")
.or_else(|_| row.try_get::<i32>("", "used").map(i64::from))
.unwrap()
}
async fn naive_reserve_one(db: &Db, resource_id: i64, capacity: i64) -> Result<bool> {
let row = db
.conn()
.query_one(Statement::from_sql_and_values(
db.backend_db(),
db.sql("SELECT COUNT(*) AS n FROM bookings WHERE resource_id = ?"),
[resource_id.into()],
))
.await
.map_err(db_error)?
.expect("count row");
let count = row
.try_get::<i64>("", "n")
.or_else(|_| row.try_get::<i32>("", "n").map(i64::from))
.unwrap();
if count >= capacity {
return Ok(false);
}
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
db.conn()
.execute(Statement::from_sql_and_values(
db.backend_db(),
db.sql("INSERT INTO bookings (resource_id) VALUES (?)"),
[resource_id.into()],
))
.await
.map_err(db_error)?;
Ok(true)
}
async fn booking_count(db: &Db, resource_id: i64) -> i64 {
let row = db
.conn()
.query_one(Statement::from_sql_and_values(
db.backend_db(),
db.sql("SELECT COUNT(*) AS n FROM bookings WHERE resource_id = ?"),
[resource_id.into()],
))
.await
.unwrap()
.expect("count row");
row.try_get::<i64>("", "n")
.or_else(|_| row.try_get::<i32>("", "n").map(i64::from))
.unwrap()
}
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
async fn atomic_reservation_reserves_exactly_capacity_on_sqlite() {
let db = Db::connect("sqlite::memory:").await.unwrap();
db.conn()
.execute_unprepared(
"CREATE TABLE seats (id INTEGER PRIMARY KEY, \
used INTEGER NOT NULL DEFAULT 0, capacity INTEGER NOT NULL)",
)
.await
.unwrap();
db.conn()
.execute_unprepared("INSERT INTO seats (id, used, capacity) VALUES (1, 0, 5)")
.await
.unwrap();
let capacity = 5i64;
let contenders = 40;
let mut handles = Vec::new();
for _ in 0..contenders {
let db = db.clone();
handles.push(tokio::spawn(async move {
atomic_reserve_one(&db, 1).await.unwrap()
}));
}
let mut reserved = 0i64;
for h in handles {
if h.await.unwrap() {
reserved += 1;
}
}
assert_eq!(reserved, capacity, "atomic reserve grants exactly capacity");
assert_eq!(
seats_used(&db, 1).await,
capacity,
"used must never exceed capacity"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 8)]
#[ignore = "needs a local postgres (set JERRYCAN_TEST_PG_URL)"]
async fn atomic_reservation_beats_the_oversell_race_on_postgres() {
let Ok(url) = std::env::var("JERRYCAN_TEST_PG_URL") else {
eprintln!("SKIP: JERRYCAN_TEST_PG_URL not set");
return;
};
let capacity = 5i64;
let contenders = 40;
let db = Db::connect(&url).await.unwrap();
db.conn()
.execute_unprepared(
"CREATE TABLE bookings (id BIGSERIAL PRIMARY KEY, resource_id BIGINT NOT NULL)",
)
.await
.unwrap();
db.conn()
.execute_unprepared(
"CREATE TABLE seats (id BIGINT PRIMARY KEY, \
used BIGINT NOT NULL DEFAULT 0, capacity BIGINT NOT NULL)",
)
.await
.unwrap();
db.conn()
.execute_unprepared("INSERT INTO seats (id, used, capacity) VALUES (1, 0, 5)")
.await
.unwrap();
let mut handles = Vec::new();
for _ in 0..contenders {
let db = db.clone();
handles.push(tokio::spawn(async move {
naive_reserve_one(&db, 1, capacity).await.unwrap()
}));
}
for h in handles {
let _ = h.await.unwrap();
}
let oversold = booking_count(&db, 1).await;
assert!(
oversold > capacity,
"naive read-then-insert must oversell on Postgres: got {oversold} bookings \
for capacity {capacity}"
);
let mut handles = Vec::new();
for _ in 0..contenders {
let db = db.clone();
handles.push(tokio::spawn(async move {
atomic_reserve_one(&db, 1).await.unwrap()
}));
}
let mut reserved = 0i64;
for h in handles {
if h.await.unwrap() {
reserved += 1;
}
}
assert_eq!(
reserved, capacity,
"atomic reserve grants exactly capacity on Postgres"
);
assert_eq!(
seats_used(&db, 1).await,
capacity,
"used must never exceed capacity"
);
db.conn()
.execute_unprepared("DROP TABLE bookings")
.await
.unwrap();
db.conn()
.execute_unprepared("DROP TABLE seats")
.await
.unwrap();
}
}