krz/hutch-stats

Server-side utility for calculating contributions for sourcehut users.

clone: git clone https://gitbay.org/krz/hutch-stats.git

main: src/srht_contrib/db.py · raw

 1from __future__ import annotations
 2
 3from collections.abc import Generator
 4
 5from fastapi import HTTPException, Request, status
 6from sqlalchemy import Engine, create_engine, event, text
 7from sqlalchemy.pool import StaticPool
 8from sqlalchemy.orm import Session, declarative_base, sessionmaker
 9
10from srht_contrib.config import Settings
11
12Base = declarative_base()
13
14
15def make_engine(settings: Settings) -> Engine:
16    connect_args = {}
17    if settings.database_url.startswith("sqlite"):
18        connect_args = {
19            "check_same_thread": False,
20            "timeout": settings.sqlite_busy_timeout_seconds,
21        }
22    engine_kwargs = {"future": True, "connect_args": connect_args}
23    if settings.database_url in {"sqlite://", "sqlite:///:memory:"}:
24        engine_kwargs["poolclass"] = StaticPool
25    engine = create_engine(settings.database_url, **engine_kwargs)
26
27    if settings.database_url.startswith("sqlite"):
28        @event.listens_for(engine, "connect")
29        def _configure_sqlite(dbapi_connection, connection_record) -> None:  # type: ignore[unused-ignore]
30            cursor = dbapi_connection.cursor()
31            cursor.execute(f"PRAGMA busy_timeout = {int(settings.sqlite_busy_timeout_seconds * 1000)}")
32            if settings.database_url not in {"sqlite://", "sqlite:///:memory:"}:
33                cursor.execute("PRAGMA journal_mode = WAL")
34                cursor.execute("PRAGMA synchronous = NORMAL")
35            cursor.close()
36
37    return engine
38
39
40def make_session_factory(settings: Settings) -> sessionmaker[Session]:
41    engine = make_engine(settings)
42    return sessionmaker(bind=engine, autoflush=False, autocommit=False, expire_on_commit=False)
43
44
45def validate_db(bind: Engine) -> None:
46    with bind.connect() as connection:
47        connection.execute(text("SELECT 1"))
48
49
50def get_db() -> Generator[Session, None, None]:
51    raise RuntimeError("Use get_db(request) dependency injection with a Request parameter.")
52
53
54def get_session_factory(request: Request) -> sessionmaker[Session]:
55    session_factory = getattr(request.app.state, "session_factory", None)
56    if session_factory is None:
57        raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database not configured.")
58    return session_factory
59
60
61def get_db_session(request: Request) -> Generator[Session, None, None]:
62    session_factory = get_session_factory(request)
63    db = session_factory()
64    try:
65        yield db
66    finally:
67        db.close()