use std::collections::{HashMap, HashSet};
use super::{Db, Dialect, now, quote, script, sql};
use anyhow::{Context, anyhow, bail};
const TABLE: &str = "renox_migrations";
macro_rules! framework_migration {
($dir:literal, $name:literal) => {
$crate::db::Migration::new(
$name,
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/migrations/",
$dir,
"/",
$name,
".up.sql"
)),
Some(include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/migrations/",
$dir,
"/",
$name,
".down.sql"
))),
)
.postgres(
include_str!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/migrations/",
$dir,
"/",
$name,
".postgres.up.sql"
)),
None,
)
};
}
pub(crate) use framework_migration;
#[derive(Debug, Clone, Copy)]
#[non_exhaustive]
pub struct Migration {
name: &'static str,
up: &'static str,
down: Option<&'static str>,
sqlite: Option<Scripts>,
postgres: Option<Scripts>,
}
#[derive(Debug, Clone, Copy)]
struct Scripts {
up: &'static str,
down: Option<&'static str>,
}
impl Migration {
pub const fn new(name: &'static str, up: &'static str, down: Option<&'static str>) -> Self {
Self {
name,
up,
down,
sqlite: None,
postgres: None,
}
}
pub const fn sqlite(mut self, up: &'static str, down: Option<&'static str>) -> Self {
self.sqlite = Some(Scripts { up, down });
self
}
pub const fn postgres(mut self, up: &'static str, down: Option<&'static str>) -> Self {
self.postgres = Some(Scripts { up, down });
self
}
pub const fn name(&self) -> &'static str {
self.name
}
fn own(&self, dialect: Dialect) -> Option<&Scripts> {
match dialect {
Dialect::Sqlite => self.sqlite.as_ref(),
Dialect::Postgres => self.postgres.as_ref(),
}
}
pub fn up_for(&self, dialect: Dialect) -> &'static str {
self.own(dialect).map_or(self.up, |own| own.up)
}
pub fn down_for(&self, dialect: Dialect) -> Option<&'static str> {
self.own(dialect).and_then(|own| own.down).or(self.down)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct MigrationStatus {
pub name: String,
pub batch: Option<i64>,
pub missing: bool,
pub changed: bool,
}
const NO_TRANSACTION: &str = "-- renox:no-transaction";
fn runs_outside_transaction(sql: &str) -> bool {
if sql.lines().any(|line| line.trim() == NO_TRANSACTION) {
return true;
}
let upper = sql.to_ascii_uppercase();
uses_concurrently(sql)
|| upper.split(';').any(|statement| {
let statement = statement.trim();
statement == "BEGIN"
|| statement.starts_with("BEGIN TRANSACTION")
|| statement.starts_with("BEGIN IMMEDIATE")
})
}
fn uses_concurrently(sql: &str) -> bool {
sql.split(|c: char| !(c.is_ascii_alphanumeric() || c == '_'))
.any(|word| word.eq_ignore_ascii_case("CONCURRENTLY"))
}
async fn run_each(db: &Db, sql: &str) -> Result<(), super::DbError> {
if db.dialect() == Dialect::Postgres && uses_concurrently(sql) {
for statement in statements(sql) {
script(db, statement).await?;
}
Ok(())
} else {
script(db, sql).await.map(|_| ())
}
}
fn statements(sql: &str) -> Vec<&str> {
let bytes = sql.as_bytes();
let mut out = Vec::new();
let (mut start, mut i) = (0, 0);
while i < bytes.len() {
match bytes[i] {
quote @ (b'\'' | b'"') => {
i += 1;
while i < bytes.len() && bytes[i] != quote {
i += 1;
}
}
b'-' if bytes.get(i + 1) == Some(&b'-') => {
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
}
b'/' if bytes.get(i + 1) == Some(&b'*') => {
i = sql[i + 2..]
.find("*/")
.map_or(bytes.len(), |end| i + 2 + end + 1);
}
b'$' => {
let tag_end = sql[i + 1..]
.find(|c: char| !(c.is_alphanumeric() || c == '_'))
.map(|n| i + 1 + n);
if let Some(end) = tag_end.filter(|&end| bytes[end] == b'$') {
let tag = &sql[i..=end];
i = sql[end + 1..]
.find(tag)
.map_or(bytes.len(), |close| end + 1 + close + tag.len() - 1);
}
}
b';' => {
out.push(&sql[start..i]);
start = i + 1;
}
_ => {}
}
i += 1;
}
out.push(&sql[start..]);
out.into_iter()
.filter(|statement| {
statement
.lines()
.any(|line| !line.trim().is_empty() && !line.trim().starts_with("--"))
})
.collect()
}
fn checksum(sql: &str) -> String {
crate::webhook::sha256_hex(sql)
}
pub(crate) struct Migrator {
migrations: Vec<Migration>,
}
struct Applied {
batch: i64,
checksum: Option<String>,
}
struct MigrationLock {
_local: tokio::sync::MutexGuard<'static, ()>,
#[cfg(feature = "postgres")]
_postgres: Option<sqlx::postgres::PgConnection>,
}
#[cfg(feature = "postgres")]
const ADVISORY_KEY: i64 = 0x7265_6e6f_786d_6967;
impl Migrator {
pub fn new(mut migrations: Vec<Migration>) -> anyhow::Result<Self> {
migrations.sort_by_key(|m| m.name);
for pair in migrations.windows(2) {
if pair[0].name == pair[1].name {
bail!("migration `{}` is registered twice", pair[0].name);
}
}
Ok(Self { migrations })
}
async fn lock(db: &Db) -> anyhow::Result<MigrationLock> {
static LOCAL: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
let local = LOCAL.lock().await;
#[cfg(feature = "postgres")]
let postgres = match db.postgres() {
Some(pool) => {
let mut conn = pool.acquire().await?.detach();
sqlx::query("SELECT pg_advisory_lock($1)")
.bind(ADVISORY_KEY)
.execute(&mut conn)
.await?;
Some(conn)
}
None => None,
};
#[cfg(not(feature = "postgres"))]
let _ = db;
Ok(MigrationLock {
_local: local,
#[cfg(feature = "postgres")]
_postgres: postgres,
})
}
async fn ensure_table(db: &Db) -> anyhow::Result<()> {
sql(format!(
"CREATE TABLE IF NOT EXISTS {TABLE} (
name TEXT PRIMARY KEY NOT NULL,
batch BIGINT NOT NULL,
applied_at TEXT NOT NULL,
checksum TEXT
)"
))
.execute(db)
.await?;
if sql(format!("SELECT checksum FROM {TABLE} WHERE 1 = 0"))
.execute(db)
.await
.is_err()
{
sql(format!("ALTER TABLE {TABLE} ADD COLUMN checksum TEXT"))
.execute(db)
.await?;
db.schema_changed();
}
Ok(())
}
async fn applied(db: &Db) -> anyhow::Result<HashMap<String, Applied>> {
Self::ensure_table(db).await?;
let rows = sql(format!("SELECT name, batch, checksum FROM {TABLE}"))
.fetch_all(db)
.await?;
rows.iter()
.map(|row| {
Ok((
row.try_get("name")?,
Applied {
batch: row.try_get("batch")?,
checksum: row.try_get("checksum")?,
},
))
})
.collect()
}
pub async fn run(&self, db: &Db) -> anyhow::Result<Vec<&'static str>> {
let _lock = Self::lock(db).await?;
self.run_locked(db).await
}
async fn run_locked(&self, db: &Db) -> anyhow::Result<Vec<&'static str>> {
let dialect = db.dialect();
let applied = Self::applied(db).await?;
for migration in &self.migrations {
let changed = applied.get(migration.name).is_some_and(|a| {
a.checksum
.as_ref()
.is_some_and(|sum| *sum != checksum(migration.up_for(dialect)))
});
if changed {
tracing::warn!(
migration = migration.name,
"an applied migration was edited; the change won't run (add a new migration)"
);
}
}
let batch = applied.values().map(|a| a.batch).max().unwrap_or(0) + 1;
let mut done = Vec::new();
let result = self.run_pending(db, &applied, batch, &mut done).await;
if !done.is_empty() {
db.schema_changed();
}
result.map(|()| done)
}
async fn run_pending(
&self,
db: &Db,
applied: &HashMap<String, Applied>,
batch: i64,
done: &mut Vec<&'static str>,
) -> anyhow::Result<()> {
let dialect = db.dialect();
for migration in self
.migrations
.iter()
.filter(|m| !applied.contains_key(m.name))
{
let up = migration.up_for(dialect);
let failed = || format!("migration `{}` failed", migration.name);
let record = sql(format!(
"INSERT INTO {TABLE} (name, batch, applied_at, checksum) VALUES (?, ?, ?, ?)"
))
.bind(migration.name)
.bind(batch)
.bind(now().to_rfc3339())
.bind(checksum(up));
if runs_outside_transaction(up) {
run_each(db, up).await.with_context(failed)?;
record.execute(db).await?;
} else {
let mut tx = db.begin_immediate().await?;
let already: i64 = sql(format!("SELECT COUNT(*) FROM {TABLE} WHERE name = ?"))
.bind(migration.name)
.scalar(&mut tx)
.await?;
if already > 0 {
continue;
}
script(&mut tx, up).await.with_context(failed)?;
record.execute(&mut tx).await?;
tx.commit().await?;
}
done.push(migration.name);
}
Ok(())
}
pub async fn rollback(&self, db: &Db, batches: u32) -> anyhow::Result<Vec<String>> {
let _lock = Self::lock(db).await?;
let dialect = db.dialect();
let applied = Self::applied(db).await?;
let mut numbers: Vec<i64> = applied
.values()
.map(|a| a.batch)
.collect::<HashSet<_>>()
.into_iter()
.collect();
numbers.sort_unstable_by(|a, b| b.cmp(a));
let targets: HashSet<i64> = numbers.into_iter().take(batches as usize).collect();
let mut names: Vec<&String> = applied
.iter()
.filter(|(_, a)| targets.contains(&a.batch))
.map(|(name, _)| name)
.collect();
names.sort_unstable_by(|a, b| b.cmp(a));
let mut steps = Vec::new();
for name in names {
let down = match self.migrations.iter().find(|m| m.name == name) {
None => None,
Some(migration) => Some(migration.down_for(dialect).ok_or_else(|| {
anyhow!(
"migration `{name}` has no .down.sql, so its batch can't be rolled back; \
nothing was rolled back"
)
})?),
};
steps.push((name, down));
}
let mut done = Vec::new();
for (name, down) in steps {
let forget = sql(format!("DELETE FROM {TABLE} WHERE name = ?")).bind(name);
match down {
None => {
tracing::warn!(
migration = %name,
"not registered any more; forgotten without undoing its changes"
);
forget.execute(db).await?;
}
Some(down) if runs_outside_transaction(down) => {
run_each(db, down)
.await
.with_context(|| format!("rolling back `{name}` failed"))?;
forget.execute(db).await?;
}
Some(down) => {
let mut tx = db.begin().await?;
script(&mut tx, down)
.await
.with_context(|| format!("rolling back `{name}` failed"))?;
forget.execute(&mut tx).await?;
tx.commit().await?;
}
}
done.push(name.clone());
db.schema_changed();
}
Ok(done)
}
pub async fn fresh(&self, db: &Db) -> anyhow::Result<Vec<&'static str>> {
let _lock = Self::lock(db).await?;
if let Some(pool) = db.sqlite() {
drop_all_sqlite(pool).await?;
}
#[cfg(feature = "postgres")]
if let Some(pool) = db.postgres() {
drop_all_postgres(pool).await?;
}
db.schema_changed();
self.run_locked(db).await
}
pub async fn status(&self, db: &Db) -> anyhow::Result<Vec<MigrationStatus>> {
let dialect = db.dialect();
let applied = Self::applied(db).await?;
let mut status: Vec<MigrationStatus> = self
.migrations
.iter()
.map(|m| {
let found = applied.get(m.name);
MigrationStatus {
name: m.name.to_owned(),
batch: found.map(|a| a.batch),
missing: false,
changed: found.is_some_and(|a| {
a.checksum
.as_ref()
.is_some_and(|sum| *sum != checksum(m.up_for(dialect)))
}),
}
})
.collect();
let mut missing: Vec<MigrationStatus> = applied
.iter()
.filter(|(name, _)| !self.migrations.iter().any(|m| m.name == name.as_str()))
.map(|(name, a)| MigrationStatus {
name: name.clone(),
batch: Some(a.batch),
missing: true,
changed: false,
})
.collect();
missing.sort_by(|a, b| a.name.cmp(&b.name));
status.extend(missing);
Ok(status)
}
}
async fn drop_all_sqlite(pool: &sqlx::SqlitePool) -> anyhow::Result<()> {
use sqlx::{AssertSqlSafe, Row};
let mut conn = pool.acquire().await?;
let objects = sqlx::query(
"SELECT type, name FROM sqlite_master \
WHERE type IN ('table', 'view') AND name NOT LIKE 'sqlite_%' \
ORDER BY CASE WHEN sql LIKE 'CREATE VIRTUAL TABLE%' THEN 0 ELSE 1 END, rowid",
)
.fetch_all(&mut *conn)
.await?;
sqlx::query("PRAGMA foreign_keys = OFF")
.execute(&mut *conn)
.await?;
for object in &objects {
let kind: String = object.try_get("type")?;
let name: String = object.try_get("name")?;
let sql = format!("DROP {} IF EXISTS {}", kind.to_uppercase(), quote(&name));
sqlx::query(AssertSqlSafe(sql)).execute(&mut *conn).await?;
}
sqlx::query("PRAGMA foreign_keys = ON")
.execute(&mut *conn)
.await?;
Ok(())
}
#[cfg(feature = "postgres")]
async fn drop_all_postgres(pool: &sqlx::PgPool) -> anyhow::Result<()> {
use sqlx::{AssertSqlSafe, Row};
const NOT_FROM_EXTENSION: &str = "NOT EXISTS (SELECT 1 FROM pg_depend d \
WHERE d.objid = {oid} AND d.deptype = 'e')";
let owned = |oid: &str| NOT_FROM_EXTENSION.replace("{oid}", oid);
let queries = [
format!(
"SELECT format('DROP MATERIALIZED VIEW IF EXISTS %I CASCADE', c.relname) \
FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = current_schema() AND c.relkind = 'm' AND {}",
owned("c.oid")
),
format!(
"SELECT format('DROP VIEW IF EXISTS %I CASCADE', c.relname) \
FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = current_schema() AND c.relkind = 'v' AND {}",
owned("c.oid")
),
format!(
"SELECT format('DROP TABLE IF EXISTS %I CASCADE', c.relname) \
FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = current_schema() AND c.relkind IN ('r', 'p') AND {}",
owned("c.oid")
),
format!(
"SELECT format('DROP SEQUENCE IF EXISTS %I CASCADE', c.relname) \
FROM pg_class c JOIN pg_namespace n ON n.oid = c.relnamespace \
WHERE n.nspname = current_schema() AND c.relkind = 'S' AND {}",
owned("c.oid")
),
format!(
"SELECT format('DROP ROUTINE IF EXISTS %s CASCADE', p.oid::regprocedure) \
FROM pg_proc p JOIN pg_namespace n ON n.oid = p.pronamespace \
WHERE n.nspname = current_schema() AND p.prokind IN ('f', 'p') AND {}",
owned("p.oid")
),
format!(
"SELECT format(CASE t.typtype WHEN 'd' THEN 'DROP DOMAIN IF EXISTS %I CASCADE' \
ELSE 'DROP TYPE IF EXISTS %I CASCADE' END, t.typname) \
FROM pg_type t JOIN pg_namespace n ON n.oid = t.typnamespace \
LEFT JOIN pg_class c ON c.oid = t.typrelid \
WHERE n.nspname = current_schema() AND t.typtype IN ('e', 'd', 'r', 'c') \
AND (t.typtype <> 'c' OR c.relkind = 'c') AND {}",
owned("t.oid")
),
];
for query in queries {
let statements = sqlx::query(AssertSqlSafe(query)).fetch_all(pool).await?;
for statement in &statements {
let statement: String = statement.try_get(0)?;
sqlx::query(AssertSqlSafe(statement)).execute(pool).await?;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn splits_statements_outside_quotes_comments_and_bodies() {
let sql = "-- renox:no-transaction\nCREATE TABLE t (s TEXT DEFAULT ';');\n\
/* a; b */ CREATE FUNCTION f() RETURNS INT AS $body$ SELECT 1; $body$ LANGUAGE SQL;\n\
CREATE INDEX CONCURRENTLY i ON t (s);\n";
let parts = statements(sql);
assert_eq!(parts.len(), 3, "{parts:?}");
assert!(parts[1].contains("SELECT 1; $body$"));
assert!(parts[2].trim().starts_with("CREATE INDEX CONCURRENTLY"));
assert!(runs_outside_transaction(sql));
assert!(runs_outside_transaction(
"BEGIN; CREATE TABLE a (id INT); COMMIT;"
));
assert!(!runs_outside_transaction(
"CREATE TABLE begin_log (id INT);"
));
}
#[test]
fn concurrently_is_found_whatever_follows_it() {
assert!(runs_outside_transaction(
"CREATE INDEX CONCURRENTLY\n users_email ON users (email);"
));
assert!(runs_outside_transaction(
"create index\tconcurrently\tx ON t (a);"
));
assert!(!runs_outside_transaction(
"CREATE TABLE concurrently_log (id INT);"
));
}
async fn tables(db: &Db) -> Vec<String> {
let list = match db.dialect() {
Dialect::Sqlite => "SELECT name FROM sqlite_master WHERE type = 'table' ORDER BY name",
Dialect::Postgres => {
"SELECT tablename AS name FROM pg_tables \
WHERE schemaname = current_schema() ORDER BY tablename"
}
};
sql(list)
.fetch_all(db)
.await
.unwrap()
.iter()
.map(|row| row.try_get::<String>("name").unwrap())
.collect()
}
#[tokio::test]
async fn scripts_outside_a_transaction_run_every_statement_both_ways() {
let db = super::super::connect(&crate::Config::default())
.await
.unwrap();
let outside = Migration::new(
"1_outside",
"-- renox:no-transaction\nCREATE TABLE a (id INT);\nCREATE TABLE b (id INT);",
Some("-- renox:no-transaction\nDROP TABLE b;\nDROP TABLE a;"),
)
.postgres(
"CREATE TABLE a (id INT);\nCREATE INDEX CONCURRENTLY a_id ON a (id);\nCREATE TABLE b (id INT);",
Some("DROP INDEX CONCURRENTLY a_id;\nDROP TABLE b;\nDROP TABLE a;"),
);
let comment = Migration::new("2_comment", "-- nothing to do yet\n", Some("-- nor here\n"));
let migrator = Migrator::new(vec![outside, comment]).unwrap();
assert_eq!(migrator.run(&db).await.unwrap(), ["1_outside", "2_comment"]);
let made = tables(&db).await;
assert!(
made.contains(&"a".to_owned()) && made.contains(&"b".to_owned()),
"{made:?}"
);
let undone = migrator.rollback(&db, 1).await.unwrap();
assert_eq!(undone, ["2_comment", "1_outside"]);
let left = tables(&db).await;
assert!(
!left.contains(&"a".to_owned()) && !left.contains(&"b".to_owned()),
"{left:?}"
);
}
#[tokio::test]
async fn an_old_migrations_table_gets_its_checksum_column() {
let db = super::super::connect(&crate::Config::default())
.await
.unwrap();
sql(format!(
"CREATE TABLE {TABLE} (name TEXT PRIMARY KEY NOT NULL, batch BIGINT NOT NULL, \
applied_at TEXT NOT NULL)"
))
.execute(&db)
.await
.unwrap();
sql(format!(
"INSERT INTO {TABLE} (name, batch, applied_at) VALUES ('1_old', 1, 'then')"
))
.execute(&db)
.await
.unwrap();
let old = Migration::new("1_old", "CREATE TABLE never (id INT);", None);
let new = Migration::new("2_new", "CREATE TABLE fresh (id INT);", None);
let migrator = Migrator::new(vec![old, new]).unwrap();
assert_eq!(migrator.run(&db).await.unwrap(), ["2_new"]);
let status = migrator.status(&db).await.unwrap();
assert_eq!(status.len(), 2);
assert!(!tables(&db).await.contains(&"never".to_owned()));
}
#[tokio::test]
async fn a_migration_recorded_meanwhile_is_skipped() {
let db = super::super::connect(&crate::Config::default())
.await
.unwrap();
let first = Migration::new(
"1_first",
"INSERT INTO renox_migrations (name, batch, applied_at) VALUES ('2_second', 1, 'meanwhile');",
None,
);
let second = Migration::new("2_second", "CREATE TABLE twice (id INT);", None);
let done = Migrator::new(vec![first, second])
.unwrap()
.run(&db)
.await
.unwrap();
assert_eq!(done, ["1_first"]);
assert!(!tables(&db).await.contains(&"twice".to_owned()));
}
#[tokio::test]
async fn pooled_connections_see_columns_added_by_a_migration() {
let dir = tempfile::tempdir().unwrap();
let config = crate::Config {
database_url: format!("sqlite://{}/app.db", dir.path().display()),
database_pool_size: 4,
..crate::Config::default()
};
let db = super::super::connect(&config).await.unwrap();
let create = Migration::new(
"1_notes",
"CREATE TABLE notes (id INTEGER PRIMARY KEY, a TEXT);",
None,
);
Migrator::new(vec![create]).unwrap().run(&db).await.unwrap();
sql("INSERT INTO notes (a) VALUES ('x')")
.execute(&db)
.await
.unwrap();
let mut open = Vec::new();
for _ in 0..4 {
let mut tx = db.begin().await.unwrap();
sql("SELECT * FROM notes").fetch_all(&mut tx).await.unwrap();
open.push(tx);
}
drop(open);
let alter = Migration::new("2_notes_b", "ALTER TABLE notes ADD COLUMN b TEXT;", None);
Migrator::new(vec![create, alter])
.unwrap()
.run(&db)
.await
.unwrap();
for _ in 0..8 {
let rows = sql("SELECT * FROM notes").fetch_all(&db).await.unwrap();
assert_eq!(rows[0].try_get::<Option<String>>("b").unwrap(), None);
}
}
}