mirror of
https://github.com/d3vyce/fastapi-toolsets.git
synced 2026-08-04 23:54:09 +00:00
Compare commits
4
Commits
c15aee22b0
...
651326d54b
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
651326d54b | ||
|
|
169bf710f0 | ||
|
|
adc0ff14c1 | ||
|
|
19a823dc8b |
@@ -491,6 +491,16 @@ The distinct values for each facet field are returned in the `filter_attributes`
|
||||
!!! info "Key format uses `__` as a separator for relationship chains."
|
||||
A direct column `User.status` produces `"status"`. A relationship tuple `(User.role, Role.name)` produces `"role__name"`. A deeper chain `(User.role, Role.permission, Permission.name)` produces `"role__permission__name"`. An unknown `filter_by` key raises [`InvalidFacetFilterError`](../reference/exceptions.md#fastapi_toolsets.exceptions.exceptions.InvalidFacetFilterError) (HTTP 422).
|
||||
|
||||
#### Skipping facet queries
|
||||
|
||||
!!! info "Added in `v5.1.0`"
|
||||
|
||||
Facet values only change with the filters, not with the page. Pass `include_facets=False` to `offset_paginate_params()` / `cursor_paginate_params()` on pages 2..N to skip the facet queries entirely (`filter_attributes` will be `None`):
|
||||
|
||||
```python
|
||||
params: Annotated[dict, Depends(UserCrud.offset_paginate_params(include_facets=False))]
|
||||
```
|
||||
|
||||
## Sorting
|
||||
|
||||
!!! info "Added in `v1.3`"
|
||||
|
||||
@@ -44,6 +44,7 @@ from ..types import (
|
||||
)
|
||||
from .search import (
|
||||
SearchConfig,
|
||||
apply_search_joins,
|
||||
build_facets,
|
||||
build_filter_by,
|
||||
build_search_filters,
|
||||
@@ -51,7 +52,6 @@ from .search import (
|
||||
search_field_keys,
|
||||
)
|
||||
|
||||
|
||||
_ForUpdateMode: TypeAlias = bool | Literal["nowait", "skip_locked"]
|
||||
|
||||
|
||||
@@ -128,17 +128,6 @@ def _apply_joins(q: Any, joins: JoinType | None, outer_join: bool) -> Any:
|
||||
return q
|
||||
|
||||
|
||||
def _apply_search_joins(q: Any, search_joins: list[Any]) -> Any:
|
||||
"""Apply relationship-based outer joins (from search/filter_by) to a query."""
|
||||
seen: set[str] = set()
|
||||
for join_rel in search_joins:
|
||||
key = str(join_rel)
|
||||
if key not in seen:
|
||||
seen.add(key)
|
||||
q = q.outerjoin(join_rel)
|
||||
return q
|
||||
|
||||
|
||||
class AsyncCrud(Generic[ModelType]):
|
||||
"""Generic async CRUD operations for SQLAlchemy models.
|
||||
|
||||
@@ -184,17 +173,30 @@ class AsyncCrud(Generic[ModelType]):
|
||||
return cls.default_load_options
|
||||
|
||||
@classmethod
|
||||
async def _reload_with_options(
|
||||
cls: type[Self], session: AsyncSession, instance: ModelType
|
||||
def _capture_pk_values(cls: type[Self], instance: ModelType) -> 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]
|
||||
) -> ModelType:
|
||||
"""Re-query instance by PK with default_load_options applied."""
|
||||
mapper = cls.model.__mapper__
|
||||
"""Re-query by previously captured PK values, with default_load_options applied."""
|
||||
# Only called when cls.default_load_options is set (see call sites).
|
||||
pk_filters = [
|
||||
getattr(cls.model, cast(str, col.key))
|
||||
== getattr(instance, cast(str, col.key))
|
||||
for col in mapper.primary_key
|
||||
getattr(cls.model, key) == value for key, value in pk_values.items()
|
||||
]
|
||||
return await cls.get(session, filters=pk_filters)
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
async def _resolve_m2m(
|
||||
@@ -265,12 +267,12 @@ class AsyncCrud(Generic[ModelType]):
|
||||
cls: type[Self],
|
||||
filter_by: dict[str, Any] | BaseModel | None,
|
||||
facet_fields: Sequence[FacetFieldType] | None,
|
||||
) -> tuple[list[Any], list[Any]]:
|
||||
"""Normalize filter_by and return (filters, joins) to apply to the query."""
|
||||
) -> tuple[dict[str, Any], list[Any]]:
|
||||
"""Normalize filter_by and return ({facet_key: filter}, joins) to apply to the query."""
|
||||
if isinstance(filter_by, BaseModel):
|
||||
filter_by = filter_by.model_dump(exclude_none=True)
|
||||
if not filter_by:
|
||||
return [], []
|
||||
return {}, []
|
||||
resolved = cls._resolve_facet_fields(facet_fields)
|
||||
return build_filter_by(filter_by, resolved or [])
|
||||
|
||||
@@ -281,8 +283,13 @@ class AsyncCrud(Generic[ModelType]):
|
||||
facet_fields: Sequence[FacetFieldType] | None,
|
||||
filters: list[Any],
|
||||
search_joins: list[Any],
|
||||
*,
|
||||
include_facets: bool = True,
|
||||
own_filters: dict[str, Any] | None = None,
|
||||
) -> dict[str, list[Any]] | None:
|
||||
"""Build facet filter_attributes, or return None if no facet fields configured."""
|
||||
"""Build facet filter_attributes, or None if disabled/no facet fields configured."""
|
||||
if not include_facets:
|
||||
return None
|
||||
resolved = cls._resolve_facet_fields(facet_fields)
|
||||
if not resolved:
|
||||
return None
|
||||
@@ -292,6 +299,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
resolved,
|
||||
base_filters=filters,
|
||||
base_joins=search_joins,
|
||||
own_filters=own_filters,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -475,6 +483,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
default_page_size: int = 20,
|
||||
max_page_size: int = 100,
|
||||
include_total: bool = True,
|
||||
include_facets: bool = True,
|
||||
search: bool = True,
|
||||
filter: bool = True,
|
||||
order: bool = True,
|
||||
@@ -490,6 +499,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
default_page_size: Default ``items_per_page`` value.
|
||||
max_page_size: Maximum ``items_per_page`` value.
|
||||
include_total: Whether to include total count (not a query param).
|
||||
include_facets: Whether to run facet queries (not a query param).
|
||||
search: Enable search query parameters.
|
||||
filter: Enable facet filter query parameters.
|
||||
order: Enable order query parameters.
|
||||
@@ -519,7 +529,10 @@ class AsyncCrud(Generic[ModelType]):
|
||||
]
|
||||
return cls._build_paginate_params(
|
||||
pagination_params=pagination_params,
|
||||
pagination_fixed={"include_total": include_total},
|
||||
pagination_fixed={
|
||||
"include_total": include_total,
|
||||
"include_facets": include_facets,
|
||||
},
|
||||
dep_name=f"{cls.model.__name__}OffsetPaginateParams",
|
||||
search=search,
|
||||
filter=filter,
|
||||
@@ -537,6 +550,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
*,
|
||||
default_page_size: int = 20,
|
||||
max_page_size: int = 100,
|
||||
include_facets: bool = True,
|
||||
search: bool = True,
|
||||
filter: bool = True,
|
||||
order: bool = True,
|
||||
@@ -551,6 +565,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
Args:
|
||||
default_page_size: Default ``items_per_page`` value.
|
||||
max_page_size: Maximum ``items_per_page`` value.
|
||||
include_facets: Whether to run facet queries (not a query param).
|
||||
search: Enable search query parameters.
|
||||
filter: Enable facet filter query parameters.
|
||||
order: Enable order query parameters.
|
||||
@@ -582,7 +597,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
]
|
||||
return cls._build_paginate_params(
|
||||
pagination_params=pagination_params,
|
||||
pagination_fixed={},
|
||||
pagination_fixed={"include_facets": include_facets},
|
||||
dep_name=f"{cls.model.__name__}CursorPaginateParams",
|
||||
search=search,
|
||||
filter=filter,
|
||||
@@ -602,6 +617,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
max_page_size: int = 100,
|
||||
default_pagination_type: PaginationType = PaginationType.OFFSET,
|
||||
include_total: bool = True,
|
||||
include_facets: bool = True,
|
||||
search: bool = True,
|
||||
filter: bool = True,
|
||||
order: bool = True,
|
||||
@@ -618,6 +634,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
max_page_size: Maximum ``items_per_page`` value.
|
||||
default_pagination_type: Default pagination strategy.
|
||||
include_total: Whether to include total count (not a query param).
|
||||
include_facets: Whether to run facet queries (not a query param).
|
||||
search: Enable search query parameters.
|
||||
filter: Enable facet filter query parameters.
|
||||
order: Enable order query parameters.
|
||||
@@ -666,7 +683,10 @@ class AsyncCrud(Generic[ModelType]):
|
||||
]
|
||||
return cls._build_paginate_params(
|
||||
pagination_params=pagination_params,
|
||||
pagination_fixed={"include_total": include_total},
|
||||
pagination_fixed={
|
||||
"include_total": include_total,
|
||||
"include_facets": include_facets,
|
||||
},
|
||||
dep_name=f"{cls.model.__name__}PaginateParams",
|
||||
search=search,
|
||||
filter=filter,
|
||||
@@ -730,9 +750,14 @@ class AsyncCrud(Generic[ModelType]):
|
||||
setattr(db_model, rel_attr, related_instances)
|
||||
|
||||
session.add(db_model)
|
||||
await session.refresh(db_model)
|
||||
pk_values: dict[str, Any] | None = None
|
||||
if cls.default_load_options:
|
||||
db_model = await cls._reload_with_options(session, db_model)
|
||||
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)
|
||||
result = cast(ModelType, db_model)
|
||||
if schema:
|
||||
return Response(data=schema.model_validate(result))
|
||||
@@ -1097,9 +1122,15 @@ 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)
|
||||
await session.refresh(db_model)
|
||||
|
||||
pk_values: dict[str, Any] | None = None
|
||||
if cls.default_load_options:
|
||||
db_model = await cls._reload_with_options(session, db_model)
|
||||
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)
|
||||
if schema:
|
||||
return Response(data=schema.model_validate(db_model))
|
||||
return db_model
|
||||
@@ -1271,6 +1302,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: str | None = None,
|
||||
order_fields: Sequence[OrderFieldType] | None = None,
|
||||
facet_fields: Sequence[FacetFieldType] | None = None,
|
||||
include_facets: bool = True,
|
||||
filter_by: dict[str, Any] | BaseModel | None = None,
|
||||
schema: type[BaseModel],
|
||||
) -> OffsetPaginatedResponse[Any]:
|
||||
@@ -1292,6 +1324,9 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: Restrict search to a single column key.
|
||||
order_fields: Fields allowed for sorting (overrides class default).
|
||||
facet_fields: Columns to compute distinct values for (overrides class default)
|
||||
include_facets: When ``False``, skip facet queries entirely;
|
||||
``filter_attributes`` will be ``None``. Useful on pages 2..N
|
||||
where the facet counts were already fetched on page 1.
|
||||
filter_by: Dict of {column_key: value} to filter by declared facet fields.
|
||||
Keys must match the column.key of a facet field. Scalar → equality,
|
||||
list → IN clause. Raises InvalidFacetFilterError for unknown keys.
|
||||
@@ -1304,7 +1339,6 @@ class AsyncCrud(Generic[ModelType]):
|
||||
offset = (page - 1) * items_per_page
|
||||
|
||||
fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields)
|
||||
filters.extend(fb_filters)
|
||||
|
||||
# Build search filters
|
||||
if search:
|
||||
@@ -1318,6 +1352,11 @@ class AsyncCrud(Generic[ModelType]):
|
||||
filters.extend(search_filters)
|
||||
search_joins.extend(new_search_joins)
|
||||
|
||||
# Facets combine these with each facet's own filter individually, so
|
||||
# fb_filters is applied to the query below but excluded here.
|
||||
facet_base_filters = list(filters)
|
||||
filters.extend(fb_filters.values())
|
||||
|
||||
# Build query with joins
|
||||
q = select(cls.model)
|
||||
|
||||
@@ -1325,11 +1364,11 @@ class AsyncCrud(Generic[ModelType]):
|
||||
q = _apply_joins(q, joins, outer_join)
|
||||
|
||||
# Apply search joins (always outer joins for search)
|
||||
q = _apply_search_joins(q, search_joins)
|
||||
q = apply_search_joins(q, search_joins)
|
||||
|
||||
# Apply order joins (relation joins required for order_by field)
|
||||
if order_joins:
|
||||
q = _apply_search_joins(q, order_joins)
|
||||
q = apply_search_joins(q, order_joins)
|
||||
|
||||
if filters:
|
||||
q = q.where(and_(*filters))
|
||||
@@ -1352,7 +1391,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
count_q = _apply_joins(count_q, joins, outer_join)
|
||||
|
||||
# Apply search joins to count query
|
||||
count_q = _apply_search_joins(count_q, search_joins)
|
||||
count_q = apply_search_joins(count_q, search_joins)
|
||||
|
||||
if filters:
|
||||
count_q = count_q.where(and_(*filters))
|
||||
@@ -1372,7 +1411,12 @@ class AsyncCrud(Generic[ModelType]):
|
||||
items: list[Any] = [schema.model_validate(item) for item in raw_items]
|
||||
|
||||
filter_attributes = await cls._build_filter_attributes(
|
||||
session, facet_fields, filters, search_joins
|
||||
session,
|
||||
facet_fields,
|
||||
facet_base_filters,
|
||||
search_joins,
|
||||
include_facets=include_facets,
|
||||
own_filters=fb_filters,
|
||||
)
|
||||
search_columns = cls._resolve_search_columns(search_fields)
|
||||
order_columns = cls._resolve_order_columns(order_fields)
|
||||
@@ -1408,6 +1452,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: str | None = None,
|
||||
order_fields: Sequence[OrderFieldType] | None = None,
|
||||
facet_fields: Sequence[FacetFieldType] | None = None,
|
||||
include_facets: bool = True,
|
||||
filter_by: dict[str, Any] | BaseModel | None = None,
|
||||
schema: type[BaseModel],
|
||||
) -> CursorPaginatedResponse[Any]:
|
||||
@@ -1430,6 +1475,8 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: Restrict search to a single column key.
|
||||
order_fields: Fields allowed for sorting (overrides class default).
|
||||
facet_fields: Columns to compute distinct values for (overrides class default).
|
||||
include_facets: When ``False``, skip facet queries entirely;
|
||||
``filter_attributes`` will be ``None``.
|
||||
filter_by: Dict of {column_key: value} to filter by declared facet fields.
|
||||
Keys must match the column.key of a facet field. Scalar → equality,
|
||||
list → IN clause. Raises InvalidFacetFilterError for unknown keys.
|
||||
@@ -1441,7 +1488,6 @@ class AsyncCrud(Generic[ModelType]):
|
||||
filters = list(filters) if filters else []
|
||||
|
||||
fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields)
|
||||
filters.extend(fb_filters)
|
||||
|
||||
if cls.cursor_column is None:
|
||||
raise ValueError(
|
||||
@@ -1473,6 +1519,11 @@ class AsyncCrud(Generic[ModelType]):
|
||||
filters.extend(search_filters)
|
||||
search_joins.extend(new_search_joins)
|
||||
|
||||
# Facets combine these with each facet's own filter individually, so
|
||||
# fb_filters is applied to the query below but excluded here.
|
||||
facet_base_filters = list(filters)
|
||||
filters.extend(fb_filters.values())
|
||||
|
||||
# Build query
|
||||
q = select(cls.model)
|
||||
|
||||
@@ -1480,11 +1531,11 @@ class AsyncCrud(Generic[ModelType]):
|
||||
q = _apply_joins(q, joins, outer_join)
|
||||
|
||||
# Apply search joins (always outer joins)
|
||||
q = _apply_search_joins(q, search_joins)
|
||||
q = apply_search_joins(q, search_joins)
|
||||
|
||||
# Apply order joins (relation joins required for order_by field)
|
||||
if order_joins:
|
||||
q = _apply_search_joins(q, order_joins)
|
||||
q = apply_search_joins(q, order_joins)
|
||||
|
||||
if filters:
|
||||
q = q.where(and_(*filters))
|
||||
@@ -1541,7 +1592,12 @@ class AsyncCrud(Generic[ModelType]):
|
||||
items: list[Any] = [schema.model_validate(item) for item in items_page]
|
||||
|
||||
filter_attributes = await cls._build_filter_attributes(
|
||||
session, facet_fields, filters, search_joins
|
||||
session,
|
||||
facet_fields,
|
||||
facet_base_filters,
|
||||
search_joins,
|
||||
include_facets=include_facets,
|
||||
own_filters=fb_filters,
|
||||
)
|
||||
search_columns = cls._resolve_search_columns(search_fields)
|
||||
order_columns = cls._resolve_order_columns(order_fields)
|
||||
@@ -1581,6 +1637,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: str | None = ...,
|
||||
order_fields: Sequence[OrderFieldType] | None = ...,
|
||||
facet_fields: Sequence[FacetFieldType] | None = ...,
|
||||
include_facets: bool = ...,
|
||||
filter_by: dict[str, Any] | BaseModel | None = ...,
|
||||
schema: type[BaseModel],
|
||||
) -> OffsetPaginatedResponse[Any]: ...
|
||||
@@ -1607,6 +1664,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: str | None = ...,
|
||||
order_fields: Sequence[OrderFieldType] | None = ...,
|
||||
facet_fields: Sequence[FacetFieldType] | None = ...,
|
||||
include_facets: bool = ...,
|
||||
filter_by: dict[str, Any] | BaseModel | None = ...,
|
||||
schema: type[BaseModel],
|
||||
) -> CursorPaginatedResponse[Any]: ...
|
||||
@@ -1632,6 +1690,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column: str | None = None,
|
||||
order_fields: Sequence[OrderFieldType] | None = None,
|
||||
facet_fields: Sequence[FacetFieldType] | None = None,
|
||||
include_facets: bool = True,
|
||||
filter_by: dict[str, Any] | BaseModel | None = None,
|
||||
schema: type[BaseModel],
|
||||
) -> OffsetPaginatedResponse[Any] | CursorPaginatedResponse[Any]:
|
||||
@@ -1662,6 +1721,8 @@ class AsyncCrud(Generic[ModelType]):
|
||||
order_fields: Fields allowed for sorting (overrides class default).
|
||||
facet_fields: Columns to compute distinct values for (overrides
|
||||
class default).
|
||||
include_facets: When ``False``, skip facet queries entirely;
|
||||
``filter_attributes`` will be ``None``.
|
||||
filter_by: Dict of ``{column_key: value}`` to filter by declared
|
||||
facet fields. Keys must match the ``column.key`` of a facet
|
||||
field. Scalar → equality, list → IN clause. Raises
|
||||
@@ -1692,6 +1753,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column=search_column,
|
||||
order_fields=order_fields,
|
||||
facet_fields=facet_fields,
|
||||
include_facets=include_facets,
|
||||
filter_by=filter_by,
|
||||
schema=schema,
|
||||
)
|
||||
@@ -1714,6 +1776,7 @@ class AsyncCrud(Generic[ModelType]):
|
||||
search_column=search_column,
|
||||
order_fields=order_fields,
|
||||
facet_fields=facet_fields,
|
||||
include_facets=include_facets,
|
||||
filter_by=filter_by,
|
||||
schema=schema,
|
||||
)
|
||||
|
||||
@@ -1,12 +1,12 @@
|
||||
"""Search utilities for AsyncCrud."""
|
||||
|
||||
import asyncio
|
||||
import functools
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
|
||||
from sqlalchemy import String, and_, func, or_, select
|
||||
from sqlalchemy import String, and_, distinct, func, or_, select
|
||||
from sqlalchemy.dialects.postgresql import aggregate_order_by
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
from sqlalchemy.orm.attributes import InstrumentedAttribute
|
||||
@@ -159,8 +159,11 @@ def build_search_filters(
|
||||
else:
|
||||
column = field
|
||||
|
||||
# Build the filter (cast to String for non-text columns)
|
||||
column_as_string = column.cast(String)
|
||||
# Build the filter (cast to String only when needed, to preserve
|
||||
# pg_trgm GIN index usability on already-String columns)
|
||||
column_as_string = (
|
||||
column if isinstance(column.type, String) else column.cast(String)
|
||||
)
|
||||
if config.case_sensitive:
|
||||
filters.append(column_as_string.like(f"%{query}%"))
|
||||
else:
|
||||
@@ -181,6 +184,21 @@ def search_field_keys(fields: Sequence[SearchFieldType]) -> list[str]:
|
||||
return facet_keys(fields)
|
||||
|
||||
|
||||
def apply_search_joins(q: Any, joins: Sequence[Any]) -> Any:
|
||||
"""Apply relationship-based outer joins (from search/filter_by/facets) to a query.
|
||||
|
||||
Deduplicates by relationship identity so a join used by several fields
|
||||
(e.g. search + a facet on the same relation) is only applied once.
|
||||
"""
|
||||
seen: set[str] = set()
|
||||
for rel in joins:
|
||||
rel_key = str(rel)
|
||||
if rel_key not in seen:
|
||||
seen.add(rel_key)
|
||||
q = q.outerjoin(rel)
|
||||
return q
|
||||
|
||||
|
||||
def facet_keys(facet_fields: Sequence[FacetFieldType]) -> list[str]:
|
||||
"""Return a key for each facet field.
|
||||
|
||||
@@ -207,6 +225,7 @@ async def build_facets(
|
||||
*,
|
||||
base_filters: "list[ColumnElement[bool]] | None" = None,
|
||||
base_joins: list[InstrumentedAttribute[Any]] | None = None,
|
||||
own_filters: "dict[str, ColumnElement[bool]] | None" = None,
|
||||
) -> dict[str, list[Any]]:
|
||||
"""Return distinct values for each facet field, respecting current filters.
|
||||
|
||||
@@ -216,15 +235,24 @@ async def build_facets(
|
||||
facet_fields: Columns or relationship tuples to facet on
|
||||
base_filters: Filter conditions already applied to the main query (search + caller filters)
|
||||
base_joins: Relationship joins already applied to the main query
|
||||
own_filters: Map of facet key -> the ``filter_by`` condition for that
|
||||
same key (if any). Excluded from that facet's own subquery so
|
||||
filtering on a facet doesn't collapse its own value list down to
|
||||
just the filtered value.
|
||||
|
||||
Returns:
|
||||
Dict mapping column key to sorted list of distinct non-None values
|
||||
"""
|
||||
existing_join_keys: set[str] = {str(j) for j in (base_joins or [])}
|
||||
if not facet_fields:
|
||||
return {}
|
||||
|
||||
keys = facet_keys(facet_fields)
|
||||
own_filters = own_filters or {}
|
||||
|
||||
async def _query_facet(field: FacetFieldType, key: str) -> tuple[str, list[Any]]:
|
||||
scalars: list[Any] = []
|
||||
enum_classes: dict[str, Any] = {}
|
||||
|
||||
for field, key in zip(facet_fields, keys):
|
||||
if isinstance(field, tuple):
|
||||
# Relationship chain: (User.role, Role.name) — last element is the column
|
||||
rels = field[:-1]
|
||||
@@ -235,51 +263,48 @@ async def build_facets(
|
||||
|
||||
col_type = column.property.columns[0].type
|
||||
is_array = isinstance(col_type, ARRAY)
|
||||
enum_classes[key] = getattr(col_type, "enum_class", None)
|
||||
|
||||
if is_array:
|
||||
unnested = func.unnest(column).label(column.key)
|
||||
q = select(unnested).select_from(model).distinct()
|
||||
else:
|
||||
q = select(column).select_from(model).distinct()
|
||||
|
||||
# Apply base joins (deduplicated) — needed here independently
|
||||
seen_joins: set[str] = set()
|
||||
for rel in base_joins or []:
|
||||
rel_key = str(rel)
|
||||
if rel_key not in seen_joins:
|
||||
seen_joins.add(rel_key)
|
||||
q = q.outerjoin(rel)
|
||||
|
||||
# Add any extra joins required by this facet field that aren't already applied
|
||||
for rel in rels:
|
||||
rel_key = str(rel)
|
||||
if rel_key not in existing_join_keys and rel_key not in seen_joins:
|
||||
seen_joins.add(rel_key)
|
||||
q = q.outerjoin(rel)
|
||||
|
||||
if base_filters:
|
||||
q = q.where(and_(*base_filters))
|
||||
|
||||
if is_array:
|
||||
q = q.order_by(unnested)
|
||||
else:
|
||||
q = q.order_by(column)
|
||||
result = await session.execute(q)
|
||||
col_type = column.property.columns[0].type
|
||||
enum_class = getattr(col_type, "enum_class", None)
|
||||
values = [
|
||||
row[0].name
|
||||
if (enum_class is not None and isinstance(row[0], enum_class))
|
||||
else row[0]
|
||||
for row in result.all()
|
||||
if row[0] is not None
|
||||
filters = [
|
||||
*(base_filters or []),
|
||||
*(f for k, f in own_filters.items() if k != key),
|
||||
]
|
||||
return key, values
|
||||
joins = [*(base_joins or []), *rels]
|
||||
|
||||
pairs = await asyncio.gather(
|
||||
*[_query_facet(f, k) for f, k in zip(facet_fields, keys)]
|
||||
if is_array:
|
||||
unnested = apply_search_joins(
|
||||
select(func.unnest(column).label("v")).select_from(model), joins
|
||||
)
|
||||
return dict(pairs)
|
||||
if filters:
|
||||
unnested = unnested.where(and_(*filters))
|
||||
unnested_sq = unnested.subquery()
|
||||
v = unnested_sq.c.v
|
||||
agg = (
|
||||
select(func.array_agg(aggregate_order_by(distinct(v), v)))
|
||||
.select_from(unnested_sq)
|
||||
.where(v.isnot(None))
|
||||
)
|
||||
else:
|
||||
agg = apply_search_joins(
|
||||
select(
|
||||
func.array_agg(aggregate_order_by(distinct(column), column))
|
||||
).select_from(model),
|
||||
joins,
|
||||
)
|
||||
agg = agg.where(and_(*filters, column.isnot(None)))
|
||||
|
||||
scalars.append(agg.scalar_subquery().label(key))
|
||||
|
||||
row = (await session.execute(select(*scalars))).one()
|
||||
|
||||
facets: dict[str, list[Any]] = {}
|
||||
for key, values in zip(keys, row):
|
||||
enum_class = enum_classes[key]
|
||||
facets[key] = [
|
||||
v.name if (enum_class is not None and isinstance(v, enum_class)) else v
|
||||
for v in (values or [])
|
||||
]
|
||||
return facets
|
||||
|
||||
|
||||
_EQUALITY_TYPES = (String, Integer, Numeric, Date, DateTime, Time, Enum, Uuid)
|
||||
@@ -301,7 +326,7 @@ def _coerce_bool(value: Any) -> bool:
|
||||
def build_filter_by(
|
||||
filter_by: dict[str, Any],
|
||||
facet_fields: Sequence[FacetFieldType],
|
||||
) -> tuple["list[ColumnElement[bool]]", list[InstrumentedAttribute[Any]]]:
|
||||
) -> tuple["dict[str, ColumnElement[bool]]", list[InstrumentedAttribute[Any]]]:
|
||||
"""Translate a {column_key: value} dict into SQLAlchemy filter conditions.
|
||||
|
||||
Args:
|
||||
@@ -309,7 +334,9 @@ def build_filter_by(
|
||||
facet_fields: Declared facet fields to validate keys against
|
||||
|
||||
Returns:
|
||||
Tuple of (filter_conditions, joins_needed)
|
||||
Tuple of ({facet_key: filter_condition}, joins_needed). One filter
|
||||
condition per key, so callers can identify (and exclude) a facet's
|
||||
own filter when computing that facet's distinct values.
|
||||
|
||||
Raises:
|
||||
InvalidFacetFilterError: If a key in filter_by is not a declared facet field
|
||||
@@ -327,7 +354,7 @@ def build_filter_by(
|
||||
index[key] = (column, rels)
|
||||
|
||||
valid_keys = set(index)
|
||||
filters: list[ColumnElement[bool]] = []
|
||||
filters: dict[str, ColumnElement[bool]] = {}
|
||||
joins: list[InstrumentedAttribute[Any]] = []
|
||||
added_join_keys: set[str] = set()
|
||||
|
||||
@@ -347,14 +374,14 @@ def build_filter_by(
|
||||
if isinstance(col_type, Boolean):
|
||||
coerce = _coerce_bool
|
||||
if isinstance(value, list):
|
||||
filters.append(column.in_([coerce(v) for v in value]))
|
||||
filters[key] = column.in_([coerce(v) for v in value])
|
||||
else:
|
||||
filters.append(column == coerce(value))
|
||||
filters[key] = column == coerce(value)
|
||||
elif isinstance(col_type, ARRAY):
|
||||
if isinstance(value, list):
|
||||
filters.append(column.overlap(value))
|
||||
filters[key] = column.overlap(value)
|
||||
else:
|
||||
filters.append(column.any(value))
|
||||
filters[key] = column.any(value)
|
||||
elif isinstance(col_type, Enum):
|
||||
enum_class = col_type.enum_class
|
||||
if enum_class is not None:
|
||||
@@ -365,19 +392,19 @@ def build_filter_by(
|
||||
return enum_class[v] # lookup by name: "PENDING", "RED"
|
||||
|
||||
if isinstance(value, list):
|
||||
filters.append(column.in_([_coerce_enum(v) for v in value]))
|
||||
filters[key] = column.in_([_coerce_enum(v) for v in value])
|
||||
else:
|
||||
filters.append(column == _coerce_enum(value))
|
||||
filters[key] = column == _coerce_enum(value)
|
||||
else: # pragma: no cover
|
||||
if isinstance(value, list):
|
||||
filters.append(column.in_(value))
|
||||
filters[key] = column.in_(value)
|
||||
else:
|
||||
filters.append(column == value)
|
||||
filters[key] = column == value
|
||||
elif isinstance(col_type, _EQUALITY_TYPES):
|
||||
if isinstance(value, list):
|
||||
filters.append(column.in_(value))
|
||||
filters[key] = column.in_(value)
|
||||
else:
|
||||
filters.append(column == value)
|
||||
filters[key] = column == value
|
||||
else:
|
||||
raise UnsupportedFacetTypeError(key, type(col_type).__name__)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import event
|
||||
from sqlalchemy import event, select, tuple_
|
||||
from sqlalchemy import inspect as sa_inspect
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import set_committed_value as _sa_set_committed_value
|
||||
@@ -29,14 +29,11 @@ _SESSION_DELETES = "_ft_deletes"
|
||||
_SESSION_UPDATES = "_ft_updates"
|
||||
_DEFERRED_STRATEGY_KEY = (("deferred", True), ("instrument", True))
|
||||
_EVENT_HANDLERS: dict[tuple[type, ModelEvent], list[Callable[..., Any]]] = {}
|
||||
_WATCHED_MODELS: set[type] = set()
|
||||
_WATCHED_CACHE: dict[type, bool] = {}
|
||||
_HANDLER_CACHE: dict[tuple[type, ModelEvent], list[Callable[..., Any]]] = {}
|
||||
|
||||
|
||||
def _invalidate_caches() -> None:
|
||||
"""Clear lookup caches after handler registration."""
|
||||
_WATCHED_CACHE.clear()
|
||||
_HANDLER_CACHE.clear()
|
||||
|
||||
|
||||
@@ -56,24 +53,12 @@ def listens_for(
|
||||
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
|
||||
for ev in evs:
|
||||
_EVENT_HANDLERS.setdefault((model_class, ev), []).append(fn)
|
||||
_WATCHED_MODELS.add(model_class)
|
||||
_invalidate_caches()
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _is_watched(obj: Any) -> bool:
|
||||
"""Return True if *obj*'s type (or any ancestor) has registered handlers."""
|
||||
cls = type(obj)
|
||||
try:
|
||||
return _WATCHED_CACHE[cls]
|
||||
except KeyError:
|
||||
result = any(klass in _WATCHED_MODELS for klass in cls.__mro__)
|
||||
_WATCHED_CACHE[cls] = result
|
||||
return result
|
||||
|
||||
|
||||
def _get_handlers(cls: type, ev: ModelEvent) -> list[Callable[..., Any]]:
|
||||
"""Return registered handlers for *cls* and *ev*, walking the MRO."""
|
||||
key = (cls, ev)
|
||||
@@ -144,18 +129,18 @@ def _upsert_changes(
|
||||
def _after_flush(session: Any, flush_context: Any) -> None:
|
||||
# New objects: capture reference. Attributes will be refreshed after commit.
|
||||
for obj in session.new:
|
||||
if _is_watched(obj):
|
||||
if _get_handlers(type(obj), ModelEvent.CREATE):
|
||||
session.info.setdefault(_SESSION_CREATES, []).append(obj)
|
||||
|
||||
# Deleted objects: snapshot now while attributes are still loaded.
|
||||
for obj in session.deleted:
|
||||
if _is_watched(obj):
|
||||
if _get_handlers(type(obj), ModelEvent.DELETE):
|
||||
snapshot = _snapshot_column_attrs(obj)
|
||||
session.info.setdefault(_SESSION_DELETES, []).append((obj, snapshot))
|
||||
|
||||
# Dirty objects: read old/new from SQLAlchemy attribute history.
|
||||
for obj in session.dirty:
|
||||
if not _is_watched(obj):
|
||||
if not _get_handlers(type(obj), ModelEvent.UPDATE):
|
||||
continue
|
||||
|
||||
watched = _get_watched_fields(type(obj))
|
||||
@@ -204,9 +189,18 @@ async def _invoke_callback(
|
||||
await result
|
||||
|
||||
|
||||
async def _reload_if_present(session: AsyncSession, obj: Any, state: Any) -> None:
|
||||
"""Re-populate *obj* from the DB if its row still exists."""
|
||||
await session.get(type(obj), state.key[1], populate_existing=True)
|
||||
async def _batch_reload(
|
||||
session: AsyncSession, model: type, pk_tuples: list[tuple[Any, ...]]
|
||||
) -> None:
|
||||
"""Re-populate all rows of *model* identified by *pk_tuples* in one round trip."""
|
||||
pk_cols = sa_inspect(model).primary_key
|
||||
where = (
|
||||
pk_cols[0].in_([pk[0] for pk in pk_tuples])
|
||||
if len(pk_cols) == 1
|
||||
else tuple_(*pk_cols).in_(pk_tuples)
|
||||
)
|
||||
q = select(model).where(where).execution_options(populate_existing=True)
|
||||
await session.execute(q)
|
||||
|
||||
|
||||
class EventSession(AsyncSession):
|
||||
@@ -250,15 +244,36 @@ class EventSession(AsyncSession):
|
||||
k: v for k, v in field_changes.items() if k not in create_ids
|
||||
}
|
||||
|
||||
# Dispatch CREATE callbacks.
|
||||
# Resolve reloadable state up front and group PKs by model type so
|
||||
# the post-commit reload is one query per type instead of one
|
||||
# session.get() per object.
|
||||
create_items: list[Any] = []
|
||||
update_items: list[tuple[Any, dict[str, dict[str, Any]]]] = []
|
||||
pk_by_type: dict[type, list[tuple[Any, ...]]] = {}
|
||||
|
||||
for obj in creates:
|
||||
try:
|
||||
state = sa_inspect(obj, raiseerr=False)
|
||||
if (
|
||||
state is None or state.detached or state.transient
|
||||
): # pragma: no cover
|
||||
if state is None or state.detached or state.transient: # pragma: no cover
|
||||
continue
|
||||
await _reload_if_present(self, obj, state)
|
||||
create_items.append(obj)
|
||||
pk_by_type.setdefault(type(obj), []).append(state.key[1])
|
||||
|
||||
for obj, changes in field_changes.values():
|
||||
state = sa_inspect(obj, raiseerr=False)
|
||||
if state is None or state.detached or state.transient: # pragma: no cover
|
||||
continue
|
||||
update_items.append((obj, changes))
|
||||
pk_by_type.setdefault(type(obj), []).append(state.key[1])
|
||||
|
||||
for model, pk_tuples in pk_by_type.items():
|
||||
try:
|
||||
await _batch_reload(self, model, pk_tuples)
|
||||
except Exception as exc:
|
||||
_logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
|
||||
|
||||
# Dispatch CREATE callbacks.
|
||||
for obj in create_items:
|
||||
try:
|
||||
for handler in _get_handlers(type(obj), ModelEvent.CREATE):
|
||||
await _invoke_callback(handler, obj, ModelEvent.CREATE, None)
|
||||
except Exception as exc:
|
||||
@@ -275,14 +290,8 @@ class EventSession(AsyncSession):
|
||||
_logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
|
||||
|
||||
# Dispatch UPDATE callbacks.
|
||||
for obj, changes in field_changes.values():
|
||||
for obj, changes in update_items:
|
||||
try:
|
||||
state = sa_inspect(obj, raiseerr=False)
|
||||
if (
|
||||
state is None or state.detached or state.transient
|
||||
): # pragma: no cover
|
||||
continue
|
||||
await _reload_if_present(self, obj, state)
|
||||
for handler in _get_handlers(type(obj), ModelEvent.UPDATE):
|
||||
await _invoke_callback(handler, obj, ModelEvent.UPDATE, changes)
|
||||
except Exception as exc:
|
||||
|
||||
@@ -476,3 +476,26 @@ async def db_session(engine):
|
||||
# Drop tables after test
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
async def db_session_expire_on_commit(engine):
|
||||
"""Session with expire_on_commit=True (the SQLAlchemy default).
|
||||
|
||||
Attributes read off an instance after commit are expired and trigger an
|
||||
implicit (sync) refresh under this setting — which fails under asyncio
|
||||
with MissingGreenlet. The other ``db_session`` fixture uses
|
||||
``expire_on_commit=False`` and would not catch that class of bug.
|
||||
"""
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=True)
|
||||
session = session_factory()
|
||||
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
await session.close()
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.drop_all)
|
||||
|
||||
@@ -417,6 +417,55 @@ class TestDefaultLoadOptionsIntegration:
|
||||
assert updated.role is not None
|
||||
assert updated.role.name == "admin"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_default_load_options_applied_to_create_expire_on_commit(
|
||||
self, db_session_expire_on_commit: AsyncSession
|
||||
):
|
||||
"""create()'s reload uses captured PK values, not an expired instance attribute.
|
||||
|
||||
Regression test for MissingGreenlet: reading a PK off `db_model` after
|
||||
commit under expire_on_commit=True (the SQLAlchemy default) would
|
||||
trigger an implicit sync refresh, which fails under asyncio.
|
||||
"""
|
||||
UserWithDefaultLoad = CrudFactory(
|
||||
User, default_load_options=[selectinload(User.role)]
|
||||
)
|
||||
role = await RoleCrud.create(
|
||||
db_session_expire_on_commit, RoleCreate(name="admin")
|
||||
)
|
||||
user = await UserWithDefaultLoad.create(
|
||||
db_session_expire_on_commit,
|
||||
UserCreate(username="alice", email="alice@test.com", role_id=role.id),
|
||||
)
|
||||
assert user.role is not None
|
||||
assert user.role.name == "admin"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_default_load_options_applied_to_update_expire_on_commit(
|
||||
self, db_session_expire_on_commit: AsyncSession
|
||||
):
|
||||
"""update()'s reload uses captured PK values, not an expired instance attribute.
|
||||
|
||||
Regression test for MissingGreenlet under expire_on_commit=True.
|
||||
"""
|
||||
UserWithDefaultLoad = CrudFactory(
|
||||
User, default_load_options=[selectinload(User.role)]
|
||||
)
|
||||
role = await RoleCrud.create(
|
||||
db_session_expire_on_commit, RoleCreate(name="admin")
|
||||
)
|
||||
user = await UserCrud.create(
|
||||
db_session_expire_on_commit,
|
||||
UserCreate(username="alice", email="alice@test.com"),
|
||||
)
|
||||
updated = await UserWithDefaultLoad.update(
|
||||
db_session_expire_on_commit,
|
||||
UserUpdate(role_id=role.id),
|
||||
filters=[User.id == user.id],
|
||||
)
|
||||
assert updated.role is not None
|
||||
assert updated.role.name == "admin"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_load_options_overrides_default_load_options(
|
||||
self, db_session: AsyncSession
|
||||
|
||||
+151
-7
@@ -372,6 +372,22 @@ class TestBuildSearchFilters:
|
||||
|
||||
assert len(joins) == 1
|
||||
|
||||
def test_skips_cast_on_string_column(self):
|
||||
"""String columns are filtered directly, without a CAST (keeps pg_trgm indexable)."""
|
||||
from fastapi_toolsets.crud.search import build_search_filters
|
||||
|
||||
filters, _ = build_search_filters(User, "john", search_fields=[User.username])
|
||||
|
||||
assert "CAST" not in str(filters[0])
|
||||
|
||||
def test_casts_non_string_column(self):
|
||||
"""Non-string columns (e.g. UUID) still get cast to String so ilike works."""
|
||||
from fastapi_toolsets.crud.search import build_search_filters
|
||||
|
||||
filters, _ = build_search_filters(User, "123", search_fields=[User.id])
|
||||
|
||||
assert "CAST" in str(filters[0])
|
||||
|
||||
|
||||
class TestSearchConfig:
|
||||
"""Tests for SearchConfig options."""
|
||||
@@ -531,6 +547,15 @@ class TestFacetsNotSet:
|
||||
|
||||
assert result.filter_attributes is None
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_build_facets_empty_field_list(self, db_session: AsyncSession):
|
||||
"""build_facets([]) is a no-op that returns {} without querying — the escape hatch."""
|
||||
from fastapi_toolsets.crud.search import build_facets
|
||||
|
||||
result = await build_facets(db_session, User, [])
|
||||
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestFacetsDirectColumn:
|
||||
"""Facets on direct model columns."""
|
||||
@@ -606,6 +631,91 @@ class TestFacetsDirectColumn:
|
||||
assert "username" not in result.filter_attributes
|
||||
|
||||
|
||||
class TestFacetsMixedTypes:
|
||||
"""Facet values keep their native Python type through the batched query."""
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_enum_and_integer_facets_preserve_types(
|
||||
self, db_session: AsyncSession
|
||||
):
|
||||
"""Enum facets return member names (not raw DB values); Integer facets return ints."""
|
||||
OrderMixedCrud = CrudFactory(
|
||||
Order, facet_fields=[Order.status, Order.priority, Order.color]
|
||||
)
|
||||
await OrderCrud.create(
|
||||
db_session,
|
||||
OrderCreate(
|
||||
name="order-1", status=OrderStatus.PENDING, priority=1, color=Color.RED
|
||||
),
|
||||
)
|
||||
await OrderCrud.create(
|
||||
db_session,
|
||||
OrderCreate(
|
||||
name="order-2", status=OrderStatus.SHIPPED, priority=3, color=Color.BLUE
|
||||
),
|
||||
)
|
||||
|
||||
result = await OrderMixedCrud.offset_paginate(db_session, schema=OrderRead)
|
||||
|
||||
assert result.filter_attributes is not None
|
||||
assert set(result.filter_attributes["status"]) == {"PENDING", "SHIPPED"}
|
||||
assert all(isinstance(v, str) for v in result.filter_attributes["status"])
|
||||
assert set(result.filter_attributes["priority"]) == {1, 3}
|
||||
assert all(isinstance(v, int) for v in result.filter_attributes["priority"])
|
||||
assert set(result.filter_attributes["color"]) == {"RED", "BLUE"}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_bool_facet_keeps_python_bool(self, db_session: AsyncSession):
|
||||
"""A Boolean facet returns Python bool values, not stringified 'true'/'false'."""
|
||||
UserBoolFacetCrud = CrudFactory(User, facet_fields=[User.is_active])
|
||||
await UserCrud.create(
|
||||
db_session, UserCreate(username="alice", email="a@test.com", is_active=True)
|
||||
)
|
||||
await UserCrud.create(
|
||||
db_session, UserCreate(username="bob", email="b@test.com", is_active=False)
|
||||
)
|
||||
|
||||
result = await UserBoolFacetCrud.offset_paginate(db_session, schema=UserRead)
|
||||
|
||||
assert result.filter_attributes is not None
|
||||
assert set(result.filter_attributes["is_active"]) == {True, False}
|
||||
assert all(isinstance(v, bool) for v in result.filter_attributes["is_active"])
|
||||
|
||||
|
||||
class TestIncludeFacets:
|
||||
"""include_facets=False skips facet queries entirely."""
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_offset_paginate_include_facets_false(self, db_session: AsyncSession):
|
||||
"""filter_attributes is None when include_facets=False, even with facet_fields set."""
|
||||
UserFacetCrud = CrudFactory(User, facet_fields=[User.username])
|
||||
await UserCrud.create(
|
||||
db_session, UserCreate(username="alice", email="a@test.com")
|
||||
)
|
||||
|
||||
result = await UserFacetCrud.offset_paginate(
|
||||
db_session, include_facets=False, schema=UserRead
|
||||
)
|
||||
|
||||
assert result.filter_attributes is None
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_cursor_paginate_include_facets_false(self, db_session: AsyncSession):
|
||||
"""filter_attributes is None when include_facets=False for cursor_paginate."""
|
||||
UserFacetCursorCrud = CrudFactory(
|
||||
User, cursor_column=User.id, facet_fields=[User.username]
|
||||
)
|
||||
await UserCrud.create(
|
||||
db_session, UserCreate(username="alice", email="a@test.com")
|
||||
)
|
||||
|
||||
result = await UserFacetCursorCrud.cursor_paginate(
|
||||
db_session, include_facets=False, schema=UserRead
|
||||
)
|
||||
|
||||
assert result.filter_attributes is None
|
||||
|
||||
|
||||
class TestFacetsRespectFilters:
|
||||
"""Facets reflect the active filter conditions."""
|
||||
|
||||
@@ -630,6 +740,28 @@ class TestFacetsRespectFilters:
|
||||
assert result.filter_attributes is not None
|
||||
assert result.filter_attributes["username"] == ["alice"]
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_array_facet_respects_unrelated_filter(
|
||||
self, db_session: AsyncSession
|
||||
):
|
||||
"""An ARRAY facet is scoped by a filter on a different column (not self-collapse)."""
|
||||
ArticleFacetCrud = CrudFactory(Article, facet_fields=[Article.labels])
|
||||
await ArticleCrud.create(
|
||||
db_session, ArticleCreate(title="Post 1", labels=["python", "fastapi"])
|
||||
)
|
||||
await ArticleCrud.create(
|
||||
db_session, ArticleCreate(title="Post 2", labels=["rust", "axum"])
|
||||
)
|
||||
|
||||
result = await ArticleFacetCrud.offset_paginate(
|
||||
db_session,
|
||||
filters=[Article.title == "Post 1"],
|
||||
schema=ArticleRead,
|
||||
)
|
||||
|
||||
assert result.filter_attributes is not None
|
||||
assert result.filter_attributes["labels"] == ["fastapi", "python"]
|
||||
|
||||
|
||||
class TestFacetsRelationship:
|
||||
"""Facets on relationship columns via tuple syntax."""
|
||||
@@ -785,8 +917,8 @@ class TestFilterBy:
|
||||
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0].username == "alice"
|
||||
# facet also scoped to the filter
|
||||
assert result.filter_attributes == {"username": ["alice"]}
|
||||
# facet excludes its own filter_by condition, so it isn't collapsed
|
||||
assert result.filter_attributes == {"username": ["alice", "bob"]}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_list_filter_produces_in_clause(self, db_session: AsyncSession):
|
||||
@@ -924,7 +1056,7 @@ class TestFilterBy:
|
||||
|
||||
assert len(result.data) == 1
|
||||
assert result.data[0].username == "alice"
|
||||
assert result.filter_attributes == {"username": ["alice"]}
|
||||
assert result.filter_attributes == {"username": ["alice", "bob"]}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_basemodel_filter_by_offset_paginate(self, db_session: AsyncSession):
|
||||
@@ -1085,8 +1217,10 @@ class TestFilterBy:
|
||||
assert result.pagination.total_count == 2
|
||||
titles = {a.title for a in result.data}
|
||||
assert titles == {"Post 1", "Post 3"}
|
||||
# facet returns individual unnested values, not whole arrays
|
||||
assert result.filter_attributes == {"labels": ["django", "fastapi", "python"]}
|
||||
# facet excludes its own filter_by condition (not collapsed to matching rows)
|
||||
assert result.filter_attributes == {
|
||||
"labels": ["axum", "django", "fastapi", "python", "rust"]
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_array_overlap_list_value(self, db_session: AsyncSession):
|
||||
@@ -2179,7 +2313,12 @@ class TestOffsetPaginateParamsSchema:
|
||||
include_total=False, search=False, filter=False, order=False
|
||||
)
|
||||
result = await dep(page=2, items_per_page=10)
|
||||
assert result == {"page": 2, "items_per_page": 10, "include_total": False}
|
||||
assert result == {
|
||||
"page": 2,
|
||||
"items_per_page": 10,
|
||||
"include_total": False,
|
||||
"include_facets": True,
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_integrates_with_offset_paginate(self, db_session: AsyncSession):
|
||||
@@ -2290,7 +2429,11 @@ class TestCursorPaginateParamsSchema:
|
||||
search=False, filter=False, order=False
|
||||
)
|
||||
result = await dep(cursor=None, items_per_page=5)
|
||||
assert result == {"cursor": None, "items_per_page": 5}
|
||||
assert result == {
|
||||
"cursor": None,
|
||||
"items_per_page": 5,
|
||||
"include_facets": True,
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_integrates_with_cursor_paginate(self, db_session: AsyncSession):
|
||||
@@ -2399,6 +2542,7 @@ class TestPaginateParamsSchema:
|
||||
"cursor": None,
|
||||
"items_per_page": 10,
|
||||
"include_total": True,
|
||||
"include_facets": True,
|
||||
}
|
||||
|
||||
@pytest.mark.anyio
|
||||
|
||||
@@ -199,8 +199,9 @@ class TestOffsetPagination:
|
||||
resp = await client.get("/articles/offset?status=published")
|
||||
|
||||
body = resp.json()
|
||||
# draft is filtered out → should not appear in filter_attributes
|
||||
assert "draft" not in body["filter_attributes"]["status"]
|
||||
# a facet excludes its own filter_by condition, so filtering by
|
||||
# status=published still shows every status the facet offers
|
||||
assert "draft" in body["filter_attributes"]["status"]
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_search_and_filter_combined(self, client: AsyncClient, ex_db_session):
|
||||
|
||||
+39
-40
@@ -25,13 +25,11 @@ from fastapi_toolsets.models.watched import (
|
||||
_SESSION_CREATES,
|
||||
_SESSION_DELETES,
|
||||
_SESSION_UPDATES,
|
||||
_WATCHED_MODELS,
|
||||
EventSession,
|
||||
_after_flush,
|
||||
_after_rollback,
|
||||
_get_watched_fields,
|
||||
_invalidate_caches,
|
||||
_is_watched,
|
||||
_snapshot_column_attrs,
|
||||
_upsert_changes,
|
||||
)
|
||||
@@ -658,22 +656,6 @@ class TestWatchInheritance:
|
||||
assert "other" in _watch_inherit_events[0]["changes"]
|
||||
|
||||
|
||||
class TestIsWatched:
|
||||
def test_watched_model_is_watched(self):
|
||||
"""_is_watched returns True for models with registered handlers."""
|
||||
obj = WatchedModel(status="x", other="y")
|
||||
assert _is_watched(obj) is True
|
||||
|
||||
def test_non_watched_model_is_not_watched(self):
|
||||
"""_is_watched returns False for models without registered handlers."""
|
||||
assert _is_watched(object()) is False
|
||||
|
||||
def test_subclass_of_watched_model_is_watched(self):
|
||||
"""_is_watched returns True for subclasses of watched models (via MRO)."""
|
||||
dog = PolyDog(name="Rex")
|
||||
assert _is_watched(dog) is True
|
||||
|
||||
|
||||
class TestUpsertChanges:
|
||||
def test_inserts_new_entry(self):
|
||||
"""New key is inserted with the full changes dict."""
|
||||
@@ -715,7 +697,10 @@ class TestAfterFlush:
|
||||
"""New watched objects are added to _SESSION_CREATES."""
|
||||
obj = object()
|
||||
session = SimpleNamespace(new=[obj], deleted=[], dirty=[], info={})
|
||||
with patch("fastapi_toolsets.models.watched._is_watched", return_value=True):
|
||||
with patch(
|
||||
"fastapi_toolsets.models.watched._get_handlers",
|
||||
return_value=[lambda *a: None],
|
||||
):
|
||||
_after_flush(session, None)
|
||||
assert session.info[_SESSION_CREATES] == [obj]
|
||||
|
||||
@@ -731,7 +716,10 @@ class TestAfterFlush:
|
||||
obj = object()
|
||||
session = SimpleNamespace(new=[], deleted=[obj], dirty=[], info={})
|
||||
with (
|
||||
patch("fastapi_toolsets.models.watched._is_watched", return_value=True),
|
||||
patch(
|
||||
"fastapi_toolsets.models.watched._get_handlers",
|
||||
return_value=[lambda *a: None],
|
||||
),
|
||||
patch(
|
||||
"fastapi_toolsets.models.watched._snapshot_column_attrs",
|
||||
return_value={"id": 1},
|
||||
@@ -1023,28 +1011,19 @@ class TestEventCallbacks:
|
||||
await other.commit()
|
||||
await engine.dispose()
|
||||
|
||||
real_get = mixin_session.get
|
||||
real_refresh = mixin_session.refresh
|
||||
real_batch_reload = _watched_module._batch_reload
|
||||
|
||||
def _matches_doomed(pk):
|
||||
return pk == doomed_id or (isinstance(pk, tuple) and pk[0] == doomed_id)
|
||||
|
||||
async def racing_get(model, pk, *args, **kwargs):
|
||||
if _matches_doomed(pk):
|
||||
async def racing_batch_reload(session, model, pk_tuples):
|
||||
if any(pk[0] == doomed_id for pk in pk_tuples):
|
||||
await kill_doomed_row_once()
|
||||
return await real_get(model, pk, *args, **kwargs)
|
||||
return await real_batch_reload(session, model, pk_tuples)
|
||||
|
||||
async def racing_refresh(obj, *args, **kwargs):
|
||||
if getattr(obj, "id", None) == doomed_id:
|
||||
await kill_doomed_row_once()
|
||||
return await real_refresh(obj, *args, **kwargs)
|
||||
|
||||
# Patch both possible reload mechanisms (session.get / session.refresh)
|
||||
# so this test still exercises the race regardless of which one
|
||||
# EventSession.commit() uses internally to pick up server defaults.
|
||||
mixin_session.get = racing_get
|
||||
mixin_session.refresh = racing_refresh
|
||||
with patch.object(_watched_module._logger, "error") as mock_error:
|
||||
# Patch the batched reload EventSession.commit() uses to pick up
|
||||
# server defaults, so this test still exercises the race.
|
||||
with (
|
||||
patch.object(_watched_module, "_batch_reload", racing_batch_reload),
|
||||
patch.object(_watched_module._logger, "error") as mock_error,
|
||||
):
|
||||
await mixin_session.commit()
|
||||
mock_error.assert_not_called()
|
||||
|
||||
@@ -1052,6 +1031,27 @@ class TestEventCallbacks:
|
||||
created_ids = {e["obj_id"] for e in _test_events if e["event"] == "create"}
|
||||
assert created_ids == {keep.id, doomed_id}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_batch_reload_exception_is_logged_and_dispatch_continues(
|
||||
self, mixin_session
|
||||
):
|
||||
"""A batched-reload failure is logged; CREATE handlers still fire."""
|
||||
obj = WatchedModel(status="active", other="x")
|
||||
mixin_session.add(obj)
|
||||
|
||||
async def failing_batch_reload(session, model, pk_tuples):
|
||||
raise RuntimeError("reload failed")
|
||||
|
||||
with (
|
||||
patch.object(_watched_module, "_batch_reload", failing_batch_reload),
|
||||
patch.object(_watched_module._logger, "error") as mock_error,
|
||||
):
|
||||
await mixin_session.commit()
|
||||
|
||||
mock_error.assert_called_once()
|
||||
creates = [e for e in _test_events if e["event"] == "create"]
|
||||
assert len(creates) == 1
|
||||
|
||||
|
||||
class TestTransientObject:
|
||||
"""Create + delete within the same transaction should fire no events."""
|
||||
@@ -1421,7 +1421,6 @@ class TestListensFor:
|
||||
for key in list(_EVENT_HANDLERS):
|
||||
if key[0] is ListenerModel:
|
||||
del _EVENT_HANDLERS[key]
|
||||
_WATCHED_MODELS.discard(ListenerModel)
|
||||
_invalidate_caches()
|
||||
|
||||
@pytest.mark.anyio
|
||||
|
||||
Reference in New Issue
Block a user