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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 8 additions & 3 deletions docs/development/ADRs/next/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,11 @@ Writing a new ADR is simple:
### General architecture #general

- [0005 - Extending Iterator IR](0005-Extending_Iterator_IR.md)
- [0019 - Connectivities](0019-Connectivities.md)
- [0020 - Runtime domains](0020-Runtime-domains.md)
- [0021 - Argument Descriptors](0021-Argument-Descriptors.md)
- [0023 - Fingerprinting](0023-Fingerprinting.md)
- [0026 - Staggered Dimensions](0024-Staggered_Dimensions.md)
- [0026 - Staggered Dimensions](0026-Staggered_Dimensions.md)

### Frontend and Parsing #frontend

Expand All @@ -45,15 +48,17 @@ Writing a new ADR is simple:
- [0006 - C++ Backend](0006-Cpp-Backend.md)
- [0007 - Fencil Processors](0007-Fencil-Processors.md)
- [0008 - Mapping Domain to Cpp Backend](0008-Mapping_Domain_to_Cpp-Backend.md)
- [0014 - DaCe backend](0014-DaCe_backend.md)
- [0016 - Multiple Backends and Build Systems](0016-Multiple-Backends-and-Build-Systems.md)
- [0017 - Toolchain Configuration](0017-Toolchain-Configuration.md)
- [0018 - Canonical Form of an SDFG in GT4Py (Especially for Optimizations)](0018-Canonical_SDFG_in_GT4Py_Transformations.md)
- [0027 - External Workspace Memory for DaCe Transients](0027-External_Workspace_Memory.md)

### Python Integration

- [0011 - On The Fly Compilation](0011-On_The_Fly_Compilation.md)
- [0012 - GridTools C++ OTF](0011-_GridTools_Cpp_OTF.md)
- [0012 - GridTools C++ OTF Steps](0012-GridTools_Cpp_OTF_Steps.md)
- [0024 - Compilation Runners](0024-Compilation-Runners.md)
- [0025 - Crash Consistent Build Caches](0025-Crash_Consistent_Build_Caches.md)

### Testing

Expand Down
55 changes: 55 additions & 0 deletions src/gt4py/eve/type_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -302,6 +302,49 @@ def __call__(
)
return self.combine_optional(name, validator) if has_none else validator

if origin_type is type:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

PR description correction — the body says two fields become strictly validated, ts.DeferredType.constraint and ffront.stages.node_class. Only the first one does.

DSLFieldOperatorDef is a plain @dataclasses.dataclass(frozen=True) (stages.py:35,74), not an eve.datamodels.DataModel, and eve.type_validation is imported by exactly one module — eve/datamodels/core.py. A plain dataclass field never reaches this factory, so node_class is unaffected.

I imported all of gt4py (0 import failures) and walked every DataModel subclass's __datamodel_fields__ looking for a type[...] annotation. There is exactly one in the tree:

gt4py.next.type_system.type_specifications.DeferredType.constraint
  Union[type[TypeSpec], tuple[type[TypeSpec], ...], None]

That one behaves as the body claims, verified end to end:

DeferredType(constraint=ts.ScalarType)                  ACCEPTED
DeferredType(constraint=None)                           ACCEPTED
DeferredType(constraint=(ts.ScalarType, ts.FieldType))  ACCEPTED
DeferredType(constraint=int)                            TypeError   <- the bug, now caught
DeferredType(constraint=ts.ScalarType(...))  (instance) TypeError

and every in-tree constraint= call site passes a ts.*Type class, None, or a tuple of them, so nothing regresses. Only the node_class half of the sentence needs dropping.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You are right, and thank you for the receipt. I reproduced it independently: walking __datamodel_fields__ over every is_datamodel class in eve/_core/storage/cartesian/next (322 classes, recursing into get_args so nested type[...] counts) turns up exactly one field — DeferredType.constraint. And grep -rn "type_validation" src/gt4py --include=*.py has a single importer, eve/datamodels/core.py, so a plain frozen dataclass like DSLFieldOperatorDef never reaches the factory.

PR description corrected to name only DeferredType.constraint.

# `type[X]`. Without this case the annotation falls through to the
# generic-collection branch below and degrades to `isinstance(value, type)`,
# i.e. "is any class at all".
if len(type_args) != 1:

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Regression: shapes that validated fine on main now raise at class-creation time.

