import logging from sqlalchemy import text from sqlalchemy.ext.asyncio import ( AsyncConnection, AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine, ) from app.db.base import Base from app.db import models as _models # noqa: F401 logger = logging.getLogger(__name__) PAYMENT_ORDERS_TABLE_NAME = "payment_orders" PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN = "provision_username" PAYMENT_ORDERS_PROVISION_USERNAME_INDEX = "ix_payment_orders_provision_username" def create_engine_and_session_factory( database_url: str, *, echo: bool = False, ) -> tuple[AsyncEngine, async_sessionmaker[AsyncSession]]: engine = create_async_engine( database_url, echo=echo, pool_pre_ping=True, pool_recycle=3600, ) session_factory = async_sessionmaker(engine, expire_on_commit=False, class_=AsyncSession) return engine, session_factory def _quote_mysql_identifier(identifier: str) -> str: return f"`{identifier.replace('`', '``')}`" async def _get_mysql_single_column_indexes( connection: AsyncConnection, *, table_name: str, column_name: str, non_unique: int, ) -> list[str]: result = await connection.execute( text( """ SELECT index_name FROM information_schema.statistics WHERE table_schema = DATABASE() AND table_name = :table_name GROUP BY index_name, non_unique HAVING non_unique = :non_unique AND COUNT(*) = 1 AND MAX(CASE WHEN column_name = :column_name THEN 1 ELSE 0 END) = 1 """ ), { "table_name": table_name, "column_name": column_name, "non_unique": non_unique, }, ) return [row[0] for row in result] async def _normalize_payment_order_provision_username_index( connection: AsyncConnection, ) -> None: if connection.dialect.name != "mysql": return unique_index_names = await _get_mysql_single_column_indexes( connection, table_name=PAYMENT_ORDERS_TABLE_NAME, column_name=PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN, non_unique=0, ) quoted_table_name = _quote_mysql_identifier(PAYMENT_ORDERS_TABLE_NAME) for index_name in unique_index_names: logger.info( "Dropping stale UNIQUE index %s on %s.%s", index_name, PAYMENT_ORDERS_TABLE_NAME, PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN, ) await connection.execute( text( "DROP INDEX " f"{_quote_mysql_identifier(index_name)} " f"ON {quoted_table_name}" ) ) non_unique_index_names = await _get_mysql_single_column_indexes( connection, table_name=PAYMENT_ORDERS_TABLE_NAME, column_name=PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN, non_unique=1, ) if non_unique_index_names: return logger.info( "Creating non-unique index %s on %s.%s", PAYMENT_ORDERS_PROVISION_USERNAME_INDEX, PAYMENT_ORDERS_TABLE_NAME, PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN, ) await connection.execute( text( "CREATE INDEX " f"{_quote_mysql_identifier(PAYMENT_ORDERS_PROVISION_USERNAME_INDEX)} " f"ON {quoted_table_name} " f"({_quote_mysql_identifier(PAYMENT_ORDERS_PROVISION_USERNAME_COLUMN)})" ) ) async def ensure_database_exists( server_url: str, database_name: str, *, echo: bool = False, ) -> None: engine = create_async_engine( server_url, echo=echo, pool_pre_ping=True, pool_recycle=3600, ) try: async with engine.begin() as connection: await connection.execute( text( "CREATE DATABASE IF NOT EXISTS " f"{_quote_mysql_identifier(database_name)} " "CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci" ) ) finally: await engine.dispose() async def init_db(engine: AsyncEngine) -> None: async with engine.begin() as connection: await connection.run_sync(Base.metadata.create_all) await _normalize_payment_order_provision_username_index(connection)