Fix database migration

This commit is contained in:
2025-10-08 17:09:04 +03:00
Unverified
parent eebe548d88
commit 14a35bc18d
2 changed files with 60 additions and 14 deletions
+54 -8
View File
@@ -104,8 +104,23 @@ def run_migrations():
# Run the upgrade command # Run the upgrade command
logger.info("Running database migrations...") logger.info("Running database migrations...")
try:
command.upgrade(alembic_cfg, "head") command.upgrade(alembic_cfg, "head")
logger.info("Database migrations completed successfully.") logger.info("Database migrations completed successfully.")
except Exception as upgrade_error:
if "Can't locate revision identified by 'direct_creation'" in str(upgrade_error):
logger.info("Found 'direct_creation' revision - resetting migration state...")
# Clear the alembic_version table and start fresh
engine = create_engine(DATABASE_URL)
with engine.connect() as connection:
from sqlalchemy import text
connection.execute(text("DELETE FROM alembic_version"))
connection.commit()
# Try upgrade again
command.upgrade(alembic_cfg, "head")
logger.info("Database migrations completed successfully after reset.")
else:
raise upgrade_error
except Exception as e: except Exception as e:
logger.error(f"Error running database migrations: {e}") logger.error(f"Error running database migrations: {e}")
@@ -165,6 +180,16 @@ def _populate_migration_file(migration_path):
with open(migration_path, 'r') as f: with open(migration_path, 'r') as f:
content = f.read() content = f.read()
# Add datetime import if needed
if "datetime.now" in migration_content and "from datetime import datetime" not in content:
# Insert the import after the existing imports
import re
content = re.sub(
r'(from alembic import op\nimport sqlalchemy as sa\n)',
r'\1from datetime import datetime\n',
content
)
# Replace the empty upgrade/downgrade functions # Replace the empty upgrade/downgrade functions
import re import re
# More flexible regex to match the actual content # More flexible regex to match the actual content
@@ -184,6 +209,7 @@ def _generate_migration_from_models():
"""Generate migration content dynamically from SQLAlchemy models.""" """Generate migration content dynamically from SQLAlchemy models."""
from models import Base from models import Base
import sqlalchemy as sa import sqlalchemy as sa
from datetime import datetime
# Generate migration content using Alembic's op functions # Generate migration content using Alembic's op functions
upgrade_statements = [] upgrade_statements = []
@@ -202,15 +228,18 @@ def _generate_migration_from_models():
for statement in schema_diff['alter_statements']: for statement in schema_diff['alter_statements']:
upgrade_statements.append(f" {statement}") upgrade_statements.append(f" {statement}")
else: else:
# Table exists and is up to date # Table exists and is up to date - skip creating it
upgrade_statements.append(f" # Table {table_name} is already up to date") upgrade_statements.append(f" # Table {table_name} already exists and is up to date")
else: else:
# Generate CREATE TABLE for new table # Generate CREATE TABLE for new table
table_code = _generate_table_creation_code(table_name, table) table_code = _generate_table_creation_code(table_name, table)
upgrade_statements.append(f" # Create {table_name} table") upgrade_statements.append(f" # Create {table_name} table")
upgrade_statements.append(table_code) upgrade_statements.append(table_code)
# Generate DROP TABLE statement for downgrade # Only add to downgrade if table actually exists
if schema_diff['table_exists']:
downgrade_statements.append(f" # op.drop_table('{table_name}') # Skipped - table exists")
else:
downgrade_statements.append(f" op.drop_table('{table_name}')") downgrade_statements.append(f" op.drop_table('{table_name}')")
# Combine all statements # Combine all statements
@@ -294,34 +323,51 @@ def _generate_table_creation_code(table_name, table):
"""Generate op.create_table code for a SQLAlchemy table.""" """Generate op.create_table code for a SQLAlchemy table."""
lines = [f" op.create_table('{table_name}',"] lines = [f" op.create_table('{table_name}',"]
# Collect all table items (columns + constraints)
all_items = []
# Add columns # Add columns
for column in table.columns: for column in table.columns:
column_def = f" sa.Column('{column.name}', {_get_column_type(column)}, nullable={column.nullable}" column_def = f" sa.Column('{column.name}', {_get_column_type(column)}, nullable={column.nullable}"
if column.default is not None: if column.default is not None:
# Handle callable defaults properly
if hasattr(column.default, 'arg') and callable(column.default.arg):
column_def += f", default=datetime.now"
else:
column_def += f", default={repr(column.default)}" column_def += f", default={repr(column.default)}"
column_def += ")" column_def += ")"
lines.append(column_def) all_items.append(column_def)
# Add constraints # Add constraints
for constraint in table.constraints: for constraint in table.constraints:
if hasattr(constraint, 'columns'): if hasattr(constraint, 'columns'):
if constraint.__class__.__name__ == 'PrimaryKeyConstraint': if constraint.__class__.__name__ == 'PrimaryKeyConstraint':
lines.append(f" sa.PrimaryKeyConstraint('{constraint.columns.keys()[0]}')") all_items.append(f" sa.PrimaryKeyConstraint('{constraint.columns.keys()[0]}')")
elif constraint.__class__.__name__ == 'UniqueConstraint': elif constraint.__class__.__name__ == 'UniqueConstraint':
cols = "', '".join(constraint.columns.keys()) cols = "', '".join(constraint.columns.keys())
lines.append(f" sa.UniqueConstraint('{cols}')") all_items.append(f" sa.UniqueConstraint('{cols}')")
# Add foreign key constraints # Add foreign key constraints
for fk in table.foreign_keys: for fk in table.foreign_keys:
lines.append(f" sa.ForeignKeyConstraint(['{fk.parent.name}'], ['{fk.column.table.name}.{fk.column.name}'], )") all_items.append(f" sa.ForeignKeyConstraint(['{fk.parent.name}'], ['{fk.column.table.name}.{fk.column.name}'], )")
# Add all items with commas (except the last one)
for i, item in enumerate(all_items):
if i < len(all_items) - 1:
item += ","
lines.append(item)
lines.append(" )") lines.append(" )")
# Add indexes # Add indexes with IF NOT EXISTS equivalent using try/except
for index in table.indexes: for index in table.indexes:
if not index.unique: if not index.unique:
cols = "', '".join([col.name for col in index.columns]) cols = "', '".join([col.name for col in index.columns])
lines.append(f" # Create index for {table_name}")
lines.append(f" try:")
lines.append(f" op.create_index(op.f('ix_{table_name}_{index.name}'), '{table_name}', ['{cols}'], unique=False)") lines.append(f" op.create_index(op.f('ix_{table_name}_{index.name}'), '{table_name}', ['{cols}'], unique=False)")
lines.append(f" except Exception:")
lines.append(f" pass # Index may already exist")
return "\n".join(lines) return "\n".join(lines)
+1 -1
View File
@@ -135,7 +135,7 @@ class DMReaction(Base):
# Relationships # Relationships
user = relationship("User") user = relationship("User")
dm_envelope = relationship("DMEnvelope") dm_envelope = relationship("DMEnvelope", overlaps="reactions")
# Ensure unique combination of dm_envelope, user, and emoji # Ensure unique combination of dm_envelope, user, and emoji
__table_args__ = (UniqueConstraint('dm_envelope_id', 'user_id', 'emoji', name='unique_dm_reaction'),) __table_args__ = (UniqueConstraint('dm_envelope_id', 'user_id', 'emoji', name='unique_dm_reaction'),)