diff --git a/src/latch_data_validation/data_validation.py b/src/latch_data_validation/data_validation.py index 59ee80e..5f48d03 100644 --- a/src/latch_data_validation/data_validation.py +++ b/src/latch_data_validation/data_validation.py @@ -28,7 +28,7 @@ Sequence as SequenceOld, # pyright: ignore[reportDeprecated] ) -from typing_extensions import NotRequired, Required, TypeForm, override +from typing_extensions import NotRequired, Required, TypeAliasType, TypeForm, override __all__ = [] @@ -189,6 +189,18 @@ def _untraced_validate( x: JsonValue, cls: TypeForm[T], *, type_vars: dict[int, TypeForm[object]] ) -> T: # todo(maximsmol): improve error messages with generics + alias_origin = get_origin(cls) + if isinstance(alias_origin, TypeAliasType): + type_vars2 = {**type_vars} + for parameter, argument in zip( + alias_origin.__type_params__, get_args(cls), strict=True + ): + type_vars2[id(parameter)] = argument + return _untraced_validate(x, alias_origin.__value__, type_vars=type_vars2) + + if isinstance(cls, TypeAliasType): + return _untraced_validate(x, cls.__value__, type_vars=type_vars) + if isinstance(cls, TypeVar): ref = type_vars.get(id(cls)) if ref is None: diff --git a/tests/test_basic.py b/tests/test_basic.py index 9024f8e..8fc8613 100644 --- a/tests/test_basic.py +++ b/tests/test_basic.py @@ -1,10 +1,11 @@ import typing from collections.abc import Mapping, Sequence +from dataclasses import dataclass from enum import Enum import pytest from syrupy.extensions.json import JSONSnapshotExtension -from typing_extensions import NotRequired, Required, TypedDict, TypeVar +from typing_extensions import NotRequired, Required, TypeAliasType, TypedDict, TypeVar BrokenJsonArray: typing.TypeAlias = typing.Sequence["BrokenJsonValue"] # pyright: ignore[reportDeprecated] BrokenJsonObject: typing.TypeAlias = typing.Mapping[str, "BrokenJsonValue"] # pyright: ignore[reportDeprecated] @@ -18,6 +19,19 @@ BrokenJsonObject2 | BrokenJsonArray2 | str | int | float | bool | None ) +TestWaitingReason = TypeAliasType( + "TestWaitingReason", typing.Literal["workspace_capacity", "taiga_rate_limited"] +) + + +@dataclass +class AliasedRequest: + waiting_reason: TestWaitingReason | None + + +AliasT = TypeVar("AliasT") +TestBox = TypeAliasType("TestBox", list[AliasT], type_params=(AliasT,)) + from latch_data_validation.data_validation import ( DataValidationError, @@ -157,6 +171,24 @@ class Test: assert x.c.a == 2 +def test_type_alias() -> None: + assert validate( + {"waiting_reason": "workspace_capacity"}, AliasedRequest + ) == AliasedRequest(waiting_reason="workspace_capacity") + assert validate({"waiting_reason": None}, AliasedRequest) == AliasedRequest( + waiting_reason=None + ) + + with pytest.raises(DataValidationError): + _ = validate({"waiting_reason": "unknown"}, AliasedRequest) + + +def test_generic_type_alias() -> None: + assert validate([1, 2], TestBox[int]) == [1, 2] + with pytest.raises(DataValidationError): + _ = validate([1, "two"], TestBox[int]) + + def test_forwardref(snapshot_json) -> None: _ = validate({"a": 123}, JsonValue)