Three annotation shapes reach this branch and fall through to the EveValueError, where before they landed in the generic-collection fallback and got loose isinstance(value, type) validation:

  • type[A | B] — one arg, but it is a UnionType, so is_actual_type is false
  • bare typing.Type — no args at all, so len(type_args) != 1
  • type[SomeAlias] where SomeAlias is a PEP 695 type X = ... — the arg is never resolved

Verified by running all three against upstream/main (all create fine) and against this branch (all raise EveValueError). No in-repo field uses these shapes, so the suite stays green, but it is an import-time break for downstream code.

The PEP 695 case is the one that stings: alias resolution happens once at the top of the factory for the whole annotation, so an alias nested inside type[...] is never resolved. This branch is a fourth annotation-dispatch funnel and needs the same treatment as the other three.

Fix: handle bare type, resolve the argument, OR-combine subclass validators for unions, and fall back to make_is_instance_of(name, type) rather than raising — never regress a shape that used to work.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in 759715fad. Unrecognized type[X] shapes now fall back to make_is_instance_of(name, type) instead of raising, so nothing that validated before can break at class-creation time. Unions of classes validate as an OR of subclass checks, and a nested PEP 695 alias is resolved with eval_type_alias first.

Regression tests added for type, typing.Type, type[A | B] and type[SomeAlias] — 8 of the new parametrizations fail against the previous version of this branch and pass now.

# bare `type` / `typing.Type`
return self.make_is_instance_of(name, type)

# A PEP 695 alias nested inside `type[...]` is not resolved by the
# whole-annotation pass at the top of this function, so resolve it here.
try:
arg = xtyping.eval_type_alias(type_args[0])
except TypeError:
return self.make_is_instance_of(name, type)

if isinstance(arg, types.UnionType): # `type[A | B]`
arg = typing.Union[arg.__args__]

# `issubclass()` is not generally usable with protocol classes: it is
# rejected outright unless the protocol is `@runtime_checkable`, and also
# for `@runtime_checkable` protocols with non-method members. Protocols
# therefore keep the loose check, like any other unsupported shape.
def is_strict_arg(a: Any) -> xtyping.TypeGuard[type]:
return xtyping.is_actual_type(a) and not xtyping.is_protocol(a)

if isinstance(arg, typing.TypeVar):
# Mirror the plain-`TypeVar` branch above, which honours the bound.
if is_strict_arg(arg.__bound__):
return self.make_is_subclass_of(name, arg.__bound__)
return self.make_is_instance_of(name, type)

if is_strict_arg(arg):
return self.make_is_subclass_of(name, arg)

if xtyping.get_origin(arg) is Union: # `type[A | B]`, `type[Union[A, B]]`
members = xtyping.get_args(arg)
if members and all(is_strict_arg(m) for m in members):
return self.combine_validators_as_or(
name, *(self.make_is_subclass_of(name, m) for m in members)
)

return self.make_is_instance_of(name, type)

if isinstance(origin_type, type):
# Deal with generic collections
if issubclass(origin_type, tuple):
Expand Down Expand Up @@ -406,6 +449,18 @@ def _is_instance_of(value: Any, **kwargs: Any) -> None:

return _is_instance_of

@staticmethod
def make_is_subclass_of(name: str, type_: type) -> FixedTypeValidator:
"""Create a ``FixedTypeValidator`` validator for ``type[type_]`` annotations."""

def _is_subclass_of(value: Any, **kwargs: Any) -> None:
if not (isinstance(value, type) and issubclass(value, type_)):

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

issubclass is called unguarded here, and it raises for Protocols. A Protocol class passes xtyping.is_actual_type, so type[SomeProtocol] takes the strict path above and then fails for every value:

type[P]  (Protocol, not runtime_checkable), value = a conforming class
  -> TypeError: Instance and class checks can only be used with @runtime_checkable protocols
type[Q]  (runtime_checkable, has a data member), value = a conforming class
  -> TypeError: Protocols with non-method members don't support issubclass()

That is inside the letter of "never reject a shape that used to validate" — it fails at validation time rather than class-creation time — but not its spirit, and the message points at issubclass rather than at the field.

Nothing breaks today: no DataModel in gt4py has such a field, and I checked the icon4py side too — its only type[...] annotations are two exc_type: type[BaseException] __exit__ parameters, neither a datamodel field. But #2844/#2845 add a lot of type[...] annotations and eve is a library, so this is worth closing now.

Cheapest fix: treat a Protocol as an unrecognized shape in the type[...] branch and fall back to the loose isinstance(value, type) check, same as the other unrecognized shapes.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good catch, and confirmed — this was a real bug. Reproduced both variants against 759715fad before fixing:

