From 738b52d2552e0396b41068fe45e613ff296d2d61 Mon Sep 17 00:00:00 2001 From: tianzi Date: Wed, 29 Jul 2026 19:27:58 +0100 Subject: [PATCH 1/4] feat: add realtime websocket safe-set --- .bumpversion.cfg | 2 +- CHANGELOG.md | 11 + __init__.py | 8 +- nodes/tests/websocket_nodes_test.py | 480 +++++++++++++++++++++++++++ nodes/tests/websocket_server_test.py | 100 +++++- nodes/websocket_nodes.py | 272 +++++++++++++-- pyproject.toml | 2 +- 7 files changed, 852 insertions(+), 23 deletions(-) diff --git a/.bumpversion.cfg b/.bumpversion.cfg index 793b4b4..d05e644 100644 --- a/.bumpversion.cfg +++ b/.bumpversion.cfg @@ -1,5 +1,5 @@ [bumpversion] -current_version = 1.1.25 +current_version = 1.1.26 commit = True tag = True parse = (?P\d+)\.(?P\d+)\.(?P\d+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 059048c..d046848 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,17 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [1.1.26 - 2026-07-29] + +### Added + +- 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] ### Fixed diff --git a/__init__.py b/__init__.py index dbc64de..f598088 100644 --- a/__init__.py +++ b/__init__.py @@ -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 * @@ -12,7 +13,7 @@ from .nodes.midi_nodes import * from .nodes.audio_music2emo_node import * from .nodes.workflow_export_nodes import * -__version__ = "1.1.25" +__version__ = "1.1.26" NODE_CLASS_MAPPINGS = { "VrchAnyOSCControlNode": VrchAnyOSCControlNode, @@ -160,6 +161,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" diff --git a/nodes/tests/websocket_nodes_test.py b/nodes/tests/websocket_nodes_test.py index dafc72d..a310e13 100644 --- a/nodes/tests/websocket_nodes_test.py +++ b/nodes/tests/websocket_nodes_test.py @@ -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): diff --git a/nodes/tests/websocket_server_test.py b/nodes/tests/websocket_server_test.py index e77877d..e0c46f6 100644 --- a/nodes/tests/websocket_server_test.py +++ b/nodes/tests/websocket_server_test.py @@ -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""" diff --git a/nodes/websocket_nodes.py b/nodes/websocket_nodes.py index c70862c..2f5ea5b 100644 --- a/nodes/websocket_nodes.py +++ b/nodes/websocket_nodes.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 1a3e7a4..1b5d8ad 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"] From 26cb85d62f1a9a45ddbab12f8a823906ac217069 Mon Sep 17 00:00:00 2001 From: tianzi Date: Fri, 31 Jul 2026 18:33:59 +0100 Subject: [PATCH 2/4] feat: add TensorRT auto loader --- CHANGELOG.md | 4 + README.md | 4 + __init__.py | 3 + docs/model_nodes.md | 29 ++++ nodes/model_nodes.py | 289 ++++++++++++++++++++++++++++++++ nodes/tests/model_nodes_test.py | 182 ++++++++++++++++++++ 6 files changed, 511 insertions(+) create mode 100644 docs/model_nodes.md create mode 100644 nodes/model_nodes.py create mode 100644 nodes/tests/model_nodes_test.py diff --git a/CHANGELOG.md b/CHANGELOG.md index d046848..e7bfed3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- add a backend-only TensorRT Auto Loader node with automatic PyTorch fallback + ## [1.1.26 - 2026-07-29] ### Added diff --git a/README.md b/README.md index f3a9a5a..5724d02 100644 --- a/README.md +++ b/README.md @@ -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) diff --git a/__init__.py b/__init__.py index f598088..11c48eb 100644 --- a/__init__.py +++ b/__init__.py @@ -12,6 +12,7 @@ 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.26" @@ -68,6 +69,7 @@ NODE_CLASS_MAPPINGS = { "VrchMidiDeviceLoaderNode": VrchMidiDeviceLoaderNode, "VrchMidiWebSocketChannelLoaderNode": VrchMidiWebSocketChannelLoaderNode, "VrchModelWebViewerNode": VrchModelWebViewerNode, + "VrchTensorRTAutoLoaderNode": VrchTensorRTAutoLoaderNode, "VrchOSCControlSettingsNode": VrchOSCControlSettingsNode, "VrchQRCodeNode": VrchQRCodeNode, "VrchSwitchOSCControlNode": VrchSwitchOSCControlNode, @@ -141,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", diff --git a/docs/model_nodes.md b/docs/model_nodes.md new file mode 100644 index 0000000..a087950 --- /dev/null +++ b/docs/model_nodes.md @@ -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. diff --git a/nodes/model_nodes.py b/nodes/model_nodes.py new file mode 100644 index 0000000..e1b00b9 --- /dev/null +++ b/nodes/model_nodes.py @@ -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}") diff --git a/nodes/tests/model_nodes_test.py b/nodes/tests/model_nodes_test.py new file mode 100644 index 0000000..aff4833 --- /dev/null +++ b/nodes/tests/model_nodes_test.py @@ -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() From ac0d307e73be2a062b1d3891a420ccf604cfce95 Mon Sep 17 00:00:00 2001 From: tianzi Date: Fri, 31 Jul 2026 19:47:05 +0100 Subject: [PATCH 3/4] =?UTF-8?q?Bump=20version:=201.1.26=20=E2=86=92=201.1.?= =?UTF-8?q?27?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .bumpversion.cfg | 2 +- CHANGELOG.md | 2 +- __init__.py | 2 +- pyproject.toml | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/.bumpversion.cfg b/.bumpversion.cfg index d05e644..d5c6907 100644 --- a/.bumpversion.cfg +++ b/.bumpversion.cfg @@ -1,5 +1,5 @@ [bumpversion] -current_version = 1.1.26 +current_version = 1.1.27 commit = True tag = True parse = (?P\d+)\.(?P\d+)\.(?P\d+) diff --git a/CHANGELOG.md b/CHANGELOG.md index e7bfed3..d89851b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,7 +5,7 @@ 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.27 - 2026-07-31] ### Added diff --git a/__init__.py b/__init__.py index 11c48eb..65128da 100644 --- a/__init__.py +++ b/__init__.py @@ -14,7 +14,7 @@ from .nodes.audio_music2emo_node import * from .nodes.workflow_export_nodes import * from .nodes.model_nodes import * -__version__ = "1.1.26" +__version__ = "1.1.27" NODE_CLASS_MAPPINGS = { "VrchAnyOSCControlNode": VrchAnyOSCControlNode, diff --git a/pyproject.toml b/pyproject.toml index 1b5d8ad..7a2e6f8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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.26" +version = "1.1.27" license = {file = "LICENSE"} dependencies = ["aiohttp","ffmpeg-python","matplotlib","pydub","audioop-lts; python_version >= '3.13'","python-osc","qrcode[pil]","scikit-learn","srt","websockets"] From 1a22e50d2c64ad93b247b003b09630d52da9aa23 Mon Sep 17 00:00:00 2001 From: tianzi Date: Fri, 31 Jul 2026 20:00:31 +0100 Subject: [PATCH 4/4] fix: release 1.1.26 --- .bumpversion.cfg | 2 +- CHANGELOG.md | 7 +------ __init__.py | 2 +- pyproject.toml | 2 +- 4 files changed, 4 insertions(+), 9 deletions(-) diff --git a/.bumpversion.cfg b/.bumpversion.cfg index d5c6907..d05e644 100644 --- a/.bumpversion.cfg +++ b/.bumpversion.cfg @@ -1,5 +1,5 @@ [bumpversion] -current_version = 1.1.27 +current_version = 1.1.26 commit = True tag = True parse = (?P\d+)\.(?P\d+)\.(?P\d+) diff --git a/CHANGELOG.md b/CHANGELOG.md index d89851b..99ec2e6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,16 +5,11 @@ 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). -## [1.1.27 - 2026-07-31] +## [1.1.26 - 2026-07-31] ### Added - add a backend-only TensorRT Auto Loader node with automatic PyTorch fallback - -## [1.1.26 - 2026-07-29] - -### Added - - 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 diff --git a/__init__.py b/__init__.py index 65128da..11c48eb 100644 --- a/__init__.py +++ b/__init__.py @@ -14,7 +14,7 @@ from .nodes.audio_music2emo_node import * from .nodes.workflow_export_nodes import * from .nodes.model_nodes import * -__version__ = "1.1.27" +__version__ = "1.1.26" NODE_CLASS_MAPPINGS = { "VrchAnyOSCControlNode": VrchAnyOSCControlNode, diff --git a/pyproject.toml b/pyproject.toml index 7a2e6f8..1b5d8ad 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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.27" +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"]