Compare commits

...
4 Commits
9 changed files with 513 additions and 188 deletions
+10
View File
@@ -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." !!! 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). 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 ## Sorting
!!! info "Added in `v1.3`" !!! info "Added in `v1.3`"
+103 -40
View File
@@ -44,6 +44,7 @@ from ..types import (
) )
from .search import ( from .search import (
SearchConfig, SearchConfig,
apply_search_joins,
build_facets, build_facets,
build_filter_by, build_filter_by,
build_search_filters, build_search_filters,
@@ -51,7 +52,6 @@ from .search import (
search_field_keys, search_field_keys,
) )
_ForUpdateMode: TypeAlias = bool | Literal["nowait", "skip_locked"] _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 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]): class AsyncCrud(Generic[ModelType]):
"""Generic async CRUD operations for SQLAlchemy models. """Generic async CRUD operations for SQLAlchemy models.
@@ -184,17 +173,30 @@ class AsyncCrud(Generic[ModelType]):
return cls.default_load_options return cls.default_load_options
@classmethod @classmethod
async def _reload_with_options( def _capture_pk_values(cls: type[Self], instance: ModelType) -> dict[str, Any]:
cls: type[Self], session: AsyncSession, instance: ModelType """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: ) -> ModelType:
"""Re-query instance by PK with default_load_options applied.""" """Re-query by previously captured PK values, with default_load_options applied."""
mapper = cls.model.__mapper__ # Only called when cls.default_load_options is set (see call sites).
pk_filters = [ pk_filters = [
getattr(cls.model, cast(str, col.key)) getattr(cls.model, key) == value for key, value in pk_values.items()
== getattr(instance, cast(str, col.key))
for col in mapper.primary_key
] ]
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 @classmethod
async def _resolve_m2m( async def _resolve_m2m(
@@ -265,12 +267,12 @@ class AsyncCrud(Generic[ModelType]):
cls: type[Self], cls: type[Self],
filter_by: dict[str, Any] | BaseModel | None, filter_by: dict[str, Any] | BaseModel | None,
facet_fields: Sequence[FacetFieldType] | None, facet_fields: Sequence[FacetFieldType] | None,
) -> tuple[list[Any], list[Any]]: ) -> tuple[dict[str, Any], list[Any]]:
"""Normalize filter_by and return (filters, joins) to apply to the query.""" """Normalize filter_by and return ({facet_key: filter}, joins) to apply to the query."""
if isinstance(filter_by, BaseModel): if isinstance(filter_by, BaseModel):
filter_by = filter_by.model_dump(exclude_none=True) filter_by = filter_by.model_dump(exclude_none=True)
if not filter_by: if not filter_by:
return [], [] return {}, []
resolved = cls._resolve_facet_fields(facet_fields) resolved = cls._resolve_facet_fields(facet_fields)
return build_filter_by(filter_by, resolved or []) return build_filter_by(filter_by, resolved or [])
@@ -281,8 +283,13 @@ class AsyncCrud(Generic[ModelType]):
facet_fields: Sequence[FacetFieldType] | None, facet_fields: Sequence[FacetFieldType] | None,
filters: list[Any], filters: list[Any],
search_joins: list[Any], search_joins: list[Any],
*,
include_facets: bool = True,
own_filters: dict[str, Any] | None = None,
) -> dict[str, list[Any]] | 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) resolved = cls._resolve_facet_fields(facet_fields)
if not resolved: if not resolved:
return None return None
@@ -292,6 +299,7 @@ class AsyncCrud(Generic[ModelType]):
resolved, resolved,
base_filters=filters, base_filters=filters,
base_joins=search_joins, base_joins=search_joins,
own_filters=own_filters,
) )
@classmethod @classmethod
@@ -475,6 +483,7 @@ class AsyncCrud(Generic[ModelType]):
default_page_size: int = 20, default_page_size: int = 20,
max_page_size: int = 100, max_page_size: int = 100,
include_total: bool = True, include_total: bool = True,
include_facets: bool = True,
search: bool = True, search: bool = True,
filter: bool = True, filter: bool = True,
order: bool = True, order: bool = True,
@@ -490,6 +499,7 @@ class AsyncCrud(Generic[ModelType]):
default_page_size: Default ``items_per_page`` value. default_page_size: Default ``items_per_page`` value.
max_page_size: Maximum ``items_per_page`` value. max_page_size: Maximum ``items_per_page`` value.
include_total: Whether to include total count (not a query param). 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. search: Enable search query parameters.
filter: Enable facet filter query parameters. filter: Enable facet filter query parameters.
order: Enable order query parameters. order: Enable order query parameters.
@@ -519,7 +529,10 @@ class AsyncCrud(Generic[ModelType]):
] ]
return cls._build_paginate_params( return cls._build_paginate_params(
pagination_params=pagination_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", dep_name=f"{cls.model.__name__}OffsetPaginateParams",
search=search, search=search,
filter=filter, filter=filter,
@@ -537,6 +550,7 @@ class AsyncCrud(Generic[ModelType]):
*, *,
default_page_size: int = 20, default_page_size: int = 20,
max_page_size: int = 100, max_page_size: int = 100,
include_facets: bool = True,
search: bool = True, search: bool = True,
filter: bool = True, filter: bool = True,
order: bool = True, order: bool = True,
@@ -551,6 +565,7 @@ class AsyncCrud(Generic[ModelType]):
Args: Args:
default_page_size: Default ``items_per_page`` value. default_page_size: Default ``items_per_page`` value.
max_page_size: Maximum ``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. search: Enable search query parameters.
filter: Enable facet filter query parameters. filter: Enable facet filter query parameters.
order: Enable order query parameters. order: Enable order query parameters.
@@ -582,7 +597,7 @@ class AsyncCrud(Generic[ModelType]):
] ]
return cls._build_paginate_params( return cls._build_paginate_params(
pagination_params=pagination_params, pagination_params=pagination_params,
pagination_fixed={}, pagination_fixed={"include_facets": include_facets},
dep_name=f"{cls.model.__name__}CursorPaginateParams", dep_name=f"{cls.model.__name__}CursorPaginateParams",
search=search, search=search,
filter=filter, filter=filter,
@@ -602,6 +617,7 @@ class AsyncCrud(Generic[ModelType]):
max_page_size: int = 100, max_page_size: int = 100,
default_pagination_type: PaginationType = PaginationType.OFFSET, default_pagination_type: PaginationType = PaginationType.OFFSET,
include_total: bool = True, include_total: bool = True,
include_facets: bool = True,
search: bool = True, search: bool = True,
filter: bool = True, filter: bool = True,
order: bool = True, order: bool = True,
@@ -618,6 +634,7 @@ class AsyncCrud(Generic[ModelType]):
max_page_size: Maximum ``items_per_page`` value. max_page_size: Maximum ``items_per_page`` value.
default_pagination_type: Default pagination strategy. default_pagination_type: Default pagination strategy.
include_total: Whether to include total count (not a query param). 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. search: Enable search query parameters.
filter: Enable facet filter query parameters. filter: Enable facet filter query parameters.
order: Enable order query parameters. order: Enable order query parameters.
@@ -666,7 +683,10 @@ class AsyncCrud(Generic[ModelType]):
] ]
return cls._build_paginate_params( return cls._build_paginate_params(
pagination_params=pagination_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", dep_name=f"{cls.model.__name__}PaginateParams",
search=search, search=search,
filter=filter, filter=filter,
@@ -730,9 +750,14 @@ class AsyncCrud(Generic[ModelType]):
setattr(db_model, rel_attr, related_instances) setattr(db_model, rel_attr, related_instances)
session.add(db_model) session.add(db_model)
await session.refresh(db_model) pk_values: dict[str, Any] | None = None
if cls.default_load_options: 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) result = cast(ModelType, db_model)
if schema: if schema:
return Response(data=schema.model_validate(result)) 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) m2m_resolved = await cls._resolve_m2m(session, obj, only_set=True)
for rel_attr, related_instances in m2m_resolved.items(): for rel_attr, related_instances in m2m_resolved.items():
setattr(db_model, rel_attr, related_instances) setattr(db_model, rel_attr, related_instances)
await session.refresh(db_model)
pk_values: dict[str, Any] | None = None
if cls.default_load_options: 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: if schema:
return Response(data=schema.model_validate(db_model)) return Response(data=schema.model_validate(db_model))
return db_model return db_model
@@ -1271,6 +1302,7 @@ class AsyncCrud(Generic[ModelType]):
search_column: str | None = None, search_column: str | None = None,
order_fields: Sequence[OrderFieldType] | None = None, order_fields: Sequence[OrderFieldType] | None = None,
facet_fields: Sequence[FacetFieldType] | None = None, facet_fields: Sequence[FacetFieldType] | None = None,
include_facets: bool = True,
filter_by: dict[str, Any] | BaseModel | None = None, filter_by: dict[str, Any] | BaseModel | None = None,
schema: type[BaseModel], schema: type[BaseModel],
) -> OffsetPaginatedResponse[Any]: ) -> OffsetPaginatedResponse[Any]:
@@ -1292,6 +1324,9 @@ class AsyncCrud(Generic[ModelType]):
search_column: Restrict search to a single column key. search_column: Restrict search to a single column key.
order_fields: Fields allowed for sorting (overrides class default). order_fields: Fields allowed for sorting (overrides class default).
facet_fields: Columns to compute distinct values for (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. 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, Keys must match the column.key of a facet field. Scalar → equality,
list → IN clause. Raises InvalidFacetFilterError for unknown keys. list → IN clause. Raises InvalidFacetFilterError for unknown keys.
@@ -1304,7 +1339,6 @@ class AsyncCrud(Generic[ModelType]):
offset = (page - 1) * items_per_page offset = (page - 1) * items_per_page
fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields) fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields)
filters.extend(fb_filters)
# Build search filters # Build search filters
if search: if search:
@@ -1318,6 +1352,11 @@ class AsyncCrud(Generic[ModelType]):
filters.extend(search_filters) filters.extend(search_filters)
search_joins.extend(new_search_joins) 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 # Build query with joins
q = select(cls.model) q = select(cls.model)
@@ -1325,11 +1364,11 @@ class AsyncCrud(Generic[ModelType]):
q = _apply_joins(q, joins, outer_join) q = _apply_joins(q, joins, outer_join)
# Apply search joins (always outer joins for search) # 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) # Apply order joins (relation joins required for order_by field)
if order_joins: if order_joins:
q = _apply_search_joins(q, order_joins) q = apply_search_joins(q, order_joins)
if filters: if filters:
q = q.where(and_(*filters)) q = q.where(and_(*filters))
@@ -1352,7 +1391,7 @@ class AsyncCrud(Generic[ModelType]):
count_q = _apply_joins(count_q, joins, outer_join) count_q = _apply_joins(count_q, joins, outer_join)
# Apply search joins to count query # 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: if filters:
count_q = count_q.where(and_(*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] items: list[Any] = [schema.model_validate(item) for item in raw_items]
filter_attributes = await cls._build_filter_attributes( 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) search_columns = cls._resolve_search_columns(search_fields)
order_columns = cls._resolve_order_columns(order_fields) order_columns = cls._resolve_order_columns(order_fields)
@@ -1408,6 +1452,7 @@ class AsyncCrud(Generic[ModelType]):
search_column: str | None = None, search_column: str | None = None,
order_fields: Sequence[OrderFieldType] | None = None, order_fields: Sequence[OrderFieldType] | None = None,
facet_fields: Sequence[FacetFieldType] | None = None, facet_fields: Sequence[FacetFieldType] | None = None,
include_facets: bool = True,
filter_by: dict[str, Any] | BaseModel | None = None, filter_by: dict[str, Any] | BaseModel | None = None,
schema: type[BaseModel], schema: type[BaseModel],
) -> CursorPaginatedResponse[Any]: ) -> CursorPaginatedResponse[Any]:
@@ -1430,6 +1475,8 @@ class AsyncCrud(Generic[ModelType]):
search_column: Restrict search to a single column key. search_column: Restrict search to a single column key.
order_fields: Fields allowed for sorting (overrides class default). order_fields: Fields allowed for sorting (overrides class default).
facet_fields: Columns to compute distinct values for (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. 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, Keys must match the column.key of a facet field. Scalar → equality,
list → IN clause. Raises InvalidFacetFilterError for unknown keys. list → IN clause. Raises InvalidFacetFilterError for unknown keys.
@@ -1441,7 +1488,6 @@ class AsyncCrud(Generic[ModelType]):
filters = list(filters) if filters else [] filters = list(filters) if filters else []
fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields) fb_filters, search_joins = cls._prepare_filter_by(filter_by, facet_fields)
filters.extend(fb_filters)
if cls.cursor_column is None: if cls.cursor_column is None:
raise ValueError( raise ValueError(
@@ -1473,6 +1519,11 @@ class AsyncCrud(Generic[ModelType]):
filters.extend(search_filters) filters.extend(search_filters)
search_joins.extend(new_search_joins) 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 # Build query
q = select(cls.model) q = select(cls.model)
@@ -1480,11 +1531,11 @@ class AsyncCrud(Generic[ModelType]):
q = _apply_joins(q, joins, outer_join) q = _apply_joins(q, joins, outer_join)
# Apply search joins (always outer joins) # 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) # Apply order joins (relation joins required for order_by field)
if order_joins: if order_joins:
q = _apply_search_joins(q, order_joins) q = apply_search_joins(q, order_joins)
if filters: if filters:
q = q.where(and_(*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] items: list[Any] = [schema.model_validate(item) for item in items_page]
filter_attributes = await cls._build_filter_attributes( 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) search_columns = cls._resolve_search_columns(search_fields)
order_columns = cls._resolve_order_columns(order_fields) order_columns = cls._resolve_order_columns(order_fields)
@@ -1581,6 +1637,7 @@ class AsyncCrud(Generic[ModelType]):
search_column: str | None = ..., search_column: str | None = ...,
order_fields: Sequence[OrderFieldType] | None = ..., order_fields: Sequence[OrderFieldType] | None = ...,
facet_fields: Sequence[FacetFieldType] | None = ..., facet_fields: Sequence[FacetFieldType] | None = ...,
include_facets: bool = ...,
filter_by: dict[str, Any] | BaseModel | None = ..., filter_by: dict[str, Any] | BaseModel | None = ...,
schema: type[BaseModel], schema: type[BaseModel],
) -> OffsetPaginatedResponse[Any]: ... ) -> OffsetPaginatedResponse[Any]: ...
@@ -1607,6 +1664,7 @@ class AsyncCrud(Generic[ModelType]):
search_column: str | None = ..., search_column: str | None = ...,
order_fields: Sequence[OrderFieldType] | None = ..., order_fields: Sequence[OrderFieldType] | None = ...,
facet_fields: Sequence[FacetFieldType] | None = ..., facet_fields: Sequence[FacetFieldType] | None = ...,
include_facets: bool = ...,
filter_by: dict[str, Any] | BaseModel | None = ..., filter_by: dict[str, Any] | BaseModel | None = ...,
schema: type[BaseModel], schema: type[BaseModel],
) -> CursorPaginatedResponse[Any]: ... ) -> CursorPaginatedResponse[Any]: ...
@@ -1632,6 +1690,7 @@ class AsyncCrud(Generic[ModelType]):
search_column: str | None = None, search_column: str | None = None,
order_fields: Sequence[OrderFieldType] | None = None, order_fields: Sequence[OrderFieldType] | None = None,
facet_fields: Sequence[FacetFieldType] | None = None, facet_fields: Sequence[FacetFieldType] | None = None,
include_facets: bool = True,
filter_by: dict[str, Any] | BaseModel | None = None, filter_by: dict[str, Any] | BaseModel | None = None,
schema: type[BaseModel], schema: type[BaseModel],
) -> OffsetPaginatedResponse[Any] | CursorPaginatedResponse[Any]: ) -> OffsetPaginatedResponse[Any] | CursorPaginatedResponse[Any]:
@@ -1662,6 +1721,8 @@ class AsyncCrud(Generic[ModelType]):
order_fields: Fields allowed for sorting (overrides class default). order_fields: Fields allowed for sorting (overrides class default).
facet_fields: Columns to compute distinct values for (overrides facet_fields: Columns to compute distinct values for (overrides
class default). 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 filter_by: Dict of ``{column_key: value}`` to filter by declared
facet fields. Keys must match the ``column.key`` of a facet facet fields. Keys must match the ``column.key`` of a facet
field. Scalar → equality, list → IN clause. Raises field. Scalar → equality, list → IN clause. Raises
@@ -1692,6 +1753,7 @@ class AsyncCrud(Generic[ModelType]):
search_column=search_column, search_column=search_column,
order_fields=order_fields, order_fields=order_fields,
facet_fields=facet_fields, facet_fields=facet_fields,
include_facets=include_facets,
filter_by=filter_by, filter_by=filter_by,
schema=schema, schema=schema,
) )
@@ -1714,6 +1776,7 @@ class AsyncCrud(Generic[ModelType]):
search_column=search_column, search_column=search_column,
order_fields=order_fields, order_fields=order_fields,
facet_fields=facet_fields, facet_fields=facet_fields,
include_facets=include_facets,
filter_by=filter_by, filter_by=filter_by,
schema=schema, schema=schema,
) )
+87 -60
View File
@@ -1,12 +1,12 @@
"""Search utilities for AsyncCrud.""" """Search utilities for AsyncCrud."""
import asyncio
import functools import functools
from collections.abc import Sequence from collections.abc import Sequence
from dataclasses import dataclass, replace from dataclasses import dataclass, replace
from typing import TYPE_CHECKING, Any, Literal 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.ext.asyncio import AsyncSession
from sqlalchemy.orm import DeclarativeBase from sqlalchemy.orm import DeclarativeBase
from sqlalchemy.orm.attributes import InstrumentedAttribute from sqlalchemy.orm.attributes import InstrumentedAttribute
@@ -159,8 +159,11 @@ def build_search_filters(
else: else:
column = field column = field
# Build the filter (cast to String for non-text columns) # Build the filter (cast to String only when needed, to preserve
column_as_string = column.cast(String) # 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: if config.case_sensitive:
filters.append(column_as_string.like(f"%{query}%")) filters.append(column_as_string.like(f"%{query}%"))
else: else:
@@ -181,6 +184,21 @@ def search_field_keys(fields: Sequence[SearchFieldType]) -> list[str]:
return facet_keys(fields) 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]: def facet_keys(facet_fields: Sequence[FacetFieldType]) -> list[str]:
"""Return a key for each facet field. """Return a key for each facet field.
@@ -207,6 +225,7 @@ async def build_facets(
*, *,
base_filters: "list[ColumnElement[bool]] | None" = None, base_filters: "list[ColumnElement[bool]] | None" = None,
base_joins: list[InstrumentedAttribute[Any]] | None = None, base_joins: list[InstrumentedAttribute[Any]] | None = None,
own_filters: "dict[str, ColumnElement[bool]] | None" = None,
) -> dict[str, list[Any]]: ) -> dict[str, list[Any]]:
"""Return distinct values for each facet field, respecting current filters. """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 facet_fields: Columns or relationship tuples to facet on
base_filters: Filter conditions already applied to the main query (search + caller filters) base_filters: Filter conditions already applied to the main query (search + caller filters)
base_joins: Relationship joins already applied to the main query 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: Returns:
Dict mapping column key to sorted list of distinct non-None values 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) 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): if isinstance(field, tuple):
# Relationship chain: (User.role, Role.name) — last element is the column # Relationship chain: (User.role, Role.name) — last element is the column
rels = field[:-1] rels = field[:-1]
@@ -235,51 +263,48 @@ async def build_facets(
col_type = column.property.columns[0].type col_type = column.property.columns[0].type
is_array = isinstance(col_type, ARRAY) is_array = isinstance(col_type, ARRAY)
enum_classes[key] = getattr(col_type, "enum_class", None)
if is_array: filters = [
unnested = func.unnest(column).label(column.key) *(base_filters or []),
q = select(unnested).select_from(model).distinct() *(f for k, f in own_filters.items() if k != key),
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
] ]
return key, values joins = [*(base_joins or []), *rels]
pairs = await asyncio.gather( if is_array:
*[_query_facet(f, k) for f, k in zip(facet_fields, keys)] 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) _EQUALITY_TYPES = (String, Integer, Numeric, Date, DateTime, Time, Enum, Uuid)
@@ -301,7 +326,7 @@ def _coerce_bool(value: Any) -> bool:
def build_filter_by( def build_filter_by(
filter_by: dict[str, Any], filter_by: dict[str, Any],
facet_fields: Sequence[FacetFieldType], 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. """Translate a {column_key: value} dict into SQLAlchemy filter conditions.
Args: Args:
@@ -309,7 +334,9 @@ def build_filter_by(
facet_fields: Declared facet fields to validate keys against facet_fields: Declared facet fields to validate keys against
Returns: 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: Raises:
InvalidFacetFilterError: If a key in filter_by is not a declared facet field 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) index[key] = (column, rels)
valid_keys = set(index) valid_keys = set(index)
filters: list[ColumnElement[bool]] = [] filters: dict[str, ColumnElement[bool]] = {}
joins: list[InstrumentedAttribute[Any]] = [] joins: list[InstrumentedAttribute[Any]] = []
added_join_keys: set[str] = set() added_join_keys: set[str] = set()
@@ -347,14 +374,14 @@ def build_filter_by(
if isinstance(col_type, Boolean): if isinstance(col_type, Boolean):
coerce = _coerce_bool coerce = _coerce_bool
if isinstance(value, list): 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: else:
filters.append(column == coerce(value)) filters[key] = column == coerce(value)
elif isinstance(col_type, ARRAY): elif isinstance(col_type, ARRAY):
if isinstance(value, list): if isinstance(value, list):
filters.append(column.overlap(value)) filters[key] = column.overlap(value)
else: else:
filters.append(column.any(value)) filters[key] = column.any(value)
elif isinstance(col_type, Enum): elif isinstance(col_type, Enum):
enum_class = col_type.enum_class enum_class = col_type.enum_class
if enum_class is not None: if enum_class is not None:
@@ -365,19 +392,19 @@ def build_filter_by(
return enum_class[v] # lookup by name: "PENDING", "RED" return enum_class[v] # lookup by name: "PENDING", "RED"
if isinstance(value, list): 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: else:
filters.append(column == _coerce_enum(value)) filters[key] = column == _coerce_enum(value)
else: # pragma: no cover else: # pragma: no cover
if isinstance(value, list): if isinstance(value, list):
filters.append(column.in_(value)) filters[key] = column.in_(value)
else: else:
filters.append(column == value) filters[key] = column == value
elif isinstance(col_type, _EQUALITY_TYPES): elif isinstance(col_type, _EQUALITY_TYPES):
if isinstance(value, list): if isinstance(value, list):
filters.append(column.in_(value)) filters[key] = column.in_(value)
else: else:
filters.append(column == value) filters[key] = column == value
else: else:
raise UnsupportedFacetTypeError(key, type(col_type).__name__) raise UnsupportedFacetTypeError(key, type(col_type).__name__)
+44 -35
View File
@@ -5,7 +5,7 @@ from collections.abc import Callable
from enum import Enum from enum import Enum
from typing import Any from typing import Any
from sqlalchemy import event from sqlalchemy import event, select, tuple_
from sqlalchemy import inspect as sa_inspect from sqlalchemy import inspect as sa_inspect
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm.attributes import set_committed_value as _sa_set_committed_value 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" _SESSION_UPDATES = "_ft_updates"
_DEFERRED_STRATEGY_KEY = (("deferred", True), ("instrument", True)) _DEFERRED_STRATEGY_KEY = (("deferred", True), ("instrument", True))
_EVENT_HANDLERS: dict[tuple[type, ModelEvent], list[Callable[..., Any]]] = {} _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]]] = {} _HANDLER_CACHE: dict[tuple[type, ModelEvent], list[Callable[..., Any]]] = {}
def _invalidate_caches() -> None: def _invalidate_caches() -> None:
"""Clear lookup caches after handler registration.""" """Clear lookup caches after handler registration."""
_WATCHED_CACHE.clear()
_HANDLER_CACHE.clear() _HANDLER_CACHE.clear()
@@ -56,24 +53,12 @@ def listens_for(
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]: def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
for ev in evs: for ev in evs:
_EVENT_HANDLERS.setdefault((model_class, ev), []).append(fn) _EVENT_HANDLERS.setdefault((model_class, ev), []).append(fn)
_WATCHED_MODELS.add(model_class)
_invalidate_caches() _invalidate_caches()
return fn return fn
return decorator 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]]: def _get_handlers(cls: type, ev: ModelEvent) -> list[Callable[..., Any]]:
"""Return registered handlers for *cls* and *ev*, walking the MRO.""" """Return registered handlers for *cls* and *ev*, walking the MRO."""
key = (cls, ev) key = (cls, ev)
@@ -144,18 +129,18 @@ def _upsert_changes(
def _after_flush(session: Any, flush_context: Any) -> None: def _after_flush(session: Any, flush_context: Any) -> None:
# New objects: capture reference. Attributes will be refreshed after commit. # New objects: capture reference. Attributes will be refreshed after commit.
for obj in session.new: for obj in session.new:
if _is_watched(obj): if _get_handlers(type(obj), ModelEvent.CREATE):
session.info.setdefault(_SESSION_CREATES, []).append(obj) session.info.setdefault(_SESSION_CREATES, []).append(obj)
# Deleted objects: snapshot now while attributes are still loaded. # Deleted objects: snapshot now while attributes are still loaded.
for obj in session.deleted: for obj in session.deleted:
if _is_watched(obj): if _get_handlers(type(obj), ModelEvent.DELETE):
snapshot = _snapshot_column_attrs(obj) snapshot = _snapshot_column_attrs(obj)
session.info.setdefault(_SESSION_DELETES, []).append((obj, snapshot)) session.info.setdefault(_SESSION_DELETES, []).append((obj, snapshot))
# Dirty objects: read old/new from SQLAlchemy attribute history. # Dirty objects: read old/new from SQLAlchemy attribute history.
for obj in session.dirty: for obj in session.dirty:
if not _is_watched(obj): if not _get_handlers(type(obj), ModelEvent.UPDATE):
continue continue
watched = _get_watched_fields(type(obj)) watched = _get_watched_fields(type(obj))
@@ -204,9 +189,18 @@ async def _invoke_callback(
await result await result
async def _reload_if_present(session: AsyncSession, obj: Any, state: Any) -> None: async def _batch_reload(
"""Re-populate *obj* from the DB if its row still exists.""" session: AsyncSession, model: type, pk_tuples: list[tuple[Any, ...]]
await session.get(type(obj), state.key[1], populate_existing=True) ) -> 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): 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 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: for obj in creates:
try:
state = sa_inspect(obj, raiseerr=False) state = sa_inspect(obj, raiseerr=False)
if ( if state is None or state.detached or state.transient: # pragma: no cover
state is None or state.detached or state.transient
): # pragma: no cover
continue 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): for handler in _get_handlers(type(obj), ModelEvent.CREATE):
await _invoke_callback(handler, obj, ModelEvent.CREATE, None) await _invoke_callback(handler, obj, ModelEvent.CREATE, None)
except Exception as exc: except Exception as exc:
@@ -275,14 +290,8 @@ class EventSession(AsyncSession):
_logger.error(_CALLBACK_ERROR_MSG, exc_info=exc) _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
# Dispatch UPDATE callbacks. # Dispatch UPDATE callbacks.
for obj, changes in field_changes.values(): for obj, changes in update_items:
try: 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): for handler in _get_handlers(type(obj), ModelEvent.UPDATE):
await _invoke_callback(handler, obj, ModelEvent.UPDATE, changes) await _invoke_callback(handler, obj, ModelEvent.UPDATE, changes)
except Exception as exc: except Exception as exc:
+23
View File
@@ -476,3 +476,26 @@ async def db_session(engine):
# Drop tables after test # Drop tables after test
async with engine.begin() as conn: async with engine.begin() as conn:
await conn.run_sync(Base.metadata.drop_all) 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)
+49
View File
@@ -417,6 +417,55 @@ class TestDefaultLoadOptionsIntegration:
assert updated.role is not None assert updated.role is not None
assert updated.role.name == "admin" 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 @pytest.mark.anyio
async def test_load_options_overrides_default_load_options( async def test_load_options_overrides_default_load_options(
self, db_session: AsyncSession self, db_session: AsyncSession
+151 -7
View File
@@ -372,6 +372,22 @@ class TestBuildSearchFilters:
assert len(joins) == 1 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: class TestSearchConfig:
"""Tests for SearchConfig options.""" """Tests for SearchConfig options."""
@@ -531,6 +547,15 @@ class TestFacetsNotSet:
assert result.filter_attributes is None 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: class TestFacetsDirectColumn:
"""Facets on direct model columns.""" """Facets on direct model columns."""
@@ -606,6 +631,91 @@ class TestFacetsDirectColumn:
assert "username" not in result.filter_attributes 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: class TestFacetsRespectFilters:
"""Facets reflect the active filter conditions.""" """Facets reflect the active filter conditions."""
@@ -630,6 +740,28 @@ class TestFacetsRespectFilters:
assert result.filter_attributes is not None assert result.filter_attributes is not None
assert result.filter_attributes["username"] == ["alice"] 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: class TestFacetsRelationship:
"""Facets on relationship columns via tuple syntax.""" """Facets on relationship columns via tuple syntax."""
@@ -785,8 +917,8 @@ class TestFilterBy:
assert len(result.data) == 1 assert len(result.data) == 1
assert result.data[0].username == "alice" assert result.data[0].username == "alice"
# facet also scoped to the filter # facet excludes its own filter_by condition, so it isn't collapsed
assert result.filter_attributes == {"username": ["alice"]} assert result.filter_attributes == {"username": ["alice", "bob"]}
@pytest.mark.anyio @pytest.mark.anyio
async def test_list_filter_produces_in_clause(self, db_session: AsyncSession): async def test_list_filter_produces_in_clause(self, db_session: AsyncSession):
@@ -924,7 +1056,7 @@ class TestFilterBy:
assert len(result.data) == 1 assert len(result.data) == 1
assert result.data[0].username == "alice" assert result.data[0].username == "alice"
assert result.filter_attributes == {"username": ["alice"]} assert result.filter_attributes == {"username": ["alice", "bob"]}
@pytest.mark.anyio @pytest.mark.anyio
async def test_basemodel_filter_by_offset_paginate(self, db_session: AsyncSession): async def test_basemodel_filter_by_offset_paginate(self, db_session: AsyncSession):
@@ -1085,8 +1217,10 @@ class TestFilterBy:
assert result.pagination.total_count == 2 assert result.pagination.total_count == 2
titles = {a.title for a in result.data} titles = {a.title for a in result.data}
assert titles == {"Post 1", "Post 3"} assert titles == {"Post 1", "Post 3"}
# facet returns individual unnested values, not whole arrays # facet excludes its own filter_by condition (not collapsed to matching rows)
assert result.filter_attributes == {"labels": ["django", "fastapi", "python"]} assert result.filter_attributes == {
"labels": ["axum", "django", "fastapi", "python", "rust"]
}
@pytest.mark.anyio @pytest.mark.anyio
async def test_array_overlap_list_value(self, db_session: AsyncSession): 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 include_total=False, search=False, filter=False, order=False
) )
result = await dep(page=2, items_per_page=10) 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 @pytest.mark.anyio
async def test_integrates_with_offset_paginate(self, db_session: AsyncSession): async def test_integrates_with_offset_paginate(self, db_session: AsyncSession):
@@ -2290,7 +2429,11 @@ class TestCursorPaginateParamsSchema:
search=False, filter=False, order=False search=False, filter=False, order=False
) )
result = await dep(cursor=None, items_per_page=5) 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 @pytest.mark.anyio
async def test_integrates_with_cursor_paginate(self, db_session: AsyncSession): async def test_integrates_with_cursor_paginate(self, db_session: AsyncSession):
@@ -2399,6 +2542,7 @@ class TestPaginateParamsSchema:
"cursor": None, "cursor": None,
"items_per_page": 10, "items_per_page": 10,
"include_total": True, "include_total": True,
"include_facets": True,
} }
@pytest.mark.anyio @pytest.mark.anyio
+3 -2
View File
@@ -199,8 +199,9 @@ class TestOffsetPagination:
resp = await client.get("/articles/offset?status=published") resp = await client.get("/articles/offset?status=published")
body = resp.json() body = resp.json()
# draft is filtered out → should not appear in filter_attributes # a facet excludes its own filter_by condition, so filtering by
assert "draft" not in body["filter_attributes"]["status"] # status=published still shows every status the facet offers
assert "draft" in body["filter_attributes"]["status"]
@pytest.mark.anyio @pytest.mark.anyio
async def test_search_and_filter_combined(self, client: AsyncClient, ex_db_session): async def test_search_and_filter_combined(self, client: AsyncClient, ex_db_session):
+39 -40
View File
@@ -25,13 +25,11 @@ from fastapi_toolsets.models.watched import (
_SESSION_CREATES, _SESSION_CREATES,
_SESSION_DELETES, _SESSION_DELETES,
_SESSION_UPDATES, _SESSION_UPDATES,
_WATCHED_MODELS,
EventSession, EventSession,
_after_flush, _after_flush,
_after_rollback, _after_rollback,
_get_watched_fields, _get_watched_fields,
_invalidate_caches, _invalidate_caches,
_is_watched,
_snapshot_column_attrs, _snapshot_column_attrs,
_upsert_changes, _upsert_changes,
) )
@@ -658,22 +656,6 @@ class TestWatchInheritance:
assert "other" in _watch_inherit_events[0]["changes"] 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: class TestUpsertChanges:
def test_inserts_new_entry(self): def test_inserts_new_entry(self):
"""New key is inserted with the full changes dict.""" """New key is inserted with the full changes dict."""
@@ -715,7 +697,10 @@ class TestAfterFlush:
"""New watched objects are added to _SESSION_CREATES.""" """New watched objects are added to _SESSION_CREATES."""
obj = object() obj = object()
session = SimpleNamespace(new=[obj], deleted=[], dirty=[], info={}) 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) _after_flush(session, None)
assert session.info[_SESSION_CREATES] == [obj] assert session.info[_SESSION_CREATES] == [obj]
@@ -731,7 +716,10 @@ class TestAfterFlush:
obj = object() obj = object()
session = SimpleNamespace(new=[], deleted=[obj], dirty=[], info={}) session = SimpleNamespace(new=[], deleted=[obj], dirty=[], info={})
with ( with (
patch("fastapi_toolsets.models.watched._is_watched", return_value=True), patch(
"fastapi_toolsets.models.watched._get_handlers",
return_value=[lambda *a: None],
),
patch( patch(
"fastapi_toolsets.models.watched._snapshot_column_attrs", "fastapi_toolsets.models.watched._snapshot_column_attrs",
return_value={"id": 1}, return_value={"id": 1},
@@ -1023,28 +1011,19 @@ class TestEventCallbacks:
await other.commit() await other.commit()
await engine.dispose() await engine.dispose()
real_get = mixin_session.get real_batch_reload = _watched_module._batch_reload
real_refresh = mixin_session.refresh
def _matches_doomed(pk): async def racing_batch_reload(session, model, pk_tuples):
return pk == doomed_id or (isinstance(pk, tuple) and pk[0] == doomed_id) if any(pk[0] == doomed_id for pk in pk_tuples):
async def racing_get(model, pk, *args, **kwargs):
if _matches_doomed(pk):
await kill_doomed_row_once() 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): # Patch the batched reload EventSession.commit() uses to pick up
if getattr(obj, "id", None) == doomed_id: # server defaults, so this test still exercises the race.
await kill_doomed_row_once() with (
return await real_refresh(obj, *args, **kwargs) patch.object(_watched_module, "_batch_reload", racing_batch_reload),
patch.object(_watched_module._logger, "error") as mock_error,
# 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:
await mixin_session.commit() await mixin_session.commit()
mock_error.assert_not_called() 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"} created_ids = {e["obj_id"] for e in _test_events if e["event"] == "create"}
assert created_ids == {keep.id, doomed_id} 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: class TestTransientObject:
"""Create + delete within the same transaction should fire no events.""" """Create + delete within the same transaction should fire no events."""
@@ -1421,7 +1421,6 @@ class TestListensFor:
for key in list(_EVENT_HANDLERS): for key in list(_EVENT_HANDLERS):
if key[0] is ListenerModel: if key[0] is ListenerModel:
del _EVENT_HANDLERS[key] del _EVENT_HANDLERS[key]
_WATCHED_MODELS.discard(ListenerModel)
_invalidate_caches() _invalidate_caches()
@pytest.mark.anyio @pytest.mark.anyio