type[P] plain Protocol        TypeError: Instance and class checks can only be used with @runtime_checkable pro...
type[Q] runtime+data member   TypeError: Protocols with non-method members do not support issubclass()

both raised for a conforming class, i.e. the field was unusable for every value. Fixed in bb848a08c by routing all three strict-path decisions through

def is_strict_arg(a: Any) -> xtyping.TypeGuard[type]:
    return xtyping.is_actual_type(a) and not xtyping.is_protocol(a)

so a protocol falls back to the loose check like any other unsupported shape. xtyping.is_protocol already existed (re-exported from typing_extensions).

One judgement call worth your sign-off: this also loosens the case where issubclass does work — a @runtime_checkable protocol whose members are all methods. I chose the blanket fallback so type[P] does not silently change behaviour when someone later adds a non-method member to P (which would turn a working datamodel into a class-creation-time error). The narrow alternative would have to reach into __non_callable_proto_members__. Happy to switch if you prefer precision here.

raise TypeError(
f"'{name}' must be a subclass of {type_} (got '{value}' which is a {type(value)})."
)

return _is_subclass_of

@staticmethod
def make_is_instance_of_int(name: str) -> FixedTypeValidator:
"""Create an ``FixedTypeValidator`` validator for ``int`` values which fails with ``bool`` values."""
Expand Down
2 changes: 1 addition & 1 deletion src/gt4py/next/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -1473,7 +1473,7 @@ def connectivity_for_cartesian_shift(dim: Dimension, offset: int | float) -> Car
`flip_staggered(dim)`).

The half-integer case encodes the convention that a staggered index sits half a cell *below*
its base index (see ADR 0024): `IHalf(0)` is the edge below `I(0)`. Because of this asymmetry,
its base index (see ADR 0026): `IHalf(0)` is the edge below `I(0)`. Because of this asymmetry,
shifting out of a non-staggered dimension needs a `+1` index correction that shifting out of a
staggered dimension does not, e.g. `I + 0.5` maps `I(i)` to `IHalf(i+1)` (position `i+½`) while
`IHalf + 0.5` maps `IHalf(i)` to `I(i)`.
Expand Down
5 changes: 5 additions & 0 deletions src/gt4py/next/type_system/mypy_plugin.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,11 @@
The documentation for how to write tests in that format is at https://github.com/typeddjango/pytest-mypy-plugins.

The documentation on mypy plugins is at https://mypy.readthedocs.io/en/latest/extending_mypy.html

