From 59dfce963cef5e8b2dfaabb794358d169eb99c34 Mon Sep 17 00:00:00 2001 From: gazorby Date: Tue, 2 Jun 2026 19:58:40 +0200 Subject: [PATCH 1/9] feat: add field groups # Conflicts: # tests/unit/mapping/test_schemas.py --- src/strawchemy/__init__.py | 5 + src/strawchemy/dto/base.py | 11 +- src/strawchemy/dto/types.py | 122 ++++++++++---- src/strawchemy/schema/factories/types.py | 11 +- src/strawchemy/utils/registry.py | 12 +- tests/unit/dto/test_dto.py | 195 +++++++++++++++++++++++ tests/unit/mapping/test_schemas.py | 22 +++ 7 files changed, 326 insertions(+), 52 deletions(-) diff --git a/src/strawchemy/__init__.py b/src/strawchemy/__init__.py index c93724b4..f9e25394 100644 --- a/src/strawchemy/__init__.py +++ b/src/strawchemy/__init__.py @@ -3,6 +3,7 @@ from __future__ import annotations from strawchemy.config.base import StrawchemyConfig +from strawchemy.dto.types import ALL, RELATIONSHIPS, SCALARS, FieldGroup from strawchemy.instance import ModelInstance from strawchemy.mapper import Strawchemy from strawchemy.repository.strawberry import StrawchemyAsyncRepository, StrawchemySyncRepository @@ -20,7 +21,11 @@ from strawchemy.validation import InputValidationError __all__ = ( + "ALL", + "RELATIONSHIPS", + "SCALARS", "ErrorType", + "FieldGroup", "Input", "InputValidationError", "ModelInstance", diff --git a/src/strawchemy/dto/base.py b/src/strawchemy/dto/base.py index f219a1b2..c6717456 100644 --- a/src/strawchemy/dto/base.py +++ b/src/strawchemy/dto/base.py @@ -37,6 +37,7 @@ IncludeFields, Purpose, PurposeConfig, + include_field, ) from strawchemy.dto.utils import config from strawchemy.exceptions import DTOError, EmptyDTOError @@ -390,11 +391,13 @@ def should_exclude_field( has_override: bool, ) -> bool: """Whether the model field should be excluded from the dto or not.""" - explictly_excluded = node.is_root and field.model_field_name in dto_config.exclude - explicitly_included = node.is_root and field.model_field_name in dto_config.include + name, is_relation = field.model_field_name, field.is_relation - globally_excluded = field.model_field_name in dto_config.global_exclude - globally_included = field.model_field_name in dto_config.global_include + explictly_excluded = node.is_root and include_field(name, is_relation, dto_config.exclude) + explicitly_included = node.is_root and include_field(name, is_relation, dto_config.include) + + globally_excluded = include_field(name, is_relation, dto_config.global_exclude) + globally_included = include_field(name, is_relation, dto_config.global_include) if dto_config.include == "all" and not explictly_excluded: explicitly_included = globally_included = True diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 771b7224..6c8f546d 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -14,8 +14,13 @@ if TYPE_CHECKING: from collections.abc import Callable, Mapping + from strawchemy.dto.base import DTOFieldDefinition + __all__ = ( + "ALL", + "RELATIONSHIPS", + "SCALARS", "DTOAuto", "DTOConfig", "DTOFieldConfig", @@ -23,17 +28,40 @@ "DTOScope", "DTOSkip", "DTOUnset", + "FieldGroup", "FieldIterable", + "FieldSelector", "IncludeFields", "Purpose", "PurposeConfig", "cast_include_fields", + "has_field_group", + "include_field", "is_fields_iterable", ) DTOScope: TypeAlias = Literal["global", "dto"] -FieldIterable: TypeAlias = "list[str] | set[str] | frozenset[str] | tuple[str, ...]" -IncludeFields: TypeAlias = "FieldIterable | Literal['all']" +FieldSelector: TypeAlias = "str | FieldGroup" +FieldIterable: TypeAlias = ( + "list[FieldSelector] | set[FieldSelector] | frozenset[FieldSelector] | tuple[FieldSelector, ...]" +) +FieldGroupStr: TypeAlias = Literal["all", "scalars", "relationships"] +IncludeFields: TypeAlias = "FieldIterable | FieldGroupStr" + + +class FieldGroup(str, Enum): + """Field-group selectors for ``include``/``exclude`` sequences.""" + + ALL = "all" + """Include all fields from model.""" + SCALARS = "scalars" + """Include everything but relationships.""" + RELATIONSHIPS = "relationships" + """Include only relationships.""" + + @staticmethod + def list_str() -> str: + return ", ".join(member.value for member in FieldGroup) @final @@ -127,11 +155,17 @@ class DTOConfig: Determines which fields from the source model are included based on their `DTOFieldConfig`. include: Explicitly include fields from the source model in the generated - DTO. Can be a list or set of field names, or the literal "all" to - include all fields not explicitly excluded. Defaults to an empty set. - exclude: Explicitly exclude fields from the source model. Can be a list - or set of field names. Defaults to an empty set. Setting this - implicitly sets `include` to "all". + DTO. Can be a list or set of field names, the literal "all" to include + all fields not explicitly excluded, or the `ALL` / `SCALARS` / + `RELATIONSHIPS` group selectors (mixable with names, e.g. + `[SCALARS, "owner"]`). `[SCALARS, RELATIONSHIPS]` is equivalent to + "all". Plain string literals "all", "scalars", "relationships" are + accepted interchangeably with the constants. Defaults to an empty set. + exclude: Explicitly exclude fields from the source model. Can be a list or + set of field names and/or the `ALL` / `SCALARS` / `RELATIONSHIPS` group + selectors (e.g. `[RELATIONSHIPS]` keeps all scalar fields and walks no + relationships). Plain string literals are accepted. A bare `exclude` (no + `include`) implies `include="all"`. Defaults to an empty set. partial: If True, makes all fields in the generated DTO optional. Defaults to None. partial_default: The default value assigned to fields when `partial` is @@ -185,15 +219,17 @@ def __post_init__(self) -> None: if self.aliases and self.alias_generator is not None: msg = "You must set `aliases` or `alias_generator`, not both" raise ValueError(msg) - if self.include and self.include != "all" and self.exclude: - msg = "When using `exclude` you must set `include='all' or leave it unset`" + if self.include and not has_field_group(self.include) and self.exclude: + msg = f"When using `exclude`, `include` must be unset or be a field group {FieldGroup.list_str()}." raise ValueError(msg) - if self.global_include and self.global_include != "all" and self.global_exclude: - msg = "When using `global_exclude` you must set `global_include='all' or leave it unset`" + if self.global_include and not has_field_group(self.global_include) and self.global_exclude: + msg = f"When using `global_exclude`, `global_include` must be unset or be a field group {FieldGroup.list_str()}." raise ValueError(msg) - if self.global_exclude: + # A bare exclude (no include) means "everything except"; promote to "all". + # If include carries a FieldGroup it is truthy, so the clobber is skipped. + if self.global_exclude and not self.global_include: self.global_include = "all" - if self.exclude: + if self.exclude and not self.include: self.include = "all" @overload @@ -320,7 +356,7 @@ def with_base_annotations(self, base: type[Any]) -> DTOConfig: 1. When include is "all" or exclude is specified: All fields from the base class are included 2. When specific fields are included: Only those fields are added to the include set """ - include: set[str] = set(self.include) if self.include != "all" else set() + include: set[FieldSelector] = set(self.include) if self.include != "all" else set() include_all = self.include == "all" or self.exclude annotation_overrides: dict[str, Any] = self.annotation_overrides try: @@ -344,45 +380,41 @@ def alias(self, name: str) -> str | None: return self.alias_generator(name) return None - def is_field_included(self, name: str) -> bool: - """Check if a field should be included based on this configuration. - - This method is used during DTO factory operations to determine which fields - from the source model should be included in the generated DTO. + def is_field_included(self, field: str | DTOFieldDefinition[Any, Any]) -> bool: + """Whether a field is included per the include/exclude rules. - Args: - name: The field name to check for inclusion. - - Returns: - True if the field should be included based on the include/exclude rules, - False otherwise. + `field` is a field name or a `DTOFieldDefinition`. A bare `str` is treated as a non-relation field name. """ + name, is_relation = (field, False) if isinstance(field, str) else (field.model_field_name, field.is_relation) if self.include == "all": - return name not in self.exclude + return not include_field(name, is_relation, self.exclude) if self.global_include == "all": - return name not in self.global_exclude - - included = set(self.include) | set(self.global_include) - excluded = set(self.exclude) | set(self.global_exclude) - return name in included and name not in excluded + return not include_field(name, is_relation, self.global_exclude) + return ( + include_field(name, is_relation, self.include) or include_field(name, is_relation, self.global_include) + ) and not ( + include_field(name, is_relation, self.exclude) or include_field(name, is_relation, self.global_exclude) + ) def __or__(self, other: DTOConfig) -> DTOConfig: return self.union(other) @overload -def cast_include_fields(value: Literal["all"]) -> Literal["all"]: ... +def cast_include_fields(value: FieldGroupStr) -> FieldGroupStr: ... @overload -def cast_include_fields(value: frozenset[str] | set[str] | list[str] | tuple[str, ...] | None) -> frozenset[str]: ... +def cast_include_fields( + value: frozenset[FieldSelector] | set[FieldSelector] | list[FieldSelector] | tuple[FieldSelector, ...] | None, +) -> frozenset[FieldSelector]: ... -def cast_include_fields(value: IncludeFields | None) -> frozenset[str] | Literal["all"]: +def cast_include_fields(value: IncludeFields | None) -> frozenset[FieldSelector] | FieldGroupStr: match value: case None: return frozenset() - case "all": + case "all" | "scalars" | "relationships": return "all" case _: return frozenset(value) @@ -395,3 +427,23 @@ def is_fields_iterable(value: Any) -> TypeIs[IncludeFields | FieldIterable | Non if isinstance(value, str): return False return isinstance(value, (frozenset, set, list, tuple)) + + +def include_field(field_name: str, is_relation: bool, fields: IncludeFields | FieldIterable) -> bool: + """Whether `fields` selects the field.""" + if fields == "all": + return True + group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS + return field_name in fields or group in fields or FieldGroup.ALL in fields + + +def has_field_group(value: IncludeFields | FieldIterable) -> bool: + """True if the selection contains a FieldGroup member.""" + if value in FieldGroup or isinstance(value, FieldGroup): + return True + return any(isinstance(item, FieldGroup) for item in value) + + +ALL = FieldGroup.ALL +SCALARS = FieldGroup.SCALARS +RELATIONSHIPS = FieldGroup.RELATIONSHIPS diff --git a/src/strawchemy/schema/factories/types.py b/src/strawchemy/schema/factories/types.py index 1f8838c7..844c7c58 100644 --- a/src/strawchemy/schema/factories/types.py +++ b/src/strawchemy/schema/factories/types.py @@ -157,16 +157,13 @@ def _relation_field( related = Self if field.related_dto is dto else field.related_dto type_annotation = list[related] if related is not None else field.type_ # ty: ignore[invalid-type-form] assert field.related_model - field_name = field.model_field_name order_by_input, distinct_on_input, pagination = None, None, False - if order_config.is_field_included(field_name) or self._mapper.config.order_config.is_field_included(field_name): + if order_config.is_field_included(field) or self._mapper.config.order_config.is_field_included(field): order_by_input = self._order_by_input_for_field(field) - if pagination_config.is_field_included(field_name) or self._mapper.config.pagination_config.is_field_included( - field_name - ): + if pagination_config.is_field_included(field) or self._mapper.config.pagination_config.is_field_included(field): pagination = default_pagination or True - if distinct_on_config.is_field_included(field_name) or self._mapper.config.distinct_on_config.is_field_included( - field_name + if distinct_on_config.is_field_included(field) or self._mapper.config.distinct_on_config.is_field_included( + field ): distinct_on_input = self._distinct_on_input_for_field(field) strawberry_field = self._mapper.field( diff --git a/src/strawchemy/utils/registry.py b/src/strawchemy/utils/registry.py index 8b254c56..0c5e5ce6 100644 --- a/src/strawchemy/utils/registry.py +++ b/src/strawchemy/utils/registry.py @@ -4,7 +4,7 @@ from collections import defaultdict from copy import copy from enum import Enum -from typing import TYPE_CHECKING, Any, ForwardRef, Literal, NewType, TypeVar, cast, overload +from typing import TYPE_CHECKING, Any, ForwardRef, NewType, TypeVar, cast, overload import strawberry from strawberry import LazyType @@ -15,7 +15,7 @@ from strawberry.types.union import StrawberryUnion from strawchemy.dto.strawberry import MappedStrawberryGraphQLDTO -from strawchemy.dto.types import cast_include_fields, is_fields_iterable +from strawchemy.dto.types import FieldGroupStr, cast_include_fields, is_fields_iterable from strawchemy.exceptions import StrawchemyError from strawchemy.utils.annotation import inner_types from strawchemy.utils.strawberry import strawberry_contained_types @@ -39,7 +39,7 @@ from strawchemy.dto import DTOConfig from strawchemy.dto.base import Node, Relation from strawchemy.dto.strawberry import EnumDTO, OrderByDTO, StrawchemyObject - from strawchemy.dto.types import DTOScope, IncludeFields + from strawchemy.dto.types import DTOScope, FieldSelector, IncludeFields from strawchemy.schema.pagination import DefaultOffsetPagination from strawchemy.typing import GraphQLType, StrawchemyObjectWithStrawberryObjectDefinition @@ -149,9 +149,9 @@ class RegistryTypeInfo: user_defined: bool = False override: bool = False pagination: DefaultOffsetPagination | None = None - order: frozenset[str] | Literal["all"] | type[OrderByDTO] = dataclasses.field(default_factory=frozenset) - distinct_on: frozenset[str] | Literal["all"] | type[EnumDTO] = dataclasses.field(default_factory=frozenset) - paginate: frozenset[str] | Literal["all"] = dataclasses.field(default_factory=frozenset) + order: frozenset[FieldSelector] | FieldGroupStr | type[OrderByDTO] = dataclasses.field(default_factory=frozenset) + distinct_on: frozenset[FieldSelector] | FieldGroupStr | type[EnumDTO] = dataclasses.field(default_factory=frozenset) + paginate: frozenset[FieldSelector] | FieldGroupStr = dataclasses.field(default_factory=frozenset) scope: DTOScope | None = None model: type[DeclarativeBase] | None = None tags: frozenset[str] = dataclasses.field(default_factory=frozenset) diff --git a/tests/unit/dto/test_dto.py b/tests/unit/dto/test_dto.py index 4bd2c26d..4ef11f9b 100644 --- a/tests/unit/dto/test_dto.py +++ b/tests/unit/dto/test_dto.py @@ -9,9 +9,13 @@ from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column from typing_extensions import Self +import strawchemy +from strawchemy import ALL, RELATIONSHIPS, SCALARS from strawchemy.dto import DTOConfig, Purpose, PurposeConfig, config, field +from strawchemy.dto.base import DTOFieldDefinition from strawchemy.dto.constants import DTO_INFO_KEY from strawchemy.dto.strawberry import DTOKey, GraphQLFieldDefinition, StrawchemyDefinition +from strawchemy.dto.types import FieldGroup, include_field from strawchemy.dto.utils import DTOFieldConfig, read_all_config, write_all_config from tests.typing import AnyFactory, MappedPydanticFactory from tests.unit.dc_models import ( @@ -26,6 +30,16 @@ from tests.utils import DTOInspect, factory_iterator +def _fruit_field(name: str, *, is_relation: bool) -> DTOFieldDefinition: # type: ignore[type-arg] + return DTOFieldDefinition( + dto_config=DTOConfig(Purpose.READ), + model=Fruit, + model_field_name=name, + type_hint=int, + is_relation=is_relation, + ) + + class _PopulateFieldsBase(DeclarativeBase): pass @@ -368,3 +382,184 @@ def test_strawchemy_definition_populate_fields(key_source: type[DeclarativeBase] assert result is definition assert definition.field_map == {DTOKey([_PopulateFieldsModel]) + "id": field_def} + + +def test_field_group_constants_are_enum_members() -> None: + """Test that the top-level SCALARS/RELATIONSHIPS constants are the enum members.""" + assert strawchemy.SCALARS is FieldGroup.SCALARS + assert strawchemy.RELATIONSHIPS is FieldGroup.RELATIONSHIPS + + +def test_field_group_constants_hashable_in_sets() -> None: + """Test that group constants are usable inside include/exclude frozensets.""" + members = frozenset([strawchemy.SCALARS, strawchemy.RELATIONSHIPS, "name"]) + assert strawchemy.SCALARS in members + assert strawchemy.RELATIONSHIPS in members + assert "name" in members + + +def test_scalars_include_allows_exclude() -> None: + """Test that a group-bearing include coexists with exclude and is not clobbered to 'all'.""" + config = DTOConfig(Purpose.READ, include=[SCALARS], exclude=["secret"]) + assert config.include != "all" + assert SCALARS in config.include + assert "secret" in config.exclude + + +def test_plain_include_with_exclude_still_raises() -> None: + """Test that a plain field-name include combined with exclude still raises.""" + with pytest.raises(ValueError, match="exclude"): + DTOConfig(Purpose.READ, include=["a", "b"], exclude=["c"]) + + +def test_bare_exclude_still_implies_all() -> None: + """Test that a bare exclude (no include) still implies include='all'.""" + config = DTOConfig(Purpose.READ, exclude=["secret"]) + assert config.include == "all" + + +def test_relationships_exclude_implies_all_include() -> None: + """Test that exclude=[RELATIONSHIPS] with no include promotes include to 'all'.""" + config = DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS]) + assert config.include == "all" + assert RELATIONSHIPS in config.exclude + + +def test_mixed_group_and_name_include_allows_exclude() -> None: + """Test that a group selector mixed with a field name coexists with exclude.""" + config = DTOConfig(Purpose.READ, include=[SCALARS, "owner"], exclude=["secret"]) + assert config.include != "all" + assert SCALARS in config.include + assert "owner" in config.include + assert "secret" in config.exclude + + +def test_global_group_include_allows_global_exclude() -> None: + """Test that a group-bearing global_include coexists with global_exclude.""" + config = DTOConfig(Purpose.READ, global_include=[SCALARS], global_exclude=["secret"]) + assert config.global_include != "all" + assert SCALARS in config.global_include + assert "secret" in config.global_exclude + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_scalars_excludes_relationships(factory: AnyFactory) -> None: + """Test that include=[SCALARS] keeps scalar fields and drops relationships.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS])) + fields = set(DTOInspect(dto).annotations()) + assert "color" not in fields + assert {"id", "name", "sweetness", "color_id"} <= fields + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_scalars_plus_named_relationship(factory: AnyFactory) -> None: + """Test that include=[SCALARS, 'color'] keeps scalars plus the named relationship.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS, "color"])) + fields = set(DTOInspect(dto).annotations()) + assert "color" in fields + assert {"id", "name", "sweetness", "color_id"} <= fields + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_relationships_only(factory: AnyFactory) -> None: + """Test that include=[RELATIONSHIPS] keeps only relationships and drops scalars.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[RELATIONSHIPS])) + fields = set(DTOInspect(dto).annotations()) + assert "color" in fields + assert "name" not in fields + assert "sweetness" not in fields + assert "color_id" not in fields + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_both_groups_equals_all(factory: AnyFactory) -> None: + """Test that include=[SCALARS, RELATIONSHIPS] is equivalent to include='all'.""" + grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS, RELATIONSHIPS])) + all_dto = factory.factory(Fruit, read_all_config) + assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_exclude_relationships_keeps_scalars(factory: AnyFactory) -> None: + """Test that exclude=[RELATIONSHIPS] keeps all scalar fields and walks no relationships.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS])) + fields = set(DTOInspect(dto).annotations()) + assert "color" not in fields + assert {"id", "name", "sweetness", "color_id"} <= fields + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_exclude_relationships_plus_named_scalar(factory: AnyFactory) -> None: + """Test that exclude=[RELATIONSHIPS, 'sweetness'] drops relationships and the named scalar.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS, "sweetness"])) + fields = set(DTOInspect(dto).annotations()) + assert "color" not in fields + assert "sweetness" not in fields + assert {"id", "name", "color_id"} <= fields + + +def test_is_field_included_relationships_group_matches_relations() -> None: + """Test that include=[RELATIONSHIPS] includes relation fields and excludes scalars.""" + config = DTOConfig(Purpose.READ, include=frozenset([RELATIONSHIPS])) + assert config.is_field_included(_fruit_field("owner", is_relation=True)) is True + assert config.is_field_included(_fruit_field("name", is_relation=False)) is False + + +def test_is_field_included_scalars_group_matches_scalars() -> None: + """Test that include=[SCALARS] includes scalar fields and excludes relations.""" + config = DTOConfig(Purpose.READ, include=frozenset([SCALARS])) + assert config.is_field_included(_fruit_field("name", is_relation=False)) is True + assert config.is_field_included(_fruit_field("owner", is_relation=True)) is False + + +def test_is_field_included_group_in_exclude() -> None: + """Test that a group selector in exclude drops matching fields.""" + config = DTOConfig(Purpose.READ, include="all", exclude=frozenset([RELATIONSHIPS])) + assert config.is_field_included(_fruit_field("owner", is_relation=True)) is False + assert config.is_field_included(_fruit_field("name", is_relation=False)) is True + + +def test_is_field_included_no_group_unchanged() -> None: + """Test that behavior is unchanged when no group constants are present.""" + config = DTOConfig(Purpose.READ, include=frozenset(["a", "b"])) + assert config.is_field_included("a") is True + assert config.is_field_included("c") is False + assert config.is_field_included(_fruit_field("a", is_relation=True)) is True + + +def test_string_literals_equal_field_group_members() -> None: + """Test that plain string literals equal their FieldGroup members.""" + assert FieldGroup.ALL == "all" + assert FieldGroup.SCALARS == "scalars" + assert FieldGroup.RELATIONSHIPS == "relationships" + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_string_scalars_literal(factory: AnyFactory) -> None: + """Test that the literal include=['scalars'] behaves like include=[SCALARS].""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=["scalars"])) + fields = set(DTOInspect(dto).annotations()) + assert "color" not in fields + assert {"id", "name", "sweetness", "color_id"} <= fields + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_all_literal_set_equals_all(factory: AnyFactory) -> None: + """Test that include={'all'} is equivalent to include='all'.""" + grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include={"all"})) + all_dto = factory.factory(Fruit, read_all_config) + assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) + + +@pytest.mark.parametrize("factory", factory_iterator()) +def test_include_all_constant_in_collection(factory: AnyFactory) -> None: + """Test that include=[ALL] is equivalent to include='all'.""" + grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[ALL])) + all_dto = factory.factory(Fruit, read_all_config) + assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) + + +def test_include_field_all_selects_everything() -> None: + """Test that include_field returns True for the 'all' selector.""" + assert include_field("anything", False, "all") is True + assert include_field("anything", True, "all") is True diff --git a/tests/unit/mapping/test_schemas.py b/tests/unit/mapping/test_schemas.py index ae873030..7f6767c3 100644 --- a/tests/unit/mapping/test_schemas.py +++ b/tests/unit/mapping/test_schemas.py @@ -16,12 +16,14 @@ from strawberry.types import get_object_definition from strawberry.types.object_type import StrawberryObjectDefinition +from strawchemy import SCALARS from strawchemy.exceptions import EmptyDTOError, QueryHookError, StrawchemyError, StrawchemyFieldError from strawchemy.schema.scalars import Interval from strawchemy.testing.pytest_plugin import MockContext from tests.fixtures import DefaultQuery from tests.unit.models import Book as BookModel from tests.unit.models import Fruit, User +from tests.utils import DTOInspect if TYPE_CHECKING: from syrupy.assertion import SnapshotAssertion @@ -639,3 +641,23 @@ def test_json_column_class_body_resolver_executes() -> None: result = schema.execute_sync("{ overriddenJson { dictCol } }") assert not result.errors assert result.data == {"overriddenJson": {"dictCol": "OVERRIDE"}} + + +def test_exclude_relationships_avoids_stub_collision(strawchemy: Strawchemy) -> None: + """Test that exclude=[RELATIONSHIPS] walks no relationships, so a later explicit type for a related model does not collide with a pre-registered walker stub (#162).""" + from strawchemy import RELATIONSHIPS + from tests.unit.models import Color, Fruit + + # First slice scopes Fruit with a relationship-free walk... + @strawchemy.type(Fruit, exclude=[RELATIONSHIPS]) + class FruitNode: + pass + + # ...so a later explicit Color type must NOT collide with a walker stub. + @strawchemy.type(Color, include=[SCALARS]) + class ColorNode: + pass + + fruit_fields = set(DTOInspect(FruitNode).annotations()) + assert "color" not in fruit_fields + assert {"id", "name", "sweetness", "color_id"} <= fruit_fields From b90eca9eeb872065cd49f52abc48973111fb224d Mon Sep 17 00:00:00 2001 From: gazorby Date: Tue, 2 Jun 2026 21:44:31 +0200 Subject: [PATCH 2/9] tests: fix python 3.10/3.111 --- src/strawchemy/dto/types.py | 8 ++++++-- tests/unit/mapping/test_schemas.py | 6 ++---- tests/utils.py | 5 +---- 3 files changed, 9 insertions(+), 10 deletions(-) diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 6c8f546d..51c60d05 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -63,6 +63,10 @@ class FieldGroup(str, Enum): def list_str() -> str: return ", ".join(member.value for member in FieldGroup) + @classmethod + def values(cls) -> frozenset[str]: + return frozenset(member.value for member in FieldGroup) + @final class DTOMissing: @@ -439,8 +443,8 @@ def include_field(field_name: str, is_relation: bool, fields: IncludeFields | Fi def has_field_group(value: IncludeFields | FieldIterable) -> bool: """True if the selection contains a FieldGroup member.""" - if value in FieldGroup or isinstance(value, FieldGroup): - return True + if isinstance(value, str): + return value in FieldGroup.values() return any(isinstance(item, FieldGroup) for item in value) diff --git a/tests/unit/mapping/test_schemas.py b/tests/unit/mapping/test_schemas.py index 7f6767c3..04754336 100644 --- a/tests/unit/mapping/test_schemas.py +++ b/tests/unit/mapping/test_schemas.py @@ -16,13 +16,13 @@ from strawberry.types import get_object_definition from strawberry.types.object_type import StrawberryObjectDefinition -from strawchemy import SCALARS +from strawchemy import RELATIONSHIPS, SCALARS from strawchemy.exceptions import EmptyDTOError, QueryHookError, StrawchemyError, StrawchemyFieldError from strawchemy.schema.scalars import Interval from strawchemy.testing.pytest_plugin import MockContext from tests.fixtures import DefaultQuery from tests.unit.models import Book as BookModel -from tests.unit.models import Fruit, User +from tests.unit.models import Color, Fruit, User from tests.utils import DTOInspect if TYPE_CHECKING: @@ -645,8 +645,6 @@ def test_json_column_class_body_resolver_executes() -> None: def test_exclude_relationships_avoids_stub_collision(strawchemy: Strawchemy) -> None: """Test that exclude=[RELATIONSHIPS] walks no relationships, so a later explicit type for a related model does not collide with a pre-registered walker stub (#162).""" - from strawchemy import RELATIONSHIPS - from tests.unit.models import Color, Fruit # First slice scopes Fruit with a relationship-free walk... @strawchemy.type(Fruit, exclude=[RELATIONSHIPS]) diff --git a/tests/utils.py b/tests/utils.py index ccb69798..09d6e6cc 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -144,9 +144,6 @@ class FactoryType(Enum): class DTOInspectProtocol(Protocol): dto: type[Any] - def __init__(self, dto: type[Any]) -> None: - self.dto = dto - @classmethod def is_class(cls, dto: type[Any]) -> bool: ... @@ -161,7 +158,7 @@ class DataclassInspect(DTOInspectProtocol): dto: type[DataclassProtocol] def __init__(self, dto: type[DataclassProtocol]) -> None: - super().__init__(dto) + self.dto = dto self._dataclass_fields = {field.name: field for field in dataclasses.fields(self.dto)} @classmethod From 129fb89092503bab4d98d07cb12edfa562dfde98 Mon Sep 17 00:00:00 2001 From: gazorby Date: Fri, 5 Jun 2026 22:21:04 +0200 Subject: [PATCH 3/9] ref: encapsulate include field logic # Conflicts: # src/strawchemy/config/base.py --- src/strawchemy/config/base.py | 12 +- src/strawchemy/dto/base.py | 37 +-- src/strawchemy/dto/types.py | 248 +++++++++------ src/strawchemy/dto/utils.py | 49 ++- src/strawchemy/mapper.py | 14 +- src/strawchemy/schema/factories/_kwargs.py | 12 +- src/strawchemy/schema/factories/base.py | 31 +- src/strawchemy/schema/factories/types.py | 8 +- src/strawchemy/schema/field.py | 10 +- src/strawchemy/utils/registry.py | 28 +- src/strawchemy/validation/pydantic.py | 6 +- tests/unit/dto/test_dto.py | 347 ++++++++------------- 12 files changed, 382 insertions(+), 420 deletions(-) diff --git a/src/strawchemy/config/base.py b/src/strawchemy/config/base.py index 2fb6a2ce..9163730e 100644 --- a/src/strawchemy/config/base.py +++ b/src/strawchemy/config/base.py @@ -7,7 +7,7 @@ from strawchemy.dto import Purpose from strawchemy.dto.inspectors import SQLAlchemyGraphQLInspector -from strawchemy.dto.types import DTOConfig, FieldIterable, IncludeFields +from strawchemy.dto.types import DTOConfig, FieldSpec from strawchemy.repository.strawberry import StrawchemySyncRepository from strawchemy.utils.strawberry import default_session_getter @@ -59,15 +59,15 @@ class StrawchemyConfig: output object types so they work as GraphQL Union / interface members without boilerplate. A user-defined is_type_of on the decorated class is always respected. Set to False to disable globally.""" - include: IncludeFields = "all" + include: FieldSpec = "all" """Globally included fields.""" - exclude: FieldIterable | None = None + exclude: FieldSpec | None = None """Globally included fields.""" - pagination: IncludeFields | None = None + pagination: FieldSpec | None = None """Enable/disable pagination on list resolvers.""" - order_by: IncludeFields | None = None + order_by: FieldSpec | None = None """Enable/disable order by on list resolvers.""" - distinct_on: IncludeFields | None = None + distinct_on: FieldSpec | None = None """Enable/disable order by on list resolvers.""" pagination_default_limit: int = 100 """Default pagination limit when `pagination=True`.""" diff --git a/src/strawchemy/dto/base.py b/src/strawchemy/dto/base.py index c6717456..467d80c6 100644 --- a/src/strawchemy/dto/base.py +++ b/src/strawchemy/dto/base.py @@ -33,11 +33,10 @@ DTOMissing, DTOSkip, DTOUnset, - FieldIterable, - IncludeFields, + FieldGroup, + FieldSpec, Purpose, PurposeConfig, - include_field, ) from strawchemy.dto.utils import config from strawchemy.exceptions import DTOError, EmptyDTOError @@ -391,30 +390,16 @@ def should_exclude_field( has_override: bool, ) -> bool: """Whether the model field should be excluded from the dto or not.""" - name, is_relation = field.model_field_name, field.is_relation - - explictly_excluded = node.is_root and include_field(name, is_relation, dto_config.exclude) - explicitly_included = node.is_root and include_field(name, is_relation, dto_config.include) - - globally_excluded = include_field(name, is_relation, dto_config.global_exclude) - globally_included = include_field(name, is_relation, dto_config.global_include) - - if dto_config.include == "all" and not explictly_excluded: - explicitly_included = globally_included = True - - if dto_config.global_include == "all" and not globally_excluded: - globally_included = True + included_by_config = node.is_root and dto_config.is_field_included(field, scope="local") + included_globally = dto_config.is_field_included(field, scope="global") excluded = dto_config.purpose not in field.allowed_purposes # Exclude fields not present in init if purpose is write - if dto_config.purpose is Purpose.WRITE and not (explicitly_included or globally_included): + if dto_config.purpose is Purpose.WRITE and not (included_by_config or included_globally): excluded = excluded or not field.init - if node.is_root: - excluded = excluded or (explictly_excluded or not explicitly_included) - else: - excluded = excluded or (globally_excluded or not globally_included) + excluded = excluded or not included_by_config if node.is_root else excluded or not included_globally return not has_override and excluded @@ -503,8 +488,10 @@ def _base_cache_key(self, dto_config: DTOConfig) -> Hashable: def _root_cache_key(self, dto_config: DTOConfig) -> Hashable: root_key = [ - frozenset(dto_config.include if dto_config.include != "all" else ()), - frozenset(dto_config.exclude), + # ALL is dropped so a root "include all" config shares its cache entry with + # nested defaults, letting relations reuse user-defined types. + frozenset(dto_config.included_fields) - {FieldGroup.ALL}, + frozenset(dto_config.excluded_fields), frozenset(dto_config.aliases.items()), frozenset(dto_config.annotation_overrides.items()), ] @@ -703,8 +690,8 @@ def decorator( model: type[ModelT], purpose: Purpose, *, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 51c60d05..3acd1552 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -3,16 +3,17 @@ from __future__ import annotations import dataclasses -from dataclasses import dataclass, field +import functools +from dataclasses import InitVar, dataclass, field from enum import Enum -from typing import TYPE_CHECKING, Any, Literal, TypeAlias, final, get_type_hints, overload +from typing import TYPE_CHECKING, Any, Literal, TypeAlias, final, get_type_hints from typing_extensions import Self, TypeIs, override from strawchemy.utils.annotation import get_annotations if TYPE_CHECKING: - from collections.abc import Callable, Mapping + from collections.abc import Callable, Iterator, Mapping from strawchemy.dto.base import DTOFieldDefinition @@ -29,14 +30,10 @@ "DTOSkip", "DTOUnset", "FieldGroup", - "FieldIterable", "FieldSelector", - "IncludeFields", + "FieldSpec", "Purpose", "PurposeConfig", - "cast_include_fields", - "has_field_group", - "include_field", "is_fields_iterable", ) @@ -46,7 +43,8 @@ "list[FieldSelector] | set[FieldSelector] | frozenset[FieldSelector] | tuple[FieldSelector, ...]" ) FieldGroupStr: TypeAlias = Literal["all", "scalars", "relationships"] -IncludeFields: TypeAlias = "FieldIterable | FieldGroupStr" +FieldSpec: TypeAlias = "FieldIterable | FieldGroupStr" +ConfigScope: TypeAlias = Literal["local", "global"] class FieldGroup(str, Enum): @@ -59,14 +57,107 @@ class FieldGroup(str, Enum): RELATIONSHIPS = "relationships" """Include only relationships.""" + def __bool__(self) -> bool: + return True + @staticmethod def list_str() -> str: return ", ".join(member.value for member in FieldGroup) @classmethod + @functools.cache def values(cls) -> frozenset[str]: return frozenset(member.value for member in FieldGroup) + @classmethod + def is_group(cls, value: str) -> TypeIs[FieldGroupStr]: + return value in cls.values() + + +@dataclass(slots=True) +class FieldSet: + """Normalized, immutable view over a field selection. + + Wraps a `FieldSpec` (a group string such as "all", an iterable of field + names and/or `FieldGroup` selectors, or `None`) into a uniform + `frozenset[FieldSelector]` so selections can be compared, hashed, and + combined regardless of how they were originally expressed. + + Set-like semantics: + - `item in field_set`: matches by field name, by the field's group + (`SCALARS`/`RELATIONSHIPS`), or by `ALL`. `item` can be a field name, + a `DTOFieldDefinition` (its `is_relation` flag picks the group), or a + `FieldGroup` — group items only match themselves or `ALL`, never via + group matching. + - `self & other`: intersection of two selections; `ALL` subsumes the + other side. + - `self | other`: union, returned as a `FieldSpec | None` — collapses + to "all" when `ALL` is present, `None` when both sides are empty. + """ + + value: InitVar[FieldSpec | None] + + field_set: frozenset[FieldSelector] = field(init=False, default_factory=frozenset) + + def __post_init__(self, value: FieldSpec | None) -> None: + self.field_set = self.normalize(value) + + def __next__(self) -> FieldSelector: + return next(iter(self.field_set)) + + def __iter__(self) -> Iterator[FieldSelector]: + return iter(self.field_set) + + def __contains__(self, item: FieldSelector | DTOFieldDefinition[Any, Any]) -> bool: + # A FieldGroup is only selected by itself or by ALL, never by group matching. + if isinstance(item, FieldGroup): + return item in self.field_set or FieldGroup.ALL in self.field_set + name, is_relation = (item, False) if isinstance(item, str) else (item.model_field_name, item.is_relation) + group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS + return name in self.field_set or group in self.field_set or FieldGroup.ALL in self.field_set + + def __and__(self, other: FieldSpec) -> FieldIterable: + other_set = FieldSet(other) + # ALL subsumes any selection: intersecting with it yields the other side. + if FieldGroup.ALL in self.field_set: + return other_set.field_set + if FieldGroup.ALL in other_set.field_set: + return self.field_set + return frozenset(field for field in self.field_set if field in other_set) | frozenset( + field for field in other_set.field_set if field in self + ) + + def __or__(self, other: FieldSpec | None) -> FieldSpec | None: + union = self.field_set | FieldSet(other).field_set + if FieldGroup.ALL in union: + return "all" + return union or None + + def __bool__(self) -> bool: + return bool(self.field_set) + + def __hash__(self) -> int: + return hash(self.field_set) + + @classmethod + def normalize(cls, value: FieldSpec | None) -> frozenset[FieldSelector]: + """Normalize a field selection into a frozenset of selectors. + + Args: + value: A group string ("all", "scalars", "relationships"), an + iterable of field names and/or `FieldGroup` members, or `None`. + + Returns: + A group string as a single-member frozenset of the matching + `FieldGroup`, `None` as an empty frozenset, and any iterable as a + frozenset of its items. + """ + if isinstance(value, str) and FieldGroup.is_group(value): + return frozenset((FieldGroup(value),)) + if value is None: + return frozenset() + return frozenset(value) + @final class DTOMissing: @@ -121,7 +212,7 @@ class Purpose(str, Enum): clients. Fields marked as TO_COMPLETE must not be null.""" -@dataclass +@dataclass(slots=True) class PurposeConfig: """Mark the field as read-only, or private.""" @@ -145,7 +236,7 @@ def purpose_config(self, dto_config: DTOConfig) -> PurposeConfig: return self.configs.get(dto_config.purpose, self.default_config) -@dataclass +@dataclass(slots=True) class DTOConfig: """Control the generated DTO. @@ -198,13 +289,13 @@ class DTOConfig: purpose: Purpose """Configure the DTO for "read" or "write" operations.""" - include: IncludeFields = field(default_factory=set) + include: FieldSpec | None = None """Explicitly include fields from the generated DTO.""" - global_include: IncludeFields = field(default_factory=set) - """Explicitly include fields from the generated DTO and all its children.""" - exclude: FieldIterable = field(default_factory=set) + exclude: FieldSpec | None = None """Explicitly exclude fields from the generated DTO. Implies `include="all"`.""" - global_exclude: FieldIterable = field(default_factory=set) + global_include: FieldSpec | None = None + """Explicitly include fields from the generated DTO and all its children.""" + global_exclude: FieldSpec | None = None """Explicitly exclude fields from the generated DTO and all its children. Implies `global_include="all"`.""" partial: bool | None = None """Make all field optional.""" @@ -219,42 +310,54 @@ class DTOConfig: exclude_from_scope: bool = False tags: set[str] = field(default_factory=set) + included_fields: FieldSet = field(init=False) + excluded_fields: FieldSet = field(init=False) + def __post_init__(self) -> None: if self.aliases and self.alias_generator is not None: msg = "You must set `aliases` or `alias_generator`, not both" raise ValueError(msg) - if self.include and not has_field_group(self.include) and self.exclude: + if self.include and not self._has_field_group(self.include) and self.exclude: msg = f"When using `exclude`, `include` must be unset or be a field group {FieldGroup.list_str()}." raise ValueError(msg) - if self.global_include and not has_field_group(self.global_include) and self.global_exclude: + if self.global_include and not self._has_field_group(self.global_include) and self.global_exclude: msg = f"When using `global_exclude`, `global_include` must be unset or be a field group {FieldGroup.list_str()}." raise ValueError(msg) # A bare exclude (no include) means "everything except"; promote to "all". # If include carries a FieldGroup it is truthy, so the clobber is skipped. - if self.global_exclude and not self.global_include: + if self.global_exclude and self.global_include is None: self.global_include = "all" - if self.exclude and not self.include: + if self.exclude and self.include is None: self.include = "all" - @overload - @classmethod - def _merge_field_iterables(cls, *iterables: FieldIterable) -> FieldIterable: ... + self.included_fields = FieldSet(self.global_include) if self.include is None else FieldSet(self.include) + self.excluded_fields = FieldSet(self.global_exclude) if self.exclude is None else FieldSet(self.exclude) + + def __or__(self, other: DTOConfig) -> DTOConfig: + return self.union(other) - @overload @classmethod - def _merge_field_iterables(cls, *iterables: IncludeFields) -> IncludeFields: ... + def _has_field_group(cls, value: FieldSpec | FieldIterable) -> bool: + """True if the selection contains a FieldGroup member.""" + if isinstance(value, str): + return value in FieldGroup.values() + return any(isinstance(item, FieldGroup) for item in value) @classmethod - def _merge_field_iterables(cls, *iterables: IncludeFields | FieldIterable) -> IncludeFields | FieldIterable: - if any(iterable == "all" for iterable in iterables): - return "all" - return set().union(*iterables) + def _include_field(cls, field_name: str, is_relation: bool, fields: FieldSpec | FieldIterable | None) -> bool: + """Whether `fields` selects the field.""" + if fields is None: + return False + if fields == "all": + return True + group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS + return field_name in fields or group in fields or FieldGroup.ALL in fields def union(self, other: DTOConfig) -> DTOConfig: - include = self._merge_field_iterables(self.include, other.include) - exclude = self._merge_field_iterables(self.exclude, other.exclude) - global_include = self._merge_field_iterables(self.global_include, other.global_include) - global_exclude = self._merge_field_iterables(self.global_exclude, other.global_exclude) + include = FieldSet(self.include) | other.include + exclude = FieldSet(self.exclude) | other.exclude + global_include = FieldSet(self.global_include) | other.global_include + global_exclude = FieldSet(self.global_exclude) | other.global_exclude type_overrides = dict(self.type_overrides) | dict(other.type_overrides) annotation_overrides = self.annotation_overrides | other.annotation_overrides tags = self.tags | other.tags @@ -270,9 +373,7 @@ def union(self, other: DTOConfig) -> DTOConfig: ) @classmethod - def from_include( - cls, include: IncludeFields | Literal[False] | None = None, purpose: Purpose = Purpose.READ - ) -> Self: + def from_include(cls, include: FieldSpec | Literal[False] | None = None, purpose: Purpose = Purpose.READ) -> Self: """Create a DTOConfig from an include specification. Factory method for creating a DTOConfig with a simplified interface, converting @@ -297,10 +398,10 @@ def from_include( def copy_with( self, purpose: Purpose | type[DTOUnset] = DTOUnset, - include: IncludeFields | None = None, - global_include: IncludeFields | None = None, - exclude: FieldIterable | None = None, - global_exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + global_include: FieldSpec | None = None, + exclude: FieldSpec | None = None, + global_exclude: FieldSpec | None = None, partial: bool | None | type[DTOUnset] = DTOUnset, unset_sentinel: Any | type[DTOUnset] = DTOUnset, type_overrides: Mapping[Any, Any] | type[DTOUnset] = DTOUnset, @@ -360,8 +461,10 @@ def with_base_annotations(self, base: type[Any]) -> DTOConfig: 1. When include is "all" or exclude is specified: All fields from the base class are included 2. When specific fields are included: Only those fields are added to the include set """ - include: set[FieldSelector] = set(self.include) if self.include != "all" else set() - include_all = self.include == "all" or self.exclude + # Root-level include/exclude only: a global "all" must not pull base fields in. + include_set = FieldSet(self.include) + include = set(include_set.field_set) + include_all = FieldGroup.ALL in include_set.field_set or bool(FieldSet(self.exclude)) annotation_overrides: dict[str, Any] = self.annotation_overrides try: base_annotations = get_type_hints(base, include_extras=True) @@ -384,47 +487,23 @@ def alias(self, name: str) -> str | None: return self.alias_generator(name) return None - def is_field_included(self, field: str | DTOFieldDefinition[Any, Any]) -> bool: + def is_field_included( + self, field: FieldSelector | DTOFieldDefinition[Any, Any], scope: ConfigScope | None = None + ) -> bool: """Whether a field is included per the include/exclude rules. `field` is a field name or a `DTOFieldDefinition`. A bare `str` is treated as a non-relation field name. """ - name, is_relation = (field, False) if isinstance(field, str) else (field.model_field_name, field.is_relation) - if self.include == "all": - return not include_field(name, is_relation, self.exclude) - if self.global_include == "all": - return not include_field(name, is_relation, self.global_exclude) - return ( - include_field(name, is_relation, self.include) or include_field(name, is_relation, self.global_include) - ) and not ( - include_field(name, is_relation, self.exclude) or include_field(name, is_relation, self.global_exclude) - ) - - def __or__(self, other: DTOConfig) -> DTOConfig: - return self.union(other) - + if scope == "local": + return field in FieldSet(self.include) and field not in FieldSet(self.exclude) + if scope == "global": + # `include="all"` propagates inclusion to children, same as `global_include="all"`. + included = field in FieldSet(self.global_include) or self.include == "all" + return included and field not in FieldSet(self.global_exclude) + return field in self.included_fields and field not in self.excluded_fields -@overload -def cast_include_fields(value: FieldGroupStr) -> FieldGroupStr: ... - -@overload -def cast_include_fields( - value: frozenset[FieldSelector] | set[FieldSelector] | list[FieldSelector] | tuple[FieldSelector, ...] | None, -) -> frozenset[FieldSelector]: ... - - -def cast_include_fields(value: IncludeFields | None) -> frozenset[FieldSelector] | FieldGroupStr: - match value: - case None: - return frozenset() - case "all" | "scalars" | "relationships": - return "all" - case _: - return frozenset(value) - - -def is_fields_iterable(value: Any) -> TypeIs[IncludeFields | FieldIterable | None]: +def is_fields_iterable(value: Any) -> TypeIs[FieldSpec | FieldIterable | None]: """Test the given value is suitable to be used as either `include` or `exclude` in a DTOConfig.""" if value == "all" or value is None: return True @@ -433,21 +512,6 @@ def is_fields_iterable(value: Any) -> TypeIs[IncludeFields | FieldIterable | Non return isinstance(value, (frozenset, set, list, tuple)) -def include_field(field_name: str, is_relation: bool, fields: IncludeFields | FieldIterable) -> bool: - """Whether `fields` selects the field.""" - if fields == "all": - return True - group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS - return field_name in fields or group in fields or FieldGroup.ALL in fields - - -def has_field_group(value: IncludeFields | FieldIterable) -> bool: - """True if the selection contains a FieldGroup member.""" - if isinstance(value, str): - return value in FieldGroup.values() - return any(isinstance(item, FieldGroup) for item in value) - - ALL = FieldGroup.ALL SCALARS = FieldGroup.SCALARS RELATIONSHIPS = FieldGroup.RELATIONSHIPS diff --git a/src/strawchemy/dto/utils.py b/src/strawchemy/dto/utils.py index af815036..e14278ab 100644 --- a/src/strawchemy/dto/utils.py +++ b/src/strawchemy/dto/utils.py @@ -15,15 +15,7 @@ from typing import TYPE_CHECKING, Any from strawchemy.dto.constants import DTO_INFO_KEY -from strawchemy.dto.types import ( - DTOConfig, - DTOFieldConfig, - DTOScope, - FieldIterable, - IncludeFields, - Purpose, - PurposeConfig, -) +from strawchemy.dto.types import DTOConfig, DTOFieldConfig, DTOScope, FieldSpec, Purpose, PurposeConfig if TYPE_CHECKING: from collections.abc import Callable, Mapping @@ -44,10 +36,10 @@ def config( purpose: Purpose, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, - global_include: IncludeFields | None = None, - global_exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, + global_include: FieldSpec | None = None, + global_exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, @@ -55,24 +47,19 @@ def config( scope: DTOScope | None = None, tags: set[str] | None = None, ) -> DTOConfig: - config = DTOConfig(purpose, alias_generator=alias_generator, scope=scope) - if exclude: - config.exclude = exclude - if include: - config.include = include - if global_include: - config.global_include = global_include - if global_exclude: - config.global_exclude = global_exclude - if type_map: - config.type_overrides = type_map - if aliases: - config.aliases = aliases - if partial is not None: - config.partial = partial - if tags: - config.tags = tags - return config + return DTOConfig( + purpose, + include=include or None, + exclude=exclude or None, + global_include=global_include or None, + global_exclude=global_exclude or None, + partial=partial, + type_overrides=type_map or {}, + aliases=aliases or {}, + alias_generator=alias_generator, + scope=scope, + tags=tags or set(), + ) def field( diff --git a/src/strawchemy/mapper.py b/src/strawchemy/mapper.py index 819aa9e2..92eab240 100644 --- a/src/strawchemy/mapper.py +++ b/src/strawchemy/mapper.py @@ -44,7 +44,7 @@ from strawberry.extensions.field_extension import FieldExtension from strawberry.types.arguments import StrawberryArgument - from strawchemy.dto.types import IncludeFields + from strawchemy.dto.types import FieldSpec from strawchemy.repository.typing import QueryHookCallable from strawchemy.schema.pagination import DefaultOffsetPagination from strawchemy.transpiler.hook import QueryHook @@ -180,9 +180,9 @@ def field( resolver: Any, *, filter_input: type[BooleanFilterDTO] | bool | None = None, - order_by: IncludeFields | type[OrderByDTO] | None = None, + order_by: FieldSpec | type[OrderByDTO] | None = None, pagination: bool | DefaultOffsetPagination | None = None, - distinct_on: IncludeFields | type[EnumDTO] | None = None, + distinct_on: FieldSpec | type[EnumDTO] | None = None, arguments: list[StrawberryArgument] | None = None, id_field_name: str | None = None, root_aggregations: bool = False, @@ -208,9 +208,9 @@ def field( self, *, filter_input: type[BooleanFilterDTO] | bool | None = None, - order_by: IncludeFields | type[OrderByDTO] | None = None, + order_by: FieldSpec | type[OrderByDTO] | None = None, pagination: bool | DefaultOffsetPagination | None = None, - distinct_on: IncludeFields | type[EnumDTO] | None = None, + distinct_on: FieldSpec | type[EnumDTO] | None = None, arguments: list[StrawberryArgument] | None = None, id_field_name: str | None = None, root_aggregations: bool = False, @@ -236,9 +236,9 @@ def field( resolver: Any | None = None, *, filter_input: type[BooleanFilterDTO] | bool | None = None, - order_by: IncludeFields | type[OrderByDTO] | None = None, + order_by: FieldSpec | type[OrderByDTO] | None = None, pagination: bool | DefaultOffsetPagination | None = None, - distinct_on: IncludeFields | type[EnumDTO] | None = None, + distinct_on: FieldSpec | type[EnumDTO] | None = None, arguments: list[StrawberryArgument] | None = None, id_field_name: str | None = None, root_aggregations: bool = False, diff --git a/src/strawchemy/schema/factories/_kwargs.py b/src/strawchemy/schema/factories/_kwargs.py index 2ad37a2d..f34fcb34 100644 --- a/src/strawchemy/schema/factories/_kwargs.py +++ b/src/strawchemy/schema/factories/_kwargs.py @@ -20,7 +20,7 @@ from strawchemy.dto.base import DTOFieldDefinition from strawchemy.dto.strawberry import BooleanFilterDTO, DTOKey, GraphQLFieldDefinition, OrderByDTO - from strawchemy.dto.types import FieldIterable, IncludeFields + from strawchemy.dto.types import FieldSpec from strawchemy.schema.pagination import DefaultOffsetPagination from strawchemy.typing import GraphQLPurpose @@ -41,8 +41,8 @@ class DTOConfigKwargs(TypedDict, total=False): """Args forwarded to ``config(...)`` to build a ``DTOConfig``.""" - include: IncludeFields | None - exclude: FieldIterable | None + include: FieldSpec | None + exclude: FieldSpec | None partial: bool | None type_map: Mapping[Any, Any] | None aliases: Mapping[str, str] | None @@ -70,11 +70,11 @@ class RegistrationKwargs(TypedDict, total=False): class TypeWrapperKwargs(TypedDict, total=False): """Args specific to ``.type()`` / ``_type_wrapper``.""" - paginate: IncludeFields | None + paginate: FieldSpec | None default_pagination: DefaultOffsetPagination | None filter_input: type[BooleanFilterDTO] | None - distinct_on: IncludeFields | None - order: IncludeFields | type[OrderByDTO] | None + distinct_on: FieldSpec | None + order: FieldSpec | type[OrderByDTO] | None query_hook: Any diff --git a/src/strawchemy/schema/factories/base.py b/src/strawchemy/schema/factories/base.py index 95c49e53..7c4b51a3 100644 --- a/src/strawchemy/schema/factories/base.py +++ b/src/strawchemy/schema/factories/base.py @@ -55,7 +55,7 @@ from strawchemy import Strawchemy from strawchemy.dto.inspectors import SQLAlchemyGraphQLInspector - from strawchemy.dto.types import FieldIterable, IncludeFields + from strawchemy.dto.types import FieldSpec from strawchemy.schema.factories._kwargs import ( InputDecoratorKwargs, MakeInputKwargs, @@ -174,8 +174,8 @@ def _resolve_config(self, dto_config: DTOConfig, base: type[Any]) -> DTOConfig: def _config( self, purpose: Purpose, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, @@ -199,7 +199,7 @@ def _config( ) def _type_order_by( - self, model: type[DeclarativeBase], include: IncludeFields | type[OrderByDTO] | None = None + self, model: type[DeclarativeBase], include: FieldSpec | type[OrderByDTO] | None = None ) -> type[OrderByDTO] | None: order_include = self._mapper.config.order_by if include is None else include if is_fields_iterable(order_include) and order_include is not None: @@ -223,7 +223,7 @@ def _type_order_by( return order_by_input def _type_distinct_on( - self, model: type[DeclarativeBase], include: IncludeFields | type[EnumDTO] | None = None + self, model: type[DeclarativeBase], include: FieldSpec | type[EnumDTO] | None = None ) -> type[EnumDTO] | None: distinct_on_include = self._mapper.config.distinct_on if include is None else include if is_fields_iterable(distinct_on_include) and distinct_on_include is not None: @@ -250,17 +250,17 @@ def _type_wrapper( model: type[T], *, mode: GraphQLPurpose, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, alias_generator: Callable[[str], str] | None = None, - paginate: IncludeFields | None = None, - distinct_on: IncludeFields | None = None, + paginate: FieldSpec | None = None, + distinct_on: FieldSpec | None = None, default_pagination: None | DefaultOffsetPagination = None, filter_input: type[BooleanFilterDTO] | None = None, - order: IncludeFields | type[OrderByDTO] | None = None, + order: FieldSpec | type[OrderByDTO] | None = None, name: str | None = None, description: str | None = None, directives: Sequence[object] | None = (), @@ -323,8 +323,8 @@ def _input_wrapper( model: type[T], *, mode: GraphQLPurpose, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, @@ -485,7 +485,7 @@ def _root_input_config(self, model: builtins.type[Any], dto_config: DTOConfig, m id_fields = self.inspector.id_field_definitions(model, dto_config) # Add PKs for update/delete inputs if mode == "update_by_pk_input": - if set(dto_config.exclude) & {name for name, _ in id_fields}: + if set(dto_config.excluded_fields) & {name for name, _ in id_fields}: msg = ( "You cannot exclude primary key columns from an input type intended for create or update mutations" ) @@ -496,9 +496,10 @@ def _root_input_config(self, model: builtins.type[Any], dto_config: DTOConfig, m if mode in {"update_by_pk_input", "update_by_filter_input"}: partial = True # Exclude default generated PKs for create inputs, if not explicitly included - elif dto_config.include == "all": + else: for name, field in id_fields: - if self.inspector.has_default(field.model_field): + # Exclude rules are deliberately ignored: default-generated PKs stay as optional inputs. + if field in dto_config.included_fields and self.inspector.has_default(field.model_field): annotations_overrides[name] = Optional[field.type_hint] return dto_config.copy_with( annotation_overrides=annotations_overrides, diff --git a/src/strawchemy/schema/factories/types.py b/src/strawchemy/schema/factories/types.py index 844c7c58..4ad9e3b0 100644 --- a/src/strawchemy/schema/factories/types.py +++ b/src/strawchemy/schema/factories/types.py @@ -23,7 +23,7 @@ MappedStrawberryGraphQLDTO, OrderByDTO, ) -from strawchemy.dto.types import DTOConfig, DTOMissing, IncludeFields, Purpose, is_fields_iterable +from strawchemy.dto.types import DTOConfig, DTOMissing, FieldSpec, Purpose, is_fields_iterable from strawchemy.dto.utils import read_partial, write_all_config from strawchemy.exceptions import EmptyDTOError from strawchemy.schema.factories import ( @@ -175,9 +175,9 @@ def _add_fields_arguments( self, dto: type[GraphQLDTOT], base: type[Any] | None, - order: IncludeFields | None = None, - paginate: IncludeFields | None = None, - distinct_on: IncludeFields | None = None, + order: FieldSpec | None = None, + paginate: FieldSpec | None = None, + distinct_on: FieldSpec | None = None, default_pagination: None | DefaultOffsetPagination = None, ) -> type[GraphQLDTOT]: """Add pagination and ordering arguments to a GraphQL DTO type. diff --git a/src/strawchemy/schema/field.py b/src/strawchemy/schema/field.py index fde91138..6e939bd7 100644 --- a/src/strawchemy/schema/field.py +++ b/src/strawchemy/schema/field.py @@ -21,7 +21,7 @@ OrderByDTO, StrawchemyObject, ) -from strawchemy.dto.types import DTOConfig, IncludeFields, Purpose +from strawchemy.dto.types import DTOConfig, FieldSpec, Purpose from strawchemy.exceptions import EmptyDTOError, StrawchemyFieldError from strawchemy.schema.pagination import DefaultOffsetPagination from strawchemy.utils.annotation import is_type_hint_optional @@ -86,8 +86,8 @@ def __init__( filter_factory: BooleanFilterFactory, distinct_on_factory: DistinctOnEnumFactory, filter_type: builtins.type[BooleanFilterDTO] | bool | None = None, - order_by: IncludeFields | builtins.type[OrderByDTO] | Literal[False] | None = None, - distinct_on: IncludeFields | builtins.type[EnumDTO] | Literal[False] | None = None, + order_by: FieldSpec | builtins.type[OrderByDTO] | Literal[False] | None = None, + distinct_on: FieldSpec | builtins.type[EnumDTO] | Literal[False] | None = None, pagination: DefaultOffsetPagination | bool | None = False, repository_type: AnyRepositoryType | None = None, root_aggregations: bool = False, @@ -263,7 +263,7 @@ def distinct_on(self) -> builtins.type[EnumDTO] | None: return self._distinct_on_factory.factory( inner_type.__dto_model__, dto_config=inner_type.__dto_config__.copy_with( - include=inner_type.__dto_config__.include if distinct_on == "all" else distinct_on + include=inner_type.__dto_config__.included_fields & distinct_on ), no_cache=True, if_no_fields="raise", @@ -294,7 +294,7 @@ def order_by(self) -> builtins.type[OrderByDTO] | None: inner_type.__dto_model__, mode="order_by", dto_config=inner_type.__dto_config__.copy_with( - include=inner_type.__dto_config__.include if order_by == "all" else order_by + include=inner_type.__dto_config__.included_fields & order_by ), no_cache=True, if_no_fields="raise", diff --git a/src/strawchemy/utils/registry.py b/src/strawchemy/utils/registry.py index 0c5e5ce6..0a6d619d 100644 --- a/src/strawchemy/utils/registry.py +++ b/src/strawchemy/utils/registry.py @@ -15,7 +15,7 @@ from strawberry.types.union import StrawberryUnion from strawchemy.dto.strawberry import MappedStrawberryGraphQLDTO -from strawchemy.dto.types import FieldGroupStr, cast_include_fields, is_fields_iterable +from strawchemy.dto.types import FieldSet, is_fields_iterable from strawchemy.exceptions import StrawchemyError from strawchemy.utils.annotation import inner_types from strawchemy.utils.strawberry import strawberry_contained_types @@ -39,7 +39,7 @@ from strawchemy.dto import DTOConfig from strawchemy.dto.base import Node, Relation from strawchemy.dto.strawberry import EnumDTO, OrderByDTO, StrawchemyObject - from strawchemy.dto.types import DTOScope, FieldSelector, IncludeFields + from strawchemy.dto.types import DTOScope, FieldSpec from strawchemy.schema.pagination import DefaultOffsetPagination from strawchemy.typing import GraphQLType, StrawchemyObjectWithStrawberryObjectDefinition @@ -149,9 +149,9 @@ class RegistryTypeInfo: user_defined: bool = False override: bool = False pagination: DefaultOffsetPagination | None = None - order: frozenset[FieldSelector] | FieldGroupStr | type[OrderByDTO] = dataclasses.field(default_factory=frozenset) - distinct_on: frozenset[FieldSelector] | FieldGroupStr | type[EnumDTO] = dataclasses.field(default_factory=frozenset) - paginate: frozenset[FieldSelector] | FieldGroupStr = dataclasses.field(default_factory=frozenset) + order: FieldSet | type[OrderByDTO] = dataclasses.field(default_factory=frozenset) + distinct_on: FieldSet | type[EnumDTO] = dataclasses.field(default_factory=frozenset) + paginate: FieldSet = dataclasses.field(default_factory=frozenset) scope: DTOScope | None = None model: type[DeclarativeBase] | None = None tags: frozenset[str] = dataclasses.field(default_factory=frozenset) @@ -349,9 +349,9 @@ def _type_info( current_node: Node[Relation[Any, Any], None] | None, override: bool = False, user_defined: bool = False, - paginate: IncludeFields | None = None, - order: IncludeFields | type[OrderByDTO] | None = None, - distinct_on: IncludeFields | type[EnumDTO] | None = None, + paginate: FieldSpec | None = None, + order: FieldSpec | type[OrderByDTO] | None = None, + distinct_on: FieldSpec | type[EnumDTO] | None = None, default_pagination: DefaultOffsetPagination | None = None, default_name: str | None = None, ) -> RegistryTypeInfo: @@ -365,9 +365,9 @@ def _type_info( override=override, user_defined=user_defined, pagination=default_pagination, - order=cast_include_fields(order) if is_fields_iterable(order) else order, - distinct_on=cast_include_fields(distinct_on) if is_fields_iterable(distinct_on) else distinct_on, - paginate=cast_include_fields(paginate), + order=FieldSet(order) if is_fields_iterable(order) else order, + distinct_on=FieldSet(distinct_on) if is_fields_iterable(distinct_on) else distinct_on, + paginate=FieldSet(paginate), scope=dto_config.scope, model=model, exclude_from_scope=dto_config.exclude_from_scope, @@ -420,9 +420,9 @@ def register_type( current_node: Node[Relation[Any, Any], None] | None = None, override: bool = False, user_defined: bool = False, - paginate: IncludeFields | None = None, - order: IncludeFields | type[OrderByDTO] | None = None, - distinct_on: IncludeFields | type[EnumDTO] | None = None, + paginate: FieldSpec | None = None, + order: FieldSpec | type[OrderByDTO] | None = None, + distinct_on: FieldSpec | type[EnumDTO] | None = None, default_pagination: DefaultOffsetPagination | None = None, default_name: str | None = None, description: str | None = None, diff --git a/src/strawchemy/validation/pydantic.py b/src/strawchemy/validation/pydantic.py index f72a842c..59f69a0f 100644 --- a/src/strawchemy/validation/pydantic.py +++ b/src/strawchemy/validation/pydantic.py @@ -24,7 +24,7 @@ from strawchemy import Strawchemy from strawchemy.dto.base import DTOFieldDefinition, MappedDTO, Relation - from strawchemy.dto.types import DTOConfig, FieldIterable, IncludeFields, Purpose + from strawchemy.dto.types import DTOConfig, FieldSpec, Purpose from strawchemy.repository.typing import DeclarativeT from strawchemy.schema.factories._kwargs import FactoryMethodKwargs from strawchemy.typing import GraphQLPurpose @@ -84,8 +84,8 @@ def input( model: type[DeclarativeT], *, mode: GraphQLPurpose, - include: IncludeFields | None = None, - exclude: FieldIterable | None = None, + include: FieldSpec | None = None, + exclude: FieldSpec | None = None, partial: bool | None = None, type_map: Mapping[Any, Any] | None = None, aliases: Mapping[str, str] | None = None, diff --git a/tests/unit/dto/test_dto.py b/tests/unit/dto/test_dto.py index 4ef11f9b..87456e5e 100644 --- a/tests/unit/dto/test_dto.py +++ b/tests/unit/dto/test_dto.py @@ -1,7 +1,7 @@ from __future__ import annotations import re -from typing import Optional +from typing import Any, Optional, get_args from uuid import UUID, uuid4 import pytest @@ -12,11 +12,11 @@ import strawchemy from strawchemy import ALL, RELATIONSHIPS, SCALARS from strawchemy.dto import DTOConfig, Purpose, PurposeConfig, config, field -from strawchemy.dto.base import DTOFieldDefinition from strawchemy.dto.constants import DTO_INFO_KEY from strawchemy.dto.strawberry import DTOKey, GraphQLFieldDefinition, StrawchemyDefinition -from strawchemy.dto.types import FieldGroup, include_field +from strawchemy.dto.types import FieldGroup from strawchemy.dto.utils import DTOFieldConfig, read_all_config, write_all_config +from strawchemy.exceptions import EmptyDTOError from tests.typing import AnyFactory, MappedPydanticFactory from tests.unit.dc_models import ( AdminDataclass, @@ -30,16 +30,6 @@ from tests.utils import DTOInspect, factory_iterator -def _fruit_field(name: str, *, is_relation: bool) -> DTOFieldDefinition: # type: ignore[type-arg] - return DTOFieldDefinition( - dto_config=DTOConfig(Purpose.READ), - model=Fruit, - model_field_name=name, - type_hint=int, - is_relation=is_relation, - ) - - class _PopulateFieldsBase(DeclarativeBase): pass @@ -49,8 +39,14 @@ class _PopulateFieldsModel(_PopulateFieldsBase): id: Mapped[int] = mapped_column(Integer, primary_key=True) -def test_config_function_produces_same_default() -> None: - assert config(Purpose.READ) == DTOConfig(Purpose.READ) +@pytest.mark.parametrize("model", [Tomato, TomatoDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_config_function_produces_same_default(factory: AnyFactory, model: type[Tomato | TomatoDataclass]) -> None: + """Test that config() and DTOConfig produce DTOs with identical fields.""" + from_function = factory.factory(model, config(Purpose.READ, include="all"), name="FromFunction") + from_class = factory.factory(model, DTOConfig(Purpose.READ, include="all"), name="FromClass") + + assert DTOInspect(from_function).annotations() == DTOInspect(from_class).annotations() def test_default_field_config() -> None: @@ -294,73 +290,48 @@ def test_forward_refs_resolved(name: str, sqlalchemy_pydantic_factory: MappedPyd ) -# Tests for DTOConfig.from_include() and is_field_included() - - -def test_from_include_with_none() -> None: - """Test that from_include(None) creates a config with empty include set.""" - config = DTOConfig.from_include(None) - assert config.include == set() - assert config.purpose == Purpose.READ - - -def test_from_include_with_all() -> None: - """Test that from_include('all') creates a config with include='all'.""" - config = DTOConfig.from_include("all") - assert config.include == "all" - assert config.purpose == Purpose.READ - - -def test_from_include_with_list() -> None: - """Test that from_include() accepts a list and converts it to the include parameter.""" - config = DTOConfig.from_include(["field1", "field2"]) - assert config.include == ["field1", "field2"] - assert config.purpose == Purpose.READ - - -def test_from_include_with_set() -> None: - """Test that from_include() accepts a set for the include parameter.""" - config = DTOConfig.from_include({"field1", "field2"}) - assert config.include == {"field1", "field2"} - assert config.purpose == Purpose.READ - - -def test_from_include_with_custom_purpose() -> None: - """Test that from_include() accepts a custom purpose.""" - config = DTOConfig.from_include(["field1"], purpose=Purpose.WRITE) - assert config.include == ["field1"] - assert config.purpose == Purpose.WRITE - - -def test_is_field_included_with_all() -> None: - """Test that is_field_included() returns True for any field when include='all'.""" - config = DTOConfig.from_include("all") - assert config.is_field_included("any_field") is True - assert config.is_field_included("another_field") is True +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_from_include_empty_raises(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: + """Test that from_include(None) produces a DTO with no fields.""" + with pytest.raises(EmptyDTOError): + factory.factory(model, DTOConfig.from_include(None), if_no_fields="raise") -def test_is_field_included_with_specific_list() -> None: - """Test that is_field_included() returns True only for listed fields.""" - config = DTOConfig.from_include(["field1", "field2"]) - assert config.is_field_included("field1") is True - assert config.is_field_included("field2") is True - assert config.is_field_included("field3") is False +@pytest.mark.parametrize( + ("include_spec", "expected_fields"), + [ + pytest.param("all", {"name", "color_id", "sweetness", "id", "color"}, id="all"), + pytest.param(["name", "sweetness"], {"name", "sweetness"}, id="list"), + pytest.param({"name", "sweetness"}, {"name", "sweetness"}, id="set"), + ], +) +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_from_include_spec( + factory: AnyFactory, + model: type[Fruit | FruitDataclass], + include_spec: Any, + expected_fields: set[str], +) -> None: + """Test that from_include() accepts 'all' and field-name collections.""" + dto = factory.factory(model, DTOConfig.from_include(include_spec)) + assert set(DTOInspect(dto).annotations()) == expected_fields -def test_is_field_included_with_empty_include() -> None: - """Test that is_field_included() returns False for all fields when include is empty.""" - config = DTOConfig.from_include(None) - assert config.is_field_included("field1") is False - assert config.is_field_included("any_field") is False +@pytest.mark.parametrize("factory", factory_iterator()) +def test_from_include_with_custom_purpose(factory: AnyFactory) -> None: + """Test that from_include() honors the purpose: read-only fields are dropped from write DTOs.""" + dto = factory.factory(Book, DTOConfig.from_include(["title", "isbn"], purpose=Purpose.WRITE)) + assert set(DTOInspect(dto).annotations()) == {"title"} -def test_is_field_included_with_exclude() -> None: - """Test that excluded fields are properly excluded even when include='all'.""" - config = DTOConfig(Purpose.READ, include="all", exclude={"field2", "field3"}) - assert config.is_field_included("field1") is True - assert config.is_field_included("field2") is False - assert config.is_field_included("field3") is False - assert config.is_field_included("field4") is True +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_named_exclude_with_include_all(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: + """Test that excluded fields are dropped from the DTO even when include='all'.""" + dto = factory.factory(model, DTOConfig(Purpose.READ, include="all", exclude={"color", "color_id"})) + assert set(DTOInspect(dto).annotations()) == {"name", "sweetness", "id"} @pytest.mark.parametrize( @@ -384,182 +355,134 @@ def test_strawchemy_definition_populate_fields(key_source: type[DeclarativeBase] assert definition.field_map == {DTOKey([_PopulateFieldsModel]) + "id": field_def} -def test_field_group_constants_are_enum_members() -> None: - """Test that the top-level SCALARS/RELATIONSHIPS constants are the enum members.""" +def test_field_group_constants() -> None: + """Test that group constants are the enum members, equal their string literals, and are hashable.""" assert strawchemy.SCALARS is FieldGroup.SCALARS assert strawchemy.RELATIONSHIPS is FieldGroup.RELATIONSHIPS - - -def test_field_group_constants_hashable_in_sets() -> None: - """Test that group constants are usable inside include/exclude frozensets.""" + assert FieldGroup.ALL == "all" + assert FieldGroup.SCALARS == "scalars" + assert FieldGroup.RELATIONSHIPS == "relationships" members = frozenset([strawchemy.SCALARS, strawchemy.RELATIONSHIPS, "name"]) assert strawchemy.SCALARS in members assert strawchemy.RELATIONSHIPS in members assert "name" in members -def test_scalars_include_allows_exclude() -> None: +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_scalars_include_allows_exclude(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: """Test that a group-bearing include coexists with exclude and is not clobbered to 'all'.""" - config = DTOConfig(Purpose.READ, include=[SCALARS], exclude=["secret"]) - assert config.include != "all" - assert SCALARS in config.include - assert "secret" in config.exclude + dto = factory.factory(model, DTOConfig(Purpose.READ, include=[SCALARS], exclude=["name"])) + assert set(DTOInspect(dto).annotations()) == {"id", "color_id", "sweetness"} -def test_plain_include_with_exclude_still_raises() -> None: +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_plain_include_with_exclude_still_raises(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: """Test that a plain field-name include combined with exclude still raises.""" with pytest.raises(ValueError, match="exclude"): - DTOConfig(Purpose.READ, include=["a", "b"], exclude=["c"]) + factory.factory(model, DTOConfig(Purpose.READ, include=["name"], exclude=["sweetness"])) -def test_bare_exclude_still_implies_all() -> None: +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_bare_exclude_still_implies_all(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: """Test that a bare exclude (no include) still implies include='all'.""" - config = DTOConfig(Purpose.READ, exclude=["secret"]) - assert config.include == "all" - - -def test_relationships_exclude_implies_all_include() -> None: - """Test that exclude=[RELATIONSHIPS] with no include promotes include to 'all'.""" - config = DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS]) - assert config.include == "all" - assert RELATIONSHIPS in config.exclude + dto = factory.factory(model, DTOConfig(Purpose.READ, exclude=["name"])) + assert set(DTOInspect(dto).annotations()) == {"id", "color_id", "sweetness", "color"} -def test_mixed_group_and_name_include_allows_exclude() -> None: +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) +@pytest.mark.parametrize("factory", factory_iterator()) +def test_mixed_group_and_name_include_allows_exclude(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: """Test that a group selector mixed with a field name coexists with exclude.""" - config = DTOConfig(Purpose.READ, include=[SCALARS, "owner"], exclude=["secret"]) - assert config.include != "all" - assert SCALARS in config.include - assert "owner" in config.include - assert "secret" in config.exclude - - -def test_global_group_include_allows_global_exclude() -> None: - """Test that a group-bearing global_include coexists with global_exclude.""" - config = DTOConfig(Purpose.READ, global_include=[SCALARS], global_exclude=["secret"]) - assert config.global_include != "all" - assert SCALARS in config.global_include - assert "secret" in config.global_exclude + dto = factory.factory(model, DTOConfig(Purpose.READ, include=[SCALARS, "color"], exclude=["name"])) + assert set(DTOInspect(dto).annotations()) == {"id", "color_id", "sweetness", "color"} +@pytest.mark.parametrize("model", [Fruit, FruitDataclass]) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_scalars_excludes_relationships(factory: AnyFactory) -> None: - """Test that include=[SCALARS] keeps scalar fields and drops relationships.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS])) - fields = set(DTOInspect(dto).annotations()) - assert "color" not in fields - assert {"id", "name", "sweetness", "color_id"} <= fields +def test_global_group_include_allows_global_exclude(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: + """Test that group-bearing global_include/global_exclude shape nested DTOs.""" + dto = factory.factory( + model, DTOConfig(Purpose.READ, include=[SCALARS, "color"], global_include=[SCALARS], global_exclude=["name"]) + ) + annotations = DTOInspect(dto).annotations() + # Root fields follow `include`; global rules don't apply at the root. + assert set(annotations) == {"id", "name", "color_id", "sweetness", "color"} + # Nested DTO follows global rules: scalars only, minus the global exclude. + color_dto = next((arg for arg in get_args(annotations["color"]) if arg is not type(None)), annotations["color"]) + assert set(DTOInspect(color_dto).annotations()) == {"id"} +@pytest.mark.parametrize( + "include_spec", + [ + pytest.param({"all"}, id="literal-set"), + pytest.param([ALL], id="constant-list"), + pytest.param([SCALARS, RELATIONSHIPS], id="both-groups"), + ], +) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_scalars_plus_named_relationship(factory: AnyFactory) -> None: - """Test that include=[SCALARS, 'color'] keeps scalars plus the named relationship.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS, "color"])) - fields = set(DTOInspect(dto).annotations()) - assert "color" in fields - assert {"id", "name", "sweetness", "color_id"} <= fields +def test_include_all_equivalents(factory: AnyFactory, include_spec: Any) -> None: + """Test that {'all'}, [ALL] and [SCALARS, RELATIONSHIPS] are all equivalent to include='all'.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) + assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness", "color"} +@pytest.mark.parametrize( + "include_spec", + [ + pytest.param([SCALARS], id="constant-list"), + pytest.param(frozenset([SCALARS]), id="constant-frozenset"), + pytest.param(["scalars"], id="string-literal"), + ], +) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_relationships_only(factory: AnyFactory) -> None: - """Test that include=[RELATIONSHIPS] keeps only relationships and drops scalars.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[RELATIONSHIPS])) - fields = set(DTOInspect(dto).annotations()) - assert "color" in fields - assert "name" not in fields - assert "sweetness" not in fields - assert "color_id" not in fields +def test_include_scalars(factory: AnyFactory, include_spec: Any) -> None: + """Test that include=[SCALARS] keeps scalar fields and drops relationships.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) + assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness"} +@pytest.mark.parametrize( + "include_spec", + [ + pytest.param([RELATIONSHIPS], id="constant-list"), + pytest.param(frozenset([RELATIONSHIPS]), id="constant-frozenset"), + pytest.param(["relationships"], id="string-literal"), + ], +) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_both_groups_equals_all(factory: AnyFactory) -> None: - """Test that include=[SCALARS, RELATIONSHIPS] is equivalent to include='all'.""" - grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[SCALARS, RELATIONSHIPS])) - all_dto = factory.factory(Fruit, read_all_config) - assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) +def test_include_relationships(factory: AnyFactory, include_spec: Any) -> None: + """Test that include=[RELATIONSHIPS] keeps only relationships and drops scalars.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) + assert set(DTOInspect(dto).annotations()) == {"color"} +@pytest.mark.parametrize( + "exclude_spec", + [ + pytest.param([RELATIONSHIPS], id="constant-list"), + pytest.param(frozenset([RELATIONSHIPS]), id="constant-frozenset"), + ], +) @pytest.mark.parametrize("factory", factory_iterator()) -def test_exclude_relationships_keeps_scalars(factory: AnyFactory) -> None: - """Test that exclude=[RELATIONSHIPS] keeps all scalar fields and walks no relationships.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS])) - fields = set(DTOInspect(dto).annotations()) - assert "color" not in fields - assert {"id", "name", "sweetness", "color_id"} <= fields +def test_exclude_relationships(factory: AnyFactory, exclude_spec: Any) -> None: + """Test that a bare exclude=[RELATIONSHIPS] implies include='all' and drops relations.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=exclude_spec)) + assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness"} @pytest.mark.parametrize("factory", factory_iterator()) def test_exclude_relationships_plus_named_scalar(factory: AnyFactory) -> None: """Test that exclude=[RELATIONSHIPS, 'sweetness'] drops relationships and the named scalar.""" dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=[RELATIONSHIPS, "sweetness"])) - fields = set(DTOInspect(dto).annotations()) - assert "color" not in fields - assert "sweetness" not in fields - assert {"id", "name", "color_id"} <= fields - - -def test_is_field_included_relationships_group_matches_relations() -> None: - """Test that include=[RELATIONSHIPS] includes relation fields and excludes scalars.""" - config = DTOConfig(Purpose.READ, include=frozenset([RELATIONSHIPS])) - assert config.is_field_included(_fruit_field("owner", is_relation=True)) is True - assert config.is_field_included(_fruit_field("name", is_relation=False)) is False - - -def test_is_field_included_scalars_group_matches_scalars() -> None: - """Test that include=[SCALARS] includes scalar fields and excludes relations.""" - config = DTOConfig(Purpose.READ, include=frozenset([SCALARS])) - assert config.is_field_included(_fruit_field("name", is_relation=False)) is True - assert config.is_field_included(_fruit_field("owner", is_relation=True)) is False - - -def test_is_field_included_group_in_exclude() -> None: - """Test that a group selector in exclude drops matching fields.""" - config = DTOConfig(Purpose.READ, include="all", exclude=frozenset([RELATIONSHIPS])) - assert config.is_field_included(_fruit_field("owner", is_relation=True)) is False - assert config.is_field_included(_fruit_field("name", is_relation=False)) is True - - -def test_is_field_included_no_group_unchanged() -> None: - """Test that behavior is unchanged when no group constants are present.""" - config = DTOConfig(Purpose.READ, include=frozenset(["a", "b"])) - assert config.is_field_included("a") is True - assert config.is_field_included("c") is False - assert config.is_field_included(_fruit_field("a", is_relation=True)) is True - - -def test_string_literals_equal_field_group_members() -> None: - """Test that plain string literals equal their FieldGroup members.""" - assert FieldGroup.ALL == "all" - assert FieldGroup.SCALARS == "scalars" - assert FieldGroup.RELATIONSHIPS == "relationships" + assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id"} @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_string_scalars_literal(factory: AnyFactory) -> None: - """Test that the literal include=['scalars'] behaves like include=[SCALARS].""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=["scalars"])) - fields = set(DTOInspect(dto).annotations()) - assert "color" not in fields - assert {"id", "name", "sweetness", "color_id"} <= fields - - -@pytest.mark.parametrize("factory", factory_iterator()) -def test_include_all_literal_set_equals_all(factory: AnyFactory) -> None: - """Test that include={'all'} is equivalent to include='all'.""" - grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include={"all"})) - all_dto = factory.factory(Fruit, read_all_config) - assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) - - -@pytest.mark.parametrize("factory", factory_iterator()) -def test_include_all_constant_in_collection(factory: AnyFactory) -> None: - """Test that include=[ALL] is equivalent to include='all'.""" - grouped = factory.factory(Fruit, DTOConfig(Purpose.READ, include=[ALL])) - all_dto = factory.factory(Fruit, read_all_config) - assert set(DTOInspect(grouped).annotations()) == set(DTOInspect(all_dto).annotations()) - - -def test_include_field_all_selects_everything() -> None: - """Test that include_field returns True for the 'all' selector.""" - assert include_field("anything", False, "all") is True - assert include_field("anything", True, "all") is True +def test_include_plain_names_with_relationship(factory: AnyFactory) -> None: + """Test that plain field-name includes select scalars and relations by name, without groups.""" + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=frozenset(["name", "color"]))) + assert set(DTOInspect(dto).annotations()) == {"name", "color"} From c060b8e60df1f397bde75f04d4c4ba48e4785eb9 Mon Sep 17 00:00:00 2001 From: gazorby Date: Sat, 6 Jun 2026 23:17:40 +0200 Subject: [PATCH 4/9] refactor(dto): remove str inheritance from FieldGroup and update FieldSet logic --- src/strawchemy/dto/types.py | 33 ++++++---------- tests/unit/dto/test_dto.py | 77 +++++++++++++++++++------------------ 2 files changed, 50 insertions(+), 60 deletions(-) diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 3acd1552..976c81ba 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -47,7 +47,7 @@ ConfigScope: TypeAlias = Literal["local", "global"] -class FieldGroup(str, Enum): +class FieldGroup(Enum): """Field-group selectors for ``include``/``exclude`` sequences.""" ALL = "all" @@ -57,9 +57,6 @@ class FieldGroup(str, Enum): RELATIONSHIPS = "relationships" """Include only relationships.""" - def __bool__(self) -> bool: - return True - @staticmethod def list_str() -> str: return ", ".join(member.value for member in FieldGroup) @@ -113,8 +110,8 @@ def __contains__(self, item: FieldSelector | DTOFieldDefinition[Any, Any]) -> bo if isinstance(item, FieldGroup): return item in self.field_set or FieldGroup.ALL in self.field_set name, is_relation = (item, False) if isinstance(item, str) else (item.model_field_name, item.is_relation) - group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS - return name in self.field_set or group in self.field_set or FieldGroup.ALL in self.field_set + item_group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS + return name in self.field_set or item_group in self.field_set or FieldGroup.ALL in self.field_set def __and__(self, other: FieldSpec) -> FieldIterable: other_set = FieldSet(other) @@ -128,10 +125,10 @@ def __and__(self, other: FieldSpec) -> FieldIterable: ) def __or__(self, other: FieldSpec | None) -> FieldSpec | None: - union = self.field_set | FieldSet(other).field_set - if FieldGroup.ALL in union: + other_set = FieldSet(other) + if FieldGroup.ALL in self.field_set or FieldGroup.ALL in other_set.field_set: return "all" - return union or None + return (self.field_set | other_set.field_set) or None def __bool__(self) -> bool: return bool(self.field_set) @@ -148,10 +145,10 @@ def normalize(cls, value: FieldSpec | None) -> frozenset[FieldSelector]: iterable of field names and/or `FieldGroup` members, or `None`. Returns: - A group string as a single-member frozenset of the matching - `FieldGroup`, `None` as an empty frozenset, and any iterable as a - frozenset of its items. + Normalized field selector set """ + if isinstance(value, FieldGroup): + return frozenset((value,)) if isinstance(value, str) and FieldGroup.is_group(value): return frozenset((FieldGroup(value),)) if value is None: @@ -339,20 +336,12 @@ def __or__(self, other: DTOConfig) -> DTOConfig: @classmethod def _has_field_group(cls, value: FieldSpec | FieldIterable) -> bool: """True if the selection contains a FieldGroup member.""" + if isinstance(value, FieldGroup): + return True if isinstance(value, str): return value in FieldGroup.values() return any(isinstance(item, FieldGroup) for item in value) - @classmethod - def _include_field(cls, field_name: str, is_relation: bool, fields: FieldSpec | FieldIterable | None) -> bool: - """Whether `fields` selects the field.""" - if fields is None: - return False - if fields == "all": - return True - group = FieldGroup.RELATIONSHIPS if is_relation else FieldGroup.SCALARS - return field_name in fields or group in fields or FieldGroup.ALL in fields - def union(self, other: DTOConfig) -> DTOConfig: include = FieldSet(self.include) | other.include exclude = FieldSet(self.exclude) | other.exclude diff --git a/tests/unit/dto/test_dto.py b/tests/unit/dto/test_dto.py index 87456e5e..9741d9c7 100644 --- a/tests/unit/dto/test_dto.py +++ b/tests/unit/dto/test_dto.py @@ -9,12 +9,11 @@ from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column from typing_extensions import Self -import strawchemy from strawchemy import ALL, RELATIONSHIPS, SCALARS from strawchemy.dto import DTOConfig, Purpose, PurposeConfig, config, field from strawchemy.dto.constants import DTO_INFO_KEY from strawchemy.dto.strawberry import DTOKey, GraphQLFieldDefinition, StrawchemyDefinition -from strawchemy.dto.types import FieldGroup +from strawchemy.dto.types import FieldSpec from strawchemy.dto.utils import DTOFieldConfig, read_all_config, write_all_config from strawchemy.exceptions import EmptyDTOError from tests.typing import AnyFactory, MappedPydanticFactory @@ -355,24 +354,17 @@ def test_strawchemy_definition_populate_fields(key_source: type[DeclarativeBase] assert definition.field_map == {DTOKey([_PopulateFieldsModel]) + "id": field_def} -def test_field_group_constants() -> None: - """Test that group constants are the enum members, equal their string literals, and are hashable.""" - assert strawchemy.SCALARS is FieldGroup.SCALARS - assert strawchemy.RELATIONSHIPS is FieldGroup.RELATIONSHIPS - assert FieldGroup.ALL == "all" - assert FieldGroup.SCALARS == "scalars" - assert FieldGroup.RELATIONSHIPS == "relationships" - members = frozenset([strawchemy.SCALARS, strawchemy.RELATIONSHIPS, "name"]) - assert strawchemy.SCALARS in members - assert strawchemy.RELATIONSHIPS in members - assert "name" in members - - +@pytest.mark.parametrize( + "include_spec", + [pytest.param("scalars", id="direct-string"), pytest.param([SCALARS], id="constant-list")], +) @pytest.mark.parametrize("model", [Fruit, FruitDataclass]) @pytest.mark.parametrize("factory", factory_iterator()) -def test_scalars_include_allows_exclude(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: +def test_scalars_include_allows_exclude( + factory: AnyFactory, model: type[Fruit | FruitDataclass], include_spec: FieldSpec +) -> None: """Test that a group-bearing include coexists with exclude and is not clobbered to 'all'.""" - dto = factory.factory(model, DTOConfig(Purpose.READ, include=[SCALARS], exclude=["name"])) + dto = factory.factory(model, DTOConfig(Purpose.READ, include=include_spec, exclude=["name"])) assert set(DTOInspect(dto).annotations()) == {"id", "color_id", "sweetness"} @@ -418,14 +410,14 @@ def test_global_group_include_allows_global_exclude(factory: AnyFactory, model: @pytest.mark.parametrize( "include_spec", [ - pytest.param({"all"}, id="literal-set"), + pytest.param("all", id="all-string"), pytest.param([ALL], id="constant-list"), pytest.param([SCALARS, RELATIONSHIPS], id="both-groups"), ], ) @pytest.mark.parametrize("factory", factory_iterator()) def test_include_all_equivalents(factory: AnyFactory, include_spec: Any) -> None: - """Test that {'all'}, [ALL] and [SCALARS, RELATIONSHIPS] are all equivalent to include='all'.""" + """Test that [ALL] and [SCALARS, RELATIONSHIPS] are equivalent to include='all'.""" dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness", "color"} @@ -433,45 +425,54 @@ def test_include_all_equivalents(factory: AnyFactory, include_spec: Any) -> None @pytest.mark.parametrize( "include_spec", [ - pytest.param([SCALARS], id="constant-list"), - pytest.param(frozenset([SCALARS]), id="constant-frozenset"), - pytest.param(["scalars"], id="string-literal"), + pytest.param(("scalars", ()), id="include-scalars-string"), + pytest.param(([SCALARS], ()), id="include-scalars-list"), + pytest.param(([ALL], [RELATIONSHIPS]), id="exclude-relationships-constant-list"), + pytest.param(("all", "relationships"), id="exclude-relationships-string"), + pytest.param((None, "relationships"), id="bare-exclude-relationships-string"), + pytest.param((None, [RELATIONSHIPS]), id="bare-exclude-relationships-constant-list"), ], ) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_scalars(factory: AnyFactory, include_spec: Any) -> None: - """Test that include=[SCALARS] keeps scalar fields and drops relationships.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) +def test_include_scalars(factory: AnyFactory, include_spec: tuple[FieldSpec | None, FieldSpec]) -> None: + """Test that scalars group selectors keep scalar fields and drop relationships.""" + include, exclude = include_spec + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include, exclude=exclude)) assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness"} @pytest.mark.parametrize( "include_spec", [ - pytest.param([RELATIONSHIPS], id="constant-list"), - pytest.param(frozenset([RELATIONSHIPS]), id="constant-frozenset"), - pytest.param(["relationships"], id="string-literal"), + pytest.param(("relationships", ()), id="include-relationships-string"), + pytest.param(([RELATIONSHIPS], ()), id="include-relationships-list"), + pytest.param(([ALL], [SCALARS]), id="exclude-scalars-constant-list"), + pytest.param(("all", "scalars"), id="exclude-scalars-string"), + pytest.param((None, "scalars"), id="bare-exclude-scalars-string"), + pytest.param((None, [SCALARS]), id="bare-exclude-scalars-constant-list"), ], ) @pytest.mark.parametrize("factory", factory_iterator()) -def test_include_relationships(factory: AnyFactory, include_spec: Any) -> None: - """Test that include=[RELATIONSHIPS] keeps only relationships and drops scalars.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec)) +def test_include_relationships(factory: AnyFactory, include_spec: tuple[FieldSpec | None, FieldSpec]) -> None: + """Test that relationships group selectors keep only relationships and drop scalars.""" + include, exclude = include_spec + dto = factory.factory(Fruit, DTOConfig(Purpose.READ, include=include, exclude=exclude)) assert set(DTOInspect(dto).annotations()) == {"color"} @pytest.mark.parametrize( - "exclude_spec", + "include_spec", [ - pytest.param([RELATIONSHIPS], id="constant-list"), - pytest.param(frozenset([RELATIONSHIPS]), id="constant-frozenset"), + pytest.param(["scalars"], id="scalars-list"), + pytest.param(["relationships"], id="relationships-list"), + pytest.param({"all"}, id="all-set"), ], ) @pytest.mark.parametrize("factory", factory_iterator()) -def test_exclude_relationships(factory: AnyFactory, exclude_spec: Any) -> None: - """Test that a bare exclude=[RELATIONSHIPS] implies include='all' and drops relations.""" - dto = factory.factory(Fruit, DTOConfig(Purpose.READ, exclude=exclude_spec)) - assert set(DTOInspect(dto).annotations()) == {"id", "name", "color_id", "sweetness"} +def test_group_string_in_iterable_is_field_name(factory: AnyFactory, include_spec: Any) -> None: + """Test that group string literals inside iterables are treated as field names, not group selectors.""" + with pytest.raises(EmptyDTOError): + factory.factory(Fruit, DTOConfig(Purpose.READ, include=include_spec), if_no_fields="raise") @pytest.mark.parametrize("factory", factory_iterator()) From ba461f832391b3781dcde31ea769a468c3eb0783 Mon Sep 17 00:00:00 2001 From: gazorby Date: Sat, 6 Jun 2026 23:48:18 +0200 Subject: [PATCH 5/9] feat(dto): add support for bare FieldGroup constants in include and exclude specs --- README.md | 64 +++++++++++++++++++++++++++++++++++++ src/strawchemy/dto/types.py | 42 ++++++++---------------- tests/unit/dto/test_dto.py | 5 +++ 3 files changed, 83 insertions(+), 28 deletions(-) diff --git a/README.md b/README.md index 084f04f0..bc70a6ec 100644 --- a/README.md +++ b/README.md @@ -276,6 +276,70 @@ See the [custom resolvers](#custom-resolvers) for more details +### Field Groups + +Instead of listing field names one by one, `include` and `exclude` accept the `SCALARS` (column fields), `RELATIONSHIPS` (relation fields) and `ALL` group selectors, importable from `strawchemy`. They can be assigned directly (`include=SCALARS`), used inside an iterable, or mixed with field names — plain strings are always treated as field names. + +
+Field group examples + +```python +from strawchemy import ALL, RELATIONSHIPS, SCALARS + + +class User(Base): + __tablename__ = "user" + + id: Mapped[int] = mapped_column(primary_key=True) + name: Mapped[str] + password: Mapped[str] + posts: Mapped[list["Post"]] = relationship("Post", back_populates="author") + + +# Only column fields: id, name, password +@strawchemy.type(User, include=[SCALARS]) +class UserType: + pass + + +# Equivalent: a bare constant assigned directly +@strawchemy.type(User, include=SCALARS) +class UserType: + pass + + +# Groups mix with field names: all columns plus the `posts` relationship +@strawchemy.type(User, include=[SCALARS, "posts"]) +class UserType: + pass + + +# Both groups together are equivalent to include=ALL +@strawchemy.type(User, include=[SCALARS, RELATIONSHIPS]) +class UserType: + pass + + +# Groups work in exclude too: a bare exclude implies everything else +# is included, so this keeps only the column fields +@strawchemy.type(User, exclude=[RELATIONSHIPS]) +class UserType: + pass +``` + +A group-bearing `include` can be combined with `exclude` to subtract fields from the group: + +```python +# All columns except `password` +@strawchemy.type(User, include=[SCALARS], exclude=["password"]) +class UserType: + pass +``` + +Combining a plain field-name `include` with `exclude` is ambiguous and raises an error — `include` must be unset or contain a field group when `exclude` is used. + +
+ ### Type Override When generating types for relationships, Strawchemy creates default names (e.g., `Type`). If you have already diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 976c81ba..d9069916 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -43,7 +43,7 @@ "list[FieldSelector] | set[FieldSelector] | frozenset[FieldSelector] | tuple[FieldSelector, ...]" ) FieldGroupStr: TypeAlias = Literal["all", "scalars", "relationships"] -FieldSpec: TypeAlias = "FieldIterable | FieldGroupStr" +FieldSpec: TypeAlias = "FieldIterable | FieldGroupStr | FieldGroup" ConfigScope: TypeAlias = Literal["local", "global"] @@ -75,21 +75,9 @@ def is_group(cls, value: str) -> TypeIs[FieldGroupStr]: class FieldSet: """Normalized, immutable view over a field selection. - Wraps a `FieldSpec` (a group string such as "all", an iterable of field - names and/or `FieldGroup` selectors, or `None`) into a uniform - `frozenset[FieldSelector]` so selections can be compared, hashed, and - combined regardless of how they were originally expressed. - - Set-like semantics: - - `item in field_set`: matches by field name, by the field's group - (`SCALARS`/`RELATIONSHIPS`), or by `ALL`. `item` can be a field name, - a `DTOFieldDefinition` (its `is_relation` flag picks the group), or a - `FieldGroup` — group items only match themselves or `ALL`, never via - group matching. - - `self & other`: intersection of two selections; `ALL` subsumes the - other side. - - `self | other`: union, returned as a `FieldSpec | None` — collapses - to "all" when `ALL` is present, `None` when both sides are empty. + Wraps a `FieldSpec`, into a uniform `frozenset[FieldSelector]` so selections can be + compared, hashed, and combined regardless of how they were originally + expressed. """ value: InitVar[FieldSpec | None] @@ -247,17 +235,16 @@ class DTOConfig: Determines which fields from the source model are included based on their `DTOFieldConfig`. include: Explicitly include fields from the source model in the generated - DTO. Can be a list or set of field names, the literal "all" to include - all fields not explicitly excluded, or the `ALL` / `SCALARS` / - `RELATIONSHIPS` group selectors (mixable with names, e.g. + DTO. Can be a list or set of field names, and/or the `ALL` / `SCALARS` / + `RELATIONSHIPS` group selectors, either assigned directly + (`include=SCALARS`) or mixed with names inside an iterable (e.g. `[SCALARS, "owner"]`). `[SCALARS, RELATIONSHIPS]` is equivalent to - "all". Plain string literals "all", "scalars", "relationships" are - accepted interchangeably with the constants. Defaults to an empty set. + `ALL`. Defaults to an empty set. exclude: Explicitly exclude fields from the source model. Can be a list or set of field names and/or the `ALL` / `SCALARS` / `RELATIONSHIPS` group selectors (e.g. `[RELATIONSHIPS]` keeps all scalar fields and walks no - relationships). Plain string literals are accepted. A bare `exclude` (no - `include`) implies `include="all"`. Defaults to an empty set. + relationships). A bare `exclude` (no `include`) implies everything else + is included. Defaults to an empty set. partial: If True, makes all fields in the generated DTO optional. Defaults to None. partial_default: The default value assigned to fields when `partial` is @@ -289,11 +276,11 @@ class DTOConfig: include: FieldSpec | None = None """Explicitly include fields from the generated DTO.""" exclude: FieldSpec | None = None - """Explicitly exclude fields from the generated DTO. Implies `include="all"`.""" + """Explicitly exclude fields from the generated DTO. Implies everything else is included.""" global_include: FieldSpec | None = None """Explicitly include fields from the generated DTO and all its children.""" global_exclude: FieldSpec | None = None - """Explicitly exclude fields from the generated DTO and all its children. Implies `global_include="all"`.""" + """Explicitly exclude fields from the generated DTO and all its children. Implies everything else is included.""" partial: bool | None = None """Make all field optional.""" partial_default: Any = None @@ -486,15 +473,14 @@ def is_field_included( if scope == "local": return field in FieldSet(self.include) and field not in FieldSet(self.exclude) if scope == "global": - # `include="all"` propagates inclusion to children, same as `global_include="all"`. - included = field in FieldSet(self.global_include) or self.include == "all" + included = field in FieldSet(self.global_include) or FieldGroup.ALL in self.included_fields.field_set return included and field not in FieldSet(self.global_exclude) return field in self.included_fields and field not in self.excluded_fields def is_fields_iterable(value: Any) -> TypeIs[FieldSpec | FieldIterable | None]: """Test the given value is suitable to be used as either `include` or `exclude` in a DTOConfig.""" - if value == "all" or value is None: + if value == "all" or value is None or isinstance(value, FieldGroup): return True if isinstance(value, str): return False diff --git a/tests/unit/dto/test_dto.py b/tests/unit/dto/test_dto.py index 9741d9c7..d837a51f 100644 --- a/tests/unit/dto/test_dto.py +++ b/tests/unit/dto/test_dto.py @@ -411,6 +411,7 @@ def test_global_group_include_allows_global_exclude(factory: AnyFactory, model: "include_spec", [ pytest.param("all", id="all-string"), + pytest.param(ALL, id="bare-constant"), pytest.param([ALL], id="constant-list"), pytest.param([SCALARS, RELATIONSHIPS], id="both-groups"), ], @@ -426,7 +427,9 @@ def test_include_all_equivalents(factory: AnyFactory, include_spec: Any) -> None "include_spec", [ pytest.param(("scalars", ()), id="include-scalars-string"), + pytest.param((SCALARS, ()), id="include-scalars-bare-constant"), pytest.param(([SCALARS], ()), id="include-scalars-list"), + pytest.param((None, RELATIONSHIPS), id="bare-exclude-relationships-bare-constant"), pytest.param(([ALL], [RELATIONSHIPS]), id="exclude-relationships-constant-list"), pytest.param(("all", "relationships"), id="exclude-relationships-string"), pytest.param((None, "relationships"), id="bare-exclude-relationships-string"), @@ -445,7 +448,9 @@ def test_include_scalars(factory: AnyFactory, include_spec: tuple[FieldSpec | No "include_spec", [ pytest.param(("relationships", ()), id="include-relationships-string"), + pytest.param((RELATIONSHIPS, ()), id="include-relationships-bare-constant"), pytest.param(([RELATIONSHIPS], ()), id="include-relationships-list"), + pytest.param((None, SCALARS), id="bare-exclude-scalars-bare-constant"), pytest.param(([ALL], [SCALARS]), id="exclude-scalars-constant-list"), pytest.param(("all", "scalars"), id="exclude-scalars-string"), pytest.param((None, "scalars"), id="bare-exclude-scalars-string"), From cf4e57072d1251c53fd16ffa26506ce909aad63d Mon Sep 17 00:00:00 2001 From: gazorby Date: Sun, 7 Jun 2026 00:20:09 +0200 Subject: [PATCH 6/9] refactor(dto): refine is_fields_iterable to exclude None for stricter typing --- src/strawchemy/dto/types.py | 4 ++-- src/strawchemy/schema/factories/base.py | 4 ++-- src/strawchemy/utils/registry.py | 6 +++--- 3 files changed, 7 insertions(+), 7 deletions(-) diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index d9069916..66789ac7 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -478,9 +478,9 @@ def is_field_included( return field in self.included_fields and field not in self.excluded_fields -def is_fields_iterable(value: Any) -> TypeIs[FieldSpec | FieldIterable | None]: +def is_fields_iterable(value: Any) -> TypeIs[FieldSpec]: """Test the given value is suitable to be used as either `include` or `exclude` in a DTOConfig.""" - if value == "all" or value is None or isinstance(value, FieldGroup): + if value == "all" or isinstance(value, FieldGroup): return True if isinstance(value, str): return False diff --git a/src/strawchemy/schema/factories/base.py b/src/strawchemy/schema/factories/base.py index 7c4b51a3..7496f303 100644 --- a/src/strawchemy/schema/factories/base.py +++ b/src/strawchemy/schema/factories/base.py @@ -202,7 +202,7 @@ def _type_order_by( self, model: type[DeclarativeBase], include: FieldSpec | type[OrderByDTO] | None = None ) -> type[OrderByDTO] | None: order_include = self._mapper.config.order_by if include is None else include - if is_fields_iterable(order_include) and order_include is not None: + if is_fields_iterable(order_include): try: order_by_input = self._mapper.order_by_factory.make_input( model=model, @@ -226,7 +226,7 @@ def _type_distinct_on( self, model: type[DeclarativeBase], include: FieldSpec | type[EnumDTO] | None = None ) -> type[EnumDTO] | None: distinct_on_include = self._mapper.config.distinct_on if include is None else include - if is_fields_iterable(distinct_on_include) and distinct_on_include is not None: + if is_fields_iterable(distinct_on_include): try: distinct_on_input = self._mapper.distinct_on_enum_factory.factory( model=model, diff --git a/src/strawchemy/utils/registry.py b/src/strawchemy/utils/registry.py index 0a6d619d..479c2e78 100644 --- a/src/strawchemy/utils/registry.py +++ b/src/strawchemy/utils/registry.py @@ -15,7 +15,7 @@ from strawberry.types.union import StrawberryUnion from strawchemy.dto.strawberry import MappedStrawberryGraphQLDTO -from strawchemy.dto.types import FieldSet, is_fields_iterable +from strawchemy.dto.types import FieldSet from strawchemy.exceptions import StrawchemyError from strawchemy.utils.annotation import inner_types from strawchemy.utils.strawberry import strawberry_contained_types @@ -365,8 +365,8 @@ def _type_info( override=override, user_defined=user_defined, pagination=default_pagination, - order=FieldSet(order) if is_fields_iterable(order) else order, - distinct_on=FieldSet(distinct_on) if is_fields_iterable(distinct_on) else distinct_on, + order=order if isinstance(order, type) else FieldSet(order), + distinct_on=distinct_on if isinstance(distinct_on, type) else FieldSet(distinct_on), paginate=FieldSet(paginate), scope=dto_config.scope, model=model, From 8dfc6828aa8f5c1fab10fba578721dd30e80f664 Mon Sep 17 00:00:00 2001 From: gazorby Date: Mon, 8 Jun 2026 20:26:54 +0200 Subject: [PATCH 7/9] chore: omit factories file from coverage reports to reduce noise --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 4dbe9ceb..32655857 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -142,7 +142,7 @@ ignore-words-list = "nin" branch = true source = ["src"] plugins = ["covdefaults"] -omit = ["*/tests/*"] +omit = ["*/tests/*", "src/strawchemy/schema/factories/_kwargs.py"] parallel = true relative_files = true From b867f7e63e9540aa0fc5140aceb09059e1ff8091 Mon Sep 17 00:00:00 2001 From: gazorby Date: Wed, 10 Jun 2026 12:30:57 +0200 Subject: [PATCH 8/9] feat: allow combining include and exclude with overlap warning --- README.md | 2 +- src/strawchemy/config/base.py | 8 +++++++- src/strawchemy/dto/types.py | 28 ++++++++++------------------ tests/unit/dto/test_dto.py | 14 +++++++++----- 4 files changed, 27 insertions(+), 25 deletions(-) diff --git a/README.md b/README.md index bc70a6ec..ef7d21a1 100644 --- a/README.md +++ b/README.md @@ -336,7 +336,7 @@ class UserType: pass ``` -Combining a plain field-name `include` with `exclude` is ambiguous and raises an error — `include` must be unset or contain a field group when `exclude` is used. +`include` and `exclude` can always be combined: a field is kept when it is selected by `include` and not selected by `exclude`. diff --git a/src/strawchemy/config/base.py b/src/strawchemy/config/base.py index 9163730e..52b7945d 100644 --- a/src/strawchemy/config/base.py +++ b/src/strawchemy/config/base.py @@ -2,12 +2,13 @@ from __future__ import annotations +import warnings from dataclasses import dataclass, field from typing import TYPE_CHECKING, Any from strawchemy.dto import Purpose from strawchemy.dto.inspectors import SQLAlchemyGraphQLInspector -from strawchemy.dto.types import DTOConfig, FieldSpec +from strawchemy.dto.types import DTOConfig, FieldGroup, FieldSet, FieldSpec from strawchemy.repository.strawberry import StrawchemySyncRepository from strawchemy.utils.strawberry import default_session_getter @@ -80,6 +81,11 @@ def __post_init__(self) -> None: """Initializes the SQLAlchemyGraphQLInspector after the dataclass is created.""" self.inspector = SQLAlchemyGraphQLInspector(self.dialect, filter_overrides=self.filter_overrides) + if overlap := FieldSet(self.include).overlap(self.exclude): + names = sorted(selector.value if isinstance(selector, FieldGroup) else selector for selector in overlap) + msg = f"Fields are both explicitly included and excluded; exclude takes precedence: {names}" + warnings.warn(msg, stacklevel=2) + @property def field_config(self) -> DTOConfig: return DTOConfig(purpose=Purpose.READ, global_include=self.include, global_exclude=self.exclude or set()) diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 66789ac7..60b9710e 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -4,6 +4,7 @@ import dataclasses import functools +import warnings from dataclasses import InitVar, dataclass, field from enum import Enum from typing import TYPE_CHECKING, Any, Literal, TypeAlias, final, get_type_hints @@ -143,6 +144,9 @@ def normalize(cls, value: FieldSpec | None) -> frozenset[FieldSelector]: return frozenset() return frozenset(value) + def overlap(self, other: FieldSpec | None) -> frozenset[FieldSelector]: + return self.field_set & FieldSet(other).field_set + @final class DTOMissing: @@ -266,9 +270,7 @@ class DTOConfig: with `aliases`. Raises: - ValueError: If both `aliases` and `alias_generator` are provided, or - if `exclude` is set while `include` is also set to a specific list/set - (i.e., not "all" or empty). + ValueError: If both `aliases` and `alias_generator` are provided. """ purpose: Purpose @@ -301,12 +303,6 @@ def __post_init__(self) -> None: if self.aliases and self.alias_generator is not None: msg = "You must set `aliases` or `alias_generator`, not both" raise ValueError(msg) - if self.include and not self._has_field_group(self.include) and self.exclude: - msg = f"When using `exclude`, `include` must be unset or be a field group {FieldGroup.list_str()}." - raise ValueError(msg) - if self.global_include and not self._has_field_group(self.global_include) and self.global_exclude: - msg = f"When using `global_exclude`, `global_include` must be unset or be a field group {FieldGroup.list_str()}." - raise ValueError(msg) # A bare exclude (no include) means "everything except"; promote to "all". # If include carries a FieldGroup it is truthy, so the clobber is skipped. if self.global_exclude and self.global_include is None: @@ -317,18 +313,14 @@ def __post_init__(self) -> None: self.included_fields = FieldSet(self.global_include) if self.include is None else FieldSet(self.include) self.excluded_fields = FieldSet(self.global_exclude) if self.exclude is None else FieldSet(self.exclude) + if overlap := FieldSet(self.include).overlap(self.exclude): + names = sorted(selector.value if isinstance(selector, FieldGroup) else selector for selector in overlap) + msg = f"Fields are both explicitly included and excluded; exclude takes precedence: {names}" + warnings.warn(msg, stacklevel=2) + def __or__(self, other: DTOConfig) -> DTOConfig: return self.union(other) - @classmethod - def _has_field_group(cls, value: FieldSpec | FieldIterable) -> bool: - """True if the selection contains a FieldGroup member.""" - if isinstance(value, FieldGroup): - return True - if isinstance(value, str): - return value in FieldGroup.values() - return any(isinstance(item, FieldGroup) for item in value) - def union(self, other: DTOConfig) -> DTOConfig: include = FieldSet(self.include) | other.include exclude = FieldSet(self.exclude) | other.exclude diff --git a/tests/unit/dto/test_dto.py b/tests/unit/dto/test_dto.py index d837a51f..1bacfc47 100644 --- a/tests/unit/dto/test_dto.py +++ b/tests/unit/dto/test_dto.py @@ -9,7 +9,7 @@ from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column from typing_extensions import Self -from strawchemy import ALL, RELATIONSHIPS, SCALARS +from strawchemy import ALL, RELATIONSHIPS, SCALARS, StrawchemyConfig from strawchemy.dto import DTOConfig, Purpose, PurposeConfig, config, field from strawchemy.dto.constants import DTO_INFO_KEY from strawchemy.dto.strawberry import DTOKey, GraphQLFieldDefinition, StrawchemyDefinition @@ -370,10 +370,14 @@ def test_scalars_include_allows_exclude( @pytest.mark.parametrize("model", [Fruit, FruitDataclass]) @pytest.mark.parametrize("factory", factory_iterator()) -def test_plain_include_with_exclude_still_raises(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: - """Test that a plain field-name include combined with exclude still raises.""" - with pytest.raises(ValueError, match="exclude"): - factory.factory(model, DTOConfig(Purpose.READ, include=["name"], exclude=["sweetness"])) +def test_plain_include_with_exclude(factory: AnyFactory, model: type[Fruit | FruitDataclass]) -> None: + """Test that a plain field-name include and exclude combine (exclude wins) and warn on overlap.""" + msg = "both explicitly included and excluded" + with pytest.warns(UserWarning, match=msg): + dto = factory.factory(model, DTOConfig(Purpose.READ, include=["name", "sweetness"], exclude=["sweetness"])) + assert set(DTOInspect(dto).annotations()) == {"name"} + with pytest.warns(UserWarning, match=msg): + StrawchemyConfig(dialect="postgresql", include=["name", "sweetness"], exclude=["sweetness"]) @pytest.mark.parametrize("model", [Fruit, FruitDataclass]) From 216dd21200d2affa44015bf50100e83ce147693f Mon Sep 17 00:00:00 2001 From: gazorby Date: Wed, 10 Jun 2026 19:51:16 +0200 Subject: [PATCH 9/9] refactor(dto): remove __next__ method from FieldGroup --- src/strawchemy/dto/types.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/src/strawchemy/dto/types.py b/src/strawchemy/dto/types.py index 60b9710e..75cddb11 100644 --- a/src/strawchemy/dto/types.py +++ b/src/strawchemy/dto/types.py @@ -88,9 +88,6 @@ class FieldSet: def __post_init__(self, value: FieldSpec | None) -> None: self.field_set = self.normalize(value) - def __next__(self) -> FieldSelector: - return next(iter(self.field_set)) - def __iter__(self) -> Iterator[FieldSelector]: return iter(self.field_set)