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
5 changes: 5 additions & 0 deletions hamilton/htypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import sys
import typing
from collections.abc import Iterable
from types import UnionType
from typing import Any, Literal, Protocol, TypeVar, Union

import typing_extensions
Expand Down Expand Up @@ -423,6 +424,10 @@ def check_instance(obj: Any, type_: Any) -> bool:
"""
if type_ == Any:
return True
# PEP 604 unions (e.g. list[int] | None) have no __origin__, and isinstance() rejects
# them when a member is a parameterized generic, so check each member instead.
if isinstance(type_, UnionType):
return any(check_instance(obj, t) for t in type_.__args__)
# Get the origin of the type (i.e., the base class for generic types)
origin = getattr(type_, "__origin__", None)

Expand Down
33 changes: 32 additions & 1 deletion tests/lifecycle/test_default.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

import pytest

from hamilton import driver
from hamilton import ad_hoc_utils, driver
from hamilton.lifecycle import default

from tests.resources import mismatched_types
Expand All @@ -37,3 +37,34 @@ def test_noedge_input_type_checking_with_adapter():
)
actual = dr.execute(["baz"], inputs={"a": 1.02, "number": "aaasdfdsf"})
assert actual == {"baz": "1.02 2 aaasdfdsf"}


def test_function_input_output_type_checker_handles_pep604_union_of_generics():
def evens(n: int) -> list[int] | None:
return [i * 2 for i in range(n)] if n else None

def total(evens: list[int] | None) -> int:
return sum(evens or [])

dr = (
driver.Builder()
.with_modules(ad_hoc_utils.create_temporary_module(evens, total))
.with_adapters(default.FunctionInputOutputTypeChecker())
.build()
)
assert dr.execute(["total"], inputs={"n": 3}) == {"total": 6}
assert dr.execute(["total"], inputs={"n": 0}) == {"total": 0}


def test_function_input_output_type_checker_rejects_wrong_pep604_union_result():
def evens(n: int) -> list[int] | None:
return ["not", "ints"]

dr = (
driver.Builder()
.with_modules(ad_hoc_utils.create_temporary_module(evens))
.with_adapters(default.FunctionInputOutputTypeChecker())
.build()
)
with pytest.raises(TypeError, match="Node evens returned a result"):
dr.execute(["evens"], inputs={"n": 3})
9 changes: 9 additions & 0 deletions tests/test_type_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,6 +403,15 @@ def test_check_instance_with_union_type():
assert not check_instance({"key1": 1, "key2": 2}, Union[int, str])


def test_check_instance_with_pep604_union_of_generics():
assert check_instance([1, 2], list[int] | None)
assert check_instance(None, list[int] | None)
assert not check_instance([1, "2"], list[int] | None)
assert not check_instance("12", list[int] | None)
assert check_instance({"a": 1}, dict[str, int] | list[int])
assert not check_instance({"a": "1"}, dict[str, int] | list[int])


def test_check_instance_with_union_type_and_literal():
from typing import Literal

Expand Down
Loading