diff --git a/src/parallel/_models.py b/src/parallel/_models.py index 8c5ab26..20b77a7 100644 --- a/src/parallel/_models.py +++ b/src/parallel/_models.py @@ -1,6 +1,8 @@ from __future__ import annotations import os +import sys +import types import inspect import weakref from typing import ( @@ -687,6 +689,15 @@ class CachedDiscriminatorType(Protocol): DISCRIMINATOR_CACHE: weakref.WeakKeyDictionary[type, DiscriminatorDetails] = weakref.WeakKeyDictionary() +def _discriminator_cache_key(union: type) -> type: + union_type = cast(Optional[type[object]], getattr(types, "UnionType", None)) + if (3, 10) <= sys.version_info < (3, 14) and union_type is not None and isinstance(union, union_type): + # PEP 604 unions cannot be weakly referenced before Python 3.14. The + # equivalent typing.Union can still serve as a weak cache key. + return cast(type, cast(Any, Union)[get_args(union)]) + return union + + class DiscriminatorDetails: field_name: str """The name of the discriminator field in the variant class, e.g. @@ -729,7 +740,8 @@ def __init__( def _build_discriminated_union_meta(*, union: type, meta_annotations: tuple[Any, ...]) -> DiscriminatorDetails | None: - cached = DISCRIMINATOR_CACHE.get(union) + key = _discriminator_cache_key(union) + cached = DISCRIMINATOR_CACHE.get(key) if cached is not None: return cached @@ -784,7 +796,7 @@ def _build_discriminated_union_meta(*, union: type, meta_annotations: tuple[Any, discriminator_field=discriminator_field_name, discriminator_alias=discriminator_alias, ) - DISCRIMINATOR_CACHE.setdefault(union, details) + DISCRIMINATOR_CACHE.setdefault(key, details) return details diff --git a/tests/test_models.py b/tests/test_models.py index 12db0ec..5303f1b 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -1,3 +1,4 @@ +import sys import json from typing import TYPE_CHECKING, Any, Dict, List, Union, Iterable, Optional, cast from datetime import datetime, timezone @@ -834,6 +835,63 @@ class B(BaseModel): assert DISCRIMINATOR_CACHE.get(UnionType) is discriminator +@pytest.mark.skipif(sys.version_info < (3, 10), reason="PEP 604 unions require Python 3.10") +def test_pep604_union_invalid_data() -> None: + class A(BaseModel): + type: Literal["a"] + + data: str + + class B(BaseModel): + type: Literal["b"] + + data: int + + PEP604Union = cast(Any, A).__or__(B).__or__(type(None)) + m = construct_type(value={"type": "b", "data": "foo"}, type_=PEP604Union) + + assert isinstance(m, A) + assert m.type == "b" # type: ignore[comparison-overlap] + assert m.data == "foo" + + +@pytest.mark.skipif(sys.version_info < (3, 10), reason="PEP 604 unions require Python 3.10") +def test_discriminated_pep604_union_invalid_data_uses_cache() -> None: + class A(BaseModel): + type: Literal["a"] + + data: str + + class B(BaseModel): + type: Literal["b"] + + data: int + + UnionType = cast(Any, Union[A, B, None]) + PEP604Union = cast(Any, A).__or__(B).__or__(type(None)) + + assert not DISCRIMINATOR_CACHE.get(UnionType) + + m = construct_type( + value={"type": "b", "data": "foo"}, + type_=cast(Any, Annotated[PEP604Union, PropertyInfo(discriminator="type")]), + ) + assert isinstance(m, B) + assert m.type == "b" + assert m.data == "foo" # type: ignore[comparison-overlap] + + discriminator = DISCRIMINATOR_CACHE.get(UnionType) + assert discriminator is not None + + m = construct_type( + value={"type": "b", "data": "bar"}, + type_=cast(Any, Annotated[PEP604Union, PropertyInfo(discriminator="type")]), + ) + assert isinstance(m, B) + assert m.data == "bar" # type: ignore[comparison-overlap] + assert DISCRIMINATOR_CACHE.get(UnionType) is discriminator + + @pytest.mark.skipif(PYDANTIC_V1, reason="TypeAliasType is not supported in Pydantic v1") def test_type_alias_type() -> None: Alias = TypeAliasType("Alias", str) # pyright: ignore