diff --git a/hamilton/htypes.py b/hamilton/htypes.py index 38fc1a419..d94bcb084 100644 --- a/hamilton/htypes.py +++ b/hamilton/htypes.py @@ -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_) diff --git a/tests/lifecycle/test_default.py b/tests/lifecycle/test_default.py index 30e4ef9c8..aa9cdef41 100644 --- a/tests/lifecycle/test_default.py +++ b/tests/lifecycle/test_default.py @@ -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 @@ -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"]) diff --git a/tests/test_type_utils.py b/tests/test_type_utils.py index 15bbe5a26..77a248ad2 100644 --- a/tests/test_type_utils.py +++ b/tests/test_type_utils.py @@ -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