Files
govoplan-core/src/govoplan_core/db/session.py
T

124 lines
3.9 KiB
Python

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()