mirror of
https://github.com/d3vyce/fastapi-toolsets.git
synced 2026-09-19 11:19:56 +00:00
Compare commits
4
Commits
v5.1.1
...
e09b911277
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e09b911277 | ||
|
|
f7ecb76e8d
|
||
|
|
ef269833b9
|
||
|
|
610b3e1ab4
|
@@ -373,6 +373,16 @@ Or via the dependency to narrow which fields are exposed as query parameters:
|
||||
params = UserCrud.offset_paginate_params(search_fields=[Post.title])
|
||||
```
|
||||
|
||||
`search_fields`, `facet_fields` and `order_fields` follow the same override rule
|
||||
everywhere they are accepted — `offset_paginate`, `cursor_paginate`,
|
||||
`paginate` and the matching `*_paginate_params` dependencies:
|
||||
|
||||
| Passed | Effect |
|
||||
| --- | --- |
|
||||
| omitted or `None` | Use the class-level declaration |
|
||||
| `[]` | Disable this feature for this call |
|
||||
| `[...]` | Use exactly these fields (the primary key is **not** prepended — that only happens for the class-level `searchable_fields`) |
|
||||
|
||||
This allows searching with both [`offset_paginate`](../reference/crud.md#fastapi_toolsets.crud.factory.AsyncCrud.offset_paginate) and [`cursor_paginate`](../reference/crud.md#fastapi_toolsets.crud.factory.AsyncCrud.cursor_paginate):
|
||||
|
||||
```python
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "fastapi-toolsets"
|
||||
version = "5.1.1"
|
||||
version = "5.1.2"
|
||||
description = "Production-ready utilities for FastAPI applications"
|
||||
readme = "README.md"
|
||||
license = "MIT"
|
||||
|
||||
@@ -24,4 +24,4 @@ Example usage:
|
||||
return Response(data={"user": user.username}, message="Success")
|
||||
"""
|
||||
|
||||
__version__ = "5.1.1"
|
||||
__version__ = "5.1.2"
|
||||
|
||||
@@ -270,32 +270,17 @@ class AsyncCrud(Generic[ModelType]):
|
||||
return cls.default_load_options
|
||||
|
||||
@classmethod
|
||||
def _capture_pk_values(
|
||||
cls: type[Self], instance: DeclarativeBase
|
||||
) -> dict[str, Any]:
|
||||
"""Capture PK values off instance — call before commit expires attributes."""
|
||||
return {
|
||||
cast(str, col.key): getattr(instance, cast(str, col.key))
|
||||
for col in cls.model.__mapper__.primary_key
|
||||
}
|
||||
|
||||
@classmethod
|
||||
async def _reload_with_options_by_pk(
|
||||
cls: type[Self], session: AsyncSession, pk_values: dict[str, Any]
|
||||
async def _reload_with_options(
|
||||
cls: type[Self], session: AsyncSession, instance: DeclarativeBase
|
||||
) -> ModelType:
|
||||
"""Re-query by previously captured PK values, with default_load_options applied."""
|
||||
# Only called when cls.default_load_options is set (see call sites).
|
||||
"""Re-query instance by PK with default_load_options applied."""
|
||||
mapper = cls.model.__mapper__
|
||||
pk_filters = [
|
||||
getattr(cls.model, key) == value for key, value in pk_values.items()
|
||||
getattr(cls.model, cast(str, col.key))
|
||||
== getattr(instance, cast(str, col.key))
|
||||
for col in mapper.primary_key
|
||||
]
|
||||
q = select(cls.model).where(and_(*pk_filters))
|
||||
q = q.execution_options(populate_existing=True)
|
||||
q = q.options(*cast(Sequence[ExecutableOption], cls.default_load_options))
|
||||
result = await session.execute(q)
|
||||
item = result.unique().scalar_one_or_none()
|
||||
if item is None: # pragma: no cover — row was just flushed in this transaction
|
||||
raise NotFoundError()
|
||||
return cast(ModelType, item)
|
||||
return await cls.get(session, filters=pk_filters)
|
||||
|
||||
@classmethod
|
||||
async def _resolve_m2m(
|
||||
@@ -401,13 +386,29 @@ class AsyncCrud(Generic[ModelType]):
|
||||
own_filters=own_filters,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _resolve_search_fields(
|
||||
cls: type[Self],
|
||||
search_fields: Sequence[SearchFieldType] | None,
|
||||
) -> Sequence[SearchFieldType] | None:
|
||||
"""Return search_fields if given, otherwise fall back to the class-level default."""
|
||||
return search_fields if search_fields is not None else cls.searchable_fields
|
||||
|
||||
@classmethod
|
||||
def _resolve_order_fields(
|
||||
cls: type[Self],
|
||||
order_fields: Sequence[OrderFieldType] | None,
|
||||
) -> Sequence[OrderFieldType] | None:
|
||||
"""Return order_fields if given, otherwise fall back to the class-level default."""
|
||||
return order_fields if order_fields is not None else cls.order_fields
|
||||
|
||||
@classmethod
|
||||
def _resolve_search_columns(
|
||||
cls: type[Self],
|
||||
search_fields: Sequence[SearchFieldType] | None,
|
||||
) -> list[str] | None:
|
||||
"""Return search column keys, or None if no searchable fields configured."""
|
||||
fields = search_fields if search_fields is not None else cls.searchable_fields
|
||||
fields = cls._resolve_search_fields(search_fields)
|
||||
if not fields:
|
||||
return None
|
||||
return search_field_keys(fields)
|
||||
@@ -418,7 +419,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
order_fields: Sequence[OrderFieldType] | None,
|
||||
) -> list[str] | None:
|
||||
"""Return sort column keys, or None if no order fields configured."""
|
||||
fields = order_fields if order_fields is not None else cls.order_fields
|
||||
fields = cls._resolve_order_fields(order_fields)
|
||||
if not fields:
|
||||
return None
|
||||
return sorted(facet_keys(fields))
|
||||
@@ -497,9 +498,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
order_field_map: dict[str, OrderFieldType] | None = None
|
||||
order_valid_keys: list[str] | None = None
|
||||
if order:
|
||||
resolved_order = (
|
||||
order_fields if order_fields is not None else cls.order_fields
|
||||
)
|
||||
resolved_order = cls._resolve_order_fields(order_fields)
|
||||
if resolved_order:
|
||||
keys = facet_keys(resolved_order)
|
||||
order_field_map = dict(zip(keys, resolved_order))
|
||||
@@ -525,8 +524,21 @@ class AsyncCrud(Generic[ModelType]):
|
||||
]
|
||||
)
|
||||
|
||||
fixed: dict[str, Any] = {
|
||||
**pagination_fixed,
|
||||
"search_fields": (cls._resolve_search_fields(search_fields) or [])
|
||||
if search
|
||||
else [],
|
||||
"facet_fields": (cls._resolve_facet_fields(facet_fields) or [])
|
||||
if filter
|
||||
else [],
|
||||
"order_fields": (cls._resolve_order_fields(order_fields) or [])
|
||||
if order
|
||||
else [],
|
||||
}
|
||||
|
||||
async def dependency(**kwargs: Any) -> dict[str, Any]:
|
||||
result: dict[str, Any] = dict(pagination_fixed)
|
||||
result: dict[str, Any] = dict(fixed)
|
||||
for name in pagination_param_names:
|
||||
result[name] = kwargs[name]
|
||||
|
||||
@@ -849,14 +861,9 @@ class AsyncCrud(Generic[ModelType]):
|
||||
setattr(db_model, rel_attr, related_instances)
|
||||
|
||||
session.add(db_model)
|
||||
pk_values: dict[str, Any] | None = None
|
||||
if cls.default_load_options:
|
||||
await session.flush()
|
||||
pk_values = cls._capture_pk_values(db_model)
|
||||
if pk_values is not None:
|
||||
db_model = await cls._reload_with_options_by_pk(session, pk_values)
|
||||
else:
|
||||
await session.refresh(db_model)
|
||||
await session.refresh(db_model)
|
||||
if cls.default_load_options:
|
||||
db_model = await cls._reload_with_options(session, db_model)
|
||||
result = cast(ModelType, db_model)
|
||||
if schema:
|
||||
return Response(data=schema.model_validate(result))
|
||||
@@ -1233,15 +1240,9 @@ class AsyncCrud(Generic[ModelType]):
|
||||
m2m_resolved = await cls._resolve_m2m(session, obj, only_set=True)
|
||||
for rel_attr, related_instances in m2m_resolved.items():
|
||||
setattr(db_model, rel_attr, related_instances)
|
||||
|
||||
pk_values: dict[str, Any] | None = None
|
||||
if cls.default_load_options:
|
||||
await session.flush()
|
||||
pk_values = cls._capture_pk_values(db_model)
|
||||
if pk_values is not None:
|
||||
db_model = await cls._reload_with_options_by_pk(session, pk_values)
|
||||
else:
|
||||
await session.refresh(db_model)
|
||||
await session.refresh(db_model)
|
||||
if cls.default_load_options:
|
||||
db_model = await cls._reload_with_options(session, db_model)
|
||||
if schema:
|
||||
return Response(data=schema.model_validate(db_model))
|
||||
return db_model
|
||||
|
||||
@@ -466,6 +466,30 @@ class TestDefaultLoadOptionsIntegration:
|
||||
assert updated.role is not None
|
||||
assert updated.role.name == "admin"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_create_does_not_expire_already_loaded_relationships(
|
||||
self, db_session: AsyncSession
|
||||
):
|
||||
"""create()'s reload must not blow away loaded state on related objects."""
|
||||
UserWithDefaultLoad = CrudFactory(
|
||||
User, default_load_options=[selectinload(User.role)]
|
||||
)
|
||||
role = await RoleCrud.create(db_session, RoleCreate(name="admin"))
|
||||
role = await RoleCrud.get(
|
||||
db_session,
|
||||
filters=[Role.id == role.id],
|
||||
load_options=[selectinload(Role.users)],
|
||||
)
|
||||
assert role.users == []
|
||||
|
||||
await UserWithDefaultLoad.create(
|
||||
db_session,
|
||||
UserCreate(username="alice", email="alice@test.com", role_id=role.id),
|
||||
)
|
||||
|
||||
# must not trigger a lazy load
|
||||
assert role.users == []
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_load_options_overrides_default_load_options(
|
||||
self, db_session: AsyncSession
|
||||
|
||||
@@ -2543,6 +2543,16 @@ class TestOrderParamsViaConsolidated:
|
||||
assert len(result.data) == 2
|
||||
|
||||
|
||||
def _fully_declared_user_crud():
|
||||
"""A CRUD class declaring all three field sets, as a real app would."""
|
||||
return CrudFactory(
|
||||
User,
|
||||
searchable_fields=[User.username],
|
||||
facet_fields=[User.email],
|
||||
order_fields=[User.username],
|
||||
)
|
||||
|
||||
|
||||
class TestOffsetPaginateParamsSchema:
|
||||
"""Tests for AsyncCrud.offset_paginate_params()."""
|
||||
|
||||
@@ -2612,6 +2622,9 @@ class TestOffsetPaginateParamsSchema:
|
||||
"items_per_page": 10,
|
||||
"include_total": False,
|
||||
"include_facets": True,
|
||||
"search_fields": [],
|
||||
"facet_fields": [],
|
||||
"order_fields": [],
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
@@ -2655,6 +2668,42 @@ class TestOffsetPaginateParamsSchema:
|
||||
assert "search" not in param_names
|
||||
assert "search_column" not in param_names
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs",
|
||||
[
|
||||
{"search": False, "filter": False, "order": False},
|
||||
{"search_fields": [], "facet_fields": [], "order_fields": []},
|
||||
],
|
||||
ids=["flags", "empty-overrides"],
|
||||
)
|
||||
async def test_disabled_features_clear_response_metadata(
|
||||
self, db_session: AsyncSession, kwargs
|
||||
):
|
||||
"""Disabling a feature on one endpoint also drops it from the response."""
|
||||
await UserCrud.create(db_session, UserCreate(username="bob", email="b@x.io"))
|
||||
Crud = _fully_declared_user_crud()
|
||||
dep = Crud.offset_paginate_params(**kwargs)
|
||||
params = await dep(page=1, items_per_page=10)
|
||||
result = await Crud.offset_paginate(db_session, **params, schema=UserRead)
|
||||
assert result.search_columns is None
|
||||
assert result.order_columns is None
|
||||
assert result.filter_attributes is None
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_enabled_features_keep_response_metadata(
|
||||
self, db_session: AsyncSession
|
||||
):
|
||||
"""The declared class defaults still reach the response when left enabled."""
|
||||
await UserCrud.create(db_session, UserCreate(username="bob", email="b@x.io"))
|
||||
Crud = _fully_declared_user_crud()
|
||||
dep = Crud.offset_paginate_params()
|
||||
params = await dep(page=1, items_per_page=10)
|
||||
result = await Crud.offset_paginate(db_session, **params, schema=UserRead)
|
||||
assert result.search_columns == ["id", "username"]
|
||||
assert result.order_columns == ["username"]
|
||||
assert result.filter_attributes == {"email": ["b@x.io"]}
|
||||
|
||||
def test_filter_enabled_but_no_facet_fields(self):
|
||||
"""filter=True with no facet_fields silently skips filter params."""
|
||||
dep = RoleCrud.offset_paginate_params(search=False, filter=True, order=False)
|
||||
@@ -2727,6 +2776,9 @@ class TestCursorPaginateParamsSchema:
|
||||
"cursor": None,
|
||||
"items_per_page": 5,
|
||||
"include_facets": True,
|
||||
"search_fields": [],
|
||||
"facet_fields": [],
|
||||
"order_fields": [],
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
@@ -2837,6 +2889,9 @@ class TestPaginateParamsSchema:
|
||||
"items_per_page": 10,
|
||||
"include_total": True,
|
||||
"include_facets": True,
|
||||
"search_fields": [],
|
||||
"facet_fields": [],
|
||||
"order_fields": [],
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
|
||||
Reference in New Issue
Block a user