Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions src/strawchemy/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,17 @@
from strawchemy.instance import ModelInstance
from strawchemy.mapper import Strawchemy
from strawchemy.repository.strawberry import StrawchemyAsyncRepository, StrawchemySyncRepository
from strawchemy.schema.filters import (
ArrayComparison,
DateComparison,
DateTimeComparison,
EqualityComparison,
GraphQLComparison,
OrderComparison,
TextComparison,
TimeComparison,
TimeDeltaComparison,
)
from strawchemy.schema.interfaces import ErrorType
from strawchemy.schema.mutation import (
Input,
Expand All @@ -24,18 +35,27 @@
"ALL",
"RELATIONSHIPS",
"SCALARS",
"ArrayComparison",
"DateComparison",
"DateTimeComparison",
"EqualityComparison",
"ErrorType",
"FieldGroup",
"GraphQLComparison",
"Input",
"InputValidationError",
"ModelInstance",
"OrderComparison",
"QueryHook",
"RequiredToManyUpdateInput",
"RequiredToOneInput",
"Strawchemy",
"StrawchemyAsyncRepository",
"StrawchemyConfig",
"StrawchemySyncRepository",
"TextComparison",
"TimeComparison",
"TimeDeltaComparison",
"ToManyCreateInput",
"ToManyUpdateInput",
"ToOneInput",
Expand Down
7 changes: 7 additions & 0 deletions src/strawchemy/dto/backend/strawberry.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,13 @@ def __init__(self, dto_base: type[AnnotatedDTOT], auto_is_type_of: bool = False)
}

def _construct_field_info(self, field_def: DTOFieldDefinition[ModelT, ModelFieldT]) -> FieldInfo:
# A field definition may carry a prebuilt strawberry field (e.g. from ``filter_field``);
# it already encodes the default and any explicit field config, so use it as-is.
# Local import: dto.strawberry imports this module at top, so a module-level import cycles.
from strawchemy.dto.strawberry import GraphQLFieldDefinition # noqa: PLC0415

if isinstance(field_def, GraphQLFieldDefinition) and field_def.graphql_field is not None:
return FieldInfo(field_def.name, field_def.type_, field_def.graphql_field)
strawberry_field: StrawberryField | None = None
if field_def.default_factory is not DTOMissing:
if isinstance(field_def.default_factory(), (list, tuple)):
Expand Down
14 changes: 9 additions & 5 deletions src/strawchemy/dto/inspectors/sqlalchemy.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@
make_full_json_comparison_input,
make_sqlite_json_comparison_input,
)
from strawchemy.utils.annotation import is_type_hint_optional
from strawchemy.utils.annotation import get_origin_or_self, is_type_hint_optional

