Skip to content
Open
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
16 changes: 14 additions & 2 deletions src/parallel/_models.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
from __future__ import annotations

import os
import sys
import types
import inspect
import weakref
from typing import (
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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


Expand Down
58 changes: 58 additions & 0 deletions tests/test_models.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down