summaryrefslogtreecommitdiff
path: root/src/srht_contrib/db.py
diff options
context:
space:
mode:
authorChristian Cleberg <[email protected]>2026-04-09 19:45:45 -0500
committerChristian Cleberg <[email protected]>2026-04-09 19:45:45 -0500
commitacbff854f2da96bddcaede1385e7fefeba0fb34b (patch)
tree48d4707cd3276d2825370f0793008a6ffcfff936 /src/srht_contrib/db.py
downloadhutch-stats-acbff854f2da96bddcaede1385e7fefeba0fb34b.tar.gz
hutch-stats-acbff854f2da96bddcaede1385e7fefeba0fb34b.tar.bz2
hutch-stats-acbff854f2da96bddcaede1385e7fefeba0fb34b.zip
initial commit
Diffstat (limited to 'src/srht_contrib/db.py')
-rw-r--r--src/srht_contrib/db.py50
1 files changed, 50 insertions, 0 deletions
diff --git a/src/srht_contrib/db.py b/src/srht_contrib/db.py
new file mode 100644
index 0000000..d96de57
--- /dev/null
+++ b/src/srht_contrib/db.py
@@ -0,0 +1,50 @@
+from __future__ import annotations
+
+from collections.abc import Generator
+
+from fastapi import HTTPException, Request, status
+from sqlalchemy import Engine, create_engine, text
+from sqlalchemy.pool import StaticPool
+from sqlalchemy.orm import Session, declarative_base, sessionmaker
+
+from srht_contrib.config import Settings
+
+Base = declarative_base()
+
+
+def make_engine(settings: Settings) -> Engine:
+ connect_args = {"check_same_thread": False} if settings.database_url.startswith("sqlite") else {}
+ engine_kwargs = {"future": True, "connect_args": connect_args}
+ if settings.database_url in {"sqlite://", "sqlite:///:memory:"}:
+ engine_kwargs["poolclass"] = StaticPool
+ return create_engine(settings.database_url, **engine_kwargs)
+
+
+def make_session_factory(settings: Settings) -> sessionmaker[Session]:
+ engine = make_engine(settings)
+ return sessionmaker(bind=engine, autoflush=False, autocommit=False, expire_on_commit=False)
+
+
+def validate_db(bind: Engine) -> None:
+ with bind.connect() as connection:
+ connection.execute(text("SELECT 1"))
+
+
+def get_db() -> Generator[Session, None, None]:
+ raise RuntimeError("Use get_db(request) dependency injection with a Request parameter.")
+
+
+def get_session_factory(request: Request) -> sessionmaker[Session]:
+ session_factory = getattr(request.app.state, "session_factory", None)
+ if session_factory is None:
+ raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Database not configured.")
+ return session_factory
+
+
+def get_db_session(request: Request) -> Generator[Session, None, None]:
+ session_factory = get_session_factory(request)
+ db = session_factory()
+ try:
+ yield db
+ finally:
+ db.close()