Merge branch 'dev' into 'main'
Release 1.1.26 See merge request vrch/comfyui/comfyui-web-viewer!61
This commit is contained in:
+1
-1
@@ -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
@@ -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]
|
||||
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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}")
|
||||
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user