diff options
| author | Christian Cleberg <[email protected]> | 2026-04-09 19:45:45 -0500 |
|---|---|---|
| committer | Christian Cleberg <[email protected]> | 2026-04-09 19:45:45 -0500 |
| commit | acbff854f2da96bddcaede1385e7fefeba0fb34b (patch) | |
| tree | 48d4707cd3276d2825370f0793008a6ffcfff936 /src/srht_contrib/db.py | |
| download | hutch-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.py | 50 |
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() |
