Skip to content
Draft
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
141 changes: 141 additions & 0 deletions src/node_graph/link.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
46 changes: 46 additions & 0 deletions src/node_graph/socket.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand All @@ -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)
Expand Down
Loading
Loading