chore: rework DB module (#324)

This commit is contained in:
d3vyce
2026-06-25 21:11:50 +02:00
committed by GitHub
parent 22f307d0fc
commit 9698a0743b
23 changed files with 1662 additions and 931 deletions
+1 -1
View File
@@ -167,7 +167,7 @@ user = await UserCrud.update(session, UserUpdate(credits=10), [User.id == user_i
``` ```
!!! warning !!! warning
`with_for_update` requires an open transaction. Wrap your call in `async with session.begin()` or use the `get_transaction` helper if you are not already inside one. `with_for_update` requires an open transaction. Wrap your call in `async with session.begin()` or use the `transaction` helper if you are not already inside one.
!!! note !!! note
`NOWAIT` raises `sqlalchemy.exc.OperationalError` immediately if the row is locked rather than waiting. `NOWAIT` raises `sqlalchemy.exc.OperationalError` immediately if the row is locked rather than waiting.
+94 -57
View File
@@ -7,96 +7,137 @@ SQLAlchemy async session management with transactions, table locking, advisory l
## Overview ## Overview
The `db` module provides helpers to create FastAPI dependencies and context managers for `AsyncSession`, along with utilities for nested transactions, table locks, advisory locks, and polling for row changes. The `db` module is built around one object, [`Database`](../reference/db.md#fastapi_toolsets.db.Database), which owns the engine and sessionmaker and exposes the FastAPI dependency, a commit-before-response middleware, session/transaction context managers, and table locking. Free helpers cover savepoint-aware transactions, advisory locks, many-to-many association tables, and row-change polling.
## Session dependency ## Setup
Use [`create_db_dependency`](../reference/db.md#fastapi_toolsets.db.create_db_dependency) to create a FastAPI dependency that yields a session and auto-commits on success: Create one `Database` for your app. Provide a **URL** (the facade builds and disposes the engine) or pass an existing **`engine=`** you own (e.g. for Alembic or `event.listen`). The session factory is built internally with `expire_on_commit=False`.
```python ```python
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker from fastapi import Depends, FastAPI
from fastapi_toolsets.db import create_db_dependency from sqlalchemy.ext.asyncio import AsyncSession
engine = create_async_engine(url="postgresql+asyncpg://...", future=True) from fastapi_toolsets.db import Database
session_maker = async_sessionmaker(bind=engine, expire_on_commit=False)
get_db = create_db_dependency(session_maker=session_maker) db = Database("postgresql+asyncpg://postgres:postgres@localhost/app")
@router.get("/users") app = FastAPI()
async def list_users(session: AsyncSession = Depends(get_db)): db.install(app) # commit middleware + engine disposal on shutdown
@app.get("/users")
async def list_users(session: AsyncSession = Depends(db)):
... ...
``` ```
## Session context manager The `Database` instance **is** the dependency: use it directly as `Depends(db)`. The whole request runs as a single transaction (CRUD writes use savepoints under it).
Use [`create_db_context`](../reference/db.md#fastapi_toolsets.db.create_db_context) for sessions outside request handlers (e.g. background tasks, CLI commands): ## Committing before the response
[`db.install(app)`](../reference/db.md#fastapi_toolsets.db.Database) adds a middleware that commits the request's session when the response starts, after the endpoint returns and before the body is sent. With the middleware installed, the dependency does not commit again.
The request is committed as a single transaction:
- **Read-after-write**: a follow-up request sees the write.
- **Atomicity**: multi-write endpoints roll back as a unit on failure.
- **Errors roll back**: on a raised exception the session rolls back and nothing is committed.
Without `install`, the session commits in the dependency teardown, which runs after the response has been sent.
!!! warning "Streaming / SSE endpoints"
For a `StreamingResponse` / `EventSourceResponse`, the commit fires at the **start** of the stream. A stream that **writes** must open a short-lived session per write with [`db.session()`](#session-context-manager); the start-time commit will not flush writes made later during the stream.
## Lifespan
`db.install(app)` disposes the engine on shutdown, composing around your own lifespan:
```python ```python
from fastapi_toolsets.db import create_db_context from contextlib import asynccontextmanager
db_context = create_db_context(session_maker=session_maker) @asynccontextmanager
async def lifespan(app):
await warm_cache() # your startup
yield
await flush_metrics() # your shutdown
app = FastAPI(lifespan=lifespan)
db.install(app) # your shutdown runs first, then the engine is disposed
```
If you have no lifespan of your own, [`db.lifespan`](../reference/db.md#fastapi_toolsets.db.Database) works standalone as `FastAPI(lifespan=db.lifespan)`. Engine disposal is idempotent and is a no-op when you passed your own `engine=`.
## Session context manager
Use [`db.session()`](../reference/db.md#fastapi_toolsets.db.Database) for sessions outside request handlers (e.g. background tasks, CLI commands). It commits on clean exit and rolls back on exception:
```python
async def seed(): async def seed():
async with db_context() as session: async with db.session() as session:
... ...
``` ```
## Nested transactions ## Transactions
[`get_transaction`](../reference/db.md#fastapi_toolsets.db.get_transaction) handles savepoints automatically, allowing safe nesting: [`transaction`](../reference/db.md#fastapi_toolsets.db.transaction) opens a transaction on a session, using a savepoint when one is already open so it nests safely:
```python ```python
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
async def create_user_with_role(session=session): async def create_user_with_role(session):
async with get_transaction(session=session): async with transaction(session):
... ...
async with get_transaction(session=session): # uses savepoint async with transaction(session): # uses a savepoint
... ...
``` ```
When you have a `Database`, [`db.begin()`](../reference/db.md#fastapi_toolsets.db.Database) opens a session already inside a transaction:
```python
async with db.begin() as session:
session.add(User(name="ada")) # commits on exit, rolls back on exception
```
## Table locking ## Table locking
[`lock_tables`](../reference/db.md#fastapi_toolsets.db.lock_tables) acquires PostgreSQL table-level locks before executing critical sections. It opens a **dedicated session** internally and yields it to the caller, so the lock is guaranteed to be released when the context exits: [`db.lock_tables`](../reference/db.md#fastapi_toolsets.db.Database) acquires PostgreSQL table-level locks for a critical section. It opens a dedicated session internally and releases the lock when the context exits:
```python ```python
from fastapi_toolsets.db import lock_tables, LockMode from fastapi_toolsets.db import LockMode
async with lock_tables(session_maker=session_maker, tables=[User], mode=LockMode.EXCLUSIVE) as session: async with db.lock_tables([User], mode=LockMode.EXCLUSIVE) as session:
# No other transaction can modify User until this block exits # No other transaction can modify User until this block exits
... ...
``` ```
Available lock modes are defined in [`LockMode`](../reference/db.md#fastapi_toolsets.db.LockMode): `ACCESS_SHARE`, `ROW_SHARE`, `ROW_EXCLUSIVE`, `SHARE_UPDATE_EXCLUSIVE`, `SHARE`, `SHARE_ROW_EXCLUSIVE`, `EXCLUSIVE`, `ACCESS_EXCLUSIVE`. Available lock modes are defined in [`LockMode`](../reference/db.md#fastapi_toolsets.db.LockMode): `ACCESS_SHARE`, `ROW_SHARE`, `ROW_EXCLUSIVE`, `SHARE_UPDATE_EXCLUSIVE`, `SHARE`, `SHARE_ROW_EXCLUSIVE`, `EXCLUSIVE`, `ACCESS_EXCLUSIVE`.
Pass `timeout` to limit how long the lock waits before giving up. On timeout, a [`LockTimeoutError`](../reference/exceptions.md#fastapi_toolsets.exceptions.exceptions.LockTimeoutError) is raised instead of a raw database error: Pass `timeout` to limit how long the lock waits. On timeout, a [`LockTimeoutError`](../reference/exceptions.md#fastapi_toolsets.exceptions.exceptions.LockTimeoutError) is raised instead of a raw database error:
```python ```python
async with lock_tables(session_maker, [Order], timeout="2s") as session: async with db.lock_tables([Order], timeout="2s") as session:
... ...
``` ```
## Advisory locking ## Advisory locking
[`advisory_lock`](../reference/db.md#fastapi_toolsets.db.advisory_lock) acquires a PostgreSQL session-level advisory lock. The lock is released explicitly when the context exits, regardless of whether the transaction has committed. [`advisory_lock`](../reference/db.md#fastapi_toolsets.db.advisory_lock) acquires a PostgreSQL session-level advisory lock on a session you provide. The lock is released when the context exits:
```python ```python
from fastapi_toolsets.db import advisory_lock from fastapi_toolsets.db import advisory_lock
# Blocking exclusive lock waits until the lock is free # Blocking exclusive lock: waits until the lock is free
async with advisory_lock(session=session, key=42): async with advisory_lock(session=session, key=42):
... ...
# Non-blocking yields False immediately if already held # Non-blocking: yields False immediately if already held
async with advisory_lock(session=session, key=42, nowait=True) as acquired: async with advisory_lock(session=session, key=42, nowait=True) as acquired:
if not acquired: if not acquired:
raise HTTPException(409, "Resource is locked") raise HTTPException(409, "Resource is locked")
# Blocking with a timeout raises LockTimeoutError if not acquired in time # Blocking with a timeout: raises LockTimeoutError if not acquired in time
async with advisory_lock(session=session, key=42, timeout="5s"): async with advisory_lock(session=session, key=42, timeout="5s"):
... ...
# Shared multiple readers allowed simultaneously, blocks exclusive writers # Shared lock: multiple readers allowed simultaneously, blocks exclusive writers
async with advisory_lock(session=session, key=42, shared=True): async with advisory_lock(session=session, key=42, shared=True):
... ...
@@ -106,11 +147,11 @@ async with advisory_lock(session=session, key=(1, user_id)):
``` ```
!!! note !!! note
Advisory locks use PostgreSQL session-level functions (`pg_advisory_lock` / `pg_advisory_unlock`). The lock is tied to the database connection, not the SQLAlchemy transaction it is released when the context exits, even if the surrounding transaction is still open. Advisory locks use PostgreSQL session-level functions (`pg_advisory_lock` / `pg_advisory_unlock`). The lock is tied to the database connection, not the SQLAlchemy transaction, so it is released when the context exits even if the surrounding transaction is still open.
## Row-change polling ## Row-change polling
[`wait_for_row_change`](../reference/db.md#fastapi_toolsets.db.wait_for_row_change) polls a row until a specific column changes value, useful for waiting on async side effects: [`wait_for_row_change`](../reference/db.md#fastapi_toolsets.db.wait_for_row_change) polls a row until a specific column changes value:
```python ```python
from fastapi_toolsets.db import wait_for_row_change from fastapi_toolsets.db import wait_for_row_change
@@ -120,7 +161,7 @@ await wait_for_row_change(
session=session, session=session,
model=Order, model=Order,
pk_value=order_id, pk_value=order_id,
columns=[Order.status], columns=["status"],
interval=1.0, interval=1.0,
timeout=30.0, timeout=30.0,
) )
@@ -128,28 +169,24 @@ await wait_for_row_change(
## Creating a database ## Creating a database
!!! info "Added in `v2.1`" [`create_database`](../reference/db.md#fastapi_toolsets.db.testing.create_database) (in `fastapi_toolsets.db.testing`) connects to *server_url* and issues a `CREATE DATABASE` statement:
[`create_database`](../reference/db.md#fastapi_toolsets.db.create_database) creates a database at a given URL. It connects to *server_url* and issues a `CREATE DATABASE` statement:
```python ```python
from fastapi_toolsets.db import create_database from fastapi_toolsets.db.testing import create_database
SERVER_URL = "postgresql+asyncpg://postgres:postgres@localhost/postgres" SERVER_URL = "postgresql+asyncpg://postgres:postgres@localhost/postgres"
await create_database(db_name="myapp_test", server_url=SERVER_URL) await create_database(db_name="myapp_test", server_url=SERVER_URL)
``` ```
For test isolation with automatic cleanup, use [`create_worker_database`](../reference/pytest.md#fastapi_toolsets.pytest.utils.create_worker_database) from the `pytest` module instead — it handles drop-before, create, and drop-after automatically. For test isolation with automatic cleanup, use [`create_worker_database`](../reference/pytest.md#fastapi_toolsets.pytest.utils.create_worker_database) from the `pytest` module, which handles drop-before, create, and drop-after.
## Cleaning up tables ## Cleaning up tables
!!! info "Added in `v2.1`" [`cleanup_tables`](../reference/db.md#fastapi_toolsets.db.testing.cleanup_tables) (in `fastapi_toolsets.db.testing`) truncates all tables:
[`cleanup_tables`](../reference/db.md#fastapi_toolsets.db.cleanup_tables) truncates all tables:
```python ```python
from fastapi_toolsets.db import cleanup_tables from fastapi_toolsets.db.testing import cleanup_tables
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
async def clean(db_session): async def clean(db_session):
@@ -159,50 +196,50 @@ async def clean(db_session):
## Many-to-Many helpers ## Many-to-Many helpers
SQLAlchemy's ORM collection API triggers lazy-loads when you append to a relationship inside a savepoint (e.g. inside `lock_tables` or a nested `get_transaction`). The three `m2m_*` helpers bypass the ORM collection entirely and issue direct SQL against the association table. The three `m2m_*` helpers modify a many-to-many association table with direct SQL, without loading the ORM collection.
### `m2m_add` insert associations ### `m2m_add`: insert associations
[`m2m_add`](../reference/db.md#fastapi_toolsets.db.m2m_add) inserts one or more rows into a secondary table without touching the ORM collection: [`m2m_add`](../reference/db.md#fastapi_toolsets.db.m2m_add) inserts one or more rows into a secondary table:
```python ```python
from fastapi_toolsets.db import lock_tables, m2m_add from fastapi_toolsets.db import m2m_add
async with lock_tables(session_maker, [Tag]) as session: async with db.lock_tables([Tag]) as session:
tag = await TagCrud.create(session, TagCreate(name="python")) tag = await TagCrud.create(session, TagCreate(name="python"))
await m2m_add(session, post, Post.tags, tag) await m2m_add(session, post, Post.tags, tag)
``` ```
Pass `ignore_conflicts=True` to silently skip associations that already exist: Pass `ignore_conflicts=True` to skip associations that already exist:
```python ```python
await m2m_add(session, post, Post.tags, tag, ignore_conflicts=True) await m2m_add(session, post, Post.tags, tag, ignore_conflicts=True)
``` ```
### `m2m_remove` delete associations ### `m2m_remove`: delete associations
[`m2m_remove`](../reference/db.md#fastapi_toolsets.db.m2m_remove) deletes specific association rows. Removing a non-existent association is a no-op: [`m2m_remove`](../reference/db.md#fastapi_toolsets.db.m2m_remove) deletes specific association rows. Removing a non-existent association is a no-op:
```python ```python
from fastapi_toolsets.db import get_transaction, m2m_remove from fastapi_toolsets.db import m2m_remove, transaction
async with get_transaction(session): async with transaction(session):
await m2m_remove(session, post, Post.tags, tag1, tag2) await m2m_remove(session, post, Post.tags, tag1, tag2)
``` ```
### `m2m_set` replace the full set ### `m2m_set`: replace the full set
[`m2m_set`](../reference/db.md#fastapi_toolsets.db.m2m_set) atomically replaces all associations: it deletes every existing row for the owner instance then inserts the new set. Passing no related instances clears the association entirely: [`m2m_set`](../reference/db.md#fastapi_toolsets.db.m2m_set) replaces all associations: it deletes every existing row for the owner instance then inserts the new set. Passing no related instances clears the association:
```python ```python
from fastapi_toolsets.db import get_transaction, m2m_set from fastapi_toolsets.db import m2m_set, transaction
# Replace all tags # Replace all tags
async with get_transaction(session): async with transaction(session):
await m2m_set(session, post, Post.tags, tag_a, tag_b) await m2m_set(session, post, Post.tags, tag_a, tag_b)
# Clear all tags # Clear all tags
async with get_transaction(session): async with transaction(session):
await m2m_set(session, post, Post.tags) await m2m_set(session, post, Post.tags)
``` ```
+1 -1
View File
@@ -134,7 +134,7 @@ SessionLocal = async_sessionmaker(engine, expire_on_commit=False, class_=EventSe
``` ```
!!! info "Callbacks fire on `session.commit()` only — not on savepoints." !!! info "Callbacks fire on `session.commit()` only — not on savepoints."
Savepoints created by [`get_transaction`](db.md) or `begin_nested()` do **not** Savepoints created by [`transaction`](db.md) or `begin_nested()` do **not**
trigger callbacks. All events accumulated across flushes are dispatched once trigger callbacks. All events accumulated across flushes are dispatched once
when the outermost `commit()` is called. when the outermost `commit()` is called.
+2 -2
View File
@@ -107,10 +107,10 @@ url = worker_database_url("postgresql+asyncpg://user:pass@localhost/myapp", defa
## Manual table cleanup ## Manual table cleanup
[`cleanup_tables`](../reference/db.md#fastapi_toolsets.db.cleanup_tables) truncates all tables in a single statement and can be called directly when you need more control: [`cleanup_tables`](../reference/db.md#fastapi_toolsets.db.testing.cleanup_tables) truncates all tables in a single statement and can be called directly when you need more control:
```python ```python
from fastapi_toolsets.db import cleanup_tables from fastapi_toolsets.pytest import cleanup_tables
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
async def clean(db_session): async def clean(db_session):
+20 -18
View File
@@ -1,46 +1,48 @@
# `db` # `db`
Here's the reference for all database session utilities, transaction helpers, and locking functions. Here's the reference for the `Database` facade, the transaction helper, locking
functions, many-to-many helpers, and row-watching utilities.
You can import them directly from `fastapi_toolsets.db`: You can import them directly from `fastapi_toolsets.db`:
```python ```python
from fastapi_toolsets.db import ( from fastapi_toolsets.db import (
Database,
LockMode, LockMode,
advisory_lock, advisory_lock,
cleanup_tables,
create_database,
create_db_dependency,
create_db_context,
get_transaction,
lock_tables, lock_tables,
m2m_add, m2m_add,
m2m_remove, m2m_remove,
m2m_set, m2m_set,
transaction,
wait_for_row_change, wait_for_row_change,
) )
``` ```
## ::: fastapi_toolsets.db.Database
## ::: fastapi_toolsets.db.transaction
## ::: fastapi_toolsets.db.LockMode ## ::: fastapi_toolsets.db.LockMode
## ::: fastapi_toolsets.db.create_db_dependency
## ::: fastapi_toolsets.db.create_db_context
## ::: fastapi_toolsets.db.get_transaction
## ::: fastapi_toolsets.db.lock_tables ## ::: fastapi_toolsets.db.lock_tables
## ::: fastapi_toolsets.db.advisory_lock ## ::: fastapi_toolsets.db.advisory_lock
## ::: fastapi_toolsets.db.wait_for_row_change
## ::: fastapi_toolsets.db.create_database
## ::: fastapi_toolsets.db.cleanup_tables
## ::: fastapi_toolsets.db.m2m_add ## ::: fastapi_toolsets.db.m2m_add
## ::: fastapi_toolsets.db.m2m_remove ## ::: fastapi_toolsets.db.m2m_remove
## ::: fastapi_toolsets.db.m2m_set ## ::: fastapi_toolsets.db.m2m_set
## ::: fastapi_toolsets.db.wait_for_row_change
Admin and test helpers live in `fastapi_toolsets.db.testing`:
```python
from fastapi_toolsets.db.testing import cleanup_tables, create_database
```
## ::: fastapi_toolsets.db.testing.create_database
## ::: fastapi_toolsets.db.testing.cleanup_tables
@@ -2,8 +2,10 @@ from fastapi import FastAPI
from fastapi_toolsets.exceptions import init_exceptions_handlers from fastapi_toolsets.exceptions import init_exceptions_handlers
from .db import db
from .routes import router from .routes import router
app = FastAPI() app = FastAPI()
db.install(app=app)
init_exceptions_handlers(app=app) init_exceptions_handlers(app=app)
app.include_router(router=router) app.include_router(router=router)
+5 -8
View File
@@ -1,17 +1,14 @@
from typing import Annotated from typing import Annotated
from fastapi import Depends from fastapi import Depends
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.ext.asyncio import AsyncSession
from fastapi_toolsets.db import create_db_context, create_db_dependency from fastapi_toolsets.db import Database
DATABASE_URL = "postgresql+asyncpg://postgres:postgres@localhost:5432/postgres" DATABASE_URL = "postgresql+asyncpg://postgres:postgres@localhost:5432/postgres"
engine = create_async_engine(url=DATABASE_URL, future=True) db = Database(url=DATABASE_URL)
async_session_maker = async_sessionmaker(bind=engine, expire_on_commit=False)
get_db = create_db_dependency(session_maker=async_session_maker) get_db = db
get_db_context = create_db_context(session_maker=async_session_maker)
SessionDep = Annotated[AsyncSession, Depends(db)]
SessionDep = Annotated[AsyncSession, Depends(get_db)]
+5 -2
View File
@@ -7,16 +7,19 @@ Example usage:
from fastapi import FastAPI, Depends from fastapi import FastAPI, Depends
from fastapi_toolsets.exceptions import init_exceptions_handlers from fastapi_toolsets.exceptions import init_exceptions_handlers
from fastapi_toolsets.crud import CrudFactory from fastapi_toolsets.crud import CrudFactory
from fastapi_toolsets.db import create_db_dependency from fastapi_toolsets.db import Database
from fastapi_toolsets.schemas import Response from fastapi_toolsets.schemas import Response
db = Database("postgresql+asyncpg://postgres:postgres@localhost/app")
app = FastAPI() app = FastAPI()
db.install(app)
init_exceptions_handlers(app) init_exceptions_handlers(app)
UserCrud = CrudFactory(User) UserCrud = CrudFactory(User)
@app.get("/users/{user_id}", response_model=Response[dict]) @app.get("/users/{user_id}", response_model=Response[dict])
async def get_user(user_id: int, session = Depends(get_db)): async def get_user(user_id: int, session = Depends(db)):
user = await UserCrud.get(session, [User.id == user_id]) user = await UserCrud.get(session, [User.id == user_id])
return Response(data={"user": user.username}, message="Success") return Response(data={"user": user.username}, message="Success")
""" """
+5 -5
View File
@@ -22,7 +22,7 @@ from sqlalchemy.orm import DeclarativeBase, QueryableAttribute, selectinload
from sqlalchemy.sql.base import ExecutableOption from sqlalchemy.sql.base import ExecutableOption
from sqlalchemy.sql.roles import WhereHavingRole from sqlalchemy.sql.roles import WhereHavingRole
from ..db import get_transaction from ..db import transaction
from ..exceptions import InvalidOrderFieldError, NotFoundError from ..exceptions import InvalidOrderFieldError, NotFoundError
from ..schemas import ( from ..schemas import (
CursorPaginatedResponse, CursorPaginatedResponse,
@@ -716,7 +716,7 @@ class AsyncCrud(Generic[ModelType]):
Returns: Returns:
Created model instance, or ``Response[schema]`` when ``schema`` is given. Created model instance, or ``Response[schema]`` when ``schema`` is given.
""" """
async with get_transaction(session): async with transaction(session):
m2m_exclude = cls._m2m_schema_fields() m2m_exclude = cls._m2m_schema_fields()
data = ( data = (
obj.model_dump(exclude=m2m_exclude) if m2m_exclude else obj.model_dump() obj.model_dump(exclude=m2m_exclude) if m2m_exclude else obj.model_dump()
@@ -1067,7 +1067,7 @@ class AsyncCrud(Generic[ModelType]):
Raises: Raises:
NotFoundError: If no record found NotFoundError: If no record found
""" """
async with get_transaction(session): async with transaction(session):
m2m_exclude = cls._m2m_schema_fields() m2m_exclude = cls._m2m_schema_fields()
# Eagerly load M2M relationships that will be updated so that # Eagerly load M2M relationships that will be updated so that
@@ -1127,7 +1127,7 @@ class AsyncCrud(Generic[ModelType]):
Returns: Returns:
Model instance Model instance
""" """
async with get_transaction(session): async with transaction(session):
values = obj.model_dump(exclude_unset=True) values = obj.model_dump(exclude_unset=True)
q = insert(cls.model).values(**values) q = insert(cls.model).values(**values)
if set_: if set_:
@@ -1189,7 +1189,7 @@ class AsyncCrud(Generic[ModelType]):
Returns: Returns:
``None``, or ``Response[None]`` when ``return_response=True``. ``None``, or ``Response[None]`` when ``return_response=True``.
""" """
async with get_transaction(session): async with transaction(session):
result = await session.execute(select(cls.model).where(and_(*filters))) result = await session.execute(select(cls.model).where(and_(*filters)))
objects = result.scalars().all() objects = result.scalars().all()
for obj in objects: for obj in objects:
-591
View File
@@ -1,591 +0,0 @@
"""Database utilities: sessions, transactions, and locks."""
import asyncio
from collections.abc import AsyncGenerator, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from enum import Enum
from typing import Any, TypeVar, cast
import asyncpg
from sqlalchemy import Table, delete, text, tuple_
from sqlalchemy import exc as sa_exc
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase, QueryableAttribute
from sqlalchemy.orm.relationships import RelationshipProperty
from .exceptions import LockTimeoutError, NotFoundError, PoolExhaustedError
def _is_lock_not_available(e: sa_exc.DBAPIError) -> bool:
return e.orig is not None and isinstance(
e.orig.__cause__, asyncpg.exceptions.LockNotAvailableError
)
__all__ = [
"LockMode",
"advisory_lock",
"cleanup_tables",
"create_database",
"create_db_context",
"create_db_dependency",
"get_transaction",
"lock_tables",
"m2m_add",
"m2m_remove",
"m2m_set",
"wait_for_row_change",
]
_SessionT = TypeVar("_SessionT", bound=AsyncSession)
def create_db_dependency(
session_maker: async_sessionmaker[_SessionT],
) -> Callable[[], AsyncGenerator[_SessionT, None]]:
"""Create a FastAPI dependency for database sessions.
Creates a dependency function that yields a session and auto-commits
if a transaction is active when the request completes.
Args:
session_maker: Async session factory from create_session_factory()
Returns:
An async generator function usable with FastAPI's Depends()
Example:
```python
from fastapi import Depends
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
from fastapi_toolsets.db import create_db_dependency
engine = create_async_engine("postgresql+asyncpg://...")
SessionLocal = async_sessionmaker(engine, expire_on_commit=False)
get_db = create_db_dependency(SessionLocal)
@app.get("/users")
async def list_users(session: AsyncSession = Depends(get_db)):
...
```
"""
async def get_db() -> AsyncGenerator[_SessionT, None]:
async with session_maker() as session:
try:
await session.connection()
except sa_exc.TimeoutError as e:
raise PoolExhaustedError() from e
yield session
if session.in_transaction():
await session.commit()
return get_db
def create_db_context(
session_maker: async_sessionmaker[_SessionT],
) -> Callable[[], AbstractAsyncContextManager[_SessionT]]:
"""Create a context manager for database sessions.
Creates a context manager for use outside of FastAPI request handlers,
such as in background tasks, CLI commands, or tests.
Args:
session_maker: Async session factory from create_session_factory()
Returns:
An async context manager function
Example:
```python
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
from fastapi_toolsets.db import create_db_context
engine = create_async_engine("postgresql+asyncpg://...")
SessionLocal = async_sessionmaker(engine, expire_on_commit=False)
get_db_context = create_db_context(SessionLocal)
async def background_task():
async with get_db_context() as session:
user = await UserCrud.get(session, [User.id == 1])
...
```
"""
get_db = create_db_dependency(session_maker)
return asynccontextmanager(get_db)
@asynccontextmanager
async def get_transaction(
session: AsyncSession,
) -> AsyncGenerator[AsyncSession, None]:
"""Get a transaction context, handling nested transactions.
If already in a transaction, creates a savepoint (nested transaction).
Otherwise, starts a new transaction.
Args:
session: AsyncSession instance
Yields:
The session within the transaction context
Example:
```python
async with get_transaction(session):
session.add(model)
# Auto-commits on exit, rolls back on exception
```
"""
if session.in_transaction():
async with session.begin_nested():
yield session
else:
async with session.begin():
yield session
class LockMode(str, Enum):
"""PostgreSQL table lock modes.
See: https://www.postgresql.org/docs/current/explicit-locking.html
"""
ACCESS_SHARE = "ACCESS SHARE"
ROW_SHARE = "ROW SHARE"
ROW_EXCLUSIVE = "ROW EXCLUSIVE"
SHARE_UPDATE_EXCLUSIVE = "SHARE UPDATE EXCLUSIVE"
SHARE = "SHARE"
SHARE_ROW_EXCLUSIVE = "SHARE ROW EXCLUSIVE"
EXCLUSIVE = "EXCLUSIVE"
ACCESS_EXCLUSIVE = "ACCESS EXCLUSIVE"
def lock_tables(
session_maker: async_sessionmaker[_SessionT],
tables: list[type[DeclarativeBase]],
*,
mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
timeout: str = "5s",
) -> AbstractAsyncContextManager[_SessionT]:
"""Lock PostgreSQL tables for the duration of a transaction.
Args:
session_maker: Async session factory used to create the dedicated
session.
tables: List of SQLAlchemy model classes to lock.
mode: Lock mode (default: SHARE UPDATE EXCLUSIVE).
timeout: Lock timeout (default: "5s").
Yields:
The dedicated session, open within the locked transaction.
Raises:
SQLAlchemyError: If the lock cannot be acquired within *timeout*.
Example:
```python
from fastapi_toolsets.db import lock_tables, LockMode
async with lock_tables(session_maker, [User, Account]) as session:
# Tables are locked; changes are committed when the context exits.
user = await UserCrud.get(session, [User.id == 1])
user.balance += 100
# With custom lock mode
async with lock_tables(session_maker, [Order], mode=LockMode.EXCLUSIVE) as session:
await process_order(session, order_id)
```
"""
table_names = ",".join(table.__tablename__ for table in tables)
@asynccontextmanager
async def _lock() -> AsyncGenerator[_SessionT, None]:
async with session_maker() as session:
try:
await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
await session.execute(text(f"LOCK {table_names} IN {mode.value} MODE"))
yield session
await session.commit()
except sa_exc.TimeoutError as e:
await session.rollback()
raise PoolExhaustedError(
f"Connection pool exhausted while locking '{table_names}'. "
) from e
except sa_exc.DBAPIError as e:
await session.rollback()
if _is_lock_not_available(e):
raise LockTimeoutError(
f"Lock on '{table_names}' could not be acquired within {timeout}."
) from e
raise # pragma: no cover
except BaseException:
await session.rollback()
raise
return _lock()
@asynccontextmanager
async def advisory_lock(
session: AsyncSession,
key: int | tuple[int, int],
*,
shared: bool = False,
nowait: bool = False,
timeout: str | None = None,
) -> AsyncGenerator[bool, None]:
"""Acquire a PostgreSQL session-level advisory lock.
Args:
session: AsyncSession instance.
key: Lock key — a single ``int`` (bigint) or a ``(int, int)`` pair for namespacing.
shared: Acquire a shared lock (multiple holders allowed). Default is exclusive.
nowait: Return ``False`` immediately if the lock is unavailable instead of waiting.
timeout: Maximum wait time (e.g. ``"5s"``, ``"500ms"``). Raises ``DBAPIError``
if exceeded. Ignored when *nowait* is ``True``.
Yields:
``True`` if the lock was acquired, ``False`` if *nowait* is ``True`` and the lock
is already held.
Raises:
LockTimeoutError: If *timeout* is set and the lock cannot be acquired in time.
Example:
```python
from fastapi_toolsets.db import advisory_lock
async with advisory_lock(session, 42):
...
async with advisory_lock(session, 42, nowait=True) as acquired:
if not acquired:
raise HTTPException(409, "Resource is locked")
async with advisory_lock(session, 42, timeout="5s"):
...
async with advisory_lock(session, (1, user_id), shared=True):
...
```
"""
suffix = "_shared" if shared else ""
acquire_fn = f"{'pg_try_advisory_lock' if nowait else 'pg_advisory_lock'}{suffix}"
release_fn = f"pg_advisory_unlock{suffix}"
if isinstance(key, tuple):
k1, k2 = key
args = "CAST(:k1 AS integer), CAST(:k2 AS integer)"
params: dict[str, int] = {"k1": k1, "k2": k2}
else:
args = ":k"
params = {"k": key}
acquire_sql = text(f"SELECT {acquire_fn}({args})")
release_sql = text(f"SELECT {release_fn}({args})")
if timeout is not None and not nowait:
await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
try:
result = await session.execute(acquire_sql, params)
except sa_exc.DBAPIError as e:
if _is_lock_not_available(e):
raise LockTimeoutError(
f"Advisory lock {key!r} could not be acquired within {timeout}."
) from e
raise # pragma: no cover
acquired = result.scalar() if nowait else True
try:
yield acquired
finally:
if acquired:
await session.execute(release_sql, params)
async def create_database(
db_name: str,
*,
server_url: str,
) -> None:
"""Create a database.
Connects to *server_url* using ``AUTOCOMMIT`` isolation and issues a
``CREATE DATABASE`` statement for *db_name*.
Args:
db_name: Name of the database to create.
server_url: URL used for server-level DDL (must point to an existing
database on the same server).
Example:
```python
from fastapi_toolsets.db import create_database
SERVER_URL = "postgresql+asyncpg://postgres:postgres@localhost/postgres"
await create_database("myapp_test", server_url=SERVER_URL)
```
"""
engine = create_async_engine(server_url, isolation_level="AUTOCOMMIT")
try:
async with engine.connect() as conn:
await conn.execute(text(f"CREATE DATABASE {db_name}"))
finally:
await engine.dispose()
async def cleanup_tables(
session: AsyncSession,
base: type[DeclarativeBase],
) -> None:
"""Truncate all tables for fast between-test cleanup.
Executes a single ``TRUNCATE … RESTART IDENTITY CASCADE`` statement
across every table in *base*'s metadata, which is significantly faster
than dropping and re-creating tables between tests.
This is a no-op when the metadata contains no tables.
Args:
session: An active async database session.
base: SQLAlchemy DeclarativeBase class containing model metadata.
Example:
```python
@pytest.fixture
async def db_session(worker_db_url):
async with create_db_session(worker_db_url, Base) as session:
yield session
await cleanup_tables(session, Base)
```
"""
tables = base.metadata.sorted_tables
if not tables:
return
table_names = ", ".join(f'"{t.name}"' for t in tables)
await session.execute(text(f"TRUNCATE {table_names} RESTART IDENTITY CASCADE"))
await session.commit()
_M = TypeVar("_M", bound=DeclarativeBase)
async def wait_for_row_change(
session: AsyncSession,
model: type[_M],
pk_value: Any,
*,
columns: list[str] | None = None,
interval: float = 0.5,
timeout: float | None = None,
) -> _M:
"""Poll a database row until a change is detected.
Queries the row every ``interval`` seconds and returns the model instance
once a change is detected in any column (or only the specified ``columns``).
Args:
session: AsyncSession instance
model: SQLAlchemy model class
pk_value: Primary key value of the row to watch
columns: Optional list of column names to watch. If None, all columns
are watched.
interval: Polling interval in seconds (default: 0.5)
timeout: Maximum time to wait in seconds. None means wait forever.
Returns:
The refreshed model instance with updated values
Raises:
NotFoundError: If the row does not exist or is deleted during polling
TimeoutError: If timeout expires before a change is detected
Example:
```python
from fastapi_toolsets.db import wait_for_row_change
# Wait for any column to change
updated = await wait_for_row_change(session, User, user_id)
# Watch specific columns with a timeout
updated = await wait_for_row_change(
session, User, user_id,
columns=["status", "email"],
interval=1.0,
timeout=30.0,
)
```
"""
instance = await session.get(model, pk_value)
if instance is None:
raise NotFoundError(f"{model.__name__} with pk={pk_value!r} not found")
if columns is not None:
watch_cols = columns
else:
watch_cols = [attr.key for attr in model.__mapper__.column_attrs]
initial = {col: getattr(instance, col) for col in watch_cols}
elapsed = 0.0
while True:
await asyncio.sleep(interval)
elapsed += interval
if timeout is not None and elapsed >= timeout:
raise TimeoutError(
f"No change detected on {model.__name__} "
f"with pk={pk_value!r} within {timeout}s"
)
session.expunge(instance)
instance = await session.get(model, pk_value)
if instance is None:
raise NotFoundError(f"{model.__name__} with pk={pk_value!r} was deleted")
current = {col: getattr(instance, col) for col in watch_cols}
if current != initial:
return instance
def _m2m_prop(rel_attr: QueryableAttribute) -> RelationshipProperty: # type: ignore[type-arg]
"""Return the validated M2M RelationshipProperty for *rel_attr*.
Raises TypeError if *rel_attr* is not a Many-to-Many relationship.
"""
prop = rel_attr.property
if not isinstance(prop, RelationshipProperty) or prop.secondary is None:
raise TypeError(
f"m2m helpers require a Many-to-Many relationship attribute, "
f"got {rel_attr!r}. Use a relationship with a secondary table."
)
return prop
async def m2m_add(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
ignore_conflicts: bool = False,
) -> None:
"""Insert rows into a Many-to-Many association table without loading the ORM collection.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: One or more related instances to associate with ``instance``.
ignore_conflicts: When ``True``, silently skip rows that already exist
in the association table (``ON CONFLICT DO NOTHING``).
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
"""
prop = _m2m_prop(rel_attr)
if not related:
return
secondary = cast(Table, prop.secondary)
assert secondary is not None # guaranteed by _m2m_prop
sync_pairs = prop.secondary_synchronize_pairs
assert sync_pairs is not None # set whenever secondary is set
# synchronize_pairs: [(parent_col, assoc_col), ...]
# secondary_synchronize_pairs: [(related_col, assoc_col), ...]
rows: list[dict[str, Any]] = []
for rel_instance in related:
row: dict[str, Any] = {}
for parent_col, assoc_col in prop.synchronize_pairs:
row[assoc_col.name] = getattr(instance, cast(str, parent_col.key))
for related_col, assoc_col in sync_pairs:
row[assoc_col.name] = getattr(rel_instance, cast(str, related_col.key))
rows.append(row)
stmt = pg_insert(secondary).values(rows)
if ignore_conflicts:
stmt = stmt.on_conflict_do_nothing()
await session.execute(stmt)
async def m2m_remove(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
) -> None:
"""Remove rows from a Many-to-Many association table without loading the ORM collection.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: One or more related instances to disassociate from ``instance``.
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
"""
prop = _m2m_prop(rel_attr)
if not related:
return
secondary = cast(Table, prop.secondary)
assert secondary is not None # guaranteed by _m2m_prop
related_pairs = prop.secondary_synchronize_pairs
assert related_pairs is not None # set whenever secondary is set
parent_where = [
assoc_col == getattr(instance, cast(str, parent_col.key))
for parent_col, assoc_col in prop.synchronize_pairs
]
if len(related_pairs) == 1:
related_col, assoc_col = related_pairs[0]
related_values = [getattr(r, cast(str, related_col.key)) for r in related]
related_where = assoc_col.in_(related_values)
else:
assoc_cols = [ac for _, ac in related_pairs]
rel_cols = [rc for rc, _ in related_pairs]
related_values_t = [
tuple(getattr(r, cast(str, rc.key)) for rc in rel_cols) for r in related
]
related_where = tuple_(*assoc_cols).in_(related_values_t)
await session.execute(delete(secondary).where(*parent_where, related_where))
async def m2m_set(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
) -> None:
"""Replace the entire Many-to-Many association set atomically.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: The new complete set of related instances.
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
"""
prop = _m2m_prop(rel_attr)
secondary = cast(Table, prop.secondary)
assert secondary is not None # guaranteed by _m2m_prop
parent_where = [
assoc_col == getattr(instance, cast(str, parent_col.key))
for parent_col, assoc_col in prop.synchronize_pairs
]
await session.execute(delete(secondary).where(*parent_where))
if related:
await m2m_add(session, instance, rel_attr, *related)
+18
View File
@@ -0,0 +1,18 @@
"""Database package: the ``Database`` facade plus PostgreSQL power-tools."""
from .core import Database, transaction
from .locks import LockMode, advisory_lock, lock_tables
from .m2m import m2m_add, m2m_remove, m2m_set
from .watch import wait_for_row_change
__all__ = [
"Database",
"LockMode",
"advisory_lock",
"lock_tables",
"m2m_add",
"m2m_remove",
"m2m_set",
"transaction",
"wait_for_row_change",
]
+315
View File
@@ -0,0 +1,315 @@
"""The ``Database`` facade: session lifecycle, dependency, middleware, transactions."""
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Any
from sqlalchemy import exc as sa_exc
from sqlalchemy.ext.asyncio import (
AsyncEngine,
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from sqlalchemy.orm import DeclarativeBase
from starlette.requests import Request
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from ..exceptions import PoolExhaustedError
from .locks import LockMode, lock_tables
@asynccontextmanager
async def transaction(
session: AsyncSession,
) -> AsyncGenerator[AsyncSession, None]:
"""Run a block inside a savepoint-aware transaction.
If *session* is already in a transaction, a nested transaction (savepoint)
is opened so the block can roll back independently. Otherwise a top-level
transaction is started. Commits on clean exit, rolls back on exception.
Args:
session: AsyncSession instance.
Yields:
The session within the transaction context.
Example:
```python
from fastapi_toolsets.db import transaction
async with transaction(session):
session.add(model)
```
"""
if session.in_transaction():
async with session.begin_nested():
yield session
else:
async with session.begin():
yield session
class _CommitOnResponseMiddleware:
"""Commit the request's DB session before the response is sent."""
def __init__(self, app: ASGIApp, *, state_attr: str) -> None:
self.app = app
self.state_attr = state_attr
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
async def send_wrapper(message: Message) -> None:
if message["type"] == "http.response.start":
# ``scope["state"]`` is the same dict ``request.state`` writes
# to, so this is the session stashed by the dependency.
state = scope.get("state")
session = state.get(self.state_attr) if state else None
if session is not None and session.in_transaction():
await session.commit()
await send(message)
await self.app(scope, receive, send_wrapper)
class Database:
"""One object that owns the engine, sessions, dependency, and middleware.
Provide exactly one of *url* (the facade builds and disposes the engine) or
*engine* (an engine you own, e.g. for Alembic or ``event.listen``, left
untouched).
Args:
url: Database connection URL (e.g. ``"postgresql+asyncpg://..."``).
engine: An existing :class:`AsyncEngine` to reuse instead of *url*.
session_class: Session class for the sessionmaker (e.g. ``EventSession``).
expire_on_commit: Expire attributes after commit. Defaults to ``False``.
autoflush: Autoflush the session before queries. Defaults to ``True``.
**engine_options: Extra keyword arguments forwarded to
:func:`create_async_engine` (URL mode only, e.g. ``pool_size``,
``echo``, ``connect_args``).
Raises:
TypeError: If neither or both of *url* and *engine* are given, or if
*engine_options* are passed together with *engine*.
Example:
```python
from fastapi import Depends, FastAPI
from fastapi_toolsets.db import Database
db = Database("postgresql+asyncpg://postgres:postgres@localhost/app")
app = FastAPI()
db.install(app)
@app.get("/users/{user_id}")
async def get_user(user_id: int, session=Depends(db)):
return await UserCrud.get(session, [User.id == user_id])
```
"""
def __init__(
self,
url: str | None = None,
*,
engine: AsyncEngine | None = None,
session_class: type[AsyncSession] = AsyncSession,
expire_on_commit: bool = False,
autoflush: bool = True,
**engine_options: Any,
) -> None:
if (url is None) == (engine is None):
raise TypeError(
"Database requires exactly one of 'url' or 'engine' "
"(got both or neither)."
)
if engine is not None and engine_options:
raise TypeError(
"engine_options are only valid in URL mode; configure the "
"engine you pass via 'engine=' yourself."
)
if engine is not None:
self._owns_engine = False
self.engine: AsyncEngine = engine
else:
assert url is not None # guaranteed by the XOR check above
self._owns_engine = True
self.engine = create_async_engine(url, **engine_options)
self._sessionmaker: async_sessionmaker[AsyncSession] = async_sessionmaker(
self.engine,
class_=session_class,
expire_on_commit=expire_on_commit,
autoflush=autoflush,
)
# Private, per-instance state attribute; cannot collide with another
# Database or be mismatched against the middleware.
self._state_attr = f"_ft_db_session_{id(self):x}"
self._middleware_installed = False
self._disposed = False
async def _dispose(self) -> None:
"""Dispose the engine once, only if we own it (idempotent)."""
if self._owns_engine and not self._disposed:
self._disposed = True
await self.engine.dispose()
@asynccontextmanager
async def lifespan(self, app: Any) -> AsyncGenerator[None, None]:
"""Dispose the engine on shutdown; use as ``FastAPI(lifespan=db.lifespan)``.
Args:
app: The ASGI application (unused; required by the lifespan protocol).
Yields:
Control to the application for its lifetime.
Example:
```python
app = FastAPI(lifespan=db.lifespan)
```
"""
try:
yield
finally:
await self._dispose()
def install(self, app: Any) -> None:
"""Wire the commit middleware and engine disposal onto *app*.
Args:
app: The FastAPI/Starlette application to wire.
Example:
```python
@asynccontextmanager
async def lifespan(app):
... # your startup
yield
... # your shutdown
app = FastAPI(lifespan=lifespan)
db.install(app)
```
"""
app.add_middleware(_CommitOnResponseMiddleware, state_attr=self._state_attr)
self._middleware_installed = True
inner_lifespan = app.router.lifespan_context
@asynccontextmanager
async def _composed(app_: Any) -> AsyncGenerator[None, None]:
async with self.lifespan(app_):
async with inner_lifespan(app_):
yield
app.router.lifespan_context = _composed
@asynccontextmanager
async def _open(self) -> AsyncGenerator[AsyncSession, None]:
"""Open a session and eagerly acquire a connection (fail-fast on pool)."""
async with self._sessionmaker() as session:
try:
await session.connection()
except sa_exc.TimeoutError as e:
raise PoolExhaustedError() from e
yield session
async def __call__(self, request: Request) -> AsyncGenerator[AsyncSession, None]:
"""FastAPI dependency: yield a session and commit once at the right time.
Args:
request: The incoming request (injected by FastAPI).
Yields:
An AsyncSession for the duration of the request.
Example:
```python
@app.get("/users/{user_id}")
async def get_user(user_id: int, session=Depends(db)):
return await UserCrud.get(session, [User.id == user_id])
```
"""
async with self._open() as session:
setattr(request.state, self._state_attr, session)
yield session
if not self._middleware_installed and session.in_transaction():
await session.commit()
@asynccontextmanager
async def session(self) -> AsyncGenerator[AsyncSession, None]:
"""Open a session outside request handlers (background tasks, CLI, tests).
Commits on clean exit, rolls back on exception.
Yields:
An AsyncSession ready for database operations.
Example:
```python
async with db.session() as session:
user = await UserCrud.get(session, [User.id == 1])
```
"""
async with self._open() as session:
yield session
if session.in_transaction():
await session.commit()
@asynccontextmanager
async def begin(self) -> AsyncGenerator[AsyncSession, None]:
"""Open a session already inside a transaction (sugar for the common case).
Equivalent to ``session()`` + :func:`transaction`. Commits on clean exit,
rolls back on exception.
Yields:
An AsyncSession open within a transaction.
Example:
```python
async with db.begin() as session:
session.add(User(name="ada"))
```
"""
async with self.session() as session, transaction(session):
yield session
def lock_tables(
self,
tables: list[type[DeclarativeBase]],
*,
mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
timeout: str = "5s",
) -> AbstractAsyncContextManager[AsyncSession]:
"""Lock PostgreSQL tables for the duration of a dedicated transaction.
Opens its own session from the facade's sessionmaker, changes are
committed when the context exits.
Args:
tables: List of SQLAlchemy model classes to lock.
mode: Lock mode (default: ``SHARE UPDATE EXCLUSIVE``).
timeout: Lock timeout (default: ``"5s"``).
Yields:
The dedicated session, open within the locked transaction.
Raises:
LockTimeoutError: If the lock cannot be acquired within *timeout*.
PoolExhaustedError: If the connection pool is exhausted.
Example:
```python
async with db.lock_tables([User, Account]) as session:
user = await UserCrud.get(session, [User.id == 1])
user.balance += 100
```
"""
return lock_tables(self._sessionmaker, tables, mode=mode, timeout=timeout)
+185
View File
@@ -0,0 +1,185 @@
"""PostgreSQL locking helpers: table locks and advisory locks."""
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from enum import Enum
from typing import TypeVar
import asyncpg
from sqlalchemy import exc as sa_exc
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
from sqlalchemy.orm import DeclarativeBase
from ..exceptions import LockTimeoutError, PoolExhaustedError
_SessionT = TypeVar("_SessionT", bound=AsyncSession)
def _is_lock_not_available(e: sa_exc.DBAPIError) -> bool:
return e.orig is not None and isinstance(
e.orig.__cause__, asyncpg.exceptions.LockNotAvailableError
)
class LockMode(str, Enum):
"""PostgreSQL table lock modes.
See: https://www.postgresql.org/docs/current/explicit-locking.html
"""
ACCESS_SHARE = "ACCESS SHARE"
ROW_SHARE = "ROW SHARE"
ROW_EXCLUSIVE = "ROW EXCLUSIVE"
SHARE_UPDATE_EXCLUSIVE = "SHARE UPDATE EXCLUSIVE"
SHARE = "SHARE"
SHARE_ROW_EXCLUSIVE = "SHARE ROW EXCLUSIVE"
EXCLUSIVE = "EXCLUSIVE"
ACCESS_EXCLUSIVE = "ACCESS EXCLUSIVE"
def lock_tables(
session_maker: async_sessionmaker[_SessionT],
tables: list[type[DeclarativeBase]],
*,
mode: LockMode = LockMode.SHARE_UPDATE_EXCLUSIVE,
timeout: str = "5s",
) -> AbstractAsyncContextManager[_SessionT]:
"""Lock PostgreSQL tables for the duration of a transaction.
Prefer the method on a :class:`Database` instance; use this
directly only when you manage your own session factory.
Args:
session_maker: Async session factory used to create the dedicated
session.
tables: List of SQLAlchemy model classes to lock.
mode: Lock mode (default: SHARE UPDATE EXCLUSIVE).
timeout: Lock timeout (default: "5s").
Yields:
The dedicated session, open within the locked transaction.
Raises:
LockTimeoutError: If the lock cannot be acquired within *timeout*.
PoolExhaustedError: If the connection pool is exhausted.
Example:
```python
from fastapi_toolsets.db import lock_tables
async with lock_tables(session_maker, [User, Account]) as session:
user = await UserCrud.get(session, [User.id == 1])
user.balance += 100
```
"""
table_names = ",".join(table.__tablename__ for table in tables)
@asynccontextmanager
async def _lock() -> AsyncGenerator[_SessionT, None]:
async with session_maker() as session:
try:
await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
await session.execute(text(f"LOCK {table_names} IN {mode.value} MODE"))
yield session
await session.commit()
except sa_exc.TimeoutError as e:
await session.rollback()
raise PoolExhaustedError(
f"Connection pool exhausted while locking '{table_names}'. "
) from e
except sa_exc.DBAPIError as e:
await session.rollback()
if _is_lock_not_available(e):
raise LockTimeoutError(
f"Lock on '{table_names}' could not be acquired within {timeout}."
) from e
raise # pragma: no cover
except BaseException:
await session.rollback()
raise
return _lock()
@asynccontextmanager
async def advisory_lock(
session: AsyncSession,
key: int | tuple[int, int],
*,
shared: bool = False,
nowait: bool = False,
timeout: str | None = None,
) -> AsyncGenerator[bool, None]:
"""Acquire a PostgreSQL session-level advisory lock.
Args:
session: AsyncSession instance.
key: Lock key, either a single ``int`` (bigint) or a ``(int, int)`` pair for namespacing.
shared: Acquire a shared lock (multiple holders allowed). Default is exclusive.
nowait: Return ``False`` immediately if the lock is unavailable instead of waiting.
timeout: Maximum wait time (e.g. ``"5s"``, ``"500ms"``). Raises ``DBAPIError``
if exceeded. Ignored when *nowait* is ``True``.
Yields:
``True`` if the lock was acquired, ``False`` if *nowait* is ``True`` and the lock
is already held.
Raises:
LockTimeoutError: If *timeout* is set and the lock cannot be acquired in time.
Example:
```python
from fastapi_toolsets.db import advisory_lock
async with advisory_lock(session, 42):
...
async with advisory_lock(session, 42, nowait=True) as acquired:
if not acquired:
raise HTTPException(409, "Resource is locked")
async with advisory_lock(session, 42, timeout="5s"):
...
async with advisory_lock(session, (1, user_id), shared=True):
...
```
"""
suffix = "_shared" if shared else ""
acquire_fn = f"{'pg_try_advisory_lock' if nowait else 'pg_advisory_lock'}{suffix}"
release_fn = f"pg_advisory_unlock{suffix}"
if isinstance(key, tuple):
k1, k2 = key
args = "CAST(:k1 AS integer), CAST(:k2 AS integer)"
params: dict[str, int] = {"k1": k1, "k2": k2}
else:
args = ":k"
params = {"k": key}
acquire_sql = text(f"SELECT {acquire_fn}({args})")
release_sql = text(f"SELECT {release_fn}({args})")
# Lock management runs raw SQL on the caller's session. Guard it with
# ``no_autoflush`` so acquiring or releasing the lock never flushes the
# caller's pending ORM changes; SQLAlchemy 2.1 autoflushes on raw
# ``text()`` too, where 2.0 did not.
try:
with session.no_autoflush:
if timeout is not None and not nowait:
await session.execute(text(f"SET LOCAL lock_timeout='{timeout}'"))
result = await session.execute(acquire_sql, params)
except sa_exc.DBAPIError as e:
if _is_lock_not_available(e):
raise LockTimeoutError(
f"Advisory lock {key!r} could not be acquired within {timeout}."
) from e
raise # pragma: no cover
acquired = result.scalar() if nowait else True
try:
yield acquired
finally:
if acquired:
with session.no_autoflush:
await session.execute(release_sql, params)
+170
View File
@@ -0,0 +1,170 @@
"""Many-to-Many association-table helpers (direct, without loading collections)."""
from typing import Any, TypeVar, cast
from sqlalchemy import ColumnElement, Table, delete, tuple_
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import DeclarativeBase, QueryableAttribute
from sqlalchemy.orm.relationships import RelationshipProperty
_M = TypeVar("_M", bound=DeclarativeBase)
def _m2m_prop(rel_attr: QueryableAttribute) -> tuple[RelationshipProperty, Table]: # type: ignore[type-arg]
"""Return the validated M2M RelationshipProperty and its secondary table.
Raises TypeError if *rel_attr* is not a Many-to-Many relationship.
"""
prop = rel_attr.property
if not isinstance(prop, RelationshipProperty) or prop.secondary is None:
raise TypeError(
f"m2m helpers require a Many-to-Many relationship attribute, "
f"got {rel_attr!r}. Use a relationship with a secondary table."
)
return prop, cast(Table, prop.secondary)
def _parent_where(
prop: RelationshipProperty, # type: ignore[type-arg]
instance: DeclarativeBase,
) -> list[ColumnElement[bool]]:
"""Build the WHERE clauses matching the owner side of *instance*."""
return [
assoc_col == getattr(instance, cast(str, parent_col.key))
for parent_col, assoc_col in prop.synchronize_pairs
]
async def m2m_add(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
ignore_conflicts: bool = False,
) -> None:
"""Insert rows into a Many-to-Many association table without loading the ORM collection.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: One or more related instances to associate with ``instance``.
ignore_conflicts: When ``True``, silently skip rows that already exist
in the association table (``ON CONFLICT DO NOTHING``).
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
Example:
```python
from fastapi_toolsets.db import m2m_add, transaction
async with transaction(session):
await m2m_add(session, post, Post.tags, tag1, tag2)
```
"""
prop, secondary = _m2m_prop(rel_attr)
if not related:
return
sync_pairs = prop.secondary_synchronize_pairs
assert sync_pairs is not None # set whenever secondary is set
# synchronize_pairs: [(parent_col, assoc_col), ...]
# secondary_synchronize_pairs: [(related_col, assoc_col), ...]
rows: list[dict[str, Any]] = []
for rel_instance in related:
row: dict[str, Any] = {}
for parent_col, assoc_col in prop.synchronize_pairs:
row[assoc_col.name] = getattr(instance, cast(str, parent_col.key))
for related_col, assoc_col in sync_pairs:
row[assoc_col.name] = getattr(rel_instance, cast(str, related_col.key))
rows.append(row)
stmt = pg_insert(secondary).values(rows)
if ignore_conflicts:
stmt = stmt.on_conflict_do_nothing()
await session.execute(stmt)
async def m2m_remove(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
) -> None:
"""Remove rows from a Many-to-Many association table without loading the ORM collection.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: One or more related instances to disassociate from ``instance``.
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
Example:
```python
from fastapi_toolsets.db import m2m_remove, transaction
async with transaction(session):
await m2m_remove(session, post, Post.tags, tag1)
```
"""
prop, secondary = _m2m_prop(rel_attr)
if not related:
return
related_pairs = prop.secondary_synchronize_pairs
assert related_pairs is not None # set whenever secondary is set
parent_where = _parent_where(prop, instance)
if len(related_pairs) == 1:
related_col, assoc_col = related_pairs[0]
related_values = [getattr(r, cast(str, related_col.key)) for r in related]
related_where = assoc_col.in_(related_values)
else:
assoc_cols = [ac for _, ac in related_pairs]
rel_cols = [rc for rc, _ in related_pairs]
related_values_t = [
tuple(getattr(r, cast(str, rc.key)) for rc in rel_cols) for r in related
]
related_where = tuple_(*assoc_cols).in_(related_values_t)
await session.execute(delete(secondary).where(*parent_where, related_where))
async def m2m_set(
session: AsyncSession,
instance: DeclarativeBase,
rel_attr: QueryableAttribute,
*related: DeclarativeBase,
) -> None:
"""Replace the entire Many-to-Many association set atomically.
Args:
session: DB async session.
instance: The "owner" side model instance (e.g. the ``A`` in ``A.b_list``).
rel_attr: The M2M relationship attribute on the model class (e.g. ``A.b_list``).
*related: The new complete set of related instances.
Raises:
TypeError: If ``rel_attr`` is not a Many-to-Many relationship.
Example:
```python
from fastapi_toolsets.db import m2m_set, transaction
async with transaction(session):
await m2m_set(session, post, Post.tags, tag1, tag2) # replaces all
```
"""
prop, secondary = _m2m_prop(rel_attr)
await session.execute(delete(secondary).where(*_parent_where(prop, instance)))
if related:
await m2m_add(session, instance, rel_attr, *related)
+69
View File
@@ -0,0 +1,69 @@
"""Database admin and test helpers: DDL and truncation."""
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
from sqlalchemy.orm import DeclarativeBase
async def create_database(
db_name: str,
*,
server_url: str,
) -> None:
"""Create a database.
Connects to *server_url* using ``AUTOCOMMIT`` isolation and issues a
``CREATE DATABASE`` statement for *db_name*.
Args:
db_name: Name of the database to create.
server_url: URL used for server-level DDL (must point to an existing
database on the same server).
Example:
```python
from fastapi_toolsets.db.testing import create_database
SERVER_URL = "postgresql+asyncpg://postgres:postgres@localhost/postgres"
await create_database("myapp_test", server_url=SERVER_URL)
```
"""
engine = create_async_engine(server_url, isolation_level="AUTOCOMMIT")
try:
async with engine.connect() as conn:
await conn.execute(text(f"CREATE DATABASE {db_name}"))
finally:
await engine.dispose()
async def cleanup_tables(
session: AsyncSession,
base: type[DeclarativeBase],
) -> None:
"""Truncate all tables for fast between-test cleanup.
Executes a single ``TRUNCATE … RESTART IDENTITY CASCADE`` statement
across every table in *base*'s metadata.
This is a no-op when the metadata contains no tables.
Args:
session: An active async database session.
base: SQLAlchemy DeclarativeBase class containing model metadata.
Example:
```python
@pytest.fixture
async def db_session(worker_db_url):
async with create_db_session(worker_db_url, Base) as session:
yield session
await cleanup_tables(session, Base)
```
"""
tables = base.metadata.sorted_tables
if not tables:
return
table_names = ", ".join(f'"{t.name}"' for t in tables)
await session.execute(text(f"TRUNCATE {table_names} RESTART IDENTITY CASCADE"))
await session.commit()
+90
View File
@@ -0,0 +1,90 @@
"""Row-watching helpers: poll a database row until it changes."""
import asyncio
from typing import Any, TypeVar
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import DeclarativeBase
from ..exceptions import NotFoundError
_M = TypeVar("_M", bound=DeclarativeBase)
async def wait_for_row_change(
session: AsyncSession,
model: type[_M],
pk_value: Any,
*,
columns: list[str] | None = None,
interval: float = 0.5,
timeout: float | None = None,
) -> _M:
"""Poll a database row until a change is detected.
Queries the row every ``interval`` seconds and returns the model instance
once a change is detected in any column (or only the specified ``columns``).
Args:
session: AsyncSession instance.
model: SQLAlchemy model class.
pk_value: Primary key value of the row to watch.
columns: Optional list of column names to watch. If None, all columns
are watched.
interval: Polling interval in seconds (default: 0.5).
timeout: Maximum time to wait in seconds. None means wait forever.
Returns:
The refreshed model instance with updated values.
Raises:
NotFoundError: If the row does not exist or is deleted during polling.
TimeoutError: If timeout expires before a change is detected.
Example:
```python
from fastapi_toolsets.db import wait_for_row_change
# Wait for any column to change
updated = await wait_for_row_change(session, User, user_id)
# Watch specific columns with a timeout
updated = await wait_for_row_change(
session, User, user_id,
columns=["status", "email"],
interval=1.0,
timeout=30.0,
)
```
"""
instance = await session.get(model, pk_value)
if instance is None:
raise NotFoundError(f"{model.__name__} with pk={pk_value!r} not found")
if columns is not None:
watch_cols = columns
else:
watch_cols = [attr.key for attr in model.__mapper__.column_attrs]
initial = {col: getattr(instance, col) for col in watch_cols}
elapsed = 0.0
while True:
await asyncio.sleep(interval)
elapsed += interval
if timeout is not None and elapsed >= timeout:
raise TimeoutError(
f"No change detected on {model.__name__} "
f"with pk={pk_value!r} within {timeout}s"
)
session.expunge(instance)
instance = await session.get(model, pk_value)
if instance is None:
raise NotFoundError(f"{model.__name__} with pk={pk_value!r} was deleted")
current = {col: getattr(instance, col) for col in watch_cols}
if current != initial:
return instance
+2 -2
View File
@@ -9,7 +9,7 @@ from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm import DeclarativeBase
from ..db import get_transaction from ..db import transaction
from ..logger import get_logger from ..logger import get_logger
from ..types import ModelType from ..types import ModelType
from .enum import LoadStrategy from .enum import LoadStrategy
@@ -229,7 +229,7 @@ async def _load_ordered(
model_name = type(instances[0]).__name__ model_name = type(instances[0]).__name__
loaded: list[DeclarativeBase] = [] loaded: list[DeclarativeBase] = []
async with get_transaction(session): async with transaction(session):
for model_cls, group in _group_by_type(instances): for model_cls, group in _group_by_type(instances):
match strategy: match strategy:
case LoadStrategy.INSERT: case LoadStrategy.INSERT:
+2 -2
View File
@@ -9,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import DeclarativeBase, selectinload from sqlalchemy.orm import DeclarativeBase, selectinload
from sqlalchemy.orm.interfaces import ExecutableOption, ORMOption from sqlalchemy.orm.interfaces import ExecutableOption, ORMOption
from ..db import get_transaction from ..db import transaction
from ..fixtures import FixtureRegistry, LoadStrategy from ..fixtures import FixtureRegistry, LoadStrategy
@@ -106,7 +106,7 @@ def _create_fixture_function(
loaded: list[DeclarativeBase] = [] loaded: list[DeclarativeBase] = []
async with get_transaction(session): async with transaction(session):
for instance in instances: for instance in instances:
if strategy == LoadStrategy.INSERT: if strategy == LoadStrategy.INSERT:
session.add(instance) session.add(instance)
+1 -1
View File
@@ -15,7 +15,7 @@ from sqlalchemy.ext.asyncio import (
) )
from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm import DeclarativeBase
from ..db import cleanup_tables, create_database from ..db.testing import cleanup_tables, create_database
from ..models.watched import EventSession from ..models.watched import EventSession
+628 -186
View File
File diff suppressed because it is too large Load Diff
+12 -6
View File
@@ -91,13 +91,19 @@ async def seed(session: AsyncSession):
class TestAppSessionDep: class TestAppSessionDep:
@pytest.mark.anyio @pytest.mark.anyio
async def test_get_db_yields_async_session(self): async def test_get_db_yields_async_session(self):
"""get_db yields a real AsyncSession when called directly.""" """The Database dependency yields a real AsyncSession when called directly."""
from docs_src.examples.pagination_search.db import get_db from starlette.requests import Request
gen = get_db() from fastapi_toolsets.db import Database
session = await gen.__anext__()
assert isinstance(session, AsyncSession) db = Database(DATABASE_URL)
await gen.aclose() try:
gen = db(Request({"type": "http", "headers": []}))
session = await gen.__anext__()
assert isinstance(session, AsyncSession)
await gen.aclose()
finally:
await db.engine.dispose()
class TestOffsetPagination: class TestOffsetPagination:
+28 -42
View File
@@ -1506,8 +1506,8 @@ class TestListensFor:
assert all(e["event"] == "change" for e in _listener_events) assert all(e["event"] == "change" for e in _listener_events)
class TestEventSessionWithGetTransaction: class TestEventSessionWithTransaction:
"""Verify callbacks fire correctly when using get_transaction / lock_tables.""" """Verify callbacks fire correctly when using transaction / lock_tables."""
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def clear_events(self): def clear_events(self):
@@ -1517,10 +1517,10 @@ class TestEventSessionWithGetTransaction:
@pytest.mark.anyio @pytest.mark.anyio
async def test_callbacks_fire_after_outer_commit_not_savepoint(self, mixin_session): async def test_callbacks_fire_after_outer_commit_not_savepoint(self, mixin_session):
"""get_transaction creates a savepoint; callbacks fire only on outer commit.""" """transaction creates a savepoint; callbacks fire only on outer commit."""
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
async with get_transaction(mixin_session): async with transaction(mixin_session):
obj = WatchedModel(status="active", other="x") obj = WatchedModel(status="active", other="x")
mixin_session.add(obj) mixin_session.add(obj)
@@ -1535,14 +1535,14 @@ class TestEventSessionWithGetTransaction:
@pytest.mark.anyio @pytest.mark.anyio
async def test_nested_transactions_accumulate_events(self, mixin_session): async def test_nested_transactions_accumulate_events(self, mixin_session):
"""Multiple get_transaction blocks accumulate events for a single commit.""" """Multiple transaction blocks accumulate events for a single commit."""
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
async with get_transaction(mixin_session): async with transaction(mixin_session):
obj1 = WatchedModel(status="first", other="x") obj1 = WatchedModel(status="first", other="x")
mixin_session.add(obj1) mixin_session.add(obj1)
async with get_transaction(mixin_session): async with transaction(mixin_session):
obj2 = WatchedModel(status="second", other="y") obj2 = WatchedModel(status="second", other="y")
mixin_session.add(obj2) mixin_session.add(obj2)
@@ -1556,14 +1556,14 @@ class TestEventSessionWithGetTransaction:
@pytest.mark.anyio @pytest.mark.anyio
async def test_savepoint_rollback_suppresses_events(self, mixin_session): async def test_savepoint_rollback_suppresses_events(self, mixin_session):
"""Objects from a rolled-back savepoint don't fire callbacks.""" """Objects from a rolled-back savepoint don't fire callbacks."""
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
survivor = WatchedModel(status="kept", other="x") survivor = WatchedModel(status="kept", other="x")
mixin_session.add(survivor) mixin_session.add(survivor)
await mixin_session.flush() await mixin_session.flush()
try: try:
async with get_transaction(mixin_session): async with transaction(mixin_session):
doomed = WatchedModel(status="doomed", other="y") doomed = WatchedModel(status="doomed", other="y")
mixin_session.add(doomed) mixin_session.add(doomed)
await mixin_session.flush() await mixin_session.flush()
@@ -1590,9 +1590,9 @@ class TestEventSessionWithGetTransaction:
assert len(creates) == 1 assert len(creates) == 1
@pytest.mark.anyio @pytest.mark.anyio
async def test_update_inside_get_transaction(self, mixin_session): async def test_update_inside_transaction(self, mixin_session):
"""UPDATE events fire with correct changes after get_transaction commit.""" """UPDATE events fire with correct changes after transaction commit."""
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
obj = WatchedModel(status="initial", other="x") obj = WatchedModel(status="initial", other="x")
mixin_session.add(obj) mixin_session.add(obj)
@@ -1600,7 +1600,7 @@ class TestEventSessionWithGetTransaction:
_test_events.clear() _test_events.clear()
async with get_transaction(mixin_session): async with transaction(mixin_session):
obj.status = "updated" obj.status = "updated"
await mixin_session.commit() await mixin_session.commit()
@@ -1696,7 +1696,7 @@ class TestEventSessionWithNullableFields:
class TestEventSessionWithFastAPIDependency: class TestEventSessionWithFastAPIDependency:
"""Verify EventSession works when session comes from create_db_dependency.""" """Verify EventSession works when session comes from the Database dependency."""
@pytest.fixture(autouse=True) @pytest.fixture(autouse=True)
def clear_events(self): def clear_events(self):
@@ -1706,31 +1706,24 @@ class TestEventSessionWithFastAPIDependency:
@pytest.mark.anyio @pytest.mark.anyio
async def test_create_event_fires_via_dependency(self): async def test_create_event_fires_via_dependency(self):
"""CREATE callback fires when session is provided by create_db_dependency.""" """CREATE callback fires when session is provided by the Database dependency."""
from fastapi import Depends, FastAPI from fastapi import Depends, FastAPI
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy.ext.asyncio import ( from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from fastapi_toolsets.db import create_db_dependency from fastapi_toolsets.db import Database
from fastapi_toolsets.models import EventSession from fastapi_toolsets.models import EventSession
engine = create_async_engine(DATABASE_URL, echo=False) engine = create_async_engine(DATABASE_URL, echo=False)
session_factory = async_sessionmaker(
engine, expire_on_commit=False, class_=EventSession
)
async with engine.begin() as conn: async with engine.begin() as conn:
await conn.run_sync(MixinBase.metadata.create_all) await conn.run_sync(MixinBase.metadata.create_all)
get_db = create_db_dependency(session_factory) db = Database(engine=engine, session_class=EventSession)
app = FastAPI() app = FastAPI()
@app.post("/watched") @app.post("/watched")
async def create_watched(session: AsyncSession = Depends(get_db)): async def create_watched(session: AsyncSession = Depends(db)):
obj = WatchedModel(status="from-api", other="x") obj = WatchedModel(status="from-api", other="x")
session.add(obj) session.add(obj)
return {"id": str(obj.id)} return {"id": str(obj.id)}
@@ -1753,40 +1746,33 @@ class TestEventSessionWithFastAPIDependency:
@pytest.mark.anyio @pytest.mark.anyio
async def test_update_event_fires_via_dependency(self): async def test_update_event_fires_via_dependency(self):
"""UPDATE callback fires when session is provided by create_db_dependency.""" """UPDATE callback fires when session is provided by the Database dependency."""
from fastapi import Depends, FastAPI from fastapi import Depends, FastAPI
from httpx import ASGITransport, AsyncClient from httpx import ASGITransport, AsyncClient
from sqlalchemy.ext.asyncio import ( from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from fastapi_toolsets.db import create_db_dependency from fastapi_toolsets.db import Database
from fastapi_toolsets.models import EventSession from fastapi_toolsets.models import EventSession
engine = create_async_engine(DATABASE_URL, echo=False) engine = create_async_engine(DATABASE_URL, echo=False)
session_factory = async_sessionmaker(
engine, expire_on_commit=False, class_=EventSession
)
async with engine.begin() as conn: async with engine.begin() as conn:
await conn.run_sync(MixinBase.metadata.create_all) await conn.run_sync(MixinBase.metadata.create_all)
get_db = create_db_dependency(session_factory) db = Database(engine=engine, session_class=EventSession)
app = FastAPI() app = FastAPI()
# Pre-seed an object. # Pre-seed an object.
async with session_factory() as seed_session: async with db.session() as seed_session:
obj = WatchedModel(status="initial", other="x") obj = WatchedModel(status="initial", other="x")
seed_session.add(obj) seed_session.add(obj)
await seed_session.commit() await seed_session.flush()
obj_id = obj.id obj_id = obj.id
_test_events.clear() _test_events.clear()
@app.put("/watched/{item_id}") @app.put("/watched/{item_id}")
async def update_watched(item_id: str, session: AsyncSession = Depends(get_db)): async def update_watched(item_id: str, session: AsyncSession = Depends(db)):
from sqlalchemy import select from sqlalchemy import select
stmt = select(WatchedModel).where(WatchedModel.id == item_id) stmt = select(WatchedModel).where(WatchedModel.id == item_id)
+7 -7
View File
@@ -11,7 +11,7 @@ from sqlalchemy.engine import make_url
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
from fastapi_toolsets.db import get_transaction from fastapi_toolsets.db import transaction
from fastapi_toolsets.fixtures import Context, FixtureRegistry, LoadStrategy from fastapi_toolsets.fixtures import Context, FixtureRegistry, LoadStrategy
from fastapi_toolsets.pytest import ( from fastapi_toolsets.pytest import (
create_async_client, create_async_client,
@@ -387,14 +387,14 @@ class TestCreateDbSession:
assert session.autoflush is False assert session.autoflush is False
@pytest.mark.anyio @pytest.mark.anyio
async def test_get_transaction_commits_visible_to_separate_session(self): async def test_transaction_commits_visible_to_separate_session(self):
"""Data written via get_transaction() is committed and visible to other sessions.""" """Data written via transaction() is committed and visible to other sessions."""
role_id = uuid.uuid4() role_id = uuid.uuid4()
async with create_db_session(DATABASE_URL, Base, drop_tables=False) as session: async with create_db_session(DATABASE_URL, Base, drop_tables=False) as session:
# Simulate what _create_fixture_function does: insert via get_transaction # Simulate what _create_fixture_function does: insert via transaction()
# with no explicit commit afterward. # with no explicit commit afterward.
async with get_transaction(session): async with transaction(session):
role = Role(id=role_id, name="visible_to_other_session") role = Role(id=role_id, name="visible_to_other_session")
session.add(role) session.add(role)
@@ -409,9 +409,9 @@ class TestCreateDbSession:
result = await other.execute(select(Role).where(Role.id == role_id)) result = await other.execute(select(Role).where(Role.id == role_id))
fetched = result.scalar_one_or_none() fetched = result.scalar_one_or_none()
assert fetched is not None, ( assert fetched is not None, (
"Fixture data inserted via get_transaction() must be committed " "Fixture data inserted via transaction() must be committed "
"and visible to a separate session. If create_db_session uses " "and visible to a separate session. If create_db_session uses "
"create_db_context, auto-begin forces get_transaction() into " "db.session(), auto-begin forces transaction() into "
"savepoints instead of real commits." "savepoints instead of real commits."
) )
assert fetched.name == "visible_to_other_session" assert fetched.name == "visible_to_other_session"