from __future__ import annotations from pathlib import Path from alembic import command from alembic.config import Config from alembic.runtime.migration import MigrationContext from cloud.database import create_database_engine, normalize_database_url HEAD_REVISION = "0006_host_planner_transport" class SchemaVersionError(RuntimeError): """Raised when the database schema is not at the required revision.""" def upgrade_database(database_url: str, revision: str = "head") -> None: command.upgrade(_alembic_config(database_url), revision) def downgrade_database(database_url: str, revision: str = "base") -> None: command.downgrade(_alembic_config(database_url), revision) def current_revision(database_url: str) -> str | None: engine = create_database_engine(database_url) try: with engine.connect() as connection: context = MigrationContext.configure(connection) return context.get_current_revision() finally: engine.dispose() def is_schema_current(database_url: str) -> bool: return current_revision(database_url) == HEAD_REVISION def require_current_schema(database_url: str) -> None: revision = current_revision(database_url) if revision != HEAD_REVISION: raise SchemaVersionError( f"cloud database schema is {revision or 'unversioned'}; " f"required revision is {HEAD_REVISION}" ) def _alembic_config(database_url: str) -> Config: migrations_dir = Path(__file__).resolve().parent / "migrations" config = Config(str(migrations_dir / "alembic.ini")) config.set_main_option("script_location", str(migrations_dir)) config.set_main_option( "sqlalchemy.url", normalize_database_url(database_url).replace("%", "%%"), ) return config