from __future__ import annotations from collections.abc import Generator, Mapping import os from typing import Any from sqlalchemy import create_engine from sqlalchemy.engine import Engine from sqlalchemy.orm import Session, sessionmaker from govoplan_core.db.query_metrics import instrument_engine _POOL_SETTING_BOUNDS: dict[str, tuple[int, int]] = { "GOVOPLAN_DB_POOL_SIZE": (1, 100), "GOVOPLAN_DB_MAX_OVERFLOW": (0, 200), "GOVOPLAN_DB_POOL_TIMEOUT_SECONDS": (1, 300), "GOVOPLAN_DB_POOL_RECYCLE_SECONDS": (0, 86_400), } def default_connect_args(database_url: str) -> dict[str, Any]: return {"check_same_thread": False} if database_url.startswith("sqlite") else {} def create_database_engine( database_url: str, *, pool_pre_ping: bool = True, connect_args: Mapping[str, Any] | None = None, **kwargs: Any, ) -> Engine: merged_connect_args = dict(default_connect_args(database_url)) if connect_args: merged_connect_args.update(connect_args) if not database_url.startswith("sqlite"): kwargs.setdefault( "pool_size", _pool_setting("GOVOPLAN_DB_POOL_SIZE", 5), ) kwargs.setdefault( "max_overflow", _pool_setting("GOVOPLAN_DB_MAX_OVERFLOW", 10), ) kwargs.setdefault( "pool_timeout", _pool_setting("GOVOPLAN_DB_POOL_TIMEOUT_SECONDS", 30), ) kwargs.setdefault( "pool_recycle", _pool_setting("GOVOPLAN_DB_POOL_RECYCLE_SECONDS", 1800), ) return instrument_engine(create_engine(database_url, pool_pre_ping=pool_pre_ping, connect_args=merged_connect_args, **kwargs)) def _pool_setting(name: str, default: int) -> int: raw = os.environ.get(name) if raw is None: return default try: value = int(raw) except ValueError as exc: raise ValueError(f"{name} must be an integer") from exc minimum, maximum = _POOL_SETTING_BOUNDS[name] if not minimum <= value <= maximum: raise ValueError( f"{name} must be between {minimum} and {maximum}", ) return value class DatabaseHandle: def __init__(self, database_url: str, *, engine: Engine | None = None) -> None: self.database_url = database_url self.engine = instrument_engine(engine) if engine is not None else create_database_engine(database_url) self.SessionLocal = sessionmaker(bind=self.engine, autoflush=False, autocommit=False, expire_on_commit=False) def session(self) -> Session: return self.SessionLocal() def dependency(self) -> Generator[Session, None, None]: with self.SessionLocal() as session: yield session def dispose(self) -> None: self.engine.dispose() _default_database: DatabaseHandle | None = None def configure_database(database_url: str, *, engine: Engine | None = None, dispose_previous: bool = False) -> DatabaseHandle: global _default_database if engine is None and _default_database is not None and _default_database.database_url == database_url: return _default_database previous_database = _default_database _default_database = DatabaseHandle(database_url, engine=engine) if dispose_previous and previous_database is not None and previous_database is not _default_database: previous_database.dispose() return _default_database def set_database(handle: DatabaseHandle) -> DatabaseHandle: global _default_database _default_database = handle return handle def reset_database(*, dispose: bool = False) -> None: global _default_database if dispose and _default_database is not None: _default_database.dispose() _default_database = None def get_database() -> DatabaseHandle: if _default_database is None: raise RuntimeError("GovOPlaN database is not configured") return _default_database def get_session() -> Generator[Session, None, None]: yield from get_database().dependency()