from __future__ import annotations import contextlib import os import pathlib from collections.abc import AsyncIterator, Iterator, Sequence from contextvars import ContextVar from typing import Any from sqlalchemy import TextClause, text from sqlalchemy.engine import URL from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine, ) from .jinjasql import JinjaSQL class Database: def __init__( self, url: str | URL, engine_args: dict[str, Any] | None = None, session_args: dict[str, Any] | None = None, template_searchpath: str | os.PathLike[str] | Sequence[str | os.PathLike[str]] | None = None, ) -> None: engine_args = engine_args or {} self._engine = create_async_engine(url, **engine_args) session_args = session_args or {"autoflush": False, "expire_on_commit": False} self._session_factory = async_sessionmaker(bind=self._engine, **session_args) # JinjaSQL self.jsql = JinjaSQL(template_searchpath or pathlib.Path(__file__)) @property def engine(self) -> AsyncEngine: return self._engine @property def is_connected(self) -> bool: return _has_session() async def close(self) -> None: await self.engine.dispose() @contextlib.asynccontextmanager async def session_context( self, in_transaction: bool = False ) -> AsyncIterator[None]: if self.is_connected: yield return async with self._session_factory() as session: with self.set_session(session): if in_transaction: async with session.begin(): yield else: yield @property def session(self) -> AsyncSession: return _get_session() @contextlib.contextmanager def set_session(self, session: AsyncSession) -> Iterator[None]: token = _session_ctx.set(session) try: yield finally: _session_ctx.reset(token) def query_from_template( self, template_name: str, *, context: dict[str, Any] | None = None, ) -> TextClause: query, bind_params = self.jsql.prepare_query(template_name, context=context) return text(query).bindparams(**bind_params) # Session context _session_ctx: ContextVar[AsyncSession | None] = ContextVar("session_ctx", default=None) def _get_session() -> AsyncSession: session = _session_ctx.get() if session is None: raise LookupError("Session is not started.") return session def _has_session() -> bool: session = _session_ctx.get() return session is not None