env.py 1.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263
  1. from __future__ import annotations
  2. import os
  3. from logging.config import fileConfig
  4. from alembic import context
  5. from sqlalchemy import engine_from_config, pool
  6. from supply_infra.config import get_infra_settings
  7. from supply_infra.db.base import Base
  8. import supply_infra.db.models # noqa: F401
  9. config = context.config
  10. if config.config_file_name is not None:
  11. fileConfig(config.config_file_name)
  12. # CI and migration smoke tests may supply an isolated database URL without
  13. # changing the application's runtime MySQL configuration.
  14. database_url = os.environ.get("ALEMBIC_DATABASE_URL") or get_infra_settings().mysql_url
  15. # Alembic uses ConfigParser interpolation; escaped credentials may contain "%".
  16. config.set_main_option(
  17. "sqlalchemy.url",
  18. database_url.replace("%", "%%"),
  19. )
  20. target_metadata = Base.metadata
  21. def run_migrations_offline() -> None:
  22. context.configure(
  23. url=config.get_main_option("sqlalchemy.url"),
  24. target_metadata=target_metadata,
  25. literal_binds=True,
  26. dialect_opts={"paramstyle": "named"},
  27. compare_type=True,
  28. )
  29. with context.begin_transaction():
  30. context.run_migrations()
  31. def run_migrations_online() -> None:
  32. connectable = engine_from_config(
  33. config.get_section(config.config_ini_section, {}),
  34. prefix="sqlalchemy.",
  35. poolclass=pool.NullPool,
  36. )
  37. with connectable.connect() as connection:
  38. if connection.dialect.name == "mysql":
  39. connection.exec_driver_sql("SET time_zone = '+08:00'")
  40. connection.commit()
  41. context.configure(
  42. connection=connection,
  43. target_metadata=target_metadata,
  44. compare_type=True,
  45. )
  46. with context.begin_transaction():
  47. context.run_migrations()
  48. if context.is_offline_mode():
  49. run_migrations_offline()
  50. else:
  51. run_migrations_online()