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
4 changes: 4 additions & 0 deletions hamilton/htypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -473,5 +473,9 @@ def check_instance(obj: Any, type_: Any) -> bool:
return False
return True

# Other parameterized generics (e.g. Callable[[int], int], type[int], frozenset[str],
# column[pd.Series, float]) matched their origin above; isinstance() rejects them.
return True

# If the type is not a generic type, just use isinstance
return isinstance(obj, type_)
45 changes: 44 additions & 1 deletion tests/lifecycle/test_default.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,9 +15,12 @@
# specific language governing permissions and limitations
# under the License.

import typing

import pandas as pd
import pytest

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

from tests.resources import mismatched_types
Expand Down Expand Up @@ -68,3 +71,43 @@ def evens(n: int) -> list[int] | None:
)
with pytest.raises(TypeError, match="Node evens returned a result"):
dr.execute(["evens"], inputs={"n": 3})


def test_function_input_output_type_checker_handles_other_parameterized_generics():
def make_adder(step: int) -> typing.Callable[[int], int]:
return lambda x: x + step

def labels() -> frozenset[str]:
return frozenset({"a", "b"})

def spend() -> htypes.column[pd.Series, float]:
return pd.Series([1.0, 2.0])

def summary(
make_adder: typing.Callable[[int], int],
labels: frozenset[str],
spend: htypes.column[pd.Series, float],
) -> str:
return f"{make_adder(1)} {sorted(labels)} {spend.sum()}"

dr = (
driver.Builder()
.with_modules(ad_hoc_utils.create_temporary_module(make_adder, labels, spend, summary))
.with_adapters(default.FunctionInputOutputTypeChecker())
.build()
)
assert dr.execute(["summary"], inputs={"step": 2}) == {"summary": "3 ['a', 'b'] 3.0"}


def test_function_input_output_type_checker_rejects_wrong_origin_for_parameterized_generic():
def labels() -> frozenset[str]:
return ["a", "b"]

dr = (
driver.Builder()
.with_modules(ad_hoc_utils.create_temporary_module(labels))
.with_adapters(default.FunctionInputOutputTypeChecker())
.build()
)
with pytest.raises(TypeError, match="Node labels returned a result"):
dr.execute(["labels"])
16 changes: 16 additions & 0 deletions tests/test_type_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -412,6 +412,22 @@ def test_check_instance_with_pep604_union_of_generics():
assert not check_instance({"a": "1"}, dict[str, int] | list[int])


def test_check_instance_with_other_parameterized_generics():
def add_one(x: int) -> int:
return x + 1

assert check_instance(add_one, typing.Callable[[int], int])
assert check_instance(int, type[int])
assert check_instance(frozenset({"a"}), frozenset[str])
assert check_instance(iter([1]), typing.Iterator[int])
assert check_instance(range(3), typing.Sequence[int])
assert check_instance(pd.Series([1.0]), htypes.column[pd.Series, float])
# the origin is still checked
assert not check_instance(1, typing.Callable[[int], int])
assert not check_instance([1], frozenset[int])
assert not check_instance([1.0], htypes.column[pd.Series, float])


def test_check_instance_with_union_type_and_literal():
from typing import Literal

Expand Down