diff --git a/src/node_graph/link.py b/src/node_graph/link.py index 2e2d535..1386a0d 100644 --- a/src/node_graph/link.py +++ b/src/node_graph/link.py @@ -85,6 +85,142 @@ def _annotated_type(self, sock: "Socket") -> str | None: return self._format_union(union) return self._annotated_py_type(sock) + def _socket_extras(self, sock: "Socket") -> dict: + return getattr(getattr(sock, "_metadata", None), "extras", {}) or {} + + def _allowed_values(self, sock: "Socket") -> list | None: + """Return the values this socket may carry, or None if unrestricted. + + A ``Literal`` socket lists them outright; a plain ``Enum`` socket + carries all of its members' values. + """ + extras = self._socket_extras(sock) + allowed = extras.get("allowed_values") + if isinstance(allowed, list): + return allowed + info = extras.get("structured_type") + if isinstance(info, dict) and info.get("kind") == "enum": + from node_graph.utils.struct_utils import import_structured_type + + try: + return [member.value for member in import_structured_type(info["path"])] + except Exception: + return None + return None + + def _literal_base(self, sock: "Socket") -> str | None: + base = self._socket_extras(sock).get("literal_base") + return str(base).lower() if base else None + + def check_allowed_values(self) -> bool | None: + """Decide a link where either end restricts the values it carries. + + Returns True to accept the link, None to leave the decision to the + remaining checks, and raises when the source can carry a value the + target forbids. Reached only for a source whose type is known: an + ``Any`` source is accepted earlier on type alone, which is what lets an + untyped parent pass a value on, and ``check_static_source_value`` reads + whatever value it already holds. + """ + from node_graph.utils.struct_utils import value_is_allowed + + from_allowed = self._allowed_values(self.from_socket) + to_allowed = self._allowed_values(self.to_socket) + if from_allowed is None and to_allowed is None: + return None + + if from_allowed is not None and to_allowed is not None: + offending = [ + value + for value in from_allowed + if not value_is_allowed(value, to_allowed) + ] + if not offending: + return True + self._raise_allowed_values_mismatch( + [repr(value) for value in offending], to_allowed + ) + + if to_allowed is not None: + # The source is unrestricted, so it can carry a forbidden value. + self._raise_allowed_values_mismatch(None, to_allowed) + + # Only the source restricts: widening into its own base type is safe. + base = self._literal_base(self.from_socket) + if base is not None and self._lower_id(self.to_socket) == base: + return True + return None + + def check_static_source_value(self) -> None: + """Reject a link whose source already holds a value the target forbids. + + A source socket declared ``Any`` passes every type check, so a value + already sitting on it is read here, the last point before the run. + Everything else is decided at run time, when the value exists: a source + fed by an upstream task, a source assigned after this link was made, + and a value bound for a sub-graph one hop further out, whose restricted + socket only appears as that sub-graph expands. + """ + from node_graph.socket import TaggedValue + + from node_graph.utils.struct_utils import ( + canonical_socket_value, + socket_subject, + ) + + extras = self._socket_extras(self.to_socket) + structured_type = extras.get("structured_type") or {} + if "allowed_values" not in extras and structured_type.get("kind") != "enum": + return + if self._is_namespace(self.from_socket): + # A namespace carries a mapping; the shape check below reports it. + return + if any(link.to_socket is self.from_socket for link in self.from_socket._links): + # Fed by an upstream task: the value only exists at run time. + return + value = getattr(self.from_socket, "_value", None) + if value is None: + return + if isinstance(value, TaggedValue): + tagged_socket = value._socket + # ``is`` only: comparing sockets with ``==`` builds an operator task. + if tagged_socket is not None and tagged_socket is not self.from_socket: + return + value = value.__wrapped__ + canonical_socket_value( + value, + structured_type=structured_type or None, + allowed=extras.get("allowed_values"), + subject=socket_subject(self.to_socket._full_name_with_task), + ) + + def _raise_allowed_values_mismatch( + self, offending: list[str] | None, allowed: list + ) -> None: + src = f"{self.from_task.name}.{self.from_socket._scoped_name}" + dst = f"{self.to_task.name}.{self.to_socket._scoped_name}" + rendered = ", ".join(repr(value) for value in allowed) + detail = ( + f" Source can carry {', '.join(offending)}, which {dst} forbids." + if offending + else " Source is not restricted to those values." + ) + raise TypeError( + "\n".join( + [ + "Socket value range mismatch:", + f" {src} [{self._format_socket_id(self.from_socket)}] -> {dst} " + f"[{self._format_socket_id(self.to_socket)}] is not allowed.", + f" {dst} accepts only {rendered}.", + detail, + "", + "Suggestions:", + " • Widen the target annotation to include the source's values", + " • Or insert a task that narrows the value before the link", + ] + ) + ) + def _namespace_item(self, sock: "Socket") -> dict | None: extras = getattr(getattr(sock, "_metadata", None), "extras", {}) or {} item = extras.get("item") @@ -169,6 +305,8 @@ def check_socket_match(self) -> None: from_id = self._lower_id(self.from_socket) to_id = self._lower_id(self.to_socket) + self.check_static_source_value() + # "any" accepts anything if self._is_any(self.from_socket) or self._is_any(self.to_socket): return @@ -254,6 +392,9 @@ def check_socket_match(self) -> None: ): return + if self.check_allowed_values(): + return + if self._is_annotated(self.from_socket) and self._is_annotated(self.to_socket): from_type = self._annotated_py_type(self.from_socket) to_type = self._annotated_py_type(self.to_socket) diff --git a/src/node_graph/socket.py b/src/node_graph/socket.py index 6ad9f5b..35a8947 100644 --- a/src/node_graph/socket.py +++ b/src/node_graph/socket.py @@ -708,6 +708,8 @@ def _set_socket_value(self, value: Any, *, value_source: str = "link") -> None: elif is_input and value_source != "property": self._update_updatable_meta({"value_source": "link"}) + value = self._canonical_value(value) + graph = getattr(self._task, "graph", None) if graph is not None: policy = getattr(graph, "serialization_policy", "off") @@ -722,6 +724,50 @@ def _set_socket_value(self, value: Any, *, value_source: str = "link") -> None: f"Socket '{self._name}' has no property to set a value." ) + def _canonical_value(self, value: Any) -> Any: + """Return the value this socket accepts, raising when it accepts none. + + Input sockets built from an ``Enum`` or a ``Literal`` carry the members + they admit; every other socket returns ``value`` unchanged. What is + stored for an ``Enum`` is the member's value, the one form that also + comes back out of storage, so a socket reads the same whichever side of + a process boundary filled it; ``coerce_inputs_from_spec`` rebuilds the + member for the task body. ``None`` passes through, so an optional socket + can still be cleared. Output sockets are left alone: a task reports what + it computed. + """ + from node_graph.utils.struct_utils import ( + canonical_socket_value, + literal_value, + retagged, + socket_subject, + ) + + if self._full_name.split(".")[0] != "inputs": + return value + extras = self._metadata.extras or {} + structured_type = extras.get("structured_type") + allowed = extras.get("allowed_values") + if allowed is None and ( + structured_type is None or structured_type.get("kind") != "enum" + ): + return value + if value is None: + return value + + raw = value.__wrapped__ if isinstance(value, TaggedValue) else value + canonical = literal_value( + canonical_socket_value( + raw, + structured_type=structured_type, + allowed=allowed, + subject=socket_subject(self._full_name_with_task), + ) + ) + if canonical is raw: + return value + return retagged(canonical, value) + def _serialize_value(self, store: bool = False) -> Any: """Serialize the socket value unless it's metadata (stored as raw).""" value = resolve_tagged_values(self._value) diff --git a/src/node_graph/socket_spec.py b/src/node_graph/socket_spec.py index 2a1d479..d45818a 100644 --- a/src/node_graph/socket_spec.py +++ b/src/node_graph/socket_spec.py @@ -11,9 +11,18 @@ ) import inspect from copy import deepcopy +from enum import Enum from node_graph.orm.mapping import type_mapping as DEFAULT_TM from node_graph.socket_meta import CallRole, SocketMeta, merge_meta -from node_graph.utils.struct_utils import structured_type_info +from node_graph.utils.struct_utils import ( + JSON_SAFE_LITERALS, + canonical_socket_value, + is_enum_type, + literal_value, + structured_type_info, + structured_type_path, + value_is_allowed, +) from .socket import TaskSocketNamespace import ast import textwrap @@ -24,7 +33,7 @@ MISSING as _DC_MISSING, ) -from typing import Annotated, get_args, get_origin, get_type_hints +from typing import Annotated, Literal, get_args, get_origin, get_type_hints try: from typing import Unpack @@ -68,6 +77,17 @@ def _is_union_origin(origin: Any) -> bool: return (origin is Union) or (_UNION_TYPE is not None and origin is _UNION_TYPE) +def _literal_arg_name(arg: Any) -> str: + """Render one ``Literal`` argument. + + An ``Enum`` member is written out with its class's import path, so two + same-named enums from different modules do not render alike. + """ + if isinstance(arg, Enum): + return f"{structured_type_path(type(arg))}.{arg.name}" + return repr(arg) + + def _is_annotated_type(tp: Any) -> bool: """True if tp is an Annotated[...] wrapper across versions.""" return get_origin(tp) is Annotated @@ -334,6 +354,35 @@ class SocketSpec: fields: Dict[str, "SocketSpec"] = field(default_factory=dict) meta: SocketMeta = field(default_factory=SocketMeta) + def __post_init__(self) -> None: + """Hold the default a socket assignment would have accepted. + + A default reaches ``TaskProperty.value`` without passing the socket's + own setter, so it is decided here instead: a default an ``Enum`` or + ``Literal`` socket forbids raises where the spec is built, and one it + admits is stored as the bare value the socket would have kept. + """ + if isinstance(self.default, type(MISSING)) or self.default is None: + return + extras = self.meta.extras or {} + structured_type = extras.get("structured_type") + allowed = extras.get("allowed_values") + if allowed is None and ( + structured_type is None or structured_type.get("kind") != "enum" + ): + return + canonical = literal_value( + canonical_socket_value( + self.default, + structured_type=structured_type, + allowed=allowed, + subject=f"the default of a socket annotated " + f"{extras.get('py_type', self.identifier)}", + ) + ) + if canonical is not self.default: + object.__setattr__(self, "default", canonical) + @property def dynamic(self) -> bool: return bool(self.meta.dynamic) @@ -408,7 +457,7 @@ def __getattr__(self, name: str) -> "SocketView": raise AttributeError("'.item' only valid on dynamic namespace specs") if spec.fields and name in spec.fields: return SocketView(spec.fields[name]) - raise AttributeError(f"'{name}' not found in namespace spec") + raise AttributeError(f"namespace spec has no field '{name}'") def __getitem__(self, name: str) -> "SocketView": return self.__getattr__(name) @@ -551,6 +600,9 @@ def _map_identifier(cls, tp: Any) -> str: @staticmethod def _py_type_name(tp: Any) -> str: + if get_origin(tp) is Literal: + args = ", ".join(_literal_arg_name(arg) for arg in get_args(tp)) + return f"typing.Literal[{args}]" module = getattr(tp, "__module__", None) qualname = getattr(tp, "__qualname__", None) or getattr(tp, "__name__", None) if module and qualname: @@ -772,6 +824,56 @@ def from_model(cls, model_cls: Type[Any]) -> SocketSpec: "from_model expects a Pydantic BaseModel subclass, dataclass, or TypedDict." ) + @classmethod + def _leaf_from_literal(cls, T: Any) -> SocketSpec: + """Build the leaf spec for ``Literal[...]``. + + The socket keeps the permitted values in the ``allowed_values`` extra. + ``structured_type`` is added when one ``Enum`` supplies every argument, + whether the argument was written as a member or as that member's value, + so the member is rebuilt after serialization; arguments spanning two + enums, or naming a value no member of the one enum carries, keep their + values alone. ``literal_base`` records the type such a socket widens + into; an enum socket records none, because its output side carries the + member rather than the bare value. An argument no spec round trip can + carry (``bytes``) leaves the socket unconstrained, as an unrecognized + annotation always has. + """ + args = list(get_args(T)) + py_type = cls._py_type_name(T) + if not args: + return SocketSpec(identifier=cls.DEFAULT, meta=SocketMeta()) + + enum_classes = {type(arg) for arg in args if isinstance(arg, Enum)} + values: list = [] + for arg in args: + value = literal_value(arg) + if not value_is_allowed(value, values): + values.append(value) + if not all(isinstance(value, JSON_SAFE_LITERALS) for value in values): + return SocketSpec( + identifier=cls.ANNOTATED, meta=SocketMeta(extras={"py_type": py_type}) + ) + + extras: Dict[str, Any] = {"py_type": py_type, "allowed_values": values} + enum_cls = next(iter(enum_classes)) if len(enum_classes) == 1 else None + member_values = [member.value for member in enum_cls] if enum_cls else [] + if enum_cls is not None and all( + value_is_allowed(value, member_values) for value in values + ): + extras["structured_type"] = structured_type_info(enum_cls) + else: + base = cls._literal_base_type(values) + if base is not None: + extras["literal_base"] = cls._leaf_from_type(base).identifier + return SocketSpec(identifier=cls.ANNOTATED, meta=SocketMeta(extras=extras)) + + @staticmethod + def _literal_base_type(values: list) -> Any: + """Return the single type behind ``values``, or ``None`` if they differ.""" + types_seen = {type(value) for value in values} + return next(iter(types_seen)) if len(types_seen) == 1 else None + @classmethod def _leaf_from_type(cls, T: Any) -> SocketSpec: if T is Any or T is inspect._empty or T is object: @@ -780,7 +882,20 @@ def _leaf_from_type(cls, T: Any) -> SocketSpec: if _is_typeddict_type(T): return cls._leaf_from_type(dict) + if is_enum_type(T): + return SocketSpec( + identifier=cls.ANNOTATED, + meta=SocketMeta( + extras={ + "py_type": cls._py_type_name(T), + "structured_type": structured_type_info(T), + } + ), + ) + origin = get_origin(T) + if origin is Literal: + return cls._leaf_from_literal(T) if _is_union_origin(origin): args = [] for arg in get_args(T): @@ -1050,10 +1165,16 @@ def build_inputs_from_signature( ): info = structured_type_info(base_T) if info is not None and "structured_type" not in spec.meta.extras: + # ``required=None`` so this overlay says nothing about + # requiredness: ``SocketMeta.required`` defaults to ``True`` + # and ``merge_meta`` prefers any non-``None`` overlay value, + # so a bare ``SocketMeta`` would republish ``required=True`` + # over the value computed from the parameter's default. spec = replace( spec, meta=merge_meta( - spec.meta, SocketMeta(extras={"structured_type": info}) + spec.meta, + SocketMeta(required=None, extras={"structured_type": info}), ), ) diff --git a/src/node_graph/utils/graph.py b/src/node_graph/utils/graph.py index ca344e9..df977ce 100644 --- a/src/node_graph/utils/graph.py +++ b/src/node_graph/utils/graph.py @@ -134,6 +134,41 @@ def _assign_graph_outputs(outputs: Any, graph: Graph) -> None: ) +def _deserialize_inputs(namespace: Any, values: Any, adapter: Any) -> Any: + """Recursively apply ``adapter.deserialize`` to leaves of ``values`` so a + ``@task.graph`` body receives the primitive its signature declares, even + when the stored value was wrapped by the engine for provenance.""" + from node_graph.socket import TaskSocketNamespace + + if adapter is None or not hasattr(adapter, "deserialize"): + return values + if not isinstance(values, dict): + return values + from node_graph.utils.struct_utils import is_structured_instance + + out = dict(values) + for name, item in namespace._sockets.items(): + if name not in out: + continue + value = out[name] + if isinstance(item, TaskSocketNamespace): + if isinstance(value, dict): + out[name] = _deserialize_inputs(item, value, adapter) + elif is_structured_instance(value): + # The namespace was already materialised into a dataclass / + # Pydantic instance by ``coerce_inputs_from_spec`` (which + # runs before ``_deserialize_inputs`` in ``materialize_graph``). + # Without this branch, the adapter never sees the structured + # value, so a serialiser that auto-promotes primitive fields + # to engine-typed wrappers (e.g. ``aiida-workgraph``'s + # ``orm.Int`` / ``orm.Float``) leaves them wrapped inside + # the dataclass and downstream ``int``-typed code breaks. + out[name] = adapter.deserialize(value, item) + else: + out[name] = adapter.deserialize(value, item) + return out + + def materialize_graph( func: Callable, in_spec: SocketSpec, @@ -173,6 +208,9 @@ def materialize_graph( tag_socket_value(graph.inputs) inputs = graph.inputs._collect_values(unwrap=False) inputs = coerce_inputs_from_spec(inputs, in_spec) + inputs = _deserialize_inputs( + graph.inputs, inputs, getattr(graph, "serialization", None) + ) raw = func(**inputs) _assign_graph_outputs(raw, graph) tag_socket_value(graph.inputs, only_uuid=True) diff --git a/src/node_graph/utils/struct_utils.py b/src/node_graph/utils/struct_utils.py index 3381b17..059e0f9 100644 --- a/src/node_graph/utils/struct_utils.py +++ b/src/node_graph/utils/struct_utils.py @@ -1,12 +1,18 @@ from __future__ import annotations from dataclasses import asdict, is_dataclass +from enum import Enum from typing import Any, Dict import importlib from pydantic import BaseModel +def is_enum_type(tp: Any) -> bool: + """Return True for Enum subclasses (including ``str``- and ``int``-Enum).""" + return isinstance(tp, type) and issubclass(tp, Enum) + + def is_structured_instance(value: Any) -> bool: """Return True for dataclass or Pydantic model instances.""" return (is_dataclass(value) and not isinstance(value, type)) or isinstance( @@ -29,11 +35,13 @@ def structured_to_dict(value: Any) -> Any: def structured_type_info(tp: Any) -> Dict[str, str] | None: - """Return a serializable descriptor for structured types.""" + """Return a serializable descriptor for structured types (incl. ``Enum``).""" if is_dataclass(tp): return {"kind": "dataclass", "path": structured_type_path(tp)} if isinstance(tp, type) and issubclass(tp, BaseModel): return {"kind": "pydantic", "path": structured_type_path(tp)} + if is_enum_type(tp): + return {"kind": "enum", "path": structured_type_path(tp)} return None @@ -50,16 +58,198 @@ def import_structured_type(path: str) -> Any: return current -def coerce_structured_value(value: Any, info: Dict[str, str] | None) -> Any: - """Rebuild a structured instance from a dict when spec says so.""" +#: Literal arguments a spec can carry through a to_dict/from_dict round trip. +JSON_SAFE_LITERALS = (str, bool, int, type(None)) + + +def untagged(value: Any) -> Any: + """Return the value a ``TaggedValue`` wraps, else ``value``. + + Canonicalization builds a fresh object, so the proxy has to come off first: + ``isinstance`` reports the wrapped type and would otherwise pass a proxy on + where callers expect a plain value. + """ + from node_graph.socket import TaggedValue + + return value.__wrapped__ if isinstance(value, TaggedValue) else value + + +def retagged(value: Any, original: Any) -> Any: + """Return ``value`` carrying ``original``'s tag, when it had one. + + Rebuilding a value loses the ``TaggedValue`` a socket wrapped it in, and + with it the socket the graph body needs to raise a link instead of a + literal. The uuid is carried over so provenance still points at one value. + """ + from node_graph.socket import TaggedValue + + if not isinstance(original, TaggedValue): + return value + tagged = TaggedValue(value, socket=original._socket) + tagged._self_uuid = original._uuid + return tagged + + +def literal_value(value: Any) -> Any: + """Return the bare value an ``Enum`` member stands for, else ``value``.""" + return value.value if isinstance(value, Enum) else value + + +def value_is_allowed(value: Any, allowed: Any) -> bool: + """Return True when ``value`` is one of ``allowed``. + + Two numbers match only when their types agree, which is typing's rule that + ``True`` is not ``1`` and ``1`` is not ``1.0``. Anything else matches on + equality alone, so a value that arrives wrapped, such as a storage node + holding ``'none'``, still names the value it equals. + """ + return any(_values_match(item, value) for item in allowed) + + +def _values_match(candidate: Any, value: Any) -> bool: + """Return True when ``value`` is the same value as ``candidate``.""" + numbers = (bool, int, float, complex) + if isinstance(candidate, numbers) and isinstance(value, numbers): + return type(candidate) is type(value) and candidate == value + try: + return bool(candidate == value) + except (TypeError, ValueError): + # An equality that yields an array, or none at all, decides nothing. + return False + + +def format_allowed_values(allowed: Any) -> str: + """Render allowed values as ``'a', 'b' or 'c'``.""" + rendered = [repr(item) for item in allowed] + if not rendered: + return "" + if len(rendered) == 1: + return rendered[0] + return f"{', '.join(rendered[:-1])} or {rendered[-1]}" + + +def socket_subject(name: str) -> str: + """Render the phrase an error uses to point at the socket ``name``.""" + return f"socket '{name}'" + + +def _invalid_value_error( + value: Any, allowed: Any, subject: str | None, hint: str +) -> ValueError: + return ValueError( + f"Invalid value for {subject or 'this input'}.\n" + f" Input should be {format_allowed_values(allowed)}. Got {value!r}.\n" + f" {hint}" + ) + + +def canonical_enum_member( + value: Any, + cls: type, + *, + allowed: Any = None, + subject: str | None = None, +) -> Any: + """Return the ``cls`` member ``value`` names, raising if it names none. + + A member of ``cls``, a member of any other ``Enum`` carrying the same + value, and a bare member value all name the same member: membership is + decided by ``value_is_allowed`` over the members' values, because + serialization keeps only the value. That is the rule a ``Literal`` socket + applies too, so an ``IntEnum`` whose members are ``1`` and ``2`` rejects + ``True`` exactly as ``Literal[1, 2]`` does. ``allowed``, when given, + restricts the result to that subset of member values. + """ + permitted = [member.value for member in cls] if allowed is None else list(allowed) + value = untagged(value) + if isinstance(value, cls): + member = value + else: + candidate = literal_value(value) + member = next( + (item for item in cls if value_is_allowed(item.value, [candidate])), + None, + ) + if member is None: + example = next(iter(cls.__members__), None) + named = f" ({cls.__name__}.{example})" if example else "" + raise _invalid_value_error( + value, + permitted, + subject, + f"{structured_type_path(cls)} members are accepted by " + f"member{named} or by value.", + ) + if not value_is_allowed(member.value, permitted): + raise _invalid_value_error( + value, + permitted, + subject, + f"The socket admits only part of {structured_type_path(cls)}.", + ) + return member + + +def canonical_literal_value( + value: Any, + allowed: Any, + *, + subject: str | None = None, +) -> Any: + """Return ``value`` when it is one of ``allowed``, raising otherwise. + + Matching follows typing's rule that ``True`` is not ``1``: a value matches + only a candidate of its own type. An ``Enum`` member matches by its value. + """ + candidate = literal_value(untagged(value)) + if value_is_allowed(candidate, allowed): + return candidate + raise _invalid_value_error( + value, allowed, subject, "Only the values listed above are accepted here." + ) + + +def canonical_socket_value( + value: Any, + *, + structured_type: Dict[str, str] | None = None, + allowed: Any = None, + subject: str | None = None, +) -> Any: + """Return the value a socket may store, raising when the socket forbids it. + + ``value`` comes back untouched when the socket constrains nothing. + """ + if structured_type is not None and structured_type.get("kind") == "enum": + cls = import_structured_type(structured_type["path"]) + return canonical_enum_member(value, cls, allowed=allowed, subject=subject) + if allowed is not None: + return canonical_literal_value(value, allowed, subject=subject) + return value + + +def coerce_structured_value( + value: Any, info: Dict[str, str] | None, name: str | None = None +) -> Any: + """Rebuild a structured instance from a flat value when spec says so.""" if info is None: return value + kind = info.get("kind") + # Enum check goes before the dict gate: its serialized form is a bare value, + # so an already-materialized-instance/dict short-circuit would never fire. + # Assignment already decided membership, so this only rebuilds the member. + if kind == "enum": + cls = import_structured_type(info["path"]) + subject = socket_subject(name) if name else None + return retagged(canonical_enum_member(value, cls, subject=subject), value) if is_structured_instance(value): return value if not isinstance(value, dict): return value + # Only import the target class once we know we must rebuild from a dict. + # Imports of already-materialized instances (e.g. locally-defined types + # whose qualname is not importable) must not reach here. cls = import_structured_type(info["path"]) - kind = info.get("kind") if kind == "pydantic": if contains_tagged_value(value): if hasattr(cls, "model_construct"): @@ -88,7 +278,7 @@ def coerce_inputs_from_spec(values: Any, spec: Any) -> Any: continue info = child.meta.extras.get("structured_type") if info: - out[name] = coerce_structured_value(out[name], info) + out[name] = coerce_structured_value(out[name], info, name) continue if child.is_namespace() and isinstance(out[name], dict): out[name] = coerce_inputs_from_spec(out[name], child) diff --git a/tests/test_engine_local.py b/tests/test_engine_local.py index 76a3037..e3f06f7 100644 --- a/tests/test_engine_local.py +++ b/tests/test_engine_local.py @@ -7,6 +7,7 @@ from node_graph.engine.provenance import ProvenanceRecorder from typing import Annotated, Any from dataclasses import dataclass +from enum import Enum from pydantic import BaseModel @@ -213,3 +214,56 @@ def test_local_engine_accepts_dataclass_models(): assert results["sum"] == 7 assert results["product"] == 12 + + +class Color(str, Enum): + RED = "red" + GREEN = "green" + BLUE = "blue" + + +class Priority(int, Enum): + LOW = 1 + HIGH = 9 + + +@task(outputs=ns(name=str)) +def color_name(color: Color) -> dict: + assert isinstance(color, Color), f"expected Color, got {type(color).__name__}" + return {"name": color.name} + + +@task(outputs=ns(name=str)) +def priority_name(priority: Priority) -> dict: + assert isinstance( + priority, Priority + ), f"expected Priority, got {type(priority).__name__}" + return {"name": priority.name} + + +def test_local_engine_coerces_str_enum(): + # A ``str``-Enum member flattens to a bare ``str`` across a serialization + # boundary (it *is* a str). Feed the engine that flattened form, not the + # member itself, so the coercion has something real to rebuild; passing + # ``Color.GREEN`` directly would leave the body's isinstance check trivially + # satisfied and the test unable to fail. The signature records the enum type, + # so ``coerce_inputs_from_spec`` must turn ``"green"`` back into ``Color.GREEN`` + # before the body runs. + ng = Graph(name="local-enum", outputs=ns(name=str)) + node = ng.add_task(color_name, "pick", color="green") + ng.add_link(node.outputs.name, ng.outputs.name) + + results = LocalEngine().run(ng) + + assert results["name"] == "GREEN" + + +def test_local_engine_coerces_int_enum(): + # Same for an ``int``-Enum: its flattened form is a bare ``int``. + ng = Graph(name="local-int-enum", outputs=ns(name=str)) + node = ng.add_task(priority_name, "pick", priority=9) + ng.add_link(node.outputs.name, ng.outputs.name) + + results = LocalEngine().run(ng) + + assert results["name"] == "HIGH" diff --git a/tests/test_enum_literal_sockets.py b/tests/test_enum_literal_sockets.py new file mode 100644 index 0000000..8a59971 --- /dev/null +++ b/tests/test_enum_literal_sockets.py @@ -0,0 +1,464 @@ +"""An enum or Literal socket accepts what names one of its members, nothing else.""" + +from enum import Enum, IntEnum +from typing import Literal, Optional + +import pytest + +from node_graph import Graph, task +from node_graph.engine.local import LocalEngine +from node_graph.socket_spec import SocketSpecAPI as api + + +class Spin(Enum): + NONE = "none" + COLLINEAR = "collinear" + SPIN_ORBIT = "spin_orbit" + + +class Narrow(Enum): + """A separate enum whose members repeat two of ``Spin``'s values.""" + + NONE = "none" + COLLINEAR = "collinear" + + +class NameOnly(Enum): + """Shares a member NAME with ``Narrow``, but not its value.""" + + NONE = "something_else" + + +class Priority(IntEnum): + """An enum whose members carry bare ints.""" + + LOW = 1 + HIGH = 2 + + +@task() +def consume_narrow(spin: Narrow) -> str: + assert isinstance(spin, Narrow), f"expected Narrow, got {type(spin).__name__}" + return spin.value + + +@task.graph() +def narrow_graph(spin: Narrow) -> str: + return consume_narrow(spin=spin).result + + +@task.graph() +def defer_literal(spin) -> str: + """Call the enum-typed sub-graph with whatever the parent was given.""" + return narrow_graph(spin=spin) + + +ACCEPTED = [ + pytest.param(Narrow.NONE, id="declared-member"), + pytest.param(Spin.NONE, id="foreign-member-same-value"), + pytest.param("none", id="bare-value"), +] + +REJECTED = [ + pytest.param("banana", id="unknown-value"), + pytest.param(Spin.SPIN_ORBIT, id="foreign-member-unknown-value"), + pytest.param(NameOnly.NONE, id="foreign-member-matching-name-only"), + pytest.param(42, id="wrong-type"), +] + + +@pytest.mark.parametrize("value", ACCEPTED) +def test_entry_graph_accepts_anything_naming_a_member(value): + ng = narrow_graph.build(spin=value) + + # Storage keeps the member's value, the form that survives a process + # boundary; the task body is handed the member back. + assert ng.inputs.spin.value == Narrow.NONE.value + + +@pytest.mark.parametrize("value", REJECTED) +def test_entry_graph_rejects_a_value_naming_no_member(value): + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + narrow_graph.build(spin=value) + + +@pytest.mark.parametrize("value", ACCEPTED) +def test_deferred_subgraph_accepts_anything_naming_a_member(value): + @task.graph() + def outer() -> str: + return narrow_graph(spin=value) + + ng = outer.build() + + assert ng.tasks.narrow_graph.inputs.spin.value == Narrow.NONE.value + + +@pytest.mark.parametrize("value", REJECTED) +def test_deferred_subgraph_rejects_at_build_not_at_run(value): + @task.graph() + def outer() -> str: + return narrow_graph(spin=value) + + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + outer.build() + + +@pytest.mark.parametrize("value", REJECTED) +def test_link_from_untyped_parent_socket_rejects_at_build(value): + """The parent socket is untyped, so the link is where the value is read.""" + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + defer_literal.build(spin=value) + + +def test_error_names_the_socket_and_omits_the_value_wrapper(): + with pytest.raises(ValueError) as excinfo: + narrow_graph.build(spin="banana") + + message = str(excinfo.value) + assert "graph_inputs.inputs.spin" in message + assert "TaggedValue" not in message + assert "uuid" not in message + + +def test_body_receives_a_member_rebuilt_from_the_stored_value(): + """A run stores the bare value; the body still sees the member.""" + ng = Graph(name="enum-readback") + node = ng.add_task(consume_narrow, "pick", spin=Spin.COLLINEAR) + ng.outputs.result = node.outputs.result + + results = LocalEngine().run(ng) + + assert results["result"] == "collinear" + + +@pytest.mark.parametrize("value", ACCEPTED) +def test_an_enum_graph_input_reaches_the_task_as_a_link(value): + """The rebuilt member keeps the socket, so the body wires rather than copies.""" + ng = narrow_graph.build(spin=value) + + sources = {link.from_task.name for link in ng.links} + + assert "graph_inputs" in sources + + +def test_an_enum_parameter_keeps_the_requiredness_of_its_default(): + def signature(a: Spin, b: Spin = Spin.NONE, c: Optional[Spin] = None): + ... + + fields = api.build_inputs_from_signature(signature).fields + + assert fields["a"].meta.required is True + assert fields["b"].meta.required is False + assert fields["b"].default == Spin.NONE.value + assert fields["c"].meta.required is False + + +# --- defaults -------------------------------------------------------------- + + +@task() +def consume_defaulted(spin: Narrow = Narrow.COLLINEAR) -> str: + assert isinstance(spin, Narrow), f"expected Narrow, got {type(spin).__name__}" + return spin.value + + +def test_a_default_is_stored_as_the_value_an_assignment_would_store(): + ng = Graph(name="defaults") + defaulted = ng.add_task(consume_defaulted, "defaulted") + assigned = ng.add_task(consume_defaulted, "assigned", spin=Narrow.COLLINEAR) + + assert defaulted.inputs.spin.value == assigned.inputs.spin.value + assert defaulted.inputs.spin._serialize_value() == Narrow.COLLINEAR.value + + +def test_a_defaulted_socket_reads_the_same_after_a_graph_round_trip(): + ng = Graph(name="defaults-round-trip") + ng.add_task(consume_defaulted, "defaulted") + + restored = Graph.from_dict(ng.to_dict()) + + assert restored.tasks.defaulted.inputs.spin.value == Narrow.COLLINEAR.value + + +def test_a_defaulted_socket_still_hands_the_body_a_member(): + ng = Graph(name="defaults-body") + node = ng.add_task(consume_defaulted, "defaulted") + ng.outputs.result = node.outputs.result + + assert LocalEngine().run(ng)["result"] == "collinear" + + +def test_a_default_naming_no_member_is_rejected_where_the_spec_is_built(): + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + + @task() + def bad_default(spin: Narrow = Spin.SPIN_ORBIT) -> str: + return "unreachable" + + +def test_a_literal_default_outside_the_allowed_set_is_rejected(): + """The value a build rejects is rejected as a default too.""" + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + + @task.graph() + def bad_literal_default(spin: Literal["none", "collinear"] = "banana") -> str: + return consume_narrow(spin=Narrow.NONE).result + + +def test_optional_enum_socket_still_takes_none(): + @task.graph() + def maybe(spin: Narrow = None) -> str: + return consume_narrow(spin=Narrow.NONE).result + + ng = maybe.build(spin=None) + + assert ng.inputs.spin.value is None + + +# --- Literal --------------------------------------------------------------- + + +def test_literal_of_enum_members_carries_the_enum_and_the_subset(): + spec = api._leaf_from_type(Literal[Spin.NONE, Spin.COLLINEAR]) + + assert spec.meta.extras["allowed_values"] == ["none", "collinear"] + assert spec.meta.extras["structured_type"]["kind"] == "enum" + assert spec.meta.extras["structured_type"]["path"].endswith(".Spin") + + +def test_literal_of_strings_carries_its_base_type(): + spec = api._leaf_from_type(Literal["none", "collinear"]) + + assert spec.meta.extras["allowed_values"] == ["none", "collinear"] + assert spec.meta.extras["literal_base"] == api._leaf_from_type(str).identifier + assert "structured_type" not in spec.meta.extras + + +def test_literal_of_mixed_types_keeps_the_values_without_a_base_type(): + spec = api._leaf_from_type(Literal["a", 1]) + + assert spec.meta.extras["allowed_values"] == ["a", 1] + assert "literal_base" not in spec.meta.extras + + +def test_literal_of_unrepresentable_arguments_constrains_nothing(): + spec = api._leaf_from_type(Literal[b"x"]) + + assert "allowed_values" not in spec.meta.extras + + +def test_a_literal_socket_survives_a_spec_round_trip(): + from node_graph.socket_spec import SocketSpec + + spec = api._leaf_from_type(Literal[Spin.NONE, Spin.COLLINEAR]) + + restored = SocketSpec.from_dict(spec.to_dict()) + + assert restored.meta.extras == spec.meta.extras + + +def test_literal_naming_one_enum_by_value_still_carries_the_enum(): + """``Spin.NONE`` and ``'collinear'`` name members of the same enum.""" + spec = api._leaf_from_type(Literal[Spin.NONE, "collinear"]) + + assert spec.meta.extras["allowed_values"] == ["none", "collinear"] + assert spec.meta.extras["structured_type"]["path"].endswith(".Spin") + + +def test_literal_naming_a_value_no_member_carries_keeps_the_values_alone(): + spec = api._leaf_from_type(Literal[Spin.NONE, "banana"]) + + assert spec.meta.extras["allowed_values"] == ["none", "banana"] + assert "structured_type" not in spec.meta.extras + + +def test_py_type_name_distinguishes_two_literals(): + assert api._py_type_name(Literal[1, 2]) != api._py_type_name(Literal["a", "b"]) + assert api._py_type_name(Literal["a", "b"]) == "typing.Literal['a', 'b']" + + +def test_py_type_name_distinguishes_same_named_enums_from_two_modules(): + here = Enum("Duplicate", {"NONE": "none"}, module="package_a") + there = Enum("Duplicate", {"NONE": "none"}, module="package_b") + + assert api._py_type_name(Literal[here.NONE]) != api._py_type_name( + Literal[there.NONE] + ) + + +def test_an_intenum_socket_and_a_literal_of_its_values_agree_on_true(): + """``True`` is not ``1`` on either path.""" + from node_graph.utils.struct_utils import canonical_socket_value + + info = api._leaf_from_type(Priority).meta.extras["structured_type"] + assert canonical_socket_value(1, structured_type=info) is Priority.LOW + for rejected in (True, 1.0): + with pytest.raises(ValueError, match="Input should be 1 or 2"): + canonical_socket_value(rejected, structured_type=info) + with pytest.raises(ValueError, match="Input should be 1 or 2"): + canonical_socket_value(rejected, allowed=[1, 2]) + + +def test_literal_socket_takes_a_member_of_its_subset_by_value(): + @task.graph() + def only_none(spin: Literal[Spin.NONE]) -> str: + return consume_narrow(spin=Narrow.NONE).result + + assert only_none.build(spin="none").inputs.spin.value == Spin.NONE.value + assert only_none.build(spin=Narrow.NONE).inputs.spin.value == Spin.NONE.value + with pytest.raises(ValueError, match="Input should be 'none'"): + only_none.build(spin=Spin.COLLINEAR) + + +# --- links ----------------------------------------------------------------- + + +@task() +def emit_subset() -> Literal[Spin.NONE, Spin.COLLINEAR]: + return Spin.NONE + + +@task() +def emit_spin() -> Spin: + return Spin.NONE + + +@task() +def emit_ints() -> Literal[1, 2]: + return 1 + + +@task() +def emit_one_string() -> Literal["none"]: + return "none" + + +@task() +def take_spin(x: Spin) -> int: + return 1 + + +@task() +def take_subset(x: Literal[Spin.NONE, Spin.COLLINEAR]) -> int: + return 1 + + +@task() +def take_strings(x: Literal["none", "collinear"]) -> int: + return 1 + + +@task() +def take_str(x: str) -> int: + return 1 + + +def link(producer, consumer): + ng = Graph(name="link") + src = ng.add_task(producer, "src") + dst = ng.add_task(consumer, "dst") + ng.add_link(src.outputs.result, dst.inputs.x) + + +def test_a_subset_flows_into_the_whole_enum(): + link(emit_subset, take_spin) + + +def test_the_whole_enum_does_not_flow_into_a_subset(): + with pytest.raises(TypeError, match="Socket value range mismatch"): + link(emit_spin, take_subset) + + +def test_a_subset_flows_into_another_enum_carrying_the_same_values(): + """Membership is by value, so the declaring class need not match.""" + + @task() + def take_narrow_enum(x: Narrow) -> int: + return 1 + + link(emit_subset, take_narrow_enum) + + +def test_literals_of_different_base_types_do_not_mix(): + with pytest.raises(TypeError, match="Socket value range mismatch"): + link(emit_ints, take_strings) + + +def test_a_string_subset_flows_into_a_wider_string_literal(): + link(emit_one_string, take_strings) + + +def test_a_string_literal_widens_into_its_base_type(): + link(emit_one_string, take_str) + + +def test_a_literal_of_enum_members_does_not_widen_into_the_value_type(): + """An enum output carries the member, which a ``str`` socket cannot take.""" + with pytest.raises(TypeError, match="Socket type mismatch"): + link(emit_subset, take_str) + + +def test_a_typed_source_of_the_same_base_does_not_flow_into_a_literal(): + """``str`` is not restricted to the two strings the target admits.""" + with pytest.raises(TypeError, match="Socket value range mismatch"): + link(take_str, take_strings) + + +def test_an_untyped_source_flows_into_a_literal_and_is_read_at_run(): + """An untyped hop is how a parent passes a value on, so the link stands.""" + + @task() + def emit_untyped(): + return "banana" + + link(emit_untyped, take_strings) + + +def test_a_value_two_untyped_hops_out_is_rejected_when_the_hop_expands(): + """The restricted socket appears only as the inner graph expands, at run.""" + + @task.graph() + def one_hop(spin) -> str: + return narrow_graph(spin=spin) + + @task.graph() + def two_hops(spin) -> str: + return one_hop(spin=spin) + + with pytest.raises(ValueError, match="Input should be 'none' or 'collinear'"): + LocalEngine().run(two_hops.build(spin=Spin.SPIN_ORBIT)) + + assert LocalEngine().run(two_hops.build(spin=Narrow.NONE))["result"] == "none" + + +def test_the_run_time_message_names_the_socket_that_refused_the_value(): + """An untyped output is read at run, where the name is the socket's own.""" + + @task() + def emit_foreign(): + return Spin.SPIN_ORBIT + + @task.graph() + def from_task() -> str: + return narrow_graph(spin=emit_foreign().result) + + with pytest.raises(ValueError, match="socket 'spin'"): + LocalEngine().run(from_task.build()) + + +def test_a_wrapped_value_read_back_from_storage_still_names_its_member(): + """Storage hands back a node that equals the value but is not a ``str``.""" + from node_graph.utils.struct_utils import canonical_socket_value + + class StoredString: + def __init__(self, value): + self.value = value + + def __eq__(self, other): + return self.value == other + + info = api._leaf_from_type(Narrow).meta.extras["structured_type"] + + assert canonical_socket_value(StoredString("none"), structured_type=info) is ( + Narrow.NONE + )