from collections.abc import Generator import logging import time from sqlalchemy import create_engine, event, text from sqlalchemy.orm import Session, sessionmaker try: # For module mode: `uvicorn app.main:app` from app.settings import settings except ModuleNotFoundError: # For script mode: `python app/main.py` or VS Code "Run Python File" from settings import settings def _create_engine() -> create_engine: if settings.db_type == "sqlite": return create_engine( settings.sqlite_dsn, connect_args={"check_same_thread": False}, future=True, ) return create_engine( settings.mysql_dsn, pool_size=5, max_overflow=10, pool_pre_ping=True, pool_recycle=1800, future=True, ) engine = _create_engine() @event.listens_for(engine, "connect") def set_sqlite_pragma(dbapi_conn, connection_record): if settings.db_type == "sqlite": cursor = dbapi_conn.cursor() cursor.execute("PRAGMA foreign_keys=ON") cursor.close() logger = logging.getLogger("app.db") @event.listens_for(engine, "before_cursor_execute") def before_cursor_execute( conn, cursor, statement, parameters, context, executemany, ): conn.info.setdefault("query_start_time", []).append(time.perf_counter()) @event.listens_for(engine, "after_cursor_execute") def after_cursor_execute( conn, cursor, statement, parameters, context, executemany, ): start_time = None query_start_time = conn.info.get("query_start_time") if query_start_time: start_time = query_start_time.pop(-1) if start_time is None: return elapsed_ms = (time.perf_counter() - start_time) * 1000 if elapsed_ms >= settings.slow_sql_ms: logger.warning( "slow sql detected", extra={ "event": "slow_sql", "sql_ms": round(elapsed_ms, 2), "statement": " ".join(statement.split())[:280], }, ) SessionLocal = sessionmaker( bind=engine, autocommit=False, autoflush=False, expire_on_commit=False, ) def get_db() -> Generator[Session, None, None]: db = SessionLocal() try: yield db finally: db.close() def init_db_tables() -> None: from app.models import Base Base.metadata.create_all(bind=engine) logger.info("database tables initialized", extra={"event": "db_init"}) def check_db_connection() -> None: with engine.connect() as conn: conn.execute(text("SELECT 1")) logger.info("database connection check succeeded", extra={"event": "db_check"}) def close_db_engine() -> None: engine.dispose() logger.info("database engine disposed", extra={"event": "db_close"})