diff --git a/src/strawchemy/utils/registry.py b/src/strawchemy/utils/registry.py index ab5e590f..16c8abc1 100644 --- a/src/strawchemy/utils/registry.py +++ b/src/strawchemy/utils/registry.py @@ -101,6 +101,9 @@ def update_type(self, strawberry_type: type[WithStrawberryObjectDefinition]) -> else: self._set_type(strawberry_type) + def contains_type(self, strawberry_type: type[WithStrawberryObjectDefinition]) -> bool: + return any(inner_type is strawberry_type for inner_type in strawberry_contained_types(self.ref_holder.type)) + @dataclasses.dataclass(frozen=True, eq=True) class RegistryTypeInfo: @@ -122,6 +125,20 @@ class RegistryTypeInfo: def scoped_id(self) -> Hashable: return self.model, self.graphql_type, self.tags + @property + def _is_user_defined_default_name_override(self) -> bool: + """Whether this override should take over generated refs for the model's default DTO name.""" + return self.scope is None and self.override and self.user_defined and self.default_name is not None + + @property + def resolves_scoped_references(self) -> bool: + """Whether this registration should satisfy refs to generated DTOs for the same model.""" + return bool( + self.model + and not self.exclude_from_scope + and (self.scope == "global" or self._is_user_defined_default_name_override) + ) + class StrawberryRegistry: def __init__(self, strawberry_config: StrawberryConfig) -> None: @@ -242,11 +259,33 @@ def _register(self, type_info: RegistryTypeInfo, strawberry_type: type[Any]) -> reference.update_type(strawberry_type) if type_info.graphql_type != "enum": self._track_references(strawberry_type, type_info.graphql_type, force=type_info.override) - if type_info.scope == "global" and type_info.model: + if type_info.resolves_scoped_references: + # A user-defined override can replace a generated/default DTO for + # the same model after relationship fields have already recorded + # references to it. Update only refs that still point at that + # previous default DTO; refs already resolved to another explicit + # override for the same model must keep their chosen type. + # ``scope="global"`` keeps the historical behavior and refreshes + # every scoped reference. + previous_default_type: type[WithStrawberryObjectDefinition] | None = None if type_info.default_name: + previous_type_info = self._names_map[type_info.graphql_type].get(type_info.default_name) + previous_default_type = cast( + "type[WithStrawberryObjectDefinition] | None", + self._type_map.get(previous_type_info) if previous_type_info else None, + ) self._namespaces[type_info.graphql_type][type_info.default_name] = strawberry_type + if type_info.default_name != type_info.name: + for reference in self._forward_type_refs[type_info.graphql_type][type_info.default_name]: + if previous_default_type is None or reference.contains_type(previous_default_type): + reference.update_type(strawberry_type) for reference in self._type_refs[type_info.scoped_id]: - reference.update_type(strawberry_type) + if ( + type_info.scope == "global" + or previous_default_type is None + or reference.contains_type(previous_default_type) + ): + reference.update_type(strawberry_type) self._scoped_types[type_info.scoped_id] = strawberry_type self._names_map[type_info.graphql_type][type_info.name] = type_info self._type_map[type_info] = strawberry_type diff --git a/tests/unit/mapping/__snapshots__/test_schemas/test_query_schemas[override_with_custom_name].gql b/tests/unit/mapping/__snapshots__/test_schemas/test_query_schemas[override_with_custom_name].gql index b9dc2f34..f6c2afc7 100644 --- a/tests/unit/mapping/__snapshots__/test_schemas/test_query_schemas[override_with_custom_name].gql +++ b/tests/unit/mapping/__snapshots__/test_schemas/test_query_schemas[override_with_custom_name].gql @@ -4,8 +4,8 @@ type ColorType { name: Int! fruitsAggregate: FruitAggregate! - """Fetch objects from the FruitType collection""" - fruits: [FruitType!]! + """Fetch objects from the FruitTypeCustomName collection""" + fruits: [FruitTypeCustomName!]! id: UUID! } @@ -39,15 +39,6 @@ type FruitSumFields { sweetness: Int } -"""GraphQL type""" -type FruitType { - color: ColorType! - name: Int! - colorId: UUID - sweetness: Int! - id: UUID! -} - """GraphQL type""" type FruitTypeCustomName { name: Int! diff --git a/tests/unit/mapping/test_schemas.py b/tests/unit/mapping/test_schemas.py index 6fb4d2b3..3ec7d304 100644 --- a/tests/unit/mapping/test_schemas.py +++ b/tests/unit/mapping/test_schemas.py @@ -69,6 +69,45 @@ class ColorSlim: assert slim_fields == {"id"} +def test_relationship_uses_late_explicit_related_type_registration(strawchemy: Strawchemy) -> None: + """Relationship type selection must not depend on decoration/import order.""" + from tests.unit.models import Color, Fruit + + @strawchemy.type(Fruit, name="FruitNode", include=["id", "color"], override=True) + class FruitNode: + pass + + @strawchemy.type(Color, name="ColorNode", include=["id", "name"], override=True) + class ColorNode: + @strawberry.field + def label(self) -> str: + return "label" + + @strawchemy.type(Color, name="AlternateColorNode", include=["id", "name"], override=True) + class AlternateColorNode: + @strawberry.field + def alternate_label(self) -> str: + return "alternate" + + @strawberry.type + class Query: + @strawberry.field(graphql_type=FruitNode | None) + def fruit(self) -> object | None: + return None + + @strawberry.field(graphql_type=ColorNode | None) + def color(self) -> object | None: + return None + + schema = strawberry.Schema(query=Query) + schema_sdl = str(schema) + + assert "color: ColorNode!" in schema_sdl + assert "label: String!" in schema_sdl + assert "color: ColorType!" not in schema_sdl + assert "color: AlternateColorNode!" not in schema_sdl + + def test_type_instance_auto_as_str(strawchemy: Strawchemy) -> None: @strawchemy.type(User) class UserType: