153 lines
4.3 KiB
Python
153 lines
4.3 KiB
Python
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)
|