fix: stabilize websocket image fallback and raw MIDI output

This commit is contained in:
tianzi
2026-05-29 17:23:31 +01:00
parent 878f8e81f1
commit b605fe065b
7 changed files with 93 additions and 30 deletions
+10
View File
@@ -5,6 +5,16 @@ 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]
### Updated
- update MIDI Control node `RAW_CC` outputs to use `INT` for raw MIDI/source values
### Fixed
- fix `IMAGE WebSocket Channel Loader @ vrch.ai` falling back to `default_image` after receiving WebSocket images when `placeholder` is set to `image`
## [1.1.19 - 2026-05-28]
### Added
+1 -1
View File
@@ -25,7 +25,7 @@ Maps values from `VRCH_MIDI` state to Int or Float outputs.
3. **Outputs:**
- **`VALUE`**: Remapped Int or Float value.
- **`RAW_CC`**: Raw source value as Float. When default is used, `RAW_CC` is `0.0`.
- **`RAW_CC`**: Raw source value as Int. When default is used, `RAW_CC` is `0`.
**Notes:**
- The default mode is `workflow_key`, which is the recommended user-facing setup.
+3 -3
View File
@@ -275,7 +275,7 @@ Provides adjustable CSS image filter parameters as a JSON object for composition
- **"black"**: pure black placeholder image.
- **"white"**: pure white placeholder image.
- **"grey"**: mid-grey placeholder image.
- **"image"**: use the provided **`default_image`** as placeholder. Requires supplying **`default_image`**. The node detects changes to this image and outputs it immediately once per change.
- **"image"**: use the provided **`default_image`** as placeholder until a WebSocket image is available. Requires supplying **`default_image`**.
- **`default_image`**: *(Optional)* Image to use when **`placeholder`** is set to **"image"**.
- **`debug`**: Enable this option to print detailed debug information to the console for troubleshooting.
@@ -286,8 +286,8 @@ Provides adjustable CSS image filter parameters as a JSON object for composition
4. **Receiving Images:**
- This node automatically connects to the specified WebSocket channel and listens for incoming image data.
- When an image is received, it will be processed and made available as the `IMAGE` output with `IS_DEFAULT_IMAGE` set to `False`.
- If **`placeholder`** is set to **"image"** and a new **`default_image`** was provided since the last execution, it is output immediately once (`IMAGE` + `True`).
- If no WebSocket image is received afterward, the **`default_image`** is used as the placeholder output (`IS_DEFAULT_IMAGE=False` in that case?).
- Once a WebSocket image has been received, the loader keeps returning the latest received image until a newer one arrives.
- If **`placeholder`** is set to **"image"** and no WebSocket image has been received yet, the **`default_image`** is used as the placeholder output with `IS_DEFAULT_IMAGE=True`.
**Notes:**
- This node is designed to work with the `IMAGE WebSocket Web Viewer @ vrch.ai` node, receiving the images it broadcasts.
+6 -6
View File
@@ -131,7 +131,7 @@ class VrchIntMidiControlNode:
}
}
RETURN_TYPES = ("INT", "FLOAT")
RETURN_TYPES = ("INT", "INT")
RETURN_NAMES = ("VALUE", "RAW_CC")
FUNCTION = "load_int_midi"
CATEGORY = CATEGORY
@@ -162,7 +162,7 @@ class VrchIntMidiControlNode:
if raw is None:
if debug:
print(f"[VrchIntMidiControlNode] {source}; using default value: {output_default}")
return int(output_default), 0.0
return int(output_default), 0
remap_func = VrchNodeUtils.select_remap_func(output_invert)
mapped = remap_func(float(raw), float(input_min), float(input_max), float(output_min), float(output_max))
mapped_int = int(mapped)
@@ -170,7 +170,7 @@ class VrchIntMidiControlNode:
if debug:
elapsed_ms = (time.perf_counter() - start) * 1000.0
print(f"[VrchIntMidiControlNode] {source}; mapped={mapped_int}; elapsed={elapsed_ms:.3f} ms")
return mapped_int, float(raw)
return mapped_int, int(raw)
class VrchFloatMidiControlNode:
@@ -193,7 +193,7 @@ class VrchFloatMidiControlNode:
}
}
RETURN_TYPES = ("FLOAT", "FLOAT")
RETURN_TYPES = ("FLOAT", "INT")
RETURN_NAMES = ("VALUE", "RAW_CC")
FUNCTION = "load_float_midi"
CATEGORY = CATEGORY
@@ -223,10 +223,10 @@ class VrchFloatMidiControlNode:
if raw is None:
if debug:
print(f"[VrchFloatMidiControlNode] {source}; using default value: {output_default}")
return float(output_default), 0.0
return float(output_default), 0
remap_func = VrchNodeUtils.select_remap_func(output_invert)
mapped = remap_func(float(raw), float(input_min), float(input_max), float(output_min), float(output_max))
if debug:
elapsed_ms = (time.perf_counter() - start) * 1000.0
print(f"[VrchFloatMidiControlNode] {source}; mapped={mapped}; elapsed={elapsed_ms:.3f} ms")
return float(mapped), float(raw)
return float(mapped), int(raw)
+15 -9
View File
@@ -30,13 +30,18 @@ def midi_state():
class TestMidiControlNodes(unittest.TestCase):
def test_raw_cc_output_type_is_int(self):
self.assertEqual(VrchIntMidiControlNode.RETURN_TYPES, ("INT", "INT"))
self.assertEqual(VrchFloatMidiControlNode.RETURN_TYPES, ("FLOAT", "INT"))
def test_int_workflow_key_lookup(self):
node = VrchIntMidiControlNode()
value, raw = node.load_int_midi(
midi_state(), "workflow_key", "brightness", "any", 0, 0, 127, 0, 127, False, 0, 0, False
)
self.assertEqual(value, 96)
self.assertEqual(raw, 96.0)
self.assertEqual(raw, 96)
self.assertIsInstance(raw, int)
def test_int_supports_large_output_ranges(self):
node = VrchIntMidiControlNode()
@@ -44,7 +49,7 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "workflow_key", "brightness", "any", 0, 0, 127, 0, 65535, False, 0, 0, False
)
self.assertEqual(value, int(96 / 127 * 65535))
self.assertEqual(raw, 96.0)
self.assertEqual(raw, 96)
def test_int_rounds_up_to_multiple(self):
node = VrchIntMidiControlNode()
@@ -52,7 +57,7 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "workflow_key", "brightness", "any", 0, 0, 96, 0, 510, False, 0, 64, False
)
self.assertEqual(value, 512)
self.assertEqual(raw, 96.0)
self.assertEqual(raw, 96)
def test_float_cc_number_lookup(self):
node = VrchFloatMidiControlNode()
@@ -60,7 +65,8 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "cc_number", "ignored", "1", 22, 0, 127, 0.0, 1.0, False, 0.0, False
)
self.assertAlmostEqual(value, 11 / 127)
self.assertEqual(raw, 11.0)
self.assertEqual(raw, 11)
self.assertIsInstance(raw, int)
def test_no_fallback_from_missing_key_to_valid_cc(self):
node = VrchIntMidiControlNode()
@@ -68,7 +74,7 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "workflow_key", "missing", "1", 22, 0, 127, 0, 127, False, 5, 0, False
)
self.assertEqual(value, 5)
self.assertEqual(raw, 0.0)
self.assertEqual(raw, 0)
def test_conflict_resolves_by_lookup_mode_only(self):
node = VrchIntMidiControlNode()
@@ -87,7 +93,7 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "workflow_key", "Brightness", "any", 0, 0, 127, 0, 127, False, 9, 0, False
)
self.assertEqual(value, 9)
self.assertEqual(raw, 0.0)
self.assertEqual(raw, 0)
def test_json_roundtrip_state_still_resolves_indexes(self):
state = json.loads(json.dumps(midi_state()))
@@ -99,9 +105,9 @@ class TestMidiControlNodes(unittest.TestCase):
state, "cc_number", "", "1", 22, 0, 127, 0, 127, False, 0, 0, False
)
self.assertEqual(key_value, 96)
self.assertEqual(key_raw, 96.0)
self.assertEqual(key_raw, 96)
self.assertEqual(cc_value, 11)
self.assertEqual(cc_raw, 11.0)
self.assertEqual(cc_raw, 11)
def test_reverse_mapping(self):
node = VrchIntMidiControlNode()
@@ -109,7 +115,7 @@ class TestMidiControlNodes(unittest.TestCase):
midi_state(), "workflow_key", "brightness", "any", 0, 0, 127, 0, 127, True, 0, 0, False
)
self.assertEqual(value, 31)
self.assertEqual(raw, 96.0)
self.assertEqual(raw, 96)
def test_range_validation(self):
node = VrchFloatMidiControlNode()
+45
View File
@@ -153,6 +153,51 @@ class TestWebSocketNodesUnit(unittest.TestCase):
self.assertTrue(payload["playlist"]["filename"].endswith(".webm"))
self.assertTrue(payload["playlist"]["autoplay_request"])
def test_09_image_loader_prefers_websocket_image_over_default_image(self):
received_image = torch.ones((1, 2, 2, 3), dtype=torch.float32)
class FakeClient:
def get_latest_data(self):
return received_image
original_get_client = ws_nodes.get_websocket_client
self.addCleanup(lambda: setattr(ws_nodes, "get_websocket_client", original_get_client))
ws_nodes.get_websocket_client = lambda *args, **kwargs: FakeClient()
node = ws_nodes.VrchImageWebSocketChannelLoaderNode()
default_image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
image, is_default = node.receive_image("1", "127.0.0.1:8001", "image", False, default_image)
self.assertIs(image, received_image)
self.assertFalse(is_default)
def test_10_image_loader_keeps_cached_websocket_image_for_ignored_messages(self):
received_image = torch.ones((1, 2, 2, 3), dtype=torch.float32)
class FakeClient:
def __init__(self):
self.values = [received_image, None]
def get_latest_data(self):
return self.values.pop(0) if self.values else None
fake_client = FakeClient()
original_get_client = ws_nodes.get_websocket_client
self.addCleanup(lambda: setattr(ws_nodes, "get_websocket_client", original_get_client))
ws_nodes.get_websocket_client = lambda *args, **kwargs: fake_client
node = ws_nodes.VrchImageWebSocketChannelLoaderNode()
default_image = torch.zeros((1, 2, 2, 3), dtype=torch.float32)
first_image, first_is_default = node.receive_image("1", "127.0.0.1:8001", "image", False, default_image)
second_image, second_is_default = node.receive_image("1", "127.0.0.1:8001", "image", False, default_image)
self.assertIs(first_image, received_image)
self.assertFalse(first_is_default)
self.assertIs(second_image, received_image)
self.assertFalse(second_is_default)
class TestWebSocketNodesIntegration(unittest.TestCase):
def setUp(self):
+13 -11
View File
@@ -1364,23 +1364,19 @@ class VrchImageWebSocketChannelLoaderNode:
CATEGORY = CATEGORY
def receive_image(self, channel, server, placeholder, debug, default_image=None):
if placeholder == "image" and default_image is not None:
# use tensor data_ptr to detect new image instance
cur_id = default_image.data_ptr() if hasattr(default_image, 'data_ptr') else id(default_image)
last_id = getattr(self, '_last_default_image_id', None)
if cur_id != last_id:
# update stored id and return new default_image immediately
self._last_default_image_id = cur_id
if debug:
print(f"[VrchImageWebSocketChannelLoaderNode] Detected new default_image, passing it downstream once")
return (default_image, True)
host, port = server.split(":")
cache = getattr(self, "_last_image_by_target", None)
if cache is None:
cache = {}
self._last_image_by_target = cache
cache_key = (server, str(channel))
# Ensure path is set correctly for loader
client = get_websocket_client(host, port, "/image", channel, data_handler=image_data_handler, debug=debug)
image = client.get_latest_data()
if image is not None:
cache[cache_key] = image
if debug and hasattr(image, "_metadata"):
meta = getattr(image, "_metadata", {})
print(
@@ -1393,6 +1389,12 @@ class VrchImageWebSocketChannelLoaderNode:
")",
)
return (image, False)
cached_image = cache.get(cache_key)
if cached_image is not None:
if debug:
print(f"[VrchImageWebSocketChannelLoaderNode] No new image data received, using cached websocket image")
return (cached_image, False)
# No image data, select placeholder
if placeholder == "image":