Merge branch 'dev' into 'main'

Release 1.1.26

See merge request vrch/comfyui/comfyui-web-viewer!61
This commit is contained in:
Tianzi Hou
2026-07-31 20:18:11 +01:00
11 changed files with 1359 additions and 24 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[bumpversion]
current_version = 1.1.25
current_version = 1.1.26
commit = True
tag = True
parse = (?P<major>\d+)\.(?P<minor>\d+)\.(?P<patch>\d+)
+11 -1
View File
@@ -5,7 +5,17 @@ All notable changes to this project will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [Unreleased]
## [1.1.26 - 2026-07-31]
### Added
- add a backend-only TensorRT Auto Loader node with automatic PyTorch fallback
- add the Default-off Realtime Safe-set for prompt-scoped JPEG reuse and stable external-server/cache boundaries
- publish the `node-safe-set-v1` capability for the cross-repository `vrch-realtime-v1` contract
### Changed
- fail closed to Default when the matching Docker GC capability is absent or incompatible
## [1.1.25 - 2026-07-23]
+4
View File
@@ -170,6 +170,10 @@ To use the `AUDIO Music to Emotion Detector @ vrch.ai` node, you'll need to inst
- Documentation: [Usage of Logic nodes](./docs/logic_nodes.md)
- Example workflows: n/a
### `Model Nodes`
- Documentation: [Usage of Model nodes](./docs/model_nodes.md)
### `Text Nodes`
- Documentation: [Usage of Text nodes](./docs/text_nodes.md)
+10 -1
View File
@@ -4,6 +4,7 @@ from .nodes.audio_nodes import *
from .nodes.text_nodes import *
from .nodes.key_control_nodes import *
from .nodes.osc_control_nodes import *
from .nodes import websocket_nodes as _vrch_websocket_nodes
from .nodes.websocket_nodes import *
from .nodes.midi_control_nodes import *
from .nodes.gamepad_nodes import *
@@ -11,8 +12,9 @@ from .nodes.logic_nodes import *
from .nodes.midi_nodes import *
from .nodes.audio_music2emo_node import *
from .nodes.workflow_export_nodes import *
from .nodes.model_nodes import *
__version__ = "1.1.25"
__version__ = "1.1.26"
NODE_CLASS_MAPPINGS = {
"VrchAnyOSCControlNode": VrchAnyOSCControlNode,
@@ -67,6 +69,7 @@ NODE_CLASS_MAPPINGS = {
"VrchMidiDeviceLoaderNode": VrchMidiDeviceLoaderNode,
"VrchMidiWebSocketChannelLoaderNode": VrchMidiWebSocketChannelLoaderNode,
"VrchModelWebViewerNode": VrchModelWebViewerNode,
"VrchTensorRTAutoLoaderNode": VrchTensorRTAutoLoaderNode,
"VrchOSCControlSettingsNode": VrchOSCControlSettingsNode,
"VrchQRCodeNode": VrchQRCodeNode,
"VrchSwitchOSCControlNode": VrchSwitchOSCControlNode,
@@ -140,6 +143,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"VrchMidiDeviceLoaderNode": "MIDI Device Loader @ vrch.ai",
"VrchMidiWebSocketChannelLoaderNode": "MIDI WebSocket Channel Loader @ vrch.ai",
"VrchModelWebViewerNode": "3D MODEL Web Viewer @ vrch.ai",
"VrchTensorRTAutoLoaderNode": "TensorRT Auto Loader @ vrch.ai",
"VrchOSCControlSettingsNode": "OSC Control Settings @ vrch.ai",
"VrchQRCodeNode": "QR Code Generator @ vrch.ai",
"VrchSwitchOSCControlNode": "SWITCH OSC Control @ vrch.ai",
@@ -160,6 +164,11 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"VrchXboxControllerNode": "Xbox Controller Mapper @ vrch.ai",
}
# Publish the Safe-set capability only after every owner module and mapping
# above loaded successfully. The Docker-side GC policy checks this marker
# before its prompt worker can enter Realtime.
_vrch_websocket_nodes._initialize_realtime_contract()
# WEB_DIRECTORY is the comfyui nodes directory that ComfyUI will link and auto-load.
WEB_DIRECTORY = "./web/comfyui"
+29
View File
@@ -0,0 +1,29 @@
# Model Nodes
## TensorRT Auto Loader @ vrch.ai
Loads a selected local TensorRT Engine when available while keeping the original ComfyUI `MODEL` as a fallback.
### Inputs
- **`model`** (`MODEL`): Original model used by `pytorch` mode and by `auto` fallback.
- **`load_mode`**:
- `auto`: Load the selected TensorRT Engine; use the original model if loading is unavailable or fails.
- `tensorrt`: Require the selected TensorRT Engine and stop the prompt if it cannot be loaded.
- `pytorch`: Bypass TensorRT and use the original model.
- **`engine_name`**: TensorRT Engine from the host's ComfyUI `output/tensorrt` directory or another registered TensorRT model path.
- **`debug`**: Print concise loader, cache, and fallback diagnostics to the ComfyUI server console.
### Outputs
- **`model`** (`MODEL`): TensorRT model or the original model selected by `load_mode`.
- **`backend`** (`STRING`): Actual backend, `tensorrt` or `pytorch`.
- **`status`** (`STRING`): Loading result or fallback reason.
### Behavior
The node is independent of Live Console Maintenance and does not require custom frontend JavaScript. Engine choices use ComfyUI's native dropdown. Refresh the ComfyUI frontend after adding a new Engine so the dropdown is rebuilt.
An Engine name saved by another host or removed later is accepted by the workflow. In `auto` mode it safely falls back to the input model; in `tensorrt` mode it reports an error. TensorRT inference errors that occur after the model has loaded are not retried with PyTorch.
The current implementation infers the TensorRT model family from the input `MODEL`. The installed `TensorRTLoader` remains responsible for Engine deserialization and compatibility.
+289
View File
@@ -0,0 +1,289 @@
"""Model loading and fallback nodes for ComfyUI workflows."""
from pathlib import Path
import folder_paths
CATEGORY = "vrch.ai/model"
NO_ENGINE_OPTION = "No TensorRT Engine Found"
LOAD_MODES = ["auto", "tensorrt", "pytorch"]
_TENSORRT_MODEL_TYPES = {
"SDXL": "sdxl_base",
"SDXLRefiner": "sdxl_refiner",
"SD15": "sd1.x",
"SD20": "sd2.x-768v",
"SVD_img2vid": "svd",
"SD3": "sd3",
"AuraFlow": "auraflow",
"Flux": "flux_dev",
"FluxSchnell": "flux_schnell",
}
def _register_output_engine_root():
get_output_directory = getattr(folder_paths, "get_output_directory", None)
registry = getattr(folder_paths, "folder_names_and_paths", None)
if not callable(get_output_directory) or not isinstance(registry, dict):
return
output_root = str((Path(get_output_directory()) / "tensorrt").resolve())
if "tensorrt" not in registry:
registry["tensorrt"] = ([output_root], {".engine"})
return
roots, extensions = registry["tensorrt"]
if output_root not in roots:
roots.insert(0, output_root)
extensions.add(".engine")
def _tensorrt_roots():
_register_output_engine_root()
try:
roots = folder_paths.get_folder_paths("tensorrt")
except (KeyError, ValueError):
roots = []
if not roots:
models_dir = getattr(folder_paths, "models_dir", None)
if models_dir:
roots = [str(Path(models_dir) / "tensorrt")]
return [Path(root).expanduser().resolve() for root in roots]
def _engine_names():
_register_output_engine_root()
names = []
try:
names = folder_paths.get_filename_list("tensorrt")
except (KeyError, ValueError):
for root in _tensorrt_roots():
if root.is_dir():
names.extend(path.name for path in root.glob("*.engine") if path.is_file())
safe_names = {
name
for name in names
if isinstance(name, str)
and Path(name).name == name
and Path(name).suffix.lower() == ".engine"
}
return sorted(safe_names, key=str.casefold)
def _engine_options():
engines = _engine_names()
return engines if engines else [NO_ENGINE_OPTION]
def _resolve_engine_path(engine_name):
_register_output_engine_root()
if (
not isinstance(engine_name, str)
or engine_name == NO_ENGINE_OPTION
or Path(engine_name).name != engine_name
or Path(engine_name).suffix.lower() != ".engine"
):
return None
candidate = None
try:
candidate = folder_paths.get_full_path("tensorrt", engine_name)
except (KeyError, ValueError):
pass
if candidate is None:
for root in _tensorrt_roots():
possible = root / engine_name
if possible.is_file():
candidate = str(possible)
break
if candidate is None:
return None
resolved = Path(candidate).expanduser().resolve()
if not resolved.is_file():
return None
if not any(resolved.is_relative_to(root) for root in _tensorrt_roots()):
return None
return resolved
def _engine_fingerprint(engine_path):
stat = engine_path.stat()
return (
stat.st_dev,
stat.st_ino,
stat.st_size,
stat.st_mtime_ns,
)
def _infer_tensorrt_model_type(model):
base_model = getattr(model, "model", None)
model_config = getattr(base_model, "model_config", None)
for candidate in (model_config, base_model):
if candidate is None:
continue
model_type = _TENSORRT_MODEL_TYPES.get(type(candidate).__name__)
if model_type:
return model_type
return None
def _get_tensorrt_loader_class():
# Resolve the optional node only when this node executes. This keeps the
# vrch.ai node package loadable on hosts without ComfyUI-TensorRT.
import nodes as comfy_nodes
return getattr(comfy_nodes, "NODE_CLASS_MAPPINGS", {}).get("TensorRTLoader")
def _one_line_error(error):
message = " ".join(str(error).split())
return f"{type(error).__name__}: {message}"[:320]
class VrchTensorRTAutoLoaderNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL",),
"load_mode": (LOAD_MODES, {"default": "auto"}),
"engine_name": (_engine_options(),),
"debug": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("MODEL", "STRING", "STRING")
RETURN_NAMES = ("model", "backend", "status")
FUNCTION = "load_model"
CATEGORY = CATEGORY
def __init__(self):
self._cached_key = None
self._cached_model = None
@classmethod
def VALIDATE_INPUTS(cls, engine_name):
# Engine choices are host-local. A workflow saved on another host, or
# before an Engine was removed, must reach load_model() so auto mode can
# fall back instead of failing ComfyUI's pre-execution COMBO check.
return True
@classmethod
def IS_CHANGED(cls, model, load_mode, engine_name, debug=False):
if load_mode == "pytorch":
return "pytorch"
engine_path = _resolve_engine_path(engine_name)
if engine_path is None:
return (load_mode, engine_name, "missing")
try:
fingerprint = _engine_fingerprint(engine_path)
except OSError:
return (load_mode, engine_name, "unreadable")
return (load_mode, engine_name, fingerprint)
def load_model(self, model, load_mode, engine_name, debug=False):
if load_mode == "pytorch":
return self._pytorch_result(
model,
"PyTorch selected",
debug,
)
engine_path = _resolve_engine_path(engine_name)
if engine_path is None:
return self._load_failure(
model,
load_mode,
f"TensorRT Engine is unavailable: {engine_name}",
debug,
)
model_type = _infer_tensorrt_model_type(model)
if model_type is None:
return self._load_failure(
model,
load_mode,
"the input model type is not supported by TensorRTLoader",
debug,
)
loader_class = _get_tensorrt_loader_class()
if loader_class is None:
return self._load_failure(
model,
load_mode,
"TensorRTLoader is not installed or registered",
debug,
)
try:
fingerprint = _engine_fingerprint(engine_path)
except OSError as error:
return self._load_failure(
model,
load_mode,
f"TensorRT Engine cannot be read: {_one_line_error(error)}",
debug,
cause=error,
)
cache_key = (id(model), model_type, engine_name, fingerprint)
if cache_key == self._cached_key and self._cached_model is not None:
self._debug(debug, f"cache hit engine={engine_name} model_type={model_type}")
return (
self._cached_model,
"tensorrt",
f"TensorRT active: {engine_name}",
)
self._debug(debug, f"loading engine={engine_name} model_type={model_type}")
try:
loaded = loader_class().load_unet(engine_name, model_type)
if not isinstance(loaded, tuple) or not loaded or loaded[0] is None:
raise RuntimeError("TensorRTLoader returned no MODEL")
tensorrt_model = loaded[0]
except Exception as error:
self._cached_key = None
self._cached_model = None
return self._load_failure(
model,
load_mode,
f"TensorRT load failed: {_one_line_error(error)}",
debug,
cause=error,
)
self._cached_key = cache_key
self._cached_model = tensorrt_model
self._debug(debug, f"TensorRT active engine={engine_name}")
return (
tensorrt_model,
"tensorrt",
f"TensorRT active: {engine_name}",
)
def _load_failure(self, model, load_mode, reason, debug, cause=None):
self._debug(debug, reason)
if load_mode == "tensorrt":
error = RuntimeError(f"TensorRT Auto Loader: {reason}")
if cause is not None:
raise error from cause
raise error
return self._pytorch_result(model, f"PyTorch fallback: {reason}", debug)
def _pytorch_result(self, model, status, debug):
self._debug(debug, status)
return (model, "pytorch", status)
@staticmethod
def _debug(enabled, message):
if enabled:
print(f"[VrchTensorRTAutoLoaderNode] {message}")
+182
View File
@@ -0,0 +1,182 @@
#!/usr/bin/env python3
"""Tests for VrchTensorRTAutoLoaderNode."""
import importlib.util
import inspect
import sys
import tempfile
import types
import unittest
from pathlib import Path
PROJECT_ROOT = Path(__file__).resolve().parents[2]
FOLDER_PATHS_STUB = types.ModuleType("folder_paths")
FOLDER_PATHS_STUB.models_dir = ""
FOLDER_PATHS_STUB.get_folder_paths = lambda _name: []
FOLDER_PATHS_STUB.get_filename_list = lambda _name: []
FOLDER_PATHS_STUB.get_full_path = lambda _name, _filename: None
sys.modules["folder_paths"] = FOLDER_PATHS_STUB
SPEC = importlib.util.spec_from_file_location(
"vrch_model_nodes_under_test",
PROJECT_ROOT / "nodes" / "model_nodes.py",
)
model_nodes = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(model_nodes)
class SDXL:
pass
class FakeBaseModel:
def __init__(self):
self.model_config = SDXL()
class FakeModelPatcher:
def __init__(self):
self.model = FakeBaseModel()
class FakeFolderPaths:
def __init__(self, root):
self.root = Path(root)
self.models_dir = str(self.root.parent)
self.folder_names_and_paths = {"tensorrt": ([str(self.root)], {".engine"})}
def get_output_directory(self):
return str(self.root.parent)
def get_folder_paths(self, folder_name):
if folder_name != "tensorrt":
raise KeyError(folder_name)
return [str(self.root)]
def get_filename_list(self, folder_name):
return sorted(path.name for path in self.root.glob("*.engine"))
def get_full_path(self, folder_name, filename):
candidate = self.root / filename
return str(candidate) if candidate.is_file() else None
class TestTensorRTAutoLoaderNode(unittest.TestCase):
def setUp(self):
self.temp_dir = tempfile.TemporaryDirectory()
self.addCleanup(self.temp_dir.cleanup)
self.engine_root = Path(self.temp_dir.name) / "tensorrt"
self.engine_root.mkdir()
self.engine_path = self.engine_root / "test.engine"
self.engine_path.write_bytes(b"engine")
self.original_folder_paths = model_nodes.folder_paths
self.original_get_loader = model_nodes._get_tensorrt_loader_class
model_nodes.folder_paths = FakeFolderPaths(self.engine_root)
self.addCleanup(setattr, model_nodes, "folder_paths", self.original_folder_paths)
self.addCleanup(setattr, model_nodes, "_get_tensorrt_loader_class", self.original_get_loader)
self.model = FakeModelPatcher()
self.tensorrt_model = object()
def install_loader(self, error=None):
output_model = self.tensorrt_model
class FakeLoader:
calls = []
def load_unet(self, engine_name, model_type):
self.calls.append((engine_name, model_type))
if error is not None:
raise error
return (output_model,)
model_nodes._get_tensorrt_loader_class = lambda: FakeLoader
return FakeLoader
def test_node_contract(self):
inputs = model_nodes.VrchTensorRTAutoLoaderNode.INPUT_TYPES()["required"]
self.assertEqual(inputs["model"], ("MODEL",))
self.assertEqual(inputs["load_mode"][0], ["auto", "tensorrt", "pytorch"])
self.assertEqual(inputs["engine_name"][0], ["test.engine"])
self.assertEqual(inputs["debug"], ("BOOLEAN", {"default": False}))
self.assertEqual(
model_nodes.VrchTensorRTAutoLoaderNode.RETURN_NAMES,
("model", "backend", "status"),
)
def test_stale_engine_validation_only_accepts_engine_name(self):
signature = inspect.signature(model_nodes.VrchTensorRTAutoLoaderNode.VALIDATE_INPUTS)
self.assertEqual(list(signature.parameters), ["engine_name"])
self.assertTrue(model_nodes.VrchTensorRTAutoLoaderNode.VALIDATE_INPUTS("removed.engine"))
def test_pytorch_mode_bypasses_tensorrt(self):
loader = self.install_loader()
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
self.model, "pytorch", "test.engine", False
)
self.assertIs(result[0], self.model)
self.assertEqual(result[1], "pytorch")
self.assertEqual(loader.calls, [])
def test_auto_mode_loads_tensorrt_and_reuses_cache(self):
loader = self.install_loader()
node = model_nodes.VrchTensorRTAutoLoaderNode()
first = node.load_model(self.model, "auto", "test.engine", False)
second = node.load_model(self.model, "auto", "test.engine", False)
self.assertIs(first[0], self.tensorrt_model)
self.assertEqual(first[1], "tensorrt")
self.assertIs(second[0], self.tensorrt_model)
self.assertEqual(loader.calls, [("test.engine", "sdxl_base")])
def test_auto_mode_falls_back_when_engine_is_missing(self):
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
self.model, "auto", "removed.engine", False
)
self.assertIs(result[0], self.model)
self.assertEqual(result[1], "pytorch")
self.assertIn("fallback", result[2])
def test_tensorrt_mode_fails_when_engine_is_missing(self):
with self.assertRaisesRegex(RuntimeError, "Engine is unavailable"):
model_nodes.VrchTensorRTAutoLoaderNode().load_model(
self.model, "tensorrt", "removed.engine", False
)
def test_auto_mode_falls_back_when_loader_fails(self):
self.install_loader(RuntimeError("incompatible engine"))
result = model_nodes.VrchTensorRTAutoLoaderNode().load_model(
self.model, "auto", "test.engine", False
)
self.assertIs(result[0], self.model)
self.assertEqual(result[1], "pytorch")
self.assertIn("incompatible engine", result[2])
def test_tensorrt_mode_surfaces_loader_failure(self):
self.install_loader(RuntimeError("incompatible engine"))
with self.assertRaisesRegex(RuntimeError, "incompatible engine"):
model_nodes.VrchTensorRTAutoLoaderNode().load_model(
self.model, "tensorrt", "test.engine", False
)
def test_missing_inventory_uses_placeholder(self):
self.engine_path.unlink()
options = model_nodes.VrchTensorRTAutoLoaderNode.INPUT_TYPES()["required"]["engine_name"][0]
self.assertEqual(options, [model_nodes.NO_ENGINE_OPTION])
if __name__ == "__main__":
unittest.main()
+480
View File
@@ -8,6 +8,7 @@ import asyncio
import base64
import io
import json
import os
import socket
import struct
import sys
@@ -15,6 +16,7 @@ import tempfile
import time
import unittest
from pathlib import Path
from unittest import mock
import numpy as np
import torch
@@ -35,6 +37,16 @@ from nodes.utils.websocket_server import ( # noqa: E402
)
def realtime_contract_env():
return {
ws_nodes.REALTIME_MODE_ENV: "realtime",
ws_nodes.REALTIME_REQUESTED_MODE_ENV: "realtime",
ws_nodes.REALTIME_CONTRACT_ENV: ws_nodes.REALTIME_CONTRACT,
ws_nodes.REALTIME_GC_CAPABILITY_ENV: ws_nodes.REALTIME_GC_CAPABILITY,
ws_nodes.REALTIME_NODE_CAPABILITY_ENV: ws_nodes.REALTIME_NODE_CAPABILITY,
}
class TestWebSocketNodesUnit(unittest.TestCase):
def test_01_json_state_merger(self):
merger = ws_nodes.JsonStateMerger(max_keys=2, clear_key="__clear__", debug=False)
@@ -294,6 +306,474 @@ class TestWebSocketNodesUnit(unittest.TestCase):
self.assertEqual(sequence, 2)
self.assertEqual(decoded_payloads, [second])
def test_14_realtime_cache_tokens_preserve_default_behavior(self):
with mock.patch.dict(
os.environ,
{"VRCH_COMFYUI_PERFORMANCE_CACHE": "default"},
):
server_token = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=True,
debug=False,
)
json_token = ws_nodes.VrchJsonWebSocketChannelLoaderNode.IS_CHANGED(
channel="2"
)
self.assertNotEqual(server_token, server_token)
self.assertNotEqual(json_token, json_token)
def test_15_realtime_cache_tokens_are_stable_and_input_sensitive(self):
class FakeClient:
def __init__(self):
self.host = "127.0.0.1"
self.port = 8001
self.path = "/json"
self.channel = 2
self.sequence = 0
def get_received_sequence(self):
return self.sequence
fake_client = FakeClient()
with ws_nodes._websocket_clients_lock:
original_clients = dict(ws_nodes._websocket_clients)
ws_nodes._websocket_clients.clear()
ws_nodes._websocket_clients["test-json"] = fake_client
self.addCleanup(
lambda: (
ws_nodes._websocket_clients.clear(),
ws_nodes._websocket_clients.update(original_clients),
)
)
with mock.patch.dict(
os.environ,
realtime_contract_env(),
):
server_first = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=True,
debug=False,
)
server_repeated = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=True,
debug=False,
)
server_changed = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8002,
external_server_only=True,
debug=False,
)
auto_mode_first = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=False,
debug=False,
)
auto_mode_second = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=False,
debug=False,
)
json_first = ws_nodes.VrchJsonWebSocketChannelLoaderNode.IS_CHANGED(
channel="2"
)
json_repeated = ws_nodes.VrchJsonWebSocketChannelLoaderNode.IS_CHANGED(
channel="2"
)
fake_client.sequence = 4
json_changed = ws_nodes.VrchJsonWebSocketChannelLoaderNode.IS_CHANGED(
channel="2"
)
self.assertEqual(server_first, server_repeated)
self.assertNotEqual(server_first, server_changed)
self.assertNotEqual(auto_mode_first, auto_mode_first)
self.assertNotEqual(auto_mode_second, auto_mode_second)
self.assertEqual(json_first, json_repeated)
self.assertNotEqual(json_first, json_changed)
def test_16_realtime_image_encode_cache_reuses_only_identical_tensor(self):
sent_messages = []
encode_formats = []
class FakeServer:
def send_to_channel(self, path, channel, data):
sent_messages.append((path, channel, data))
class FakeImage:
def save(self, buffer, format):
encode_formats.append(format)
buffer.write(b"encoded-image")
original_get_server = ws_nodes.get_global_server
original_fromarray = ws_nodes.Image.fromarray
self.addCleanup(
lambda: setattr(ws_nodes, "get_global_server", original_get_server)
)
self.addCleanup(
lambda: setattr(ws_nodes.Image, "fromarray", original_fromarray)
)
self.addCleanup(ws_nodes._reset_realtime_image_encode_cache)
ws_nodes.get_global_server = lambda *args, **kwargs: FakeServer()
ws_nodes.Image.fromarray = lambda *args, **kwargs: FakeImage()
ws_nodes._reset_realtime_image_encode_cache()
first_node = ws_nodes.VrchImageWebSocketSimpleWebViewerNode()
second_node = ws_nodes.VrchImageWebSocketSimpleWebViewerNode()
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
kwargs = {
"images": image,
"channel": "1",
"server": "127.0.0.1:8001",
"format": "JPEG",
"number_of_images": 1,
"image_display_duration": 50,
"fade_anim_duration": 10,
"window_width": 512,
"window_height": 512,
"show_url": False,
"dev_mode": False,
"debug": False,
"extra_params": "",
"url": "",
"prompt_scope": {},
}
with mock.patch.dict(
os.environ,
realtime_contract_env(),
):
first_node.send_images(**kwargs)
kwargs["channel"] = "3"
second_node.send_images(**kwargs)
self.assertEqual(encode_formats, ["JPEG"])
self.assertEqual(sent_messages[0][2][8:], sent_messages[1][2][8:])
self.assertEqual(
ws_nodes._realtime_image_encode_cache_stats(),
{"hits": 1, "misses": 1},
)
image.add_(1)
first_node.send_images(**kwargs)
self.assertEqual(encode_formats, ["JPEG", "JPEG"])
self.assertEqual(
ws_nodes._realtime_image_encode_cache_stats(),
{"hits": 1, "misses": 2},
)
def test_17_default_mode_encodes_each_output_node_independently(self):
encode_count = 0
def fake_encode(images, image_format):
nonlocal encode_count
encode_count += 1
return (b"encoded-image",)
original_uncached = ws_nodes._encode_image_batch_uncached
self.addCleanup(
lambda: setattr(
ws_nodes,
"_encode_image_batch_uncached",
original_uncached,
)
)
ws_nodes._encode_image_batch_uncached = fake_encode
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
with mock.patch.dict(
os.environ,
{"VRCH_COMFYUI_PERFORMANCE_CACHE": "default"},
):
ws_nodes._encode_image_batch(image, "JPEG")
ws_nodes._encode_image_batch(image, "JPEG")
self.assertEqual(encode_count, 2)
def test_18_inference_tensor_without_version_counter_is_supported(self):
with torch.inference_mode():
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
key = ws_nodes._image_batch_cache_key(image, "JPEG")
self.assertIsNone(key[-1])
def test_19_inference_tensor_cache_is_consumed_after_one_hit(self):
encoded_values = []
def fake_encode(images, image_format):
value = int(images[0, 0, 0, 0].item())
encoded_values.append(value)
return (f"encoded-{value}".encode(),)
original_uncached = ws_nodes._encode_image_batch_uncached
self.addCleanup(
lambda: setattr(
ws_nodes,
"_encode_image_batch_uncached",
original_uncached,
)
)
self.addCleanup(ws_nodes._reset_realtime_image_encode_cache)
ws_nodes._encode_image_batch_uncached = fake_encode
ws_nodes._reset_realtime_image_encode_cache()
with (
torch.inference_mode(),
mock.patch.dict(
os.environ,
realtime_contract_env(),
),
):
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
first_prompt = {}
second_prompt = {}
first = ws_nodes._encode_image_batch(
image, "JPEG", first_prompt
)
second = ws_nodes._encode_image_batch(
image, "JPEG", first_prompt
)
image.add_(1)
third = ws_nodes._encode_image_batch(
image, "JPEG", second_prompt
)
fourth = ws_nodes._encode_image_batch(
image, "JPEG", second_prompt
)
self.assertEqual(first, second)
self.assertEqual(third, fourth)
self.assertNotEqual(first, third)
self.assertEqual(encoded_values, [0, 1])
self.assertEqual(
ws_nodes._realtime_image_encode_cache_stats(),
{"hits": 2, "misses": 2},
)
def test_20_image_output_resolves_server_on_every_send(self):
get_server_calls = []
class FakeServer:
def send_to_channel(self, path, channel, data):
pass
def fake_get_server(*args, **kwargs):
get_server_calls.append((args, kwargs))
return FakeServer()
original_get_server = ws_nodes.get_global_server
original_encode = ws_nodes._encode_image_batch
self.addCleanup(
lambda: setattr(ws_nodes, "get_global_server", original_get_server)
)
self.addCleanup(
lambda: setattr(ws_nodes, "_encode_image_batch", original_encode)
)
ws_nodes.get_global_server = fake_get_server
ws_nodes._encode_image_batch = lambda *args, **kwargs: (b"encoded",)
node = ws_nodes.VrchImageWebSocketSimpleWebViewerNode()
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
with mock.patch.dict(
os.environ,
realtime_contract_env(),
):
for _ in range(2):
node.send_images(
image, "1", "127.0.0.1:8001", "JPEG", 1, 50, 10,
512, 512, False, False, False, "", "",
)
self.assertEqual(len(get_server_calls), 2)
def test_21_prompt_scope_prevents_incomplete_pair_stale_hit(self):
encoded_values = []
def fake_encode(images, image_format):
value = int(images[0, 0, 0, 0].item())
encoded_values.append(value)
return (f"encoded-{value}".encode(),)
original_uncached = ws_nodes._encode_image_batch_uncached
self.addCleanup(
lambda: setattr(
ws_nodes,
"_encode_image_batch_uncached",
original_uncached,
)
)
self.addCleanup(ws_nodes._reset_realtime_image_encode_cache)
ws_nodes._encode_image_batch_uncached = fake_encode
ws_nodes._reset_realtime_image_encode_cache()
with (
torch.inference_mode(),
mock.patch.dict(
os.environ,
realtime_contract_env(),
),
):
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
first_prompt = {}
second_prompt = {}
first = ws_nodes._encode_image_batch(
image, "JPEG", first_prompt
)
image.add_(1)
second = ws_nodes._encode_image_batch(
image, "JPEG", second_prompt
)
repeated = ws_nodes._encode_image_batch(
image, "JPEG", second_prompt
)
self.assertNotEqual(first, second)
self.assertEqual(second, repeated)
self.assertEqual(encoded_values, [0, 1])
self.assertEqual(
ws_nodes._realtime_image_encode_cache_stats(),
{"hits": 1, "misses": 2},
)
def test_22_realtime_without_prompt_scope_fails_closed(self):
encode_count = 0
def fake_encode(images, image_format):
nonlocal encode_count
encode_count += 1
return (b"encoded-image",)
original_uncached = ws_nodes._encode_image_batch_uncached
self.addCleanup(
lambda: setattr(
ws_nodes,
"_encode_image_batch_uncached",
original_uncached,
)
)
self.addCleanup(ws_nodes._reset_realtime_image_encode_cache)
ws_nodes._encode_image_batch_uncached = fake_encode
ws_nodes._reset_realtime_image_encode_cache()
image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
with mock.patch.dict(
os.environ,
realtime_contract_env(),
):
ws_nodes._encode_image_batch(image, "JPEG")
ws_nodes._encode_image_batch(image, "JPEG")
self.assertEqual(encode_count, 2)
self.assertEqual(
ws_nodes._realtime_image_encode_cache_stats(),
{"hits": 0, "misses": 0},
)
def test_23_realtime_contract_activates_with_matching_gc_capability(self):
env = realtime_contract_env()
env.pop(ws_nodes.REALTIME_NODE_CAPABILITY_ENV)
with (
mock.patch.dict(os.environ, env, clear=True),
self.assertLogs(level="INFO") as captured,
):
state = ws_nodes._initialize_realtime_contract()
enabled = ws_nodes._realtime_safe_set_enabled()
self.assertEqual(state["effective"], "realtime")
self.assertEqual(state["status"], "active")
self.assertTrue(enabled)
self.assertIn(
"component=node contract=vrch-realtime-v1 "
"capability=node-safe-set-v1 requested=realtime "
"effective=realtime",
"\n".join(captured.output),
)
def test_24_realtime_contract_fails_closed_without_gc_capability(self):
with (
mock.patch.dict(
os.environ,
{
ws_nodes.REALTIME_MODE_ENV: "realtime",
ws_nodes.REALTIME_REQUESTED_MODE_ENV: "realtime",
},
clear=True,
),
self.assertLogs(level="ERROR") as captured,
):
state = ws_nodes._initialize_realtime_contract()
effective_mode = os.environ[ws_nodes.REALTIME_MODE_ENV]
node_capability = os.environ.get(
ws_nodes.REALTIME_NODE_CAPABILITY_ENV
)
enabled = ws_nodes._realtime_safe_set_enabled()
self.assertEqual(state["effective"], "default")
self.assertEqual(state["status"], "skew-fail-closed")
self.assertEqual(state["peer"], "missing")
self.assertEqual(effective_mode, "default")
self.assertIsNone(node_capability)
self.assertFalse(enabled)
self.assertIn("status=skew-fail-closed", "\n".join(captured.output))
def test_25_realtime_contract_fails_closed_on_contract_mismatch(self):
with (
mock.patch.dict(
os.environ,
{
ws_nodes.REALTIME_MODE_ENV: "realtime",
ws_nodes.REALTIME_REQUESTED_MODE_ENV: "realtime",
ws_nodes.REALTIME_CONTRACT_ENV: "vrch-realtime-v0",
ws_nodes.REALTIME_GC_CAPABILITY_ENV:
ws_nodes.REALTIME_GC_CAPABILITY,
},
clear=True,
),
self.assertLogs(level="ERROR"),
):
state = ws_nodes._initialize_realtime_contract()
effective_mode = os.environ[ws_nodes.REALTIME_MODE_ENV]
node_capability = os.environ.get(
ws_nodes.REALTIME_NODE_CAPABILITY_ENV
)
self.assertEqual(state["effective"], "default")
self.assertEqual(state["status"], "skew-fail-closed")
self.assertEqual(effective_mode, "default")
self.assertIsNone(node_capability)
self.assertFalse(ws_nodes._realtime_safe_set_enabled())
def test_26_unset_mode_initializes_with_default_semantics(self):
with (
mock.patch.dict(os.environ, {}, clear=True),
self.assertLogs(level="INFO") as captured,
):
state = ws_nodes._initialize_realtime_contract()
enabled = ws_nodes._realtime_safe_set_enabled()
node_capability = os.environ.get(
ws_nodes.REALTIME_NODE_CAPABILITY_ENV
)
server_token = ws_nodes.VrchWebSocketServerNode.IS_CHANGED(
server="0.0.0.0",
port=8001,
external_server_only=True,
debug=False,
)
self.assertEqual(state["requested"], "default")
self.assertEqual(state["effective"], "default")
self.assertEqual(state["status"], "standby")
self.assertIsNone(node_capability)
self.assertFalse(enabled)
self.assertNotEqual(server_token, server_token)
self.assertIn("status=standby", "\n".join(captured.output))
class TestWebSocketNodesIntegration(unittest.TestCase):
def setUp(self):
+99 -1
View File
@@ -766,7 +766,10 @@ class TestWebSocketServerIntegration(unittest.TestCase):
proxy.send_to_channel("/image", 1, '{"settings":{"numberOfImages":1}}')
await asyncio.sleep(0.2)
endpoint_uri = f"ws://{self.test_host}:{port}/image?channel=1"
endpoint_uri = (
f"ws://{self.test_host}:{port}/image?channel=1"
"&client=comfyui-output&role=service"
)
queue = proxy._endpoint_queues.get(endpoint_uri)
self.assertIsNotNone(queue, "Proxy endpoint queue should exist")
self.assertEqual(queue.maxsize, 1, "Realtime endpoint queue must be bounded")
@@ -960,6 +963,101 @@ class TestWebSocketServerIntegration(unittest.TestCase):
print("✓ Proxy sender stability under downlink pressure test passed")
def test_18_proxy_recovers_when_external_endpoint_appears(self):
"""A live proxy should deliver again after its endpoint was unavailable."""
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.bind((self.test_host, 0))
port = sock.getsockname()[1]
payload = struct.pack(">II", 1, (1 << 16) | (0 << 8) | 1) + b"recovered"
received = []
received_event = threading.Event()
server_ready = threading.Event()
server_loop = asyncio.new_event_loop()
server_holder = {"instance": None}
async def handler(websocket, _path=None):
try:
received.append(
await asyncio.wait_for(websocket.recv(), timeout=4.0)
)
received_event.set()
except Exception:
pass
def run_server():
asyncio.set_event_loop(server_loop)
async def start_server():
server_holder["instance"] = await websockets.serve(
handler,
self.test_host,
port,
ping_interval=0.2,
ping_timeout=0.2,
)
server_ready.set()
server_loop.run_until_complete(start_server())
server_loop.run_forever()
proxy = WebSocketClientProxy(self.test_host, port, debug=False)
proxy.register_path("/image")
server_thread = None
try:
# The endpoint is deliberately absent for the first delivery attempt.
proxy.send_to_channel("/image", 1, payload)
time.sleep(0.75)
server_thread = threading.Thread(target=run_server, daemon=True)
server_thread.start()
self.assertTrue(
server_ready.wait(timeout=3.0),
"Recovery test endpoint did not start",
)
deadline = time.monotonic() + 5.0
while not received_event.is_set() and time.monotonic() < deadline:
proxy.send_to_channel("/image", 1, payload)
received_event.wait(timeout=0.1)
self.assertTrue(
received_event.is_set(),
"Proxy did not recover after the external endpoint appeared",
)
self.assertEqual(received[-1], payload)
self.assertTrue(proxy.is_running())
finally:
try:
proxy.stop()
except Exception:
pass
if (
server_holder["instance"] is not None
and server_loop.is_running()
):
try:
async def shutdown_server():
server_holder["instance"].close()
await server_holder["instance"].wait_closed()
future = asyncio.run_coroutine_threadsafe(
shutdown_server(),
server_loop,
)
future.result(timeout=2.0)
except Exception:
pass
try:
server_loop.call_soon_threadsafe(server_loop.stop)
except Exception:
pass
if server_thread is not None:
server_thread.join(timeout=1.0)
print("✓ Proxy unavailable-to-available recovery test passed")
def run_all_tests():
"""Run both unit tests and integration tests"""
+253 -19
View File
@@ -1,11 +1,14 @@
import hashlib
import io
import json
import logging
import os
import time
import struct
import base64
import re
import tempfile
import weakref
import numpy as np
import asyncio
import websockets
@@ -37,6 +40,203 @@ AUDIO_PLAYER_QUALITY_PRESETS_KBPS = {
"high": 192,
}
REALTIME_CONTRACT = "vrch-realtime-v1"
REALTIME_NODE_CAPABILITY = "node-safe-set-v1"
REALTIME_GC_CAPABILITY = "gc-policy-v3"
REALTIME_MODE_ENV = "VRCH_COMFYUI_PERFORMANCE_CACHE"
REALTIME_REQUESTED_MODE_ENV = "VRCH_COMFYUI_PERFORMANCE_REQUESTED_MODE"
REALTIME_CONTRACT_ENV = "VRCH_COMFYUI_REALTIME_CONTRACT"
REALTIME_NODE_CAPABILITY_ENV = "VRCH_COMFYUI_REALTIME_NODE_CAPABILITY"
REALTIME_GC_CAPABILITY_ENV = "VRCH_COMFYUI_REALTIME_GC_CAPABILITY"
def _requested_realtime_mode():
value = os.environ.get(
REALTIME_REQUESTED_MODE_ENV,
os.environ.get(REALTIME_MODE_ENV, "default"),
)
return "realtime" if str(value).strip().lower() == "realtime" else "default"
def _initialize_realtime_contract():
"""Publish the node capability only after the owner package loads."""
requested = _requested_realtime_mode()
peer_contract = os.environ.get(REALTIME_CONTRACT_ENV, "").strip()
peer_capability = os.environ.get(REALTIME_GC_CAPABILITY_ENV, "").strip()
compatible = (
peer_contract == REALTIME_CONTRACT
and peer_capability == REALTIME_GC_CAPABILITY
)
if requested == "realtime" and compatible:
effective = "realtime"
status = "active"
os.environ[REALTIME_NODE_CAPABILITY_ENV] = REALTIME_NODE_CAPABILITY
log = logging.info
elif requested == "realtime":
effective = "default"
status = "skew-fail-closed"
os.environ.pop(REALTIME_NODE_CAPABILITY_ENV, None)
os.environ[REALTIME_MODE_ENV] = "default"
log = logging.error
else:
effective = "default"
status = "standby"
os.environ.pop(REALTIME_NODE_CAPABILITY_ENV, None)
log = logging.info
peer = peer_capability or "missing"
contract = peer_contract or "missing"
log(
"[VRCH_REALTIME_CAPABILITY] component=node contract=%s "
"capability=%s requested=%s effective=%s peer=%s "
"peer_contract=%s status=%s",
REALTIME_CONTRACT,
REALTIME_NODE_CAPABILITY,
requested,
effective,
peer,
contract,
status,
)
return {
"component": "node",
"contract": REALTIME_CONTRACT,
"capability": REALTIME_NODE_CAPABILITY,
"requested": requested,
"effective": effective,
"peer": peer,
"peer_contract": contract,
"status": status,
}
def _realtime_safe_set_enabled():
return (
os.environ.get(REALTIME_MODE_ENV, "").strip().lower() == "realtime"
and os.environ.get(REALTIME_CONTRACT_ENV, "").strip()
== REALTIME_CONTRACT
and os.environ.get(REALTIME_GC_CAPABILITY_ENV, "").strip()
== REALTIME_GC_CAPABILITY
and os.environ.get(REALTIME_NODE_CAPABILITY_ENV, "").strip()
== REALTIME_NODE_CAPABILITY
)
_realtime_image_encode_cache_lock = threading.Lock()
_realtime_image_encode_cache = {
"images_ref": None,
"prompt_scope": None,
"key": None,
"payloads": None,
"hits": 0,
"misses": 0,
}
def _reset_realtime_image_encode_cache():
with _realtime_image_encode_cache_lock:
_realtime_image_encode_cache.update(
{
"images_ref": None,
"prompt_scope": None,
"key": None,
"payloads": None,
"hits": 0,
"misses": 0,
}
)
def _realtime_image_encode_cache_stats():
with _realtime_image_encode_cache_lock:
return {
"hits": int(_realtime_image_encode_cache["hits"]),
"misses": int(_realtime_image_encode_cache["misses"]),
}
def _encode_image_batch_uncached(images, image_format):
payloads = []
for tensor in images:
arr = 255.0 * tensor.cpu().numpy()
img = Image.fromarray(np.clip(arr, 0, 255).astype(np.uint8))
buf = io.BytesIO()
img.save(buf, format=image_format)
payloads.append(buf.getvalue())
return tuple(payloads)
def _image_batch_cache_key(images, image_format):
try:
version = int(images._version)
except (AttributeError, RuntimeError, TypeError, ValueError):
version = None
return (
str(image_format).upper(),
tuple(getattr(images, "shape", ())),
str(getattr(images, "dtype", "")),
str(getattr(images, "device", "")),
version,
)
def _encode_image_batch(images, image_format, prompt_scope=None):
if not _realtime_safe_set_enabled() or prompt_scope is None:
return _encode_image_batch_uncached(images, image_format)
try:
images_ref = weakref.ref(images)
except TypeError:
return _encode_image_batch_uncached(images, image_format)
cache_key = _image_batch_cache_key(images, image_format)
with _realtime_image_encode_cache_lock:
cached_ref = _realtime_image_encode_cache["images_ref"]
cached_prompt_scope = _realtime_image_encode_cache["prompt_scope"]
if (
cached_ref is not None
and cached_ref() is images
and cached_prompt_scope is prompt_scope
and _realtime_image_encode_cache["key"] == cache_key
):
_realtime_image_encode_cache["hits"] += 1
payloads = _realtime_image_encode_cache["payloads"]
# The canonical workflow has exactly two output nodes. Consume the
# cached bytes on the second node so an inference tensor without a
# mutation version can never reuse stale bytes across prompts.
_realtime_image_encode_cache.update(
{
"images_ref": None,
"prompt_scope": None,
"key": None,
"payloads": None,
}
)
else:
payloads = _encode_image_batch_uncached(images, image_format)
_realtime_image_encode_cache.update(
{
"images_ref": images_ref,
"prompt_scope": prompt_scope,
"key": cache_key,
"payloads": payloads,
"misses": _realtime_image_encode_cache["misses"] + 1,
}
)
calls = (
_realtime_image_encode_cache["hits"]
+ _realtime_image_encode_cache["misses"]
)
if calls and calls % 1024 == 0:
print(
"[VRCH_REALTIME_SAFE_SET] jpeg_encode_cache "
f"hits={_realtime_image_encode_cache['hits']} "
f"misses={_realtime_image_encode_cache['misses']}"
)
return payloads
def _describe_image_binary_payload(data):
if not isinstance(data, (bytes, bytearray)):
@@ -118,7 +318,16 @@ class VrchWebSocketServerNode:
@classmethod
def IS_CHANGED(cls, **kwargs):
return float("NaN") # Always trigger evaluation to check server status
external_server_only = bool(kwargs.get("external_server_only", False))
if not _realtime_safe_set_enabled() or not external_server_only:
return float("NaN") # Preserve status checks outside external-only.
server = str(kwargs.get("server", DEFAULT_SERVER_IP))
port = int(kwargs.get("port", DEFAULT_SERVER_PORT))
debug = bool(kwargs.get("debug", False))
return json.dumps(
[server, port, external_server_only, debug],
separators=(",", ":"),
)
class VrchImageWebSocketWebViewerNode:
@@ -148,7 +357,8 @@ class VrchImageWebSocketWebViewerNode:
"debug": ("BOOLEAN", {"default": False}),
"extra_params":("STRING", {"multiline": True, "dynamicPrompts": False}),
"url": ("STRING", {"default": "", "multiline": True}),
}
},
"hidden": {"prompt_scope": "PROMPT"},
}
RETURN_TYPES = ("IMAGE", "STRING")
RETURN_NAMES = ("IMAGES", "URL")
@@ -176,7 +386,8 @@ class VrchImageWebSocketWebViewerNode:
dev_mode,
debug,
extra_params,
url):
url,
prompt_scope=None):
results = []
host, port = server.split(":")
server = get_global_server(host, port, path="/image", debug=debug) # Ensure path is set correctly for viewer
@@ -186,12 +397,8 @@ class VrchImageWebSocketWebViewerNode:
batch_id = (batch_id + 1) % 65536
self._last_batch_id = batch_id
for index, tensor in enumerate(images):
arr = 255.0 * tensor.cpu().numpy()
img = Image.fromarray(np.clip(arr, 0, 255).astype(np.uint8))
buf = io.BytesIO()
img.save(buf, format=format)
binary_data = buf.getvalue()
encoded_images = _encode_image_batch(images, format, prompt_scope)
for index, binary_data in enumerate(encoded_images):
meta = (batch_id << 16) | ((index & 0xFF) << 8) | (batch_size & 0xFF)
header = struct.pack(">II", 1, meta)
data = header + binary_data
@@ -240,7 +447,8 @@ class VrchImageWebSocketSimpleWebViewerNode:
"debug": ("BOOLEAN", {"default": False}),
"extra_params":("STRING", {"multiline": True, "dynamicPrompts": False}),
"url": ("STRING", {"default": "", "multiline": True}),
}
},
"hidden": {"prompt_scope": "PROMPT"},
}
RETURN_TYPES = ("IMAGE", "STRING")
@@ -263,7 +471,8 @@ class VrchImageWebSocketSimpleWebViewerNode:
dev_mode,
debug,
extra_params,
url):
url,
prompt_scope=None):
results = []
host, port = server.split(":")
server = get_global_server(host, port, path="/image", debug=debug) # Ensure path is set correctly for viewer
@@ -273,12 +482,8 @@ class VrchImageWebSocketSimpleWebViewerNode:
batch_id = (batch_id + 1) % 65536
self._last_batch_id = batch_id
for index, tensor in enumerate(images):
arr = 255.0 * tensor.cpu().numpy()
img = Image.fromarray(np.clip(arr, 0, 255).astype(np.uint8))
buf = io.BytesIO()
img.save(buf, format=format)
binary_data = buf.getvalue()
encoded_images = _encode_image_batch(images, format, prompt_scope)
for index, binary_data in enumerate(encoded_images):
meta = (batch_id << 16) | ((index & 0xFF) << 8) | (batch_size & 0xFF)
header = struct.pack(">II", 1, meta)
data = header + binary_data
@@ -722,6 +927,10 @@ class WebSocketClient:
with self.lock:
return self.received_data, self.received_sequence
def get_received_sequence(self):
with self.lock:
return self.received_sequence
def _is_latest_message_candidate(self, message):
if self.path == "/image":
return isinstance(message, (bytes, bytearray)) and len(message) >= 8
@@ -802,6 +1011,29 @@ class WebSocketClient:
if self.thread and self.thread.is_alive():
self.thread.join(timeout=1.5)
def _websocket_source_sequence_token(path, channel):
normalized_path = path if str(path).startswith("/") else f"/{path}"
normalized_channel = int(channel)
with _websocket_clients_lock:
sources = [
[
str(client.host),
int(client.port),
int(client.get_received_sequence()),
]
for client in _websocket_clients.values()
if client.path == normalized_path and client.channel == normalized_channel
]
if not sources:
return f"{normalized_path}|{normalized_channel}|no-client"
sources.sort()
return json.dumps(
[normalized_path, normalized_channel, sources],
separators=(",", ":"),
)
def get_websocket_client(host, port, path, channel, data_handler=None, debug=False, latest_only=False):
key = f"{host}:{port}:{path}:{channel}"
with _websocket_clients_lock:
@@ -1413,8 +1645,10 @@ class VrchJsonWebSocketChannelLoaderNode:
@classmethod
def IS_CHANGED(cls, **kwargs):
# Always trigger evaluation to check for new data
return float("NaN")
if not _realtime_safe_set_enabled():
return float("NaN") # Preserve the original Default behavior.
channel = kwargs.get("channel", 1)
return _websocket_source_sequence_token("/json", channel)
class VrchMidiWebSocketChannelLoaderNode:
@classmethod
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-web-viewer"
description = "The ComfyUI Web Viewer by vrch.ai is a custom node collection offering a real-time AI-generated interactive art framework. This utility integrates realtime streaming into ComfyUI workflows, supporting keyboard control nodes, OSC control nodes, sound input nodes, and more. Accessible from any device with a web browser, it enables real time interaction with AI-generated content, making it ideal for interactive visual projects and enhancing ComfyUI workflows with efficient content management and display."
version = "1.1.25"
version = "1.1.26"
license = {file = "LICENSE"}
dependencies = ["aiohttp","ffmpeg-python","matplotlib","pydub","audioop-lts; python_version >= '3.13'","python-osc","qrcode[pil]","scikit-learn","srt","websockets"]