use std::collections::HashSet;
use std::time::SystemTime;
use frunk::{HCons, HNil, hlist::HList};
use sea_orm::{ActiveValue, DatabaseConnection, DbErr, EntityTrait};
use sea_orm_migration::{MigratorTrait, seaql_migrations};
use crate::{
app::{App, MountedApp},
capability::{
ApplyHooks, CapStore, Capability, FoldRegistrarHooks, apply_registrar_hook,
mount_with_hooks,
},
db::{DbState, DbTag},
tag::Tagged,
traits::{
add::{AddCapability, CapTagAbsent},
get::GetByTag,
},
};
pub struct MigrationTag;
#[derive(Clone)]
pub struct MigrationCapability<Migrators> {
pub migrators: Migrators,
}
impl MigrationCapability<HNil> {
pub fn new() -> Self {
Self { migrators: HNil }
}
}
impl Default for MigrationCapability<HNil> {
fn default() -> Self {
Self::new()
}
}
impl<Migrators> MigrationCapability<Migrators> {
pub fn prepend<Tag, M>(
self,
migrator: M,
) -> MigrationCapability<HCons<Tagged<Tag, M>, Migrators>>
where
Migrators: HList,
M: MigratorTrait + Clone,
{
MigrationCapability {
migrators: HCons {
head: Tagged::new(migrator),
tail: self.migrators,
},
}
}
pub async fn run(self, db: &DatabaseConnection) -> Result<(), DbErr>
where
Migrators: RunMigrations,
{
self.migrators.run_migrations(db).await
}
}
pub type MigrationCap<Hooks, Items> = CapStore<MigrationTag, Hooks, Items>;
impl<Hooks, Items> MigrationCap<Hooks, Items> {
pub fn resolve_hooks(
self,
) -> MigrationCap<HNil, <Hooks as FoldRegistrarHooks<MigrationTag, Items>>::Output>
where
Hooks: FoldRegistrarHooks<MigrationTag, Items>,
{
CapStore::with_items(self.hooks.fold_registrar_hooks(self.items))
}
}
pub trait MigrationRegistrar<M>: Sized {
type Output;
fn register_migrations(self, cap: MigrationCapability<M>) -> MigrationCapability<Self::Output>;
}
apply_registrar_hook! {
capability: MigrationCapability;
trait: MigrationRegistrar;
method: register_migrations;
field: migrators;
proof: crate::capability::MigrationHookProof;
tag: MigrationTag;
}
impl<Hooks, Items> Capability for MigrationCap<Hooks, Items>
where
Hooks: ApplyHooks<Items>,
{
type Value = MigrationCapability<Hooks::Output>;
type Output = Tagged<MigrationTag, MigrationCapability<Hooks::Output>>;
type Hooks = Hooks;
type Items = Items;
fn mount(self) -> Self::Output {
mount_with_hooks(self, |items| MigrationCapability { migrators: items })
}
}
pub trait RunMigrations {
fn run_migrations(
self,
db: &DatabaseConnection,
) -> impl std::future::Future<Output = Result<(), DbErr>> + Send;
}
pub trait CollectMigrations {
fn collect_migrations(self) -> Vec<Box<dyn sea_orm_migration::MigrationTrait>>;
}
impl CollectMigrations for HNil {
fn collect_migrations(self) -> Vec<Box<dyn sea_orm_migration::MigrationTrait>> {
Vec::new()
}
}
impl<Tag, M, Tail> CollectMigrations for HCons<Tagged<Tag, M>, Tail>
where
M: MigratorTrait,
Tail: CollectMigrations,
{
fn collect_migrations(self) -> Vec<Box<dyn sea_orm_migration::MigrationTrait>> {
let mut migrations = self.tail.collect_migrations();
migrations.extend(M::migrations());
migrations
}
}
thread_local! {
static COMPOSITE_MIGRATIONS: std::cell::RefCell<
Option<Vec<Box<dyn sea_orm_migration::MigrationTrait>>>,
> = const { std::cell::RefCell::new(None) };
}
struct CompositeMigrator;
impl MigratorTrait for CompositeMigrator {
fn migrations() -> Vec<Box<dyn sea_orm_migration::MigrationTrait>> {
COMPOSITE_MIGRATIONS.with(|cell| cell.borrow_mut().take().unwrap_or_default())
}
}
impl<L> RunMigrations for L
where
L: CollectMigrations + Send,
{
async fn run_migrations(self, db: &DatabaseConnection) -> Result<(), DbErr> {
let migrations = self.collect_migrations();
COMPOSITE_MIGRATIONS.with(|cell| {
*cell.borrow_mut() = Some(migrations);
});
CompositeMigrator::up(db, None).await
}
}
#[macro_export]
macro_rules! define_register_migrations {
(
plugin: $plugin:ty;
migrator: $migrator:ty;
) => {
#[derive(Clone, Copy, Default)]
pub struct Hook;
impl<M> $crate::migration::MigrationRegistrar<M> for Hook
where
M: ::frunk::hlist::HList + Clone + $crate::migration::CollectMigrations + Send,
{
type Output =
impl ::frunk::hlist::HList + $crate::migration::CollectMigrations + Clone + Send;
fn register_migrations(
self,
cap: $crate::migration::MigrationCapability<M>,
) -> $crate::migration::MigrationCapability<Self::Output> {
cap.prepend::<$plugin, _>(<$migrator>::default())
}
}
};
}
pub fn with_migrations<L, Proof>(app: App<L>) -> App<HCons<MigrationCap<HNil, HNil>, L>>
where
L: HList + CapTagAbsent<MigrationTag, Proof>,
{
app.add_capability(CapStore::with_items(HNil))
}
pub async fn run_migrations<M, MigIdx, DbIdx, Migrators>(app: &MountedApp<M>) -> Result<(), DbErr>
where
M: GetByTag<MigrationTag, MigIdx, Value = MigrationCapability<Migrators>>,
M: GetByTag<DbTag, DbIdx, Value = DbState>,
Migrators: RunMigrations + Clone,
{
let db = app.get_capability_output::<DbTag, DbIdx>().conn.clone();
let migrations = app.get_capability_output::<MigrationTag, MigIdx>().clone();
migrations.run(&db).await
}
pub async fn mark_migrations_applied<L>(
db: &DatabaseConnection,
migrators: L,
) -> Result<usize, DbErr>
where
L: CollectMigrations + Send,
{
CompositeMigrator::install(db).await?;
let existing = seaql_migrations::Entity::find().all(db).await?;
let applied: HashSet<String> = existing.into_iter().map(|row| row.version).collect();
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.expect("SystemTime before UNIX EPOCH!")
.as_secs() as i64;
let mut inserted = 0usize;
for migration in migrators.collect_migrations() {
let version = migration.name().to_owned();
if applied.contains(&version) {
continue;
}
seaql_migrations::Entity::insert(seaql_migrations::ActiveModel {
version: ActiveValue::Set(version),
applied_at: ActiveValue::Set(now),
})
.exec(db)
.await?;
inserted += 1;
}
Ok(inserted)
}
pub async fn mark_migrations<M, MigIdx, DbIdx, Migrators>(
app: &MountedApp<M>,
) -> Result<usize, DbErr>
where
M: GetByTag<MigrationTag, MigIdx, Value = MigrationCapability<Migrators>>,
M: GetByTag<DbTag, DbIdx, Value = DbState>,
Migrators: CollectMigrations + Clone + Send,
{
let db = app.get_capability_output::<DbTag, DbIdx>().conn.clone();
let migrators = app.get_capability_output::<MigrationTag, MigIdx>().clone();
mark_migrations_applied(&db, migrators.migrators).await
}