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