use burncloud_database_core::error::{DatabaseResult, DatabaseError};
use burncloud_database_core::{QueryExecutor, QueryContext, MigrationInfo};
use async_trait::async_trait;
use std::collections::HashMap;
use chrono::{DateTime, Utc};
pub struct MigrationManager {
query_executor: Box<dyn QueryExecutor>,
migrations: Vec<Migration>,
}
pub struct Migration {
pub version: String,
pub name: String,
pub up_sql: String,
pub down_sql: String,
pub checksum: String,
}
impl MigrationManager {
pub fn new(query_executor: Box<dyn QueryExecutor>) -> Self {
let mut manager = Self {
query_executor,
migrations: Vec::new(),
};
manager.register_migrations();
manager
}
fn register_migrations(&mut self) {
self.add_migration(Migration {
version: "001".to_string(),
name: "create_base_tables".to_string(),
up_sql: include_str!("../migrations/001_create_base_tables.sql").to_string(),
down_sql: include_str!("../migrations/001_create_base_tables_down.sql").to_string(),
checksum: self.calculate_checksum("001"),
});
self.add_migration(Migration {
version: "002".to_string(),
name: "create_ai_model_tables".to_string(),
up_sql: include_str!("../migrations/002_create_ai_model_tables.sql").to_string(),
down_sql: include_str!("../migrations/002_create_ai_model_tables_down.sql").to_string(),
checksum: self.calculate_checksum("002"),
});
self.add_migration(Migration {
version: "003".to_string(),
name: "create_monitoring_tables".to_string(),
up_sql: include_str!("../migrations/003_create_monitoring_tables.sql").to_string(),
down_sql: include_str!("../migrations/003_create_monitoring_tables_down.sql").to_string(),
checksum: self.calculate_checksum("003"),
});
self.add_migration(Migration {
version: "004".to_string(),
name: "create_user_security_tables".to_string(),
up_sql: include_str!("../migrations/004_create_user_security_tables.sql").to_string(),
down_sql: include_str!("../migrations/004_create_user_security_tables_down.sql").to_string(),
checksum: self.calculate_checksum("004"),
});
self.add_migration(Migration {
version: "005".to_string(),
name: "create_indexes".to_string(),
up_sql: include_str!("../migrations/005_create_indexes.sql").to_string(),
down_sql: include_str!("../migrations/005_create_indexes_down.sql").to_string(),
checksum: self.calculate_checksum("005"),
});
}
fn add_migration(&mut self, migration: Migration) {
self.migrations.push(migration);
}
fn calculate_checksum(&self, version: &str) -> String {
format!("checksum_{}", version)
}
pub async fn run_migrations(&self, context: &QueryContext) -> DatabaseResult<()> {
self.create_migration_table(context).await?;
let applied_migrations = self.get_applied_migrations(context).await?;
for migration in &self.migrations {
if !applied_migrations.iter().any(|m| m.version == migration.version) {
println!("Running migration: {} - {}", migration.version, migration.name);
self.apply_migration(migration, context).await?;
self.record_migration(migration, context).await?;
println!("✅ Migration {} completed", migration.version);
}
}
Ok(())
}
pub async fn rollback_migration(&self, version: &str, context: &QueryContext) -> DatabaseResult<()> {
if let Some(migration) = self.migrations.iter().find(|m| m.version == version) {
println!("Rolling back migration: {} - {}", migration.version, migration.name);
self.execute_sql(&migration.down_sql, context).await?;
self.remove_migration_record(version, context).await?;
println!("✅ Migration {} rolled back", version);
Ok(())
} else {
Err(DatabaseError::ConfigurationError(format!("Migration {} not found", version)))
}
}
pub async fn get_migration_status(&self, context: &QueryContext) -> DatabaseResult<Vec<MigrationInfo>> {
self.get_applied_migrations(context).await
}
async fn create_migration_table(&self, context: &QueryContext) -> DatabaseResult<()> {
let sql = "
CREATE TABLE IF NOT EXISTS schema_migrations (
version VARCHAR(255) PRIMARY KEY,
name VARCHAR(255) NOT NULL,
applied_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
checksum VARCHAR(255) NOT NULL
)
";
self.execute_sql(sql, context).await
}
async fn apply_migration(&self, migration: &Migration, context: &QueryContext) -> DatabaseResult<()> {
self.execute_sql(&migration.up_sql, context).await
}
async fn record_migration(&self, migration: &Migration, context: &QueryContext) -> DatabaseResult<()> {
let sql = "
INSERT INTO schema_migrations (version, name, checksum)
VALUES ($1, $2, $3)
";
let version_param = burncloud_database_impl::StringParam(migration.version.clone());
let name_param = burncloud_database_impl::StringParam(migration.name.clone());
let checksum_param = burncloud_database_impl::StringParam(migration.checksum.clone());
let params: Vec<&dyn burncloud_database_core::QueryParam> = vec![&version_param, &name_param, &checksum_param];
self.query_executor.execute_query(sql, ¶ms, context).await?;
Ok(())
}
async fn remove_migration_record(&self, version: &str, context: &QueryContext) -> DatabaseResult<()> {
let sql = "DELETE FROM schema_migrations WHERE version = $1";
let version_param = burncloud_database_impl::StringParam(version.to_string());
let params: Vec<&dyn burncloud_database_core::QueryParam> = vec![&version_param];
self.query_executor.execute_query(sql, ¶ms, context).await?;
Ok(())
}
async fn get_applied_migrations(&self, context: &QueryContext) -> DatabaseResult<Vec<MigrationInfo>> {
let sql = "SELECT version, name, applied_at, checksum FROM schema_migrations ORDER BY version";
let params: Vec<&dyn burncloud_database_core::QueryParam> = vec![];
let result = self.query_executor.execute_query(sql, ¶ms, context).await?;
let mut migrations = Vec::new();
for row in result.rows {
let version = row.get("version")
.and_then(|v| v.as_str())
.ok_or_else(|| DatabaseError::SerializationError("Missing version".to_string()))?
.to_string();
let name = row.get("name")
.and_then(|v| v.as_str())
.ok_or_else(|| DatabaseError::SerializationError("Missing name".to_string()))?
.to_string();
let applied_at_str = row.get("applied_at")
.and_then(|v| v.as_str())
.ok_or_else(|| DatabaseError::SerializationError("Missing applied_at".to_string()))?;
let applied_at = DateTime::parse_from_rfc3339(applied_at_str)
.map_err(|e| DatabaseError::SerializationError(e.to_string()))?
.with_timezone(&Utc);
let checksum = row.get("checksum")
.and_then(|v| v.as_str())
.ok_or_else(|| DatabaseError::SerializationError("Missing checksum".to_string()))?
.to_string();
migrations.push(MigrationInfo {
version,
name,
applied_at,
checksum,
});
}
Ok(migrations)
}
async fn execute_sql(&self, sql: &str, context: &QueryContext) -> DatabaseResult<()> {
let params: Vec<&dyn burncloud_database_core::QueryParam> = vec![];
self.query_executor.execute_query(sql, ¶ms, context).await?;
Ok(())
}
}
#[async_trait]
impl burncloud_database_core::MigrationManager for MigrationManager {
async fn run_migrations(&self) -> DatabaseResult<()> {
let context = QueryContext::default();
self.run_migrations(&context).await
}
async fn rollback_migration(&self, version: &str) -> DatabaseResult<()> {
let context = QueryContext::default();
self.rollback_migration(version, &context).await
}
async fn get_migration_status(&self) -> DatabaseResult<Vec<MigrationInfo>> {
let context = QueryContext::default();
self.get_migration_status(&context).await
}
}