use crate::pool::Connection;
use crate::DbError;
#[derive(Debug, Clone, PartialEq)]
pub struct ColumnDef {
pub name: String,
pub sql_type: String,
pub nullable: bool,
pub primary_key: bool,
pub default: Option<String>,
}
impl ColumnDef {
pub fn new(
name: impl Into<String>,
sql_type: impl Into<String>,
nullable: bool,
primary_key: bool,
default: Option<String>,
) -> Self {
Self {
name: name.into(),
sql_type: sql_type.into(),
nullable,
primary_key,
default,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct TableDef {
pub name: String,
pub columns: Vec<ColumnDef>,
}
impl TableDef {
pub fn new(name: impl Into<String>, columns: Vec<ColumnDef>) -> Self {
Self {
name: name.into(),
columns,
}
}
pub fn get_column(&self, name: &str) -> Option<&ColumnDef> {
self.columns.iter().find(|c| c.name == name)
}
}
#[derive(Debug, Clone, Default)]
pub struct SchemaDiff {
pub added_tables: Vec<TableDef>,
pub dropped_tables: Vec<String>,
pub added_columns: Vec<(String, ColumnDef)>,
pub dropped_columns: Vec<(String, String)>,
pub type_changed_columns: Vec<(String, ColumnDef, ColumnDef)>,
pub renamed_columns: Vec<(String, String, String)>,
}
impl SchemaDiff {
pub fn is_empty(&self) -> bool {
self.added_tables.is_empty()
&& self.dropped_tables.is_empty()
&& self.added_columns.is_empty()
&& self.dropped_columns.is_empty()
&& self.type_changed_columns.is_empty()
&& self.renamed_columns.is_empty()
}
pub fn has_destructive_changes(&self) -> bool {
!self.dropped_tables.is_empty() || !self.dropped_columns.is_empty()
}
}
#[derive(Debug, Clone)]
pub struct SyncResult {
pub affected_tables: Vec<String>,
pub executed_ddl: Vec<String>,
}
pub fn diff(entity: &[TableDef], db: &[TableDef]) -> SchemaDiff {
let mut result = SchemaDiff::default();
let db_map: std::collections::HashMap<&str, &TableDef> =
db.iter().map(|t| (t.name.as_str(), t)).collect();
let entity_map: std::collections::HashMap<&str, &TableDef> =
entity.iter().map(|t| (t.name.as_str(), t)).collect();
for t in entity {
if !db_map.contains_key(t.name.as_str()) {
result.added_tables.push(t.clone());
}
}
for t in db {
if !entity_map.contains_key(t.name.as_str()) {
result.dropped_tables.push(t.name.clone());
}
}
for entity_table in entity {
if let Some(db_table) = db_map.get(entity_table.name.as_str()) {
diff_columns(&mut result, entity_table, db_table);
}
}
result
}
fn diff_columns(result: &mut SchemaDiff, entity: &TableDef, db: &TableDef) {
let db_col_map: std::collections::HashMap<&str, &ColumnDef> =
db.columns.iter().map(|c| (c.name.as_str(), c)).collect();
let entity_col_map: std::collections::HashMap<&str, &ColumnDef> = entity
.columns
.iter()
.map(|c| (c.name.as_str(), c))
.collect();
for col in &entity.columns {
if !db_col_map.contains_key(col.name.as_str()) {
result
.added_columns
.push((entity.name.clone(), col.clone()));
}
}
for col in &db.columns {
if !entity_col_map.contains_key(col.name.as_str()) {
result
.dropped_columns
.push((entity.name.clone(), col.name.clone()));
}
}
for entity_col in &entity.columns {
if let Some(db_col) = db_col_map.get(entity_col.name.as_str()) {
if entity_col.sql_type != db_col.sql_type || entity_col.nullable != db_col.nullable {
result.type_changed_columns.push((
entity.name.clone(),
(*db_col).clone(),
entity_col.clone(),
));
}
}
}
}
pub trait DdlGenerator: Send + Sync {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError>;
}
pub struct MySqlDdlGenerator;
impl DdlGenerator for MySqlDdlGenerator {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError> {
let mut ddl = Vec::new();
for table in &diff.added_tables {
ddl.push(generate_create_table_mysql(table));
}
for (table, col) in &diff.added_columns {
ddl.push(format!(
"ALTER TABLE {} ADD COLUMN {} {}{}{}",
table,
col.name,
col.sql_type,
if col.nullable { "" } else { " NOT NULL" },
if col.primary_key { " PRIMARY KEY" } else { "" }
));
}
for (table, _old, new) in &diff.type_changed_columns {
ddl.push(format!(
"ALTER TABLE {} MODIFY COLUMN {} {}{}",
table,
new.name,
new.sql_type,
if new.nullable { "" } else { " NOT NULL" }
));
}
for (table, old, new) in &diff.renamed_columns {
ddl.push(format!(
"ALTER TABLE {} RENAME COLUMN {} TO {}",
table, old, new
));
}
Ok(ddl)
}
}
fn generate_create_table_mysql(table: &TableDef) -> String {
let columns: Vec<String> = table
.columns
.iter()
.map(|c| {
format!(
"{} {}{}{}{}",
c.name,
c.sql_type,
if c.nullable { "" } else { " NOT NULL" },
if c.primary_key { " PRIMARY KEY" } else { "" },
c.default
.as_ref()
.map(|d| format!(" DEFAULT {}", d))
.unwrap_or_default()
)
})
.collect();
format!("CREATE TABLE {} ({})", table.name, columns.join(", "))
}
pub struct PgDdlGenerator;
impl DdlGenerator for PgDdlGenerator {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError> {
let mut ddl = Vec::new();
for table in &diff.added_tables {
ddl.push(generate_create_table_mysql(table)); }
for (table, col) in &diff.added_columns {
ddl.push(format!(
"ALTER TABLE {} ADD COLUMN {} {}{}{}",
table,
col.name,
col.sql_type,
if col.nullable { "" } else { " NOT NULL" },
if col.primary_key { " PRIMARY KEY" } else { "" }
));
}
for (table, _old, new) in &diff.type_changed_columns {
ddl.push(format!(
"ALTER TABLE {} ALTER COLUMN {} TYPE {}",
table, new.name, new.sql_type
));
}
for (table, old, new) in &diff.renamed_columns {
ddl.push(format!(
"ALTER TABLE {} RENAME COLUMN {} TO {}",
table, old, new
));
}
Ok(ddl)
}
}
pub struct SqliteDdlGenerator;
impl DdlGenerator for SqliteDdlGenerator {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError> {
let mut ddl = Vec::new();
for table in &diff.added_tables {
ddl.push(generate_create_table_mysql(table));
}
for (table, col) in &diff.added_columns {
ddl.push(format!(
"ALTER TABLE {} ADD COLUMN {} {}{}",
table,
col.name,
col.sql_type,
col.default
.as_ref()
.map(|d| format!(" DEFAULT {}", d))
.unwrap_or_else(|| " DEFAULT NULL".to_string())
));
}
if !diff.type_changed_columns.is_empty() {
return Err(DbError::Unsupported(
"SQLite does not support altering column type; table rebuild required".to_string(),
));
}
for (table, old, new) in &diff.renamed_columns {
ddl.push(format!(
"ALTER TABLE {} RENAME COLUMN {} TO {}",
table, old, new
));
}
Ok(ddl)
}
}
pub struct OracleDdlGenerator;
impl DdlGenerator for OracleDdlGenerator {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError> {
let mut ddl = Vec::new();
for table in &diff.added_tables {
ddl.push(generate_create_table_mysql(table));
}
for (table, col) in &diff.added_columns {
ddl.push(format!(
"ALTER TABLE {} ADD ({} {}{}{})",
table,
col.name,
col.sql_type,
if col.nullable { "" } else { " NOT NULL" },
if col.primary_key { " PRIMARY KEY" } else { "" }
));
}
for (table, _old, new) in &diff.type_changed_columns {
ddl.push(format!(
"ALTER TABLE {} MODIFY ({} {}{})",
table,
new.name,
new.sql_type,
if new.nullable { "" } else { " NOT NULL" }
));
}
for (table, old, new) in &diff.renamed_columns {
ddl.push(format!(
"ALTER TABLE {} RENAME COLUMN {} TO {}",
table, old, new
));
}
Ok(ddl)
}
}
pub struct MssqlDdlGenerator;
impl DdlGenerator for MssqlDdlGenerator {
fn generate(&self, diff: &SchemaDiff) -> Result<Vec<String>, DbError> {
let mut ddl = Vec::new();
for table in &diff.added_tables {
ddl.push(generate_create_table_mysql(table));
}
for (table, col) in &diff.added_columns {
ddl.push(format!(
"ALTER TABLE {} ADD {} {}{}{}",
table,
col.name,
col.sql_type,
if col.nullable { "" } else { " NOT NULL" },
if col.primary_key { " PRIMARY KEY" } else { "" }
));
}
for (table, _old, new) in &diff.type_changed_columns {
ddl.push(format!(
"ALTER TABLE {} ALTER COLUMN {} {}{}",
table,
new.name,
new.sql_type,
if new.nullable { "" } else { " NOT NULL" }
));
}
for (table, old, new) in &diff.renamed_columns {
ddl.push(format!(
"EXEC sp_rename '{}.{}', '{}', 'COLUMN'",
table, old, new
));
}
Ok(ddl)
}
}
pub struct SchemaSync {
entity_tables: Vec<TableDef>,
ddl_generator: Box<dyn DdlGenerator>,
}
impl SchemaSync {
pub fn new(entity_tables: Vec<TableDef>) -> Self {
Self {
entity_tables,
ddl_generator: Box::new(MySqlDdlGenerator),
}
}
pub fn with_generator(
entity_tables: Vec<TableDef>,
ddl_generator: Box<dyn DdlGenerator>,
) -> Self {
Self {
entity_tables,
ddl_generator,
}
}
pub async fn sync_dry_run(&self, conn: &mut dyn Connection) -> Result<Vec<String>, DbError> {
let db_tables = introspect(conn).await?;
let diff_result = diff(&self.entity_tables, &db_tables);
if diff_result.has_destructive_changes() {
return Err(DbError::Internal(format!(
"DestructiveChangeDetected: dropped_tables={:?}, dropped_columns={:?}",
diff_result.dropped_tables, diff_result.dropped_columns
)));
}
self.ddl_generator.generate(&diff_result)
}
pub async fn sync(&self, conn: &mut dyn Connection) -> Result<SyncResult, DbError> {
let ddl = self.sync_dry_run(conn).await?;
if ddl.is_empty() {
return Ok(SyncResult {
affected_tables: Vec::new(),
executed_ddl: Vec::new(),
});
}
conn.begin_transaction().await?;
let mut executed = Vec::new();
for ddl_stmt in &ddl {
match conn.execute(ddl_stmt).await {
Ok(_) => executed.push(ddl_stmt.clone()),
Err(e) => {
let _ = conn.rollback().await;
return Err(DbError::Internal(format!(
"DDL execution failed: {} — SQL: {}",
e, ddl_stmt
)));
}
}
}
conn.commit().await?;
Ok(SyncResult {
affected_tables: self.entity_tables.iter().map(|t| t.name.clone()).collect(),
executed_ddl: executed,
})
}
pub fn diff_against(&self, db_tables: &[TableDef]) -> SchemaDiff {
diff(&self.entity_tables, db_tables)
}
}
async fn introspect(conn: &mut dyn Connection) -> Result<Vec<TableDef>, DbError> {
let _ = conn;
Ok(Vec::new())
}
#[cfg(test)]
mod tests {
use super::*;
fn make_column(name: &str, sql_type: &str) -> ColumnDef {
ColumnDef::new(name, sql_type, true, false, None)
}
fn make_table(name: &str, columns: Vec<ColumnDef>) -> TableDef {
TableDef::new(name, columns)
}
#[test]
fn test_diff_add_table() {
let entity = vec![make_table("users", vec![make_column("id", "BIGINT")])];
let db = vec![];
let result = diff(&entity, &db);
assert_eq!(result.added_tables.len(), 1);
assert_eq!(result.added_tables[0].name, "users");
}
#[test]
fn test_diff_drop_table() {
let entity = vec![];
let db = vec![make_table("legacy", vec![make_column("id", "BIGINT")])];
let result = diff(&entity, &db);
assert_eq!(result.dropped_tables.len(), 1);
assert_eq!(result.dropped_tables[0], "legacy");
assert!(result.has_destructive_changes());
}
#[test]
fn test_diff_add_column() {
let entity = vec![make_table(
"users",
vec![
make_column("id", "BIGINT"),
make_column("email", "VARCHAR(255)"),
],
)];
let db = vec![make_table("users", vec![make_column("id", "BIGINT")])];
let result = diff(&entity, &db);
assert_eq!(result.added_columns.len(), 1);
assert_eq!(result.added_columns[0].0, "users");
assert_eq!(result.added_columns[0].1.name, "email");
}
#[test]
fn test_diff_drop_column() {
let entity = vec![make_table("users", vec![make_column("id", "BIGINT")])];
let db = vec![make_table(
"users",
vec![
make_column("id", "BIGINT"),
make_column("legacy_col", "TEXT"),
],
)];
let result = diff(&entity, &db);
assert_eq!(result.dropped_columns.len(), 1);
assert_eq!(
result.dropped_columns[0],
("users".to_string(), "legacy_col".to_string())
);
assert!(result.has_destructive_changes());
}
#[test]
fn test_diff_type_change() {
let entity = vec![make_table(
"users",
vec![
make_column("id", "BIGINT"),
make_column("name", "VARCHAR(255)"),
],
)];
let db = vec![make_table(
"users",
vec![
make_column("id", "BIGINT"),
make_column("name", "VARCHAR(100)"),
],
)];
let result = diff(&entity, &db);
assert_eq!(result.type_changed_columns.len(), 1);
assert_eq!(result.type_changed_columns[0].0, "users");
assert_eq!(result.type_changed_columns[0].1.sql_type, "VARCHAR(100)");
assert_eq!(result.type_changed_columns[0].2.sql_type, "VARCHAR(255)");
}
#[test]
fn test_diff_no_change() {
let entity = vec![make_table("users", vec![make_column("id", "BIGINT")])];
let db = vec![make_table("users", vec![make_column("id", "BIGINT")])];
let result = diff(&entity, &db);
assert!(result.is_empty());
}
#[test]
fn test_mysql_ddl_add_table() {
let diff_result = SchemaDiff {
added_tables: vec![make_table(
"users",
vec![ColumnDef::new("id", "BIGINT", false, true, None)],
)],
..Default::default()
};
let ddl = MySqlDdlGenerator.generate(&diff_result).unwrap();
assert_eq!(ddl.len(), 1);
assert!(ddl[0].contains("CREATE TABLE users"));
assert!(ddl[0].contains("id BIGINT NOT NULL PRIMARY KEY"));
}
#[test]
fn test_mysql_ddl_add_column() {
let diff_result = SchemaDiff {
added_columns: vec![(
"users".to_string(),
ColumnDef::new("email", "VARCHAR(255)", false, false, None),
)],
..Default::default()
};
let ddl = MySqlDdlGenerator.generate(&diff_result).unwrap();
assert_eq!(ddl.len(), 1);
assert!(ddl[0].contains("ALTER TABLE users ADD COLUMN email VARCHAR(255) NOT NULL"));
}
#[test]
fn test_pg_ddl_type_change() {
let diff_result = SchemaDiff {
type_changed_columns: vec![(
"users".to_string(),
ColumnDef::new("name", "VARCHAR(100)", true, false, None),
ColumnDef::new("name", "VARCHAR(255)", true, false, None),
)],
..Default::default()
};
let ddl = PgDdlGenerator.generate(&diff_result).unwrap();
assert_eq!(ddl.len(), 1);
assert!(ddl[0].contains("ALTER TABLE users ALTER COLUMN name TYPE VARCHAR(255)"));
}
#[test]
fn test_sqlite_ddl_type_change_unsupported() {
let diff_result = SchemaDiff {
type_changed_columns: vec![(
"users".to_string(),
ColumnDef::new("name", "VARCHAR(100)", true, false, None),
ColumnDef::new("name", "VARCHAR(255)", true, false, None),
)],
..Default::default()
};
let result = SqliteDdlGenerator.generate(&diff_result);
assert!(result.is_err());
}
#[test]
fn test_oracle_ddl_add_column() {
let diff_result = SchemaDiff {
added_columns: vec![(
"users".to_string(),
ColumnDef::new("email", "VARCHAR2(255)", true, false, None),
)],
..Default::default()
};
let ddl = OracleDdlGenerator.generate(&diff_result).unwrap();
assert_eq!(ddl.len(), 1);
assert!(ddl[0].contains("ALTER TABLE users ADD (email VARCHAR2(255))"));
}
#[test]
fn test_mssql_ddl_rename() {
let diff_result = SchemaDiff {
renamed_columns: vec![(
"users".to_string(),
"old_name".to_string(),
"new_name".to_string(),
)],
..Default::default()
};
let ddl = MssqlDdlGenerator.generate(&diff_result).unwrap();
assert_eq!(ddl.len(), 1);
assert!(ddl[0].contains("EXEC sp_rename 'users.old_name', 'new_name', 'COLUMN'"));
}
#[test]
fn test_destructive_change_detected() {
let diff_result = SchemaDiff {
dropped_columns: vec![("users".to_string(), "legacy".to_string())],
..Default::default()
};
assert!(diff_result.has_destructive_changes());
}
#[test]
fn test_schema_diff_is_empty() {
let empty = SchemaDiff::default();
assert!(empty.is_empty());
let non_empty = SchemaDiff {
added_columns: vec![(
"users".to_string(),
ColumnDef::new("email", "VARCHAR(255)", true, false, None),
)],
..Default::default()
};
assert!(!non_empty.is_empty());
}
#[test]
fn test_sync_result() {
let result = SyncResult {
affected_tables: vec!["users".to_string()],
executed_ddl: vec!["ALTER TABLE users ADD COLUMN email VARCHAR(255)".to_string()],
};
assert_eq!(result.affected_tables.len(), 1);
assert_eq!(result.executed_ddl.len(), 1);
}
}