Known limitation: the fallback hook that rewrites stray 'Dimension' instances in type aliases
matches on the *variable name* ending in 'Dim' ('fullname.endswith("Dim")'). A dimension bound to
a name that does not end in 'Dim' -- e.g. 'I = gtx.Dimension("I")' -- silently gets no plugin
support at all.
"""

from __future__ import annotations
Expand Down
50 changes: 50 additions & 0 deletions tests/eve_tests/unit_tests/test_type_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,23 @@ class SampleDataClass:
a: int


type SampleTypeAlias = SampleEmptyClass # PEP 695, nested inside `type[...]`

type SampleRecursiveTypeAlias = SampleRecursiveTypeAlias # cannot be resolved

SampleBoundTypeVar = typing.TypeVar("SampleBoundTypeVar", bound=SampleEmptyClass)
SampleUnboundTypeVar = typing.TypeVar("SampleUnboundTypeVar")


class SampleProtocol(typing.Protocol):
def method(self) -> None: ...


@typing.runtime_checkable
class SampleRuntimeCheckableProtocol(typing.Protocol):
attribute: int


# Each item should be a tuple like:
# ( annotation: Any, valid_values: Sequence, wrong_values: Sequence,
# globalns: Optional[Dict[str, Any]], localns: Optional[Dict[str, Any]] )
Expand All @@ -75,6 +92,39 @@ class SampleDataClass:
(typing.List[int], ([1, 2, 3], []), (1, [1.0]), None, None),
(typing.Set[int], ({1, 2, 3}, set()), (1, [1], (1,), {1: None}), None, None),
(typing.Dict[int, str], ({}, {3: "three"}), ([(3, "three")], 3, "three", []), None, None),
(type[SampleEmptyClass], [SampleEmptyClass], [SampleEmptyClass(), int, 3], None, None),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The table covers the recognized shapes well. Three gaps worth closing, the first one most:

  1. Nothing asserts, at datamodel level, that DeferredType(constraint=int) now raises. That is the only behaviour change this PR actually causes anywhere in the tree, and it is the example in the PR description — but the assertion for it lives only in this annotation-level table, against SampleEmptyClass.
  2. typing.Union[A, B] is not exercised — only the A | B spelling is. They are different runtime objects and the branch handles them on different paths (types.UnionType normalisation vs. get_origin(...) is Union).
  3. The except TypeError fallback around eval_type_alias is untested. A recursive PEP 695 alias reaches it; I confirmed by hand that type[Rec] with type Rec = Rec falls back to the loose check rather than raising, but nothing in the suite pins that.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

All three closed in bb848a08c, each verified to fail without the fix:

  1. tests/next_tests/unit_tests/type_system_tests/test_type_specifications.py (new) asserts at datamodel level that DeferredType(constraint=int) raises, and that None, ts.ScalarType, ts.TypeSpec and tuples of them are still accepted. Put under next_tests rather than eve_tests because nothing in eve_tests imports gt4py.next today — tach would not have objected (source_roots is ["src"]), but the layering intent is clear.
  2. type[typing.Union[A, B]] row added — with the branch disabled it fails, so it genuinely exercises the get_origin(...) is Union path rather than the types.UnionType one.
  3. type Rec = Rec row added. Removing the try/except TypeError makes it fail with Type alias 'Rec' cannot be resolved (recursive definition) at validator-construction time.

(type[SampleDataClass], [SampleDataClass], [SampleDataClass(a=1), str], None, None),
(type[Any], [SampleEmptyClass, int, str], [3, "int", SampleEmptyClass()], None, None),
# `type[X]` shapes that are not a plain class must stay *usable*: they validated
# loosely before the `type[X]` case existed, so they fall back to "is a class"
# rather than being rejected at class-creation time.
(type, [SampleEmptyClass, int], [3, "int"], None, None),
(typing.Type, [SampleEmptyClass, int], [3, "int"], None, None),
(
type[SampleEmptyClass | SampleDataClass],
[SampleEmptyClass, SampleDataClass],
[int, SampleEmptyClass(), 3],
None,
None,
),
(
type[typing.Union[SampleEmptyClass, SampleDataClass]],
[SampleEmptyClass, SampleDataClass],
[int, SampleEmptyClass(), 3],
None,
None,
),
(type[SampleTypeAlias], [SampleEmptyClass], [int, SampleEmptyClass(), 3], None, None),
# An alias which cannot be resolved keeps the loose check instead of propagating
# the resolution error out of the class definition.
(type[SampleRecursiveTypeAlias], [SampleEmptyClass, int], [3, "int"], None, None),
# Protocols do not support 'issubclass()' in general, so they keep the loose check
# too, independently of whether they are '@runtime_checkable'.
(type[SampleProtocol], [SampleEmptyClass, int], [3, "int"], None, None),
(type[SampleRuntimeCheckableProtocol], [SampleEmptyClass, int], [3, "int"], None, None),
# a bounded TypeVar is honoured, mirroring the plain-TypeVar branch
(type[SampleBoundTypeVar], [SampleEmptyClass], [int, SampleEmptyClass()], None, None),
(type[SampleUnboundTypeVar], [SampleEmptyClass, int], [3, "int"], None, None),
(
frozendict[int, str],
(
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
# GT4Py - GridTools Framework
#
# Copyright (c) 2014-2024, ETH Zurich
# All rights reserved.
#
# Please, refer to the LICENSE file in the root directory.
# SPDX-License-Identifier: BSD-3-Clause

import pytest

from gt4py.next.type_system import type_specifications as ts


@pytest.mark.parametrize(
"constraint",
[
None,
ts.ScalarType,
ts.TypeSpec,
(ts.ScalarType,),
(ts.ScalarType, ts.FieldType),
],
)
def test_deferred_type_accepts_type_spec_constraints(constraint):
assert ts.DeferredType(constraint=constraint).constraint is constraint


@pytest.mark.parametrize("constraint", [int, ts.ScalarType(kind=ts.ScalarKind.INT32)])
def test_deferred_type_rejects_non_type_spec_constraint(constraint):
# 'constraint' is annotated as 'type[TypeSpec] | tuple[type[TypeSpec], ...] | None',
# which is checked with 'issubclass()' and not just with "is a class at all".
with pytest.raises(TypeError, match="constraint"):
ts.DeferredType(constraint=constraint)


def test_deferred_type_rejects_non_type_spec_constraint_in_tuple():
with pytest.raises(TypeError, match="constraint"):
ts.DeferredType(constraint=(ts.ScalarType, int))