From b605fe065bd65d8ffffecbeaa126d95ca03dd26d Mon Sep 17 00:00:00 2001 From: tianzi Date: Fri, 29 May 2026 17:23:31 +0100 Subject: [PATCH] fix: stabilize websocket image fallback and raw MIDI output --- CHANGELOG.md | 10 ++++++ docs/midi_control_nodes.md | 2 +- docs/websocket_nodes.md | 6 ++-- nodes/midi_control_nodes.py | 12 +++---- nodes/tests/midi_control_nodes_test.py | 24 ++++++++------ nodes/tests/websocket_nodes_test.py | 45 ++++++++++++++++++++++++++ nodes/websocket_nodes.py | 24 +++++++------- 7 files changed, 93 insertions(+), 30 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index d2603be..d90bf94 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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 diff --git a/docs/midi_control_nodes.md b/docs/midi_control_nodes.md index 0ae62b2..3b2f551 100644 --- a/docs/midi_control_nodes.md +++ b/docs/midi_control_nodes.md @@ -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. diff --git a/docs/websocket_nodes.md b/docs/websocket_nodes.md index ada64cd..d0616cd 100644 --- a/docs/websocket_nodes.md +++ b/docs/websocket_nodes.md @@ -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. diff --git a/nodes/midi_control_nodes.py b/nodes/midi_control_nodes.py index 1874d89..e3d7433 100644 --- a/nodes/midi_control_nodes.py +++ b/nodes/midi_control_nodes.py @@ -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) diff --git a/nodes/tests/midi_control_nodes_test.py b/nodes/tests/midi_control_nodes_test.py index 0949e6a..2f15678 100644 --- a/nodes/tests/midi_control_nodes_test.py +++ b/nodes/tests/midi_control_nodes_test.py @@ -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() diff --git a/nodes/tests/websocket_nodes_test.py b/nodes/tests/websocket_nodes_test.py index f8cbe74..5b80d67 100644 --- a/nodes/tests/websocket_nodes_test.py +++ b/nodes/tests/websocket_nodes_test.py @@ -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): diff --git a/nodes/websocket_nodes.py b/nodes/websocket_nodes.py index ef13599..6bb51be 100644 --- a/nodes/websocket_nodes.py +++ b/nodes/websocket_nodes.py @@ -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":