from __future__ import annotations import contextlib import functools import logging import pathlib from collections.abc import Callable, Iterator from contextvars import ContextVar from dataclasses import dataclass from types import TracebackType from typing import Any, Self, overload import jinja2 import sqlalchemy as sa from jinja2sql import Jinja2SQL from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.sql.elements import TextClause logger = logging.getLogger(__name__) class Database: def __init__( self, url: str | sa.URL, engine_args: dict[str, Any] | None = None, session_args: dict[str, Any] | None = None, jinja2sql: Jinja2SQL | None = None, ) -> None: engine_args = engine_args or {} self._engine = sa.create_engine(url, **engine_args) session_args = session_args or { "autoflush": False, "expire_on_commit": False, } self._session_factory = sessionmaker(bind=self._engine, **session_args) # Jinja2SQL self._jinja2sql = jinja2sql or Jinja2SQL( jinja2.Environment( loader=jinja2.FileSystemLoader(searchpath=pathlib.Path(__file__)) ) ) # Contexts self._session_ctx: _ScopedVar[Session | None] = _ScopedVar("session", None) self._in_transaction_ctx: _ScopedVar[bool] = _ScopedVar("in_transaction", False) self._rollback_mode_ctx: _ScopedVar[bool] = _ScopedVar("rollback_mode", False) def __enter__(self) -> Self: return self def __exit__( self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None, ) -> None: self.close() @property def engine(self) -> sa.Engine: return self._engine def close(self) -> None: self.engine.dispose() @property def session(self) -> Session: return self._get_session() def has_session(self) -> bool: """Check if the current session is started.""" return self._session_ctx.get() is not None @contextlib.contextmanager def global_context(self) -> Iterator[None]: """Replace per-task ContextVars with a single shared value for this scope.""" with ( self._session_ctx.global_mode(), self._in_transaction_ctx.global_mode(), self._rollback_mode_ctx.global_mode(), ): yield @contextlib.contextmanager def session_factory( self, *, bind: sa.Connection | sa.Engine | None = None, replace: bool = False, ) -> Iterator[Session]: """Create a new session.""" if self.has_session(): if not replace: yield self.session return self.session.close() with ( self._session_factory(bind=bind or self.engine) as session, self._set_session(session), ): yield session @overload def transaction[T, **P](self, func: Callable[P, T]) -> Callable[P, T]: ... @overload def transaction( self, func: None = None, *, commit_on_error: type[Exception] | tuple[type[Exception], ...] | None = None, ) -> _Transaction: ... def transaction[T, **P]( self, func: Callable[P, T] | None = None, *, commit_on_error: type[Exception] | tuple[type[Exception], ...] | None = None, ) -> Callable[P, T] | _Transaction: """Use as `with db.transaction():`, `@db.transaction`, or `@db.transaction(...)`.""" transaction = _Transaction( lambda: self._transaction_cm(commit_on_error=commit_on_error) ) if func is not None: return transaction(func) return transaction @overload def autocommit[T, **P](self, func: Callable[P, T]) -> Callable[P, T]: ... @overload def autocommit(self, func: None = None) -> _Transaction: ... def autocommit[T, **P]( self, func: Callable[P, T] | None = None, ) -> Callable[P, T] | _Transaction: """Like `transaction`, but each statement commits immediately (AUTOCOMMIT isolation). Use as `with db.autocommit():`, `@db.autocommit`, or `@db.autocommit()`. """ transaction = _Transaction(self._autocommit_cm) if func is not None: return transaction(func) return transaction @contextlib.contextmanager def _transaction_cm( self, *, commit_on_error: type[Exception] | tuple[type[Exception], ...] | None = None, ) -> Iterator[None]: if self._in_transaction_ctx.get(): raise RuntimeError( "Transaction already started, nested transactions are not supported." ) in_rollback_mode = self._rollback_mode_ctx.get() with self._transaction() as session: token = self._in_transaction_ctx.set(True) try: yield if in_rollback_mode: session.flush() else: session.commit() except Exception as exc: if not in_rollback_mode: if commit_on_error and isinstance(exc, commit_on_error): session.commit() else: session.rollback() raise finally: self._in_transaction_ctx.reset(token) @contextlib.contextmanager def _autocommit_cm(self) -> Iterator[None]: if self._rollback_mode_ctx.get(): yield return if self.in_transaction(): raise RuntimeError("Cannot use autocommit inside a transaction.") with self.engine.connect() as conn: conn = conn.execution_options(isolation_level="AUTOCOMMIT") with self.session_factory(bind=conn, replace=True): yield @contextlib.contextmanager def _transaction(self) -> Iterator[Session]: if self.has_session() and self._rollback_mode_ctx.get(): yield self.session return with self.session_factory(replace=True) as session: yield session @contextlib.contextmanager def rollback_transaction(self) -> Iterator[None]: if self.has_session(): raise RuntimeError( "Cannot use rollback_transaction inside an existing session." ) with self.engine.begin() as conn, self.session_factory(bind=conn): self._rollback_mode_ctx.set(True) try: yield finally: conn.rollback() self._rollback_mode_ctx.set(False) def in_transaction(self) -> bool: """Check if the current session is in a transaction.""" return self.has_session() and ( self._in_transaction_ctx.get() or self._rollback_mode_ctx.get() ) @property def jinja2sql(self) -> Jinja2SQL: return self._jinja2sql def query_from_template( self, template_name: str, *, context: dict[str, Any] | None = None, ) -> TextClause: context = context or {} context["db_dialect"] = self.engine.dialect.name query, bind_params = self.jinja2sql.from_file(template_name, context=context) return sa.text(query).bindparams(**dict(bind_params)) @contextlib.contextmanager def _set_session(self, session: Session) -> Iterator[None]: """Set the current session.""" token = self._session_ctx.set(session) try: yield finally: self._session_ctx.reset(token) def _get_session(self) -> Session: """Get the current session.""" if (session := self._session_ctx.get()) is None: raise RuntimeError("Session is not started.") return session class _Transaction: """Supports context manager and decorator usage for database transactions. Instantiate via `Database.transaction` or `Database.autocommit`. """ def __init__( self, cm_factory: Callable[[], contextlib.AbstractContextManager[None]], ) -> None: self._cm_factory = cm_factory self._ctx: contextlib.AbstractContextManager[None] | None = None def __enter__(self) -> None: self._ctx = self._cm_factory() return self._ctx.__enter__() def __exit__( self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None, ) -> bool | None: assert self._ctx is not None return self._ctx.__exit__(exc_type, exc_val, exc_tb) def __call__[T, **P](self, func: Callable[P, T]) -> Callable[P, T]: cm_factory = self._cm_factory @functools.wraps(func) def wrapper(*args: P.args, **kwargs: P.kwargs) -> T: with _Transaction(cm_factory): return func(*args, **kwargs) return wrapper @dataclass class _ScopedToken[T]: """Opaque token returned by _ScopedVar.set(); passed back to _ScopedVar.reset().""" is_global: bool inner: Any # Token[T] (ContextVar) or T (global value) class _ScopedVar[T]: """ContextVar that can switch to global (cross-task) storage.""" def __init__(self, name: str, default: T) -> None: self._context = ContextVar(name, default=default) self._value = default self._is_global = False def get(self) -> T: return self._value if self._is_global else self._context.get() def set(self, value: T) -> _ScopedToken[T]: if self._is_global: old, self._value = self._value, value return _ScopedToken(is_global=True, inner=old) return _ScopedToken(is_global=False, inner=self._context.set(value)) def reset(self, token: _ScopedToken[T]) -> None: if token.is_global: self._value = token.inner else: self._context.reset(token.inner) @contextlib.contextmanager def global_mode(self) -> Iterator[None]: saved = self._value self._is_global = True try: yield finally: self._is_global = False self._value = saved