mirror of
https://github.com/d3vyce/fastapi-toolsets.git
synced 2026-08-04 23:54:09 +00:00
feat: add NOWAIT/SKIP LOCKED row locking support and extend with_for_update to get_multi and update (#281)
This commit is contained in:
@@ -141,6 +141,37 @@ Use `first` when you only care about any one match and don't need uniqueness:
|
|||||||
user = await UserCrud.first(session=session, filters=[User.is_active == True])
|
user = await UserCrud.first(session=session, filters=[User.is_active == True])
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Row locking
|
||||||
|
|
||||||
|
`get`, `get_or_none`, `first`, `get_multi`, and `update` all accept a `with_for_update` parameter that appends a `FOR UPDATE` clause to the underlying `SELECT`, preventing concurrent transactions from modifying the matched rows until the current transaction commits.
|
||||||
|
|
||||||
|
| Value | SQL clause |
|
||||||
|
|---|---|
|
||||||
|
| `False` (default) | no locking |
|
||||||
|
| `True` | `FOR UPDATE` |
|
||||||
|
| `"nowait"` | `FOR UPDATE NOWAIT` |
|
||||||
|
| `"skip_locked"` | `FOR UPDATE SKIP LOCKED` |
|
||||||
|
|
||||||
|
```python
|
||||||
|
# Lock before reading — typical read-modify-write pattern
|
||||||
|
user = await UserCrud.get(session, [User.id == user_id], with_for_update=True)
|
||||||
|
|
||||||
|
# Raise immediately if another transaction holds the lock
|
||||||
|
user = await UserCrud.get(session, [User.id == user_id], with_for_update="nowait")
|
||||||
|
|
||||||
|
# Skip rows already locked by another transaction (e.g. job queues)
|
||||||
|
rows = await JobCrud.get_multi(session, filters=[Job.status == "pending"], with_for_update="skip_locked")
|
||||||
|
|
||||||
|
# Lock atomically as part of update (prevents race between SELECT and UPDATE)
|
||||||
|
user = await UserCrud.update(session, UserUpdate(credits=10), [User.id == user_id], with_for_update=True)
|
||||||
|
```
|
||||||
|
|
||||||
|
!!! 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.
|
||||||
|
|
||||||
|
!!! note
|
||||||
|
`NOWAIT` raises `sqlalchemy.exc.OperationalError` immediately if the row is locked rather than waiting.
|
||||||
|
|
||||||
## Pagination
|
## Pagination
|
||||||
|
|
||||||
!!! info "Added in `v1.1` (only offset_pagination via `paginate` if `<v1.1`)"
|
!!! info "Added in `v1.1` (only offset_pagination via `paginate` if `<v1.1`)"
|
||||||
|
|||||||
@@ -10,7 +10,7 @@ from collections.abc import Awaitable, Callable, Sequence
|
|||||||
from datetime import date, datetime
|
from datetime import date, datetime
|
||||||
from decimal import Decimal
|
from decimal import Decimal
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from typing import Any, ClassVar, Generic, Literal, Self, cast, overload
|
from typing import Any, ClassVar, Generic, Literal, Self, TypeAlias, cast, overload
|
||||||
|
|
||||||
from fastapi import Query
|
from fastapi import Query
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
@@ -52,6 +52,19 @@ from .search import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
_ForUpdateMode: TypeAlias = bool | Literal["nowait", "skip_locked"]
|
||||||
|
|
||||||
|
|
||||||
|
def _apply_for_update(q: Any, mode: _ForUpdateMode) -> Any:
|
||||||
|
if not mode:
|
||||||
|
return q
|
||||||
|
if mode == "nowait":
|
||||||
|
return q.with_for_update(nowait=True)
|
||||||
|
if mode == "skip_locked":
|
||||||
|
return q.with_for_update(skip_locked=True)
|
||||||
|
return q.with_for_update()
|
||||||
|
|
||||||
|
|
||||||
class _CursorDirection(str, Enum):
|
class _CursorDirection(str, Enum):
|
||||||
NEXT = "next"
|
NEXT = "next"
|
||||||
PREV = "prev"
|
PREV = "prev"
|
||||||
@@ -733,7 +746,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[SchemaType],
|
schema: type[SchemaType],
|
||||||
) -> Response[SchemaType]: ...
|
) -> Response[SchemaType]: ...
|
||||||
@@ -747,7 +760,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: None = ...,
|
schema: None = ...,
|
||||||
) -> ModelType: ...
|
) -> ModelType: ...
|
||||||
@@ -760,7 +773,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[BaseModel] | None = None,
|
schema: type[BaseModel] | None = None,
|
||||||
) -> ModelType | Response[Any]:
|
) -> ModelType | Response[Any]:
|
||||||
@@ -805,7 +818,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[SchemaType],
|
schema: type[SchemaType],
|
||||||
) -> Response[SchemaType] | None: ...
|
) -> Response[SchemaType] | None: ...
|
||||||
@@ -819,7 +832,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: None = ...,
|
schema: None = ...,
|
||||||
) -> ModelType | None: ...
|
) -> ModelType | None: ...
|
||||||
@@ -832,7 +845,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[BaseModel] | None = None,
|
schema: type[BaseModel] | None = None,
|
||||||
) -> ModelType | Response[Any] | None:
|
) -> ModelType | Response[Any] | None:
|
||||||
@@ -864,8 +877,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
q = q.where(and_(*filters))
|
q = q.where(and_(*filters))
|
||||||
if resolved := cls._resolve_load_options(load_options):
|
if resolved := cls._resolve_load_options(load_options):
|
||||||
q = q.options(*resolved)
|
q = q.options(*resolved)
|
||||||
if with_for_update:
|
q = _apply_for_update(q, with_for_update)
|
||||||
q = q.with_for_update()
|
|
||||||
result = await session.execute(q)
|
result = await session.execute(q)
|
||||||
item = result.unique().scalar_one_or_none()
|
item = result.unique().scalar_one_or_none()
|
||||||
if item is None:
|
if item is None:
|
||||||
@@ -884,7 +896,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[SchemaType],
|
schema: type[SchemaType],
|
||||||
) -> Response[SchemaType] | None: ...
|
) -> Response[SchemaType] | None: ...
|
||||||
@@ -898,7 +910,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: None = ...,
|
schema: None = ...,
|
||||||
) -> ModelType | None: ...
|
) -> ModelType | None: ...
|
||||||
@@ -911,7 +923,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
with_for_update: bool = False,
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
schema: type[BaseModel] | None = None,
|
schema: type[BaseModel] | None = None,
|
||||||
) -> ModelType | Response[Any] | None:
|
) -> ModelType | Response[Any] | None:
|
||||||
@@ -937,8 +949,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
q = q.where(and_(*filters))
|
q = q.where(and_(*filters))
|
||||||
if resolved := cls._resolve_load_options(load_options):
|
if resolved := cls._resolve_load_options(load_options):
|
||||||
q = q.options(*resolved)
|
q = q.options(*resolved)
|
||||||
if with_for_update:
|
q = _apply_for_update(q, with_for_update)
|
||||||
q = q.with_for_update()
|
|
||||||
result = await session.execute(q)
|
result = await session.execute(q)
|
||||||
item = result.unique().scalars().first()
|
item = result.unique().scalars().first()
|
||||||
if item is None:
|
if item is None:
|
||||||
@@ -956,6 +967,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
filters: list[Any] | None = None,
|
filters: list[Any] | None = None,
|
||||||
joins: JoinType | None = None,
|
joins: JoinType | None = None,
|
||||||
outer_join: bool = False,
|
outer_join: bool = False,
|
||||||
|
with_for_update: _ForUpdateMode = False,
|
||||||
load_options: Sequence[ExecutableOption] | None = None,
|
load_options: Sequence[ExecutableOption] | None = None,
|
||||||
order_by: OrderByClause | None = None,
|
order_by: OrderByClause | None = None,
|
||||||
limit: int | None = None,
|
limit: int | None = None,
|
||||||
@@ -968,6 +980,9 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
filters: List of SQLAlchemy filter conditions
|
filters: List of SQLAlchemy filter conditions
|
||||||
joins: List of (model, condition) tuples for joining related tables
|
joins: List of (model, condition) tuples for joining related tables
|
||||||
outer_join: Use LEFT OUTER JOIN instead of INNER JOIN
|
outer_join: Use LEFT OUTER JOIN instead of INNER JOIN
|
||||||
|
with_for_update: Lock rows for update. ``True`` for plain ``FOR UPDATE``,
|
||||||
|
``"nowait"`` for ``FOR UPDATE NOWAIT``, ``"skip_locked"`` for
|
||||||
|
``FOR UPDATE SKIP LOCKED``.
|
||||||
load_options: SQLAlchemy loader options
|
load_options: SQLAlchemy loader options
|
||||||
order_by: Column or list of columns to order by
|
order_by: Column or list of columns to order by
|
||||||
limit: Max number of rows to return
|
limit: Max number of rows to return
|
||||||
@@ -982,6 +997,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
q = q.where(and_(*filters))
|
q = q.where(and_(*filters))
|
||||||
if resolved := cls._resolve_load_options(load_options):
|
if resolved := cls._resolve_load_options(load_options):
|
||||||
q = q.options(*resolved)
|
q = q.options(*resolved)
|
||||||
|
q = _apply_for_update(q, with_for_update)
|
||||||
if order_by is not None:
|
if order_by is not None:
|
||||||
q = q.order_by(order_by)
|
q = q.order_by(order_by)
|
||||||
if offset is not None:
|
if offset is not None:
|
||||||
@@ -1001,6 +1017,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
exclude_unset: bool = True,
|
exclude_unset: bool = True,
|
||||||
exclude_none: bool = False,
|
exclude_none: bool = False,
|
||||||
|
with_for_update: _ForUpdateMode = False,
|
||||||
schema: type[SchemaType],
|
schema: type[SchemaType],
|
||||||
) -> Response[SchemaType]: ...
|
) -> Response[SchemaType]: ...
|
||||||
|
|
||||||
@@ -1014,6 +1031,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
exclude_unset: bool = True,
|
exclude_unset: bool = True,
|
||||||
exclude_none: bool = False,
|
exclude_none: bool = False,
|
||||||
|
with_for_update: _ForUpdateMode = False,
|
||||||
schema: None = ...,
|
schema: None = ...,
|
||||||
) -> ModelType: ...
|
) -> ModelType: ...
|
||||||
|
|
||||||
@@ -1026,6 +1044,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
*,
|
*,
|
||||||
exclude_unset: bool = True,
|
exclude_unset: bool = True,
|
||||||
exclude_none: bool = False,
|
exclude_none: bool = False,
|
||||||
|
with_for_update: _ForUpdateMode = False,
|
||||||
schema: type[BaseModel] | None = None,
|
schema: type[BaseModel] | None = None,
|
||||||
) -> ModelType | Response[Any]:
|
) -> ModelType | Response[Any]:
|
||||||
"""Update a record in the database.
|
"""Update a record in the database.
|
||||||
@@ -1036,6 +1055,9 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
filters: List of SQLAlchemy filter conditions
|
filters: List of SQLAlchemy filter conditions
|
||||||
exclude_unset: Exclude fields not explicitly set in the schema
|
exclude_unset: Exclude fields not explicitly set in the schema
|
||||||
exclude_none: Exclude fields with None value
|
exclude_none: Exclude fields with None value
|
||||||
|
with_for_update: Lock the row before updating. ``True`` for plain
|
||||||
|
``FOR UPDATE``, ``"nowait"`` for ``FOR UPDATE NOWAIT``,
|
||||||
|
``"skip_locked"`` for ``FOR UPDATE SKIP LOCKED``.
|
||||||
schema: Pydantic schema to serialize the result into. When provided,
|
schema: Pydantic schema to serialize the result into. When provided,
|
||||||
the result is automatically wrapped in a ``Response[schema]``.
|
the result is automatically wrapped in a ``Response[schema]``.
|
||||||
|
|
||||||
@@ -1059,6 +1081,7 @@ class AsyncCrud(Generic[ModelType]):
|
|||||||
db_model = await cls.get(
|
db_model = await cls.get(
|
||||||
session=session,
|
session=session,
|
||||||
filters=filters,
|
filters=filters,
|
||||||
|
with_for_update=with_for_update,
|
||||||
load_options=m2m_load_options or None,
|
load_options=m2m_load_options or None,
|
||||||
)
|
)
|
||||||
values = obj.model_dump(
|
values = obj.model_dump(
|
||||||
|
|||||||
+203
-1
@@ -670,6 +670,28 @@ class TestCrudFirst:
|
|||||||
assert role is not None
|
assert role is not None
|
||||||
assert role.name == "admin"
|
assert role.name == "admin"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_first_with_for_update_nowait(self, db_session: AsyncSession):
|
||||||
|
"""First with with_for_update='nowait' emits FOR UPDATE NOWAIT."""
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="nowait_first"))
|
||||||
|
|
||||||
|
role = await RoleCrud.first(
|
||||||
|
db_session, [Role.name == "nowait_first"], with_for_update="nowait"
|
||||||
|
)
|
||||||
|
assert role is not None
|
||||||
|
assert role.name == "nowait_first"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_first_with_for_update_skip_locked(self, db_session: AsyncSession):
|
||||||
|
"""First with with_for_update='skip_locked' emits FOR UPDATE SKIP LOCKED."""
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="skip_first"))
|
||||||
|
|
||||||
|
role = await RoleCrud.first(
|
||||||
|
db_session, [Role.name == "skip_first"], with_for_update="skip_locked"
|
||||||
|
)
|
||||||
|
assert role is not None
|
||||||
|
assert role.name == "skip_first"
|
||||||
|
|
||||||
|
|
||||||
class TestCrudGetMulti:
|
class TestCrudGetMulti:
|
||||||
"""Tests for CRUD get_multi operations."""
|
"""Tests for CRUD get_multi operations."""
|
||||||
@@ -735,6 +757,45 @@ class TestCrudGetMulti:
|
|||||||
names = [r.name for r in roles]
|
names = [r.name for r in roles]
|
||||||
assert names == ["alpha", "bravo", "charlie"]
|
assert names == ["alpha", "bravo", "charlie"]
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_multi_with_for_update(self, db_session: AsyncSession):
|
||||||
|
"""get_multi() with with_for_update=True locks the rows."""
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="lock1"))
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="lock2"))
|
||||||
|
|
||||||
|
roles = await RoleCrud.get_multi(
|
||||||
|
db_session,
|
||||||
|
filters=[Role.name.in_(["lock1", "lock2"])],
|
||||||
|
with_for_update=True,
|
||||||
|
)
|
||||||
|
assert len(roles) == 2
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_multi_with_for_update_nowait(self, db_session: AsyncSession):
|
||||||
|
"""get_multi() with with_for_update='nowait' emits FOR UPDATE NOWAIT."""
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="nowait_multi"))
|
||||||
|
|
||||||
|
roles = await RoleCrud.get_multi(
|
||||||
|
db_session,
|
||||||
|
filters=[Role.name == "nowait_multi"],
|
||||||
|
with_for_update="nowait",
|
||||||
|
)
|
||||||
|
assert len(roles) == 1
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_multi_with_for_update_skip_locked(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
):
|
||||||
|
"""get_multi() with with_for_update='skip_locked' emits FOR UPDATE SKIP LOCKED."""
|
||||||
|
await RoleCrud.create(db_session, RoleCreate(name="skip_multi"))
|
||||||
|
|
||||||
|
roles = await RoleCrud.get_multi(
|
||||||
|
db_session,
|
||||||
|
filters=[Role.name == "skip_multi"],
|
||||||
|
with_for_update="skip_locked",
|
||||||
|
)
|
||||||
|
assert len(roles) == 1
|
||||||
|
|
||||||
|
|
||||||
class TestCrudUpdate:
|
class TestCrudUpdate:
|
||||||
"""Tests for CRUD update operations."""
|
"""Tests for CRUD update operations."""
|
||||||
@@ -781,6 +842,48 @@ class TestCrudUpdate:
|
|||||||
assert updated.email == "john@test.com"
|
assert updated.email == "john@test.com"
|
||||||
assert updated.is_active is True
|
assert updated.is_active is True
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_update_with_for_update(self, db_session: AsyncSession):
|
||||||
|
"""update() with with_for_update=True locks the row before writing."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="before"))
|
||||||
|
|
||||||
|
updated = await RoleCrud.update(
|
||||||
|
db_session,
|
||||||
|
RoleUpdate(name="after"),
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert updated.name == "after"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_update_with_for_update_nowait(self, db_session: AsyncSession):
|
||||||
|
"""update() with with_for_update='nowait' locks the row with NOWAIT."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="before_nowait"))
|
||||||
|
|
||||||
|
updated = await RoleCrud.update(
|
||||||
|
db_session,
|
||||||
|
RoleUpdate(name="after_nowait"),
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update="nowait",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert updated.name == "after_nowait"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_update_with_for_update_skip_locked(self, db_session: AsyncSession):
|
||||||
|
"""update() with with_for_update='skip_locked' locks the row with SKIP LOCKED."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="before_skip"))
|
||||||
|
|
||||||
|
updated = await RoleCrud.update(
|
||||||
|
db_session,
|
||||||
|
RoleUpdate(name="after_skip"),
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update="skip_locked",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert updated.name == "after_skip"
|
||||||
|
|
||||||
|
|
||||||
class TestCrudDelete:
|
class TestCrudDelete:
|
||||||
"""Tests for CRUD delete operations."""
|
"""Tests for CRUD delete operations."""
|
||||||
@@ -2610,7 +2713,7 @@ class TestCursorPaginateSearchJoins:
|
|||||||
|
|
||||||
|
|
||||||
class TestGetWithForUpdate:
|
class TestGetWithForUpdate:
|
||||||
"""Tests for get() with with_for_update=True."""
|
"""Tests for get/get_or_none with_for_update variants."""
|
||||||
|
|
||||||
@pytest.mark.anyio
|
@pytest.mark.anyio
|
||||||
async def test_get_with_for_update(self, db_session: AsyncSession):
|
async def test_get_with_for_update(self, db_session: AsyncSession):
|
||||||
@@ -2626,6 +2729,105 @@ class TestGetWithForUpdate:
|
|||||||
assert result.id == role.id
|
assert result.id == role.id
|
||||||
assert result.name == "locked"
|
assert result.name == "locked"
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_with_for_update_nowait(self, db_session: AsyncSession):
|
||||||
|
"""get() with with_for_update='nowait' emits FOR UPDATE NOWAIT."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="nowait"))
|
||||||
|
|
||||||
|
result = await RoleCrud.get(
|
||||||
|
db_session,
|
||||||
|
filters=[Role.id == role.id],
|
||||||
|
with_for_update="nowait",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.id == role.id
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_with_for_update_skip_locked(self, db_session: AsyncSession):
|
||||||
|
"""get() with with_for_update='skip_locked' emits FOR UPDATE SKIP LOCKED."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="skip"))
|
||||||
|
|
||||||
|
result = await RoleCrud.get(
|
||||||
|
db_session,
|
||||||
|
filters=[Role.id == role.id],
|
||||||
|
with_for_update="skip_locked",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.id == role.id
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_or_none_with_for_update(self, db_session: AsyncSession):
|
||||||
|
"""get_or_none() with with_for_update=True locks the row."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="locked2"))
|
||||||
|
|
||||||
|
result = await RoleCrud.get_or_none(
|
||||||
|
db_session,
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.id == role.id
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_or_none_with_for_update_nowait(self, db_session: AsyncSession):
|
||||||
|
"""get_or_none() with with_for_update='nowait' emits FOR UPDATE NOWAIT."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="nowait2"))
|
||||||
|
|
||||||
|
result = await RoleCrud.get_or_none(
|
||||||
|
db_session,
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update="nowait",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.id == role.id
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_get_or_none_with_for_update_skip_locked(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
):
|
||||||
|
"""get_or_none() with with_for_update='skip_locked' emits FOR UPDATE SKIP LOCKED."""
|
||||||
|
role = await RoleCrud.create(db_session, RoleCreate(name="skip2"))
|
||||||
|
|
||||||
|
result = await RoleCrud.get_or_none(
|
||||||
|
db_session,
|
||||||
|
[Role.id == role.id],
|
||||||
|
with_for_update="skip_locked",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert result.id == role.id
|
||||||
|
|
||||||
|
def test_for_update_sql_clauses(self):
|
||||||
|
"""Verify _apply_for_update emits the correct SQL FOR UPDATE clauses."""
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.dialects import postgresql
|
||||||
|
|
||||||
|
from fastapi_toolsets.crud.factory import _apply_for_update
|
||||||
|
|
||||||
|
base = select(Role)
|
||||||
|
|
||||||
|
plain = str(_apply_for_update(base, True).compile(dialect=postgresql.dialect()))
|
||||||
|
assert "FOR UPDATE" in plain
|
||||||
|
assert "NOWAIT" not in plain
|
||||||
|
assert "SKIP LOCKED" not in plain
|
||||||
|
|
||||||
|
nowait = str(
|
||||||
|
_apply_for_update(base, "nowait").compile(dialect=postgresql.dialect())
|
||||||
|
)
|
||||||
|
assert "FOR UPDATE NOWAIT" in nowait
|
||||||
|
|
||||||
|
skip = str(
|
||||||
|
_apply_for_update(base, "skip_locked").compile(dialect=postgresql.dialect())
|
||||||
|
)
|
||||||
|
assert "FOR UPDATE SKIP LOCKED" in skip
|
||||||
|
|
||||||
|
no_lock = str(
|
||||||
|
_apply_for_update(base, False).compile(dialect=postgresql.dialect())
|
||||||
|
)
|
||||||
|
assert "FOR UPDATE" not in no_lock
|
||||||
|
|
||||||
|
|
||||||
class TestCursorPaginateColumnTypes:
|
class TestCursorPaginateColumnTypes:
|
||||||
"""Tests for cursor_paginate() covering DateTime, Date and Numeric column types."""
|
"""Tests for cursor_paginate() covering DateTime, Date and Numeric column types."""
|
||||||
|
|||||||
Reference in New Issue
Block a user