Files
RemnaWaveBOT/app/db/session.py
2026-04-22 18:34:13 +03:00

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)