use std::collections::HashSet;
use std::fmt;
use async_trait::async_trait;
use turso_orm::sql::{ColumnDef, Expr, Order, Query, Table};
use turso_orm::{ConnectionTrait, Database, DbErr, Statement, TransactionMode, TransactionTrait};
use crate::MigrationTrait;
use crate::manager::{SchemaManager, has_table};
const DEFAULT_TABLE: &str = "turso_migrations";
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MigrationStatus {
pub name: String,
pub applied: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum MigrationIssue {
Unknown(String),
OutOfOrder(String),
}
impl fmt::Display for MigrationIssue {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Unknown(name) => write!(f, "applied migration `{name}` is not declared"),
Self::OutOfOrder(name) => write!(
f,
"pending migration `{name}` is declared before an applied one"
),
}
}
}
#[async_trait]
pub trait MigratorTrait: Send {
fn migrations() -> Vec<Box<dyn MigrationTrait>>;
fn migration_table_name() -> &'static str {
DEFAULT_TABLE
}
fn strict() -> bool {
false
}
async fn check(db: &Database) -> Result<Vec<MigrationIssue>, DbErr> {
let migrations = Self::migrations();
ensure_unique(&migrations)?;
let applied = Self::get_applied_migrations(db).await?;
Ok(find_issues(&migrations, &applied))
}
async fn install(db: &Database) -> Result<(), DbErr> {
let stmt = Table::create()
.table(Self::migration_table_name())
.if_not_exists()
.col(ColumnDef::text("version").primary_key().not_null())
.col(ColumnDef::integer("applied_at").not_null());
db.execute(turso_orm::Build::to_statement(&stmt)).await?;
Ok(())
}
async fn get_applied_migrations(db: &Database) -> Result<Vec<String>, DbErr> {
if !has_table(db, Self::migration_table_name()).await? {
return Ok(Vec::new());
}
let stmt = Query::select()
.column("version")
.from(Self::migration_table_name())
.order_by("applied_at", Order::Asc)
.order_by("version", Order::Asc);
let rows = db.query_all(turso_orm::Build::to_statement(&stmt)).await?;
rows.iter()
.map(|r| r.get::<String>("version").map_err(DbErr::from))
.collect()
}
async fn status(db: &Database) -> Result<Vec<MigrationStatus>, DbErr> {
let migrations = Self::migrations();
ensure_unique(&migrations)?;
let applied = Self::get_applied_migrations(db).await?;
Ok(migrations
.iter()
.map(|m| MigrationStatus {
name: m.name().to_owned(),
applied: applied.iter().any(|a| a == m.name()),
})
.collect())
}
async fn up(db: &Database, steps: Option<u32>) -> Result<(), DbErr> {
let migrations = Self::migrations();
ensure_unique(&migrations)?;
Self::install(db).await?;
let applied = Self::get_applied_migrations(db).await?;
report_issues(&find_issues(&migrations, &applied), Self::strict())?;
let mut remaining = steps.map_or(usize::MAX, |s| s as usize);
for migration in migrations {
if remaining == 0 {
break;
}
if applied.iter().any(|a| a == migration.name()) {
continue;
}
tracing::info!(name = migration.name(), "applying migration");
let txn = db.begin_with_mode(TransactionMode::Immediate).await?;
if is_applied(&txn, Self::migration_table_name(), migration.name()).await? {
txn.rollback().await?;
continue;
}
{
let manager = SchemaManager::new(&txn);
migration.up(&manager).await?;
}
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| i64::try_from(d.as_secs()).unwrap_or(i64::MAX))
.unwrap_or_default();
let insert = Query::insert()
.into_table(Self::migration_table_name())
.columns(["version", "applied_at"])
.values([Expr::val(migration.name()), Expr::val(now)]);
txn.execute(turso_orm::Build::to_statement(&insert)).await?;
txn.commit().await?;
remaining -= 1;
}
Ok(())
}
async fn down(db: &Database, steps: Option<u32>) -> Result<(), DbErr> {
let migrations = Self::migrations();
ensure_unique(&migrations)?;
Self::install(db).await?;
let applied = Self::get_applied_migrations(db).await?;
let mut remaining = steps.map_or(usize::MAX, |s| s as usize);
for migration in migrations.into_iter().rev() {
if remaining == 0 {
break;
}
if !applied.iter().any(|a| a == migration.name()) {
continue;
}
tracing::info!(name = migration.name(), "reverting migration");
let txn = db.begin_with_mode(TransactionMode::Immediate).await?;
if !is_applied(&txn, Self::migration_table_name(), migration.name()).await? {
txn.rollback().await?;
continue;
}
{
let manager = SchemaManager::new(&txn);
migration.down(&manager).await?;
}
let delete = Query::delete()
.from_table(Self::migration_table_name())
.and_where(Expr::col("version").eq(Expr::val(migration.name())));
txn.execute(turso_orm::Build::to_statement(&delete)).await?;
txn.commit().await?;
remaining -= 1;
}
Ok(())
}
async fn fresh(db: &Database) -> Result<(), DbErr> {
ensure_unique(&Self::migrations())?;
let rows = db
.query_all(Statement::from_string(
"SELECT name FROM sqlite_schema WHERE type = 'table' AND name NOT LIKE 'sqlite_%' AND name NOT LIKE '__turso_%'",
))
.await?;
let txn = db.begin_with_mode(TransactionMode::Immediate).await?;
let mut pending: Vec<String> = rows
.iter()
.map(|row| row.get::<String>("name"))
.collect::<Result<_, _>>()?;
while !pending.is_empty() {
let before = pending.len();
let mut failed = Vec::new();
let mut last_error = None;
for name in pending {
let drop = Table::drop().table(name.clone()).if_exists();
if let Err(err) = txn.execute(turso_orm::Build::to_statement(&drop)).await {
failed.push(name);
last_error = Some(err);
}
}
if failed.len() == before
&& let Some(err) = last_error
{
return Err(err.into());
}
pending = failed;
}
txn.commit().await?;
Self::up(db, None).await
}
async fn refresh(db: &Database) -> Result<(), DbErr> {
if Self::strict() {
let unknown: Vec<MigrationIssue> = Self::check(db)
.await?
.into_iter()
.filter(|issue| matches!(issue, MigrationIssue::Unknown(_)))
.collect();
report_issues(&unknown, true)?;
}
Self::down(db, None).await?;
Self::up(db, None).await
}
async fn reset(db: &Database) -> Result<(), DbErr> {
Self::down(db, None).await
}
}
async fn is_applied<C: ConnectionTrait>(
conn: &C,
table: &'static str,
version: &str,
) -> Result<bool, DbErr> {
let stmt = Query::select()
.column("version")
.from(table)
.and_where(Expr::col("version").eq(Expr::val(version)));
Ok(conn
.query_one(turso_orm::Build::to_statement(&stmt))
.await?
.is_some())
}
fn ensure_unique(migrations: &[Box<dyn MigrationTrait>]) -> Result<(), DbErr> {
let mut seen = HashSet::new();
match migrations.iter().find(|m| !seen.insert(m.name())) {
Some(duplicate) => Err(DbErr::Migration(format!(
"duplicate migration name `{}`",
duplicate.name()
))),
None => Ok(()),
}
}
fn find_issues(migrations: &[Box<dyn MigrationTrait>], applied: &[String]) -> Vec<MigrationIssue> {
let declared: HashSet<&str> = migrations.iter().map(|m| m.name()).collect();
let is_applied = |name: &str| applied.iter().any(|a| a == name);
let last_applied = migrations.iter().rposition(|m| is_applied(m.name()));
let unknown = applied
.iter()
.filter(|a| !declared.contains(a.as_str()))
.map(|a| MigrationIssue::Unknown(a.clone()));
let out_of_order = migrations
.iter()
.take(last_applied.unwrap_or(0))
.filter(|m| !is_applied(m.name()))
.map(|m| MigrationIssue::OutOfOrder(m.name().to_owned()));
unknown.chain(out_of_order).collect()
}
fn report_issues(issues: &[MigrationIssue], strict: bool) -> Result<(), DbErr> {
if strict && !issues.is_empty() {
let list: Vec<String> = issues.iter().map(ToString::to_string).collect();
return Err(DbErr::Migration(list.join("; ")));
}
for issue in issues {
tracing::warn!("{issue}");
}
Ok(())
}