if TYPE_CHECKING:
from collections.abc import Callable, Generator, Iterable
Expand Down Expand Up @@ -665,8 +665,8 @@ def _filter_type(cls, type_: type[Any], sqlalchemy_filter: type[GraphQLCompariso
"""
return sqlalchemy_filter if cls._is_specialized(sqlalchemy_filter) else sqlalchemy_filter[type_] # ty: ignore[not-subscriptable] # runtime generic specialization with a dynamic type argument

def get_field_comparison(
self, field_definition: DTOFieldDefinition[DeclarativeBase, QueryableAttribute[Any]]
def get_comparison(
self, field_definition: DTOFieldDefinition[DeclarativeBase, QueryableAttribute[Any]], subscribed: bool = True
) -> type[GraphQLComparison]:
"""Determines the GraphQL comparison filter type for a DTO field.

Expand All @@ -678,14 +678,18 @@ def get_field_comparison(
Args:
field_definition: The DTO field definition, which contains information
about the model attribute and its type.
subscribed: When `False`, return the unsubscripted comparison class
(e.g. `OrderComparison` instead of `OrderComparison[int]`).

Returns:
The GraphQL comparison filter type suitable for the field.
"""
field_type = field_definition.model_field.type
if isinstance(field_type, ARRAY) and self.db_features.dialect == "postgresql":
return ArrayComparison[field_type.item_type.python_type]
return self.get_type_comparison(self.model_field_type(field_definition))
comparison: type[GraphQLComparison] = ArrayComparison[field_type.item_type.python_type]
else:
comparison = self.get_type_comparison(self.model_field_type(field_definition))
return comparison if subscribed else get_origin_or_self(comparison)

def get_type_comparison(self, type_: type[Any]) -> type[GraphQLComparison]:
"""Determines the GraphQL comparison filter type for a Python type.
Expand Down
34 changes: 32 additions & 2 deletions src/strawchemy/dto/strawberry.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,8 +63,10 @@
from collections.abc import Callable, Hashable, Iterator

from sqlalchemy import ColumnElement
from strawberry.types.field import StrawberryField

from strawchemy.schema.filters import EqualityComparison, GraphQLComparison
from strawchemy.schema.filters.fields import CustomFilterApply, JoinStrategy

T = TypeVar("T")

Expand Down Expand Up @@ -182,6 +184,8 @@ class GraphQLFieldDefinition(DTOFieldDefinition[DeclarativeBase, QueryableAttrib
is_aggregate: bool = False
is_function: bool = False
is_function_arg: bool = False
graphql_field: StrawberryField | None = None
"""Prebuilt strawberry field carrying explicit field config."""

_function: FunctionInfo | None = None

Expand Down Expand Up @@ -286,6 +290,16 @@ def __post_init__(self) -> None:
self.is_function_arg = True


@dataclass(eq=False, repr=False, kw_only=True)
class CustomFilterFieldDefinition(GraphQLFieldDefinition):
"""Field definition for a custom-apply virtual filter field."""

apply: CustomFilterApply
"""The custom filter callable; always set for this field type."""
join: JoinStrategy = "exists"
"""Fold-back strategy used by the transpiler (``"exists"`` or ``"in"``)."""


@dataclass(eq=False)
class QueryNode(Node[GraphQLFieldDefinition, QueryNodeMetadata]):
node_metadata: NodeMetadata[QueryNodeMetadata] | None = dataclasses.field(
Expand Down Expand Up @@ -354,9 +368,23 @@ class AggregationFilter:
distinct: bool | None = None


@dataclass
class CustomFilter:
"""A set custom-apply filter value, ready to be folded into the query."""

apply: CustomFilterApply
"""The custom filter callable (see ``CustomFilterApply``)."""
value: Any
"""The scalar value supplied in the GraphQL query."""
join: JoinStrategy
"""Fold-back strategy (``"exists"`` or ``"in"``)."""
field_node: QueryNodeType
"""The model's query node, used to correlate the EXISTS/IN subquery."""


@dataclass
class Filter:
and_: list[Self | GraphQLComparison | AggregationFilter] = dataclasses.field(default_factory=list)
and_: list[Self | GraphQLComparison | AggregationFilter | CustomFilter] = dataclasses.field(default_factory=list)
or_: list[Self] = dataclasses.field(default_factory=list)
not_: Self | None = None

Expand Down Expand Up @@ -531,7 +559,9 @@ def filters_tree(self, _node: QueryNodeType | None = None) -> tuple[QueryNodeTyp
for name in self.dto_set_fields:
value: EqualityComparison[Any] | BooleanFilterDTO | AggregateFilterDTO = getattr(self, name)
field = self.__dto_field_definitions__[name]
if isinstance(value, BooleanFilterDTO):
if isinstance(field, CustomFilterFieldDefinition):
query.and_.append(CustomFilter(apply=field.apply, value=value, join=field.join, field_node=node))
elif isinstance(value, BooleanFilterDTO):
child, _ = node.upsert_child(field, match_on="value_equality")
_, sub_query = value.filters_tree(child)
if sub_query:
Expand Down
Loading