diff --git a/python/packages/core/agent_framework/_workflows/_typing_utils.py b/python/packages/core/agent_framework/_workflows/_typing_utils.py index 6496716f1ec..f50f33f6397 100644 --- a/python/packages/core/agent_framework/_workflows/_typing_utils.py +++ b/python/packages/core/agent_framework/_workflows/_typing_utils.py @@ -485,6 +485,13 @@ def is_type_compatible(source_type: type | UnionType | Any, target_type: type | target_origin = get_origin(target_type) target_args = get_args(target_type) + # A bare alias from `typing` (typing.Sequence, typing.List) has an origin but no type + # arguments; treat it as the bare container class it stands for. + if source_origin is not None and not source_args and isinstance(source_origin, type): + source_type, source_origin = source_origin, None + if target_origin is not None and not target_args and isinstance(target_origin, type): + target_type, target_origin = target_origin, None + # Case 2: target is Union/Optional - source is compatible if it matches any target member if target_origin is Union or target_origin is UnionType: # Special case: if source is also a Union, check that each source member @@ -512,6 +519,13 @@ def is_type_compatible(source_type: type | UnionType | Any, target_type: type | return False return source_type == target_type + # A bare container class on one side (list[str] -> list, or list -> list[str]): + # the unparameterized side accepts any element type, so compare the containers. + if target_origin is None and isinstance(target_type, type) and isinstance(source_origin, type): + return issubclass(source_origin, target_type) + if source_origin is None and isinstance(source_type, type) and isinstance(target_origin, type): + return issubclass(source_type, target_origin) + # Case 5: different container types are not compatible if source_origin != target_origin: return False diff --git a/python/packages/core/tests/workflow/test_typing_utils.py b/python/packages/core/tests/workflow/test_typing_utils.py index 805ad197107..eb2fcce85d1 100644 --- a/python/packages/core/tests/workflow/test_typing_utils.py +++ b/python/packages/core/tests/workflow/test_typing_utils.py @@ -441,6 +441,31 @@ class Dog(Animal): assert not is_type_compatible(Animal, Dog) +def test_type_compatibility_bare_container() -> None: + """A bare container class accepts any parameterization of it, in either direction.""" + import typing + from collections.abc import Sequence + + assert is_type_compatible(list[str], list) + assert is_type_compatible(dict[str, int], dict) + assert is_type_compatible(tuple[int, int], tuple) + assert is_type_compatible(list[int], Sequence) + assert is_type_compatible(list, list[str]) + + # Bare aliases from `typing` carry an origin but no arguments. + assert is_type_compatible(list[str], typing.Sequence) + assert is_type_compatible(list[str], typing.List) # noqa: UP006 + assert is_type_compatible(dict[str, int], typing.Mapping) + assert is_type_compatible(typing.List, typing.Sequence) # noqa: UP006 + assert is_type_compatible(typing.List, list[str]) # noqa: UP006 + assert not is_type_compatible(list[str], typing.Mapping) + assert not is_type_compatible(typing.Dict, list[str]) # noqa: UP006 + + assert not is_type_compatible(list[str], dict) + assert not is_type_compatible(dict, list[str]) + assert not is_type_compatible(list[str], str) + + def test_type_compatibility_unions() -> None: """Test type compatibility with Union types.""" # Source matches target union member diff --git a/python/packages/core/tests/workflow/test_validation.py b/python/packages/core/tests/workflow/test_validation.py index 5798c5ec778..82ca127bcd0 100644 --- a/python/packages/core/tests/workflow/test_validation.py +++ b/python/packages/core/tests/workflow/test_validation.py @@ -98,6 +98,25 @@ def test_edge_duplication_validation_fails(): assert exc_info.value.validation_type == ValidationTypeEnum.EDGE_DUPLICATION +def test_parameterized_output_to_bare_container_input_builds(): + # A list[str] output going to a handler typed with plain `list` used to fail + # validation, although the message is a list at runtime. + class ListProducer(Executor): + @handler + async def produce(self, message: str, ctx: WorkflowContext[list[str]]) -> None: + await ctx.send_message([message]) + + class ListConsumer(Executor): + @handler + async def consume(self, messages: list, ctx: WorkflowContext[None, list]) -> None: # type: ignore[type-arg] + await ctx.yield_output(messages) + + producer = ListProducer(id="producer") + consumer = ListConsumer(id="consumer") + + WorkflowBuilder(start_executor=producer).add_edge(producer, consumer).build() + + def test_type_compatibility_validation_fails(): string_executor = StringExecutor(id="string_executor") int_executor = IntExecutor(id="int_executor")