Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
69af6f818d | ||
|
|
8e91b27453 | ||
|
|
11a84be06d | ||
|
|
6b66f6c5a4 | ||
|
|
9843d63569 | ||
|
|
d7e2683e9d |
@@ -28,15 +28,24 @@ from .protocol import (
|
||||
log = logging.getLogger("comfy_remote_nodes.client")
|
||||
|
||||
|
||||
# Capability tokens advertised on every outbound request. Currently
|
||||
# covers the input/output encodings the registered providers (LTXV,
|
||||
# Gemini) need, plus the ``schema:v3`` marker so descriptors come back
|
||||
# in the V3 ``get_v1_info()`` shape.
|
||||
# Capability tokens advertised on every outbound request. Covers every
|
||||
# input/output encoding ``serialization.py`` knows how to round-trip,
|
||||
# plus the ``schema:v3`` marker so descriptors come back in the V3
|
||||
# ``get_v1_info()`` shape. Audio rides as ``mp3_base64`` outbound (the
|
||||
# encoder always re-encodes a ComfyUI AUDIO dict via
|
||||
# ``audio_input_to_mp3``) and decodes any of mp3_base64/mp3_inline/
|
||||
# wav_base64/wav_inline inbound — Stability returns wav_base64 by
|
||||
# default.
|
||||
CLIENT_CAPABILITIES = [
|
||||
Capability.SCHEMA_V3,
|
||||
Capability.VIDEO_MP4_INLINE,
|
||||
Capability.IMAGE_PNG_BASE64,
|
||||
Capability.MASK_PNG_BASE64,
|
||||
Capability.AUDIO_MP3_BASE64,
|
||||
Capability.AUDIO_WAV_BASE64,
|
||||
Capability.AUDIO_MP3_INLINE,
|
||||
Capability.AUDIO_WAV_INLINE,
|
||||
Capability.SVG_XML_BASE64,
|
||||
]
|
||||
|
||||
# Reported in the X-RNP-Client-Version header.
|
||||
|
||||
+17
-2
@@ -21,7 +21,18 @@ DEFAULT_POLL_INTERVAL_S = 1.5
|
||||
# Values mirror comfy_api_nodes.util.client.poll_op_raw so RNP-routed
|
||||
# nodes behave identically to a direct upstream call when the descriptor
|
||||
# leaves a field unset.
|
||||
DEFAULT_MAX_POLL_ATTEMPTS = 160
|
||||
#
|
||||
# Lifetime budget design: the wire format declares ``soft_timeout_s``
|
||||
# (server-side stuck threshold) and ``hard_timeout_s`` (server-side
|
||||
# hard kill) as durations. The client's poll-loop cap is a derived
|
||||
# value — ``max_poll_attempts = ceil((hard_timeout_s + DEFAULT_POLL_GRACE_S)
|
||||
# / poll_interval_s)`` — not a separate wire field, so changing the
|
||||
# poll cadence can never silently shrink or stretch the contract.
|
||||
# Industry precedent: OAuth2 device flow §3.5, Kubernetes
|
||||
# ``--timeout``, gRPC deadlines all express bounds as durations.
|
||||
DEFAULT_SOFT_TIMEOUT_S = 240.0
|
||||
DEFAULT_HARD_TIMEOUT_S = 300.0
|
||||
DEFAULT_POLL_GRACE_S = 30.0
|
||||
DEFAULT_TIMEOUT_PER_POLL_S = 120.0
|
||||
DEFAULT_MAX_RETRIES_PER_POLL = 10
|
||||
DEFAULT_RETRY_DELAY_PER_POLL_S = 1.0
|
||||
@@ -95,12 +106,16 @@ class Capability:
|
||||
IMAGE_RAW_CHW_UINT8 = "image:raw_b64_chw_uint8"
|
||||
VIDEO_MP4_INLINE = "video:mp4_inline"
|
||||
AUDIO_MP3_INLINE = "audio:mp3_inline"
|
||||
AUDIO_MP3_BASE64 = "audio:mp3_base64"
|
||||
AUDIO_WAV_BASE64 = "audio:wav_base64"
|
||||
AUDIO_WAV_INLINE = "audio:wav_inline"
|
||||
MASK_PNG_BASE64 = "mask:png_base64"
|
||||
SVG_XML_BASE64 = "svg:svg_xml_base64"
|
||||
SCHEMA_V3 = "schema:v3"
|
||||
ASYNC_EXECUTE = "execute:async"
|
||||
|
||||
|
||||
HEAVY_TYPES = frozenset({"image", "video", "audio", "mask"})
|
||||
HEAVY_TYPES = frozenset({"image", "video", "audio", "mask", "svg"})
|
||||
|
||||
|
||||
def is_envelope(value: Any) -> bool:
|
||||
|
||||
+232
-22
@@ -3,9 +3,11 @@
|
||||
The descriptor wire format is V3 ``Schema.get_v1_info()`` verbatim, so
|
||||
per-input dicts are handed straight to V3's IO classes without
|
||||
reinterpretation. ``_parse_input_spec`` knows STRING / INT / FLOAT /
|
||||
BOOLEAN / COMBO / IMAGE / VIDEO / AUDIO / MASK; an input outside that
|
||||
set causes the whole node to be skipped with a warning rather than
|
||||
half-built.
|
||||
BOOLEAN / COMBO / IMAGE / VIDEO / AUDIO / MASK natively; any other
|
||||
``io_type`` string is treated as an opaque pass-through socket (see
|
||||
``IO.Custom``) so partner-defined types like ``RECRAFT_STYLE_V3`` /
|
||||
``RECRAFT_COLOR`` round-trip between RNP nodes without the descriptor
|
||||
being rejected.
|
||||
|
||||
Hidden inputs (``auth_token_comfy_org`` / ``api_key_comfy_org`` /
|
||||
``unique_id`` / ``prompt`` / ``extra_pnginfo`` / ``dynprompt``) are
|
||||
@@ -22,6 +24,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
import math
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
@@ -36,11 +39,13 @@ from . import client as rnp_client
|
||||
from . import serialization
|
||||
from .protocol import (
|
||||
DEFAULT_CANCEL_TIMEOUT_S,
|
||||
DEFAULT_MAX_POLL_ATTEMPTS,
|
||||
DEFAULT_HARD_TIMEOUT_S,
|
||||
DEFAULT_MAX_RETRIES_PER_POLL,
|
||||
DEFAULT_POLL_GRACE_S,
|
||||
DEFAULT_POLL_INTERVAL_S,
|
||||
DEFAULT_RETRY_BACKOFF_PER_POLL,
|
||||
DEFAULT_RETRY_DELAY_PER_POLL_S,
|
||||
DEFAULT_SOFT_TIMEOUT_S,
|
||||
DEFAULT_TIMEOUT_PER_POLL_S,
|
||||
ErrorCode,
|
||||
ExecutionMode,
|
||||
@@ -146,19 +151,32 @@ def _extract_execution_policy(execution: dict[str, Any]) -> dict[str, Any]:
|
||||
node with an empty ``execution`` block behaves identically to a
|
||||
direct upstream call. The schema is the wire contract from
|
||||
``comfy_rnp_protocol.constants``: ``poll_interval_s``,
|
||||
``max_poll_attempts``, ``timeout_per_poll_s``, ``cancel_timeout_s``,
|
||||
``max_retries_per_poll``, ``retry_delay_per_poll_s``,
|
||||
``retry_backoff_per_poll``, ``retry: {max, delay_s, backoff,
|
||||
retry_on}``, ``idempotency``.
|
||||
``soft_timeout_s``, ``hard_timeout_s``, ``timeout_per_poll_s``,
|
||||
``cancel_timeout_s``, ``max_retries_per_poll``,
|
||||
``retry_delay_per_poll_s``, ``retry_backoff_per_poll``,
|
||||
``retry: {max, delay_s, backoff, retry_on}``, ``idempotency``.
|
||||
|
||||
The client's poll-loop cap (``max_poll_attempts``) is **derived**
|
||||
from ``hard_timeout_s + DEFAULT_POLL_GRACE_S`` and
|
||||
``poll_interval_s`` — there is no wire field for it. Coupling a
|
||||
contract bound to a cadence value silently breaks every time you
|
||||
tune the cadence; expressing the bound as a duration keeps the
|
||||
contract invariant under cadence changes (see protocol.py
|
||||
"Lifetime budget design" comment).
|
||||
"""
|
||||
retry_block = execution.get("retry") if isinstance(execution.get("retry"), dict) else {}
|
||||
poll_interval_s = _coerce_pos(
|
||||
execution.get("poll_interval_s"), DEFAULT_POLL_INTERVAL_S,
|
||||
)
|
||||
hard_timeout_s = _coerce_pos(
|
||||
execution.get("hard_timeout_s"),
|
||||
_coerce_pos(execution.get("soft_timeout_s"), DEFAULT_HARD_TIMEOUT_S),
|
||||
)
|
||||
poll_budget_s = hard_timeout_s + DEFAULT_POLL_GRACE_S
|
||||
max_poll_attempts = max(1, math.ceil(poll_budget_s / poll_interval_s))
|
||||
return {
|
||||
"poll_interval_s": _coerce_pos(
|
||||
execution.get("poll_interval_s"), DEFAULT_POLL_INTERVAL_S,
|
||||
),
|
||||
"max_poll_attempts": _coerce_pos_int(
|
||||
execution.get("max_poll_attempts"), DEFAULT_MAX_POLL_ATTEMPTS,
|
||||
),
|
||||
"poll_interval_s": poll_interval_s,
|
||||
"max_poll_attempts": max_poll_attempts,
|
||||
"timeout_per_poll_s": _coerce_pos(
|
||||
execution.get("timeout_per_poll_s"), DEFAULT_TIMEOUT_PER_POLL_S,
|
||||
),
|
||||
@@ -306,6 +324,16 @@ def build_node_class(
|
||||
url_fetch_policy = (
|
||||
remote.get("url_fetch") if isinstance(remote.get("url_fetch"), dict) else {}
|
||||
)
|
||||
# ``descriptor.remote.hidden_forwarded`` opt-in for non-auth hidden
|
||||
# markers (``prompt`` / ``extra_pnginfo`` / ``dynprompt``). Auth
|
||||
# tokens and ``unique_id`` always forward when declared. Anything
|
||||
# else is declared on the schema (so the executor populates
|
||||
# ``cls.hidden``) but dropped before building the request body so
|
||||
# large EXTRA_PNGINFO blobs don't ship by default.
|
||||
hidden_forwarded_raw = remote.get("hidden_forwarded") or []
|
||||
hidden_forwarded: list[str] = [
|
||||
h for h in hidden_forwarded_raw if isinstance(h, str)
|
||||
] if isinstance(hidden_forwarded_raw, (list, tuple)) else []
|
||||
|
||||
display_name = descriptor.get("display_name") or node_id
|
||||
category = descriptor.get("category") or "remote"
|
||||
@@ -364,6 +392,7 @@ def build_node_class(
|
||||
timeout_s=timeout_s,
|
||||
outputs_meta=schema_outputs,
|
||||
hidden_names=hidden_names,
|
||||
hidden_forwarded=hidden_forwarded,
|
||||
inputs=kwargs,
|
||||
max_inline_bytes=max_inline_bytes,
|
||||
estimated_duration_s=estimated_duration_s,
|
||||
@@ -439,13 +468,17 @@ def _parse_input_spec(name: str, spec: list[Any], optional: bool) -> Any | None:
|
||||
if isinstance(io_type, list):
|
||||
# Legacy V1-style combo: io_type *is* the options list.
|
||||
return IO.Combo.Input(name, options=list(io_type),
|
||||
default=options.get("default"), **common)
|
||||
default=options.get("default"),
|
||||
control_after_generate=options.get("control_after_generate"),
|
||||
**common)
|
||||
if io_type == "COMBO":
|
||||
opts = options.get("options")
|
||||
if not isinstance(opts, (list, tuple)):
|
||||
return None
|
||||
return IO.Combo.Input(name, options=list(opts),
|
||||
default=options.get("default"), **common)
|
||||
default=options.get("default"),
|
||||
control_after_generate=options.get("control_after_generate"),
|
||||
**common)
|
||||
if io_type == "STRING":
|
||||
return IO.String.Input(
|
||||
name,
|
||||
@@ -460,6 +493,7 @@ def _parse_input_spec(name: str, spec: list[Any], optional: bool) -> Any | None:
|
||||
min=options.get("min", 0),
|
||||
max=options.get("max", 2147483647),
|
||||
step=options.get("step", 1),
|
||||
control_after_generate=options.get("control_after_generate"),
|
||||
**common,
|
||||
)
|
||||
if io_type == "FLOAT":
|
||||
@@ -483,6 +517,132 @@ def _parse_input_spec(name: str, spec: list[Any], optional: bool) -> Any | None:
|
||||
return IO.Audio.Input(name, **common)
|
||||
if io_type == "MASK":
|
||||
return IO.Mask.Input(name, **common)
|
||||
if io_type == "SVG":
|
||||
return IO.SVG.Input(name, **common)
|
||||
if io_type == "AUTOGROW":
|
||||
# Wire shape:
|
||||
# ["AUTOGROW", {
|
||||
# "template": ["IMAGE", {...}], # nested input spec
|
||||
# "prefix": "image",
|
||||
# "min": 1,
|
||||
# "max": 5,
|
||||
# }]
|
||||
# The descriptor declares N "<prefix>0".."<prefix>{max-1}"
|
||||
# virtual inputs that the V3 frontend renders as a single
|
||||
# auto-growing block — see ``IO.Autogrow.TemplatePrefix``.
|
||||
# The first ``min`` slots are required, the rest optional;
|
||||
# the executor packs the connected ones into a dict keyed by
|
||||
# virtual id which the server provider iterates.
|
||||
template_spec = options.get("template")
|
||||
if (
|
||||
not isinstance(template_spec, (list, tuple))
|
||||
or len(template_spec) < 1
|
||||
or not isinstance(template_spec[0], str)
|
||||
):
|
||||
log.warning(
|
||||
"AUTOGROW input %s: malformed template spec %r",
|
||||
name, template_spec,
|
||||
)
|
||||
return None
|
||||
# The template is itself an ``[io_type, {...}]`` pair — reuse
|
||||
# the parser recursively so AUTOGROW inherits every IO type the
|
||||
# parser knows about (IMAGE, MASK, STRING, custom IO …).
|
||||
# The framework forbids ``DynamicInput`` as an Autogrow template
|
||||
# (``assert(not isinstance(input, DynamicInput))`` in
|
||||
# ``_AutogrowTemplate.__init__``), so passing AUTOGROW or
|
||||
# DYNAMIC_COMBO here will raise inside ``TemplatePrefix`` —
|
||||
# caught below.
|
||||
template_input = _parse_input_spec(
|
||||
"_autogrow_template", list(template_spec), optional=False,
|
||||
)
|
||||
if template_input is None:
|
||||
log.warning(
|
||||
"AUTOGROW input %s: failed to parse template %r",
|
||||
name, template_spec,
|
||||
)
|
||||
return None
|
||||
prefix = options.get("prefix")
|
||||
if not isinstance(prefix, str) or not prefix:
|
||||
return None
|
||||
min_slots = int(options.get("min", 1))
|
||||
max_slots = int(options.get("max", 1))
|
||||
if max_slots < 1 or min_slots < 0 or min_slots > max_slots:
|
||||
return None
|
||||
try:
|
||||
template = IO.Autogrow.TemplatePrefix(
|
||||
template_input,
|
||||
prefix=prefix,
|
||||
min=min_slots,
|
||||
max=max_slots,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
log.warning(
|
||||
"AUTOGROW input %s: failed to build TemplatePrefix: %s",
|
||||
name, e,
|
||||
)
|
||||
return None
|
||||
return IO.Autogrow.Input(
|
||||
name,
|
||||
template=template,
|
||||
display_name=options.get("display_name"),
|
||||
optional=optional,
|
||||
tooltip=options.get("tooltip"),
|
||||
)
|
||||
if io_type == "DYNAMIC_COMBO":
|
||||
# Wire shape:
|
||||
# ["DYNAMIC_COMBO", {
|
||||
# "options": [
|
||||
# {"key": "recraftv4", "inputs": [
|
||||
# ["size", "COMBO", {"options": [...]}],
|
||||
# ...
|
||||
# ]},
|
||||
# {"key": "recraftv4_pro", "inputs": [...]},
|
||||
# ],
|
||||
# "default": "recraftv4",
|
||||
# "tooltip": "...",
|
||||
# }]
|
||||
# The selected key drives which sub-widget block the frontend
|
||||
# renders; the executor packs ``{"model": <key>, ...sub fields}``
|
||||
# into a dict the server consumes via
|
||||
# ``inputs["model"]["model"]`` / ``inputs["model"]["size"]``
|
||||
# idiom — same shape ``IO.DynamicCombo`` produces locally.
|
||||
opts = options.get("options")
|
||||
if not isinstance(opts, list) or not opts:
|
||||
return None
|
||||
dc_options: list[Any] = []
|
||||
for entry in opts:
|
||||
if not isinstance(entry, dict):
|
||||
return None
|
||||
key = entry.get("key")
|
||||
sub_inputs_spec = entry.get("inputs")
|
||||
if not isinstance(key, str) or not isinstance(sub_inputs_spec, list):
|
||||
return None
|
||||
sub_inputs: list[Any] = []
|
||||
for sub in sub_inputs_spec:
|
||||
if not (isinstance(sub, (list, tuple)) and len(sub) >= 2 and isinstance(sub[0], str)):
|
||||
return None
|
||||
sub_name = sub[0]
|
||||
sub_spec = list(sub[1:])
|
||||
# ``[name, io_type, options]`` flattens to the parser's
|
||||
# ``[io_type, options]`` shape with sub_name carried separately.
|
||||
sub_input = _parse_input_spec(sub_name, sub_spec, optional=False)
|
||||
if sub_input is None:
|
||||
return None
|
||||
sub_inputs.append(sub_input)
|
||||
dc_options.append(IO.DynamicCombo.Option(key, sub_inputs))
|
||||
return IO.DynamicCombo.Input(
|
||||
name,
|
||||
options=dc_options,
|
||||
display_name=options.get("display_name"),
|
||||
optional=optional,
|
||||
tooltip=options.get("tooltip"),
|
||||
)
|
||||
# Any other string io_type — partner-defined custom IO (RecraftStyle,
|
||||
# RecraftColor, Tripo3DModel, …). Built as an opaque pass-through
|
||||
# socket so the descriptor is accepted; the value flows between RNP
|
||||
# nodes verbatim and the proxy_node never inspects or serialises it.
|
||||
if isinstance(io_type, str) and io_type:
|
||||
return IO.Custom(io_type).Input(name, **common)
|
||||
return None
|
||||
|
||||
|
||||
@@ -567,10 +727,18 @@ def _parse_outputs(descriptor: dict[str, Any]) -> list[Any] | None:
|
||||
tip = output_tooltips[i] if i < len(output_tooltips) else None
|
||||
cls = _OUTPUT_CLASSES.get(io_type)
|
||||
if cls is None:
|
||||
log.warning(
|
||||
"Skipping descriptor: unsupported output type %r", io_type,
|
||||
)
|
||||
return None
|
||||
# Partner-defined custom IO output (RecraftStyle, Tripo3DModel,
|
||||
# Hunyuan3DMesh, …). Built as an opaque pass-through socket so
|
||||
# the value flows downstream verbatim — no envelope decoding.
|
||||
if not isinstance(io_type, str) or not io_type:
|
||||
log.warning(
|
||||
"Skipping descriptor: unsupported output type %r", io_type,
|
||||
)
|
||||
return None
|
||||
out.append(IO.Custom(io_type).Output(
|
||||
display_name=name, is_output_list=is_list, tooltip=tip,
|
||||
))
|
||||
continue
|
||||
out.append(cls(display_name=name, is_output_list=is_list, tooltip=tip))
|
||||
return out
|
||||
|
||||
@@ -580,6 +748,7 @@ _OUTPUT_CLASSES = {
|
||||
"IMAGE": IO.Image.Output,
|
||||
"AUDIO": IO.Audio.Output,
|
||||
"MASK": IO.Mask.Output,
|
||||
"SVG": IO.SVG.Output,
|
||||
"STRING": IO.String.Output,
|
||||
"INT": IO.Int.Output,
|
||||
"FLOAT": IO.Float.Output,
|
||||
@@ -825,6 +994,7 @@ async def _execute_remote(
|
||||
timeout_s: float,
|
||||
outputs_meta: list[Any],
|
||||
hidden_names: list[str],
|
||||
hidden_forwarded: list[str],
|
||||
inputs: dict[str, Any],
|
||||
max_inline_bytes: int | None = None,
|
||||
estimated_duration_s: int | float | None = None,
|
||||
@@ -837,7 +1007,15 @@ async def _execute_remote(
|
||||
# Strip hidden inputs from the user-supplied kwargs. Auth credentials
|
||||
# (``auth_token_comfy_org`` / ``api_key_comfy_org``) travel as standard
|
||||
# HTTP headers so the server can forward them to api.comfy.org without
|
||||
# parsing the body; everything else rides in the request ``context``.
|
||||
# parsing the body; ``unique_id`` always rides in ``context`` for
|
||||
# correlation. Other markers (``prompt`` / ``extra_pnginfo`` /
|
||||
# ``dynprompt``) only forward when the descriptor explicitly opted
|
||||
# in via ``remote.hidden_forwarded`` — otherwise they're declared on
|
||||
# the schema (so the executor populates ``cls.hidden``) but stripped
|
||||
# before we build the request body, so EXTRA_PNGINFO blobs don't
|
||||
# leak by default.
|
||||
_ALWAYS_FORWARD = {"auth_token_comfy_org", "api_key_comfy_org", "unique_id"}
|
||||
forward_set = _ALWAYS_FORWARD | set(hidden_forwarded)
|
||||
context: dict[str, Any] = {}
|
||||
auth_headers: dict[str, str] = {}
|
||||
hidden_holder = getattr(cls, "hidden", None)
|
||||
@@ -858,7 +1036,7 @@ async def _execute_remote(
|
||||
auth_headers["Authorization"] = f"Bearer {value}"
|
||||
elif hname == "api_key_comfy_org":
|
||||
auth_headers["X-API-KEY"] = str(value)
|
||||
else:
|
||||
elif hname in forward_set:
|
||||
context[hname] = value
|
||||
|
||||
# Encode heavy-typed inputs (IMAGE/MASK/AUDIO) as RNP value
|
||||
@@ -1158,9 +1336,41 @@ async def _encode_inputs(
|
||||
return out
|
||||
|
||||
|
||||
def _dict_contains_tensor(value: Any) -> bool:
|
||||
"""True iff a dict (or nested dict) holds at least one torch tensor.
|
||||
|
||||
Used to gate the recursive encode branch so we don't pointlessly
|
||||
rebuild scalar custom-IO chain dicts (RECRAFT_V3_STYLE etc.) on
|
||||
every encode pass, and so already-encoded envelope dicts pass
|
||||
straight through. Nested-dict support keeps DynamicCombo→Autogrow
|
||||
composition (and any future layered dynamic IO) round-trippable.
|
||||
"""
|
||||
if isinstance(value, dict):
|
||||
return any(_dict_contains_tensor(v) for v in value.values())
|
||||
return serialization._is_torch_tensor(value)
|
||||
|
||||
|
||||
def _encode_one(name: str, value: Any) -> Any:
|
||||
if serialization.is_audio_input(value):
|
||||
return serialization.encode_audio_input(value)
|
||||
if (
|
||||
isinstance(value, dict)
|
||||
and not serialization.is_envelope(value)
|
||||
and not serialization.is_audio_input(value)
|
||||
and _dict_contains_tensor(value)
|
||||
):
|
||||
# IO.Autogrow inputs arrive grouped as a dict (e.g. ``{"image0":
|
||||
# tensor, "image1": tensor}``) keyed by the per-slot virtual id.
|
||||
# Nested dynamic types (DynamicCombo whose option inputs include
|
||||
# an Autogrow, etc.) can produce arbitrarily nested dict-of-
|
||||
# tensor shapes — recurse so every heavy-typed leaf becomes an
|
||||
# envelope. Skip ordinary scalar dicts (custom-IO chain values
|
||||
# like RECRAFT_V3_STYLE) and already-encoded envelopes; both
|
||||
# pass through unchanged via the ``return value`` tail.
|
||||
encoded_dict: dict[str, Any] = {}
|
||||
for sub_name, sub_value in value.items():
|
||||
encoded_dict[sub_name] = _encode_one(f"{name}.{sub_name}", sub_value)
|
||||
return encoded_dict
|
||||
if not serialization._is_torch_tensor(value):
|
||||
return value
|
||||
rank = value.dim()
|
||||
|
||||
@@ -252,6 +252,69 @@ def decode_video_envelope(envelope: dict[str, Any]) -> Any:
|
||||
return InputImpl.VideoFromFile(BytesIO(decode_envelope_data(envelope)))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SVG decode (encode lands when a remote node accepts SVG inputs; today
|
||||
# Recraft's V3/V4 vector endpoints only emit SVG, so the client is
|
||||
# decode-only).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def decode_svg_envelope(envelope: dict[str, Any]) -> Any:
|
||||
"""Decode an SVG envelope into a ``comfy_extras.nodes_images.SVG`` instance.
|
||||
|
||||
The wire payload is one or more concatenated ``<svg>…</svg>`` XML
|
||||
documents; we split on the closing tag so a multi-frame Recraft
|
||||
response (n>1) hydrates a single ``SVG`` carrying a list of
|
||||
BytesIO frames — the same shape the local node produces via
|
||||
``SVG(svg_data)``.
|
||||
"""
|
||||
encoding = envelope.get("encoding")
|
||||
if encoding not in ("svg_xml_base64", "svg_xml_inline"):
|
||||
raise RnpProtocolError(
|
||||
f"Unsupported svg encoding: {encoding!r}",
|
||||
code=ErrorCode.INTERNAL,
|
||||
)
|
||||
from comfy_extras.nodes_images import SVG
|
||||
|
||||
raw = decode_envelope_data(envelope)
|
||||
text = raw.decode("utf-8", "replace")
|
||||
# The server joins per-frame SVGs with ``\n`` (see
|
||||
# ``_emit_svg_from_urls``); split on the closing tag followed by a
|
||||
# newline so an SVG document that legitimately contains ``</svg>``
|
||||
# in CDATA / a foreignObject / a comment isn't shredded mid-frame.
|
||||
# Fall back to splitting on the bare closing tag for single-frame
|
||||
# payloads or older server builds that didn't add the newline.
|
||||
delimiter = "</svg>\n"
|
||||
if delimiter not in text:
|
||||
delimiter = "</svg>"
|
||||
parts: list[BytesIO] = []
|
||||
cursor = 0
|
||||
closing = "</svg>"
|
||||
while True:
|
||||
idx = text.find(delimiter, cursor)
|
||||
if idx < 0:
|
||||
tail = text[cursor:].strip()
|
||||
if tail:
|
||||
parts.append(BytesIO(tail.encode("utf-8")))
|
||||
break
|
||||
end = idx + len(closing) # keep the closing tag on this frame
|
||||
chunk = text[cursor:end].strip()
|
||||
if chunk:
|
||||
parts.append(BytesIO(chunk.encode("utf-8")))
|
||||
cursor = idx + len(delimiter)
|
||||
declared = envelope.get("count")
|
||||
if isinstance(declared, int) and declared > 0 and declared != len(parts):
|
||||
log.warning(
|
||||
"SVG envelope count=%d but parsed %d documents — server payload may be malformed",
|
||||
declared, len(parts),
|
||||
)
|
||||
if not parts:
|
||||
raise RnpProtocolError(
|
||||
"SVG envelope payload contained no <svg> documents",
|
||||
code=ErrorCode.INTERNAL,
|
||||
)
|
||||
return SVG(parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dispatch by envelope type
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -261,6 +324,7 @@ _DECODERS = {
|
||||
"mask": decode_mask_envelope,
|
||||
"audio": decode_audio_envelope,
|
||||
"video": decode_video_envelope,
|
||||
"svg": decode_svg_envelope,
|
||||
}
|
||||
|
||||
|
||||
@@ -310,6 +374,7 @@ __all__ = [
|
||||
"encode_audio_input",
|
||||
"decode_audio_envelope",
|
||||
"decode_video_envelope",
|
||||
"decode_svg_envelope",
|
||||
"decode_envelope",
|
||||
"is_envelope",
|
||||
"is_image_tensor",
|
||||
|
||||
Reference in New Issue
Block a user