from collections.abc import Generator import logging import time from sqlalchemy import bindparam, 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") SHARED_TALKINGQ_REQUIRED_TABLES = ( "device_auth", "device_configs", "conversation_histories", "conversation_messages", "roles", "role_languages", "device_firmware_update", "system_config", "parents", "children", "parent_child_relations", "device_bindings", "device_bind_sessions", "device_bind_history", "cards", "device_settings", "im_conversations", "im_messages", "child_location_current", "child_location_history", ) @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 check_shared_schema_tables() -> None: if not settings.uses_shared_talkingq_db: return query = text( """ SELECT table_name FROM information_schema.tables WHERE table_schema = :table_schema AND table_name IN :table_names """ ).bindparams(bindparam("table_names", expanding=True)) with engine.connect() as conn: rows = conn.execute( query, { "table_schema": settings.db_name, "table_names": list(SHARED_TALKINGQ_REQUIRED_TABLES), }, ) existing_tables = {row[0] for row in rows} missing_tables = sorted(set(SHARED_TALKINGQ_REQUIRED_TABLES) - existing_tables) if missing_tables: raise RuntimeError( "shared talkingq schema is incomplete; missing tables: " + ", ".join(missing_tables) ) logger.info( "shared talkingq schema check succeeded", extra={ "event": "db_shared_schema_check", "db_name": settings.db_name, "table_count": len(SHARED_TALKINGQ_REQUIRED_TABLES), }, ) def close_db_engine() -> None: engine.dispose() logger.info("database engine disposed", extra={"event": "db_close"})