From 314c62374ade94b2d8146a32d95c3a949656a9f6 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 24 Sep 2026 00:24:06 -0400 Subject: [PATCH 1/2] Resolve generic stream types in graph metadata --- src/ezmsg/core/backend.py | 9 ++++- src/ezmsg/core/type_resolution.py | 55 ++++++++++++++++++++++++++++ tests/test_stream_type_resolution.py | 46 +++++++++++++++++++++++ 3 files changed, 108 insertions(+), 2 deletions(-) create mode 100644 src/ezmsg/core/type_resolution.py create mode 100644 tests/test_stream_type_resolution.py diff --git a/src/ezmsg/core/backend.py b/src/ezmsg/core/backend.py index 084282ec..21497b98 100644 --- a/src/ezmsg/core/backend.py +++ b/src/ezmsg/core/backend.py @@ -31,6 +31,7 @@ OutputRelay, ) from .unit import Unit, PROCESS_ATTR, SUBSCRIBES_ATTR, PUBLISHES_ATTR +from .type_resolution import resolve_stream_type from .settings import Settings from .graphmeta import ( CollectionMetadata, @@ -421,7 +422,9 @@ def _component_metadata(self) -> GraphMetadata: else None ), settings_type=( - self._stream_type_name(input_settings.msg_type) + self._stream_type_name( + resolve_stream_type(type(comp), input_settings.msg_type) + ) if isinstance(input_settings, InputStream) else None ), @@ -431,7 +434,9 @@ def _component_metadata(self) -> GraphMetadata: topic_entries: dict[str, TopicMetadataType] = {} relay_entries: dict[str, RelayMetadataType] = {} for stream_name, stream in comp.streams.items(): - msg_type = self._stream_type_name(stream.msg_type) + msg_type = self._stream_type_name( + resolve_stream_type(type(comp), stream.msg_type) + ) if isinstance(stream, InputRelay): runtime = _relay_runtime_info(stream) relay_entries[stream_name] = InputRelayMetadata( diff --git a/src/ezmsg/core/type_resolution.py b/src/ezmsg/core/type_resolution.py new file mode 100644 index 00000000..f570b73d --- /dev/null +++ b/src/ezmsg/core/type_resolution.py @@ -0,0 +1,55 @@ +"""Resolve inherited generic annotations without changing shared stream objects.""" + +import types +from functools import reduce +from operator import or_ +from typing import TypeVar, Union, get_args, get_origin + + +def _substitute(annotation, bindings): + if isinstance(annotation, TypeVar): + seen = set() + while isinstance(annotation, TypeVar) and annotation in bindings and annotation not in seen: + seen.add(annotation) + annotation = bindings[annotation] + return annotation + args = get_args(annotation) + if not args: + return annotation + resolved = tuple(_substitute(arg, bindings) for arg in args) + if resolved == args: + return annotation + if hasattr(annotation, "copy_with"): + return annotation.copy_with(resolved) + origin = get_origin(annotation) + if origin in (Union, types.UnionType): + return reduce(or_, resolved) + try: + return origin[resolved[0] if len(resolved) == 1 else resolved] + except (TypeError, AttributeError): + return annotation + + +def resolve_stream_type(component_type: type, annotation): + """Substitute explicit generic base arguments; leave unbound types unresolved.""" + candidates = {} + + def walk(cls, inherited): + for parent in cls.__dict__.get("__orig_bases__", cls.__bases__): + origin = get_origin(parent) or parent + args = tuple(_substitute(arg, inherited) for arg in get_args(parent)) + bindings = {**inherited, **dict(zip(getattr(origin, "__parameters__", ()), args))} + for variable, value in bindings.items(): + if variable != value: + candidates.setdefault(variable, []).append(value) + if isinstance(origin, type) and origin is not object: + walk(origin, bindings) + + walk(component_type, {}) + # Ambiguous multiple inheritance must not advertise an invented concrete type. + bindings = { + variable: values[0] + for variable, values in candidates.items() + if all(value == values[0] for value in values) + } + return _substitute(annotation, bindings) diff --git a/tests/test_stream_type_resolution.py b/tests/test_stream_type_resolution.py new file mode 100644 index 00000000..a9f3dc76 --- /dev/null +++ b/tests/test_stream_type_resolution.py @@ -0,0 +1,46 @@ +from typing import Generic, TypeVar + +import ezmsg.core as ez +from ezmsg.core.backend import ExecutionContext, GraphRunner +from ezmsg.core.type_resolution import resolve_stream_type + +T = TypeVar("T") +U = TypeVar("U") +S = TypeVar("S") + + +class GenericUnit(ez.Unit, Generic[S, T, U]): + INPUT_SETTINGS = ez.InputStream(S) + INPUT = ez.InputStream(T) + OUTPUT = ez.OutputStream(U) + + +class Intermediate(GenericUnit[ez.Settings, T, list[T]], Generic[T]): + pass + + +class Concrete(Intermediate[int]): + pass + + +def test_specialized_stream_and_settings_metadata(): + unit = Concrete() + ExecutionContext.setup({"UNIT": unit}) + metadata = GraphRunner(components={"UNIT": unit})._component_metadata().components["UNIT"] + assert metadata.streams["INPUT"].msg_type == "builtins.int" + assert metadata.streams["OUTPUT"].msg_type == "list[int]" + assert metadata.dynamic_settings.settings_type == "ezmsg.core.settings.Settings" + assert GenericUnit.__streams__["INPUT"].msg_type is T + assert unit.INPUT.msg_type is T # Metadata does not mutate stream declarations. + + +def test_unspecialized_variables_remain_unresolved(): + assert resolve_stream_type(GenericUnit, T) is T + + +def test_independent_specializations(): + class Text(Intermediate[str]): + pass + + assert resolve_stream_type(Text, U) == list[str] + assert resolve_stream_type(Concrete, U) == list[int] From b52fb5a784ad261c306675ae109f60883f85b3df Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 24 Sep 2026 01:20:24 -0400 Subject: [PATCH 2/2] Preserve parameterized stream metadata on Python 3.10 --- src/ezmsg/core/backend.py | 5 ++++- tests/test_stream_type_resolution.py | 22 ++++++++++++++++++++++ 2 files changed, 26 insertions(+), 1 deletion(-) diff --git a/src/ezmsg/core/backend.py b/src/ezmsg/core/backend.py index 21497b98..e6a5a04e 100644 --- a/src/ezmsg/core/backend.py +++ b/src/ezmsg/core/backend.py @@ -14,6 +14,7 @@ from multiprocessing.synchronize import Barrier as BarrierType from multiprocessing.connection import wait, Connection from socket import socket +from typing import get_origin from .netprotocol import DEFAULT_SHM_SIZE, AddressType @@ -392,7 +393,9 @@ def _type_name(self, tp: type) -> str: return f"{tp.__module__}.{tp.__qualname__}" def _stream_type_name(self, stream_type: object) -> str: - if inspect.isclass(stream_type): + # Python 3.10 considers GenericAlias objects (e.g. list[int]) classes. + # Preserve their parameters instead of naming only the origin class. + if get_origin(stream_type) is None and inspect.isclass(stream_type): return self._type_name(stream_type) return repr(stream_type) diff --git a/tests/test_stream_type_resolution.py b/tests/test_stream_type_resolution.py index a9f3dc76..861aa743 100644 --- a/tests/test_stream_type_resolution.py +++ b/tests/test_stream_type_resolution.py @@ -1,5 +1,7 @@ from typing import Generic, TypeVar +import pytest + import ezmsg.core as ez from ezmsg.core.backend import ExecutionContext, GraphRunner from ezmsg.core.type_resolution import resolve_stream_type @@ -44,3 +46,23 @@ class Text(Intermediate[str]): assert resolve_stream_type(Text, U) == list[str] assert resolve_stream_type(Concrete, U) == list[int] + + +@pytest.mark.parametrize( + "annotation, expected", + [ + (int, "builtins.int"), + (list, "builtins.list"), + (list[int], "list[int]"), + (dict[str, list[int]], "dict[str, list[int]]"), + (tuple[int, str], "tuple[int, str]"), + ], +) +def test_parameterized_stream_metadata_preserves_arguments(annotation, expected): + class TypedUnit(ez.Unit): + OUTPUT = ez.OutputStream(annotation) + + unit = TypedUnit() + ExecutionContext.setup({"UNIT": unit}) + metadata = GraphRunner(components={"UNIT": unit})._component_metadata().components["UNIT"] + assert metadata.streams["OUTPUT"].msg_type == expected