From bd78a3455fc87eb59a393ea3ebca88e67111339f Mon Sep 17 00:00:00 2001 From: Jacob Segal Date: Tue, 8 Jul 2025 00:14:20 -0700 Subject: [PATCH] Add a feature flags message to reduce bandwidth We now only send 1 preview message of the latest type the client can support. We'll add a console warning when the client fails to send a feature flags message at some point in the future. --- comfy_api/feature_flags.py | 69 +++++++++++ comfy_execution/progress.py | 109 +++++++++++++----- main.py | 15 ++- server.py | 47 +++++++- tests-unit/feature_flags_test.py | 98 ++++++++++++++++ tests-unit/websocket_feature_flags_test.py | 128 +++++++++++++++++++++ 6 files changed, 431 insertions(+), 35 deletions(-) create mode 100644 comfy_api/feature_flags.py create mode 100644 tests-unit/feature_flags_test.py create mode 100644 tests-unit/websocket_feature_flags_test.py diff --git a/comfy_api/feature_flags.py b/comfy_api/feature_flags.py new file mode 100644 index 000000000..0d4389a6e --- /dev/null +++ b/comfy_api/feature_flags.py @@ -0,0 +1,69 @@ +""" +Feature flags module for ComfyUI WebSocket protocol negotiation. + +This module handles capability negotiation between frontend and backend, +allowing graceful protocol evolution while maintaining backward compatibility. +""" + +from typing import Any, Dict + +from comfy.cli_args import args + +# Default server capabilities +SERVER_FEATURE_FLAGS: Dict[str, Any] = { + "supports_preview_metadata": True, + "max_upload_size": args.max_upload_size * 1024 * 1024, # Convert MB to bytes +} + + +def get_connection_feature( + sockets_metadata: Dict[str, Dict[str, Any]], + sid: str, + feature_name: str, + default: Any = False +) -> Any: + """ + Get a feature flag value for a specific connection. + + Args: + sockets_metadata: Dictionary of socket metadata + sid: Session ID of the connection + feature_name: Name of the feature to check + default: Default value if feature not found + + Returns: + Feature value or default if not found + """ + if sid not in sockets_metadata: + return default + + return sockets_metadata[sid].get("feature_flags", {}).get(feature_name, default) + + +def supports_feature( + sockets_metadata: Dict[str, Dict[str, Any]], + sid: str, + feature_name: str +) -> bool: + """ + Check if a connection supports a specific feature. + + Args: + sockets_metadata: Dictionary of socket metadata + sid: Session ID of the connection + feature_name: Name of the feature to check + + Returns: + Boolean indicating if feature is supported + """ + return get_connection_feature(sockets_metadata, sid, feature_name, False) is True + + +def get_server_features() -> Dict[str, Any]: + """ + Get the server's feature flags. + + Returns: + Dictionary of server feature flags + """ + return SERVER_FEATURE_FLAGS.copy() diff --git a/comfy_execution/progress.py b/comfy_execution/progress.py index 3a14969dd..68e7d792d 100644 --- a/comfy_execution/progress.py +++ b/comfy_execution/progress.py @@ -6,6 +6,8 @@ from abc import ABC from tqdm import tqdm from comfy_execution.graph import DynamicPrompt from protocol import BinaryEventTypes +from comfy_api import feature_flags + class NodeState(Enum): Pending = "pending" @@ -13,19 +15,23 @@ class NodeState(Enum): Finished = "finished" Error = "error" + class NodeProgressState(TypedDict): """ A class to represent the state of a node's progress. """ + state: NodeState value: float max: float + class ProgressHandler(ABC): """ Abstract base class for progress handlers. Progress handlers receive progress updates and display them in various ways. """ + def __init__(self, name: str): self.name = name self.enabled = True @@ -37,8 +43,15 @@ class ProgressHandler(ABC): """Called when a node starts processing""" pass - def update_handler(self, node_id: str, value: float, max_value: float, - state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None): + def update_handler( + self, + node_id: str, + value: float, + max_value: float, + state: NodeProgressState, + prompt_id: str, + image: Optional[Image.Image] = None, + ): """Called when a node's progress is updated""" pass @@ -58,10 +71,12 @@ class ProgressHandler(ABC): """Disable this handler""" self.enabled = False + class CLIProgressHandler(ProgressHandler): """ Handler that displays progress using tqdm progress bars in the CLI. """ + def __init__(self): super().__init__("cli") self.progress_bars: Dict[str, tqdm] = {} @@ -75,12 +90,19 @@ class CLIProgressHandler(ProgressHandler): desc=f"Node {node_id}", unit="steps", leave=True, - position=len(self.progress_bars) + position=len(self.progress_bars), ) @override - def update_handler(self, node_id: str, value: float, max_value: float, - state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None): + def update_handler( + self, + node_id: str, + value: float, + max_value: float, + state: NodeProgressState, + prompt_id: str, + image: Optional[Image.Image] = None, + ): # Handle case where start_handler wasn't called if node_id not in self.progress_bars: self.progress_bars[node_id] = tqdm( @@ -88,7 +110,7 @@ class CLIProgressHandler(ProgressHandler): desc=f"Node {node_id}", unit="steps", leave=True, - position=len(self.progress_bars) + position=len(self.progress_bars), ) self.progress_bars[node_id].update(value) else: @@ -119,10 +141,12 @@ class CLIProgressHandler(ProgressHandler): bar.close() self.progress_bars.clear() + class WebUIProgressHandler(ProgressHandler): """ Handler that sends progress updates to the WebUI via WebSockets. """ + def __init__(self, server_instance): super().__init__("webui") self.server_instance = server_instance @@ -145,17 +169,16 @@ class WebUIProgressHandler(ProgressHandler): "prompt_id": prompt_id, "display_node_id": self.registry.dynprompt.get_display_node_id(node_id), "parent_node_id": self.registry.dynprompt.get_parent_node_id(node_id), - "real_node_id": self.registry.dynprompt.get_real_node_id(node_id) + "real_node_id": self.registry.dynprompt.get_real_node_id(node_id), } for node_id, state in nodes.items() if state["state"] != NodeState.Pending } # Send a combined progress_state message with all node states - self.server_instance.send_sync("progress_state", { - "prompt_id": prompt_id, - "nodes": active_nodes - }) + self.server_instance.send_sync( + "progress_state", {"prompt_id": prompt_id, "nodes": active_nodes} + ) @override def start_handler(self, node_id: str, state: NodeProgressState, prompt_id: str): @@ -164,21 +187,41 @@ class WebUIProgressHandler(ProgressHandler): self._send_progress_state(prompt_id, self.registry.nodes) @override - def update_handler(self, node_id: str, value: float, max_value: float, - state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None): + def update_handler( + self, + node_id: str, + value: float, + max_value: float, + state: NodeProgressState, + prompt_id: str, + image: Optional[Image.Image] = None, + ): # Send progress state of all nodes if self.registry: self._send_progress_state(prompt_id, self.registry.nodes) if image: - metadata = { - "node_id": node_id, - "prompt_id": prompt_id, - "display_node_id": self.registry.dynprompt.get_display_node_id(node_id), - "parent_node_id": self.registry.dynprompt.get_parent_node_id(node_id), - "real_node_id": self.registry.dynprompt.get_real_node_id(node_id) - } - self.server_instance.send_sync(BinaryEventTypes.PREVIEW_IMAGE_WITH_METADATA, (image, metadata), self.server_instance.client_id) - + # Only send new format if client supports it + if feature_flags.supports_feature( + self.server_instance.sockets_metadata, + self.server_instance.client_id, + "supports_preview_metadata", + ): + metadata = { + "node_id": node_id, + "prompt_id": prompt_id, + "display_node_id": self.registry.dynprompt.get_display_node_id( + node_id + ), + "parent_node_id": self.registry.dynprompt.get_parent_node_id( + node_id + ), + "real_node_id": self.registry.dynprompt.get_real_node_id(node_id), + } + self.server_instance.send_sync( + BinaryEventTypes.PREVIEW_IMAGE_WITH_METADATA, + (image, metadata), + self.server_instance.client_id, + ) @override def finish_handler(self, node_id: str, state: NodeProgressState, prompt_id: str): @@ -186,10 +229,12 @@ class WebUIProgressHandler(ProgressHandler): if self.registry: self._send_progress_state(prompt_id, self.registry.nodes) + class ProgressRegistry: """ Registry that maintains node progress state and notifies registered handlers. """ + def __init__(self, prompt_id: str, dynprompt: DynamicPrompt): self.prompt_id = prompt_id self.dynprompt = dynprompt @@ -221,9 +266,7 @@ class ProgressRegistry: """Ensure a node entry exists""" if node_id not in self.nodes: self.nodes[node_id] = NodeProgressState( - state = NodeState.Pending, - value = 0, - max = 1 + state=NodeState.Pending, value=0, max=1 ) return self.nodes[node_id] @@ -239,7 +282,9 @@ class ProgressRegistry: if handler.enabled: handler.start_handler(node_id, entry, self.prompt_id) - def update_progress(self, node_id: str, value: float, max_value: float, image: Optional[Image.Image]) -> None: + def update_progress( + self, node_id: str, value: float, max_value: float, image: Optional[Image.Image] + ) -> None: """Update progress for a node""" entry = self.ensure_entry(node_id) entry["state"] = NodeState.Running @@ -249,7 +294,9 @@ class ProgressRegistry: # Notify all enabled handlers for handler in self.handlers.values(): if handler.enabled: - handler.update_handler(node_id, value, max_value, entry, self.prompt_id, image) + handler.update_handler( + node_id, value, max_value, entry, self.prompt_id, image + ) def finish_progress(self, node_id: str) -> None: """Finish progress tracking for a node""" @@ -267,8 +314,12 @@ class ProgressRegistry: for handler in self.handlers.values(): handler.reset() + # Global registry instance -global_progress_registry: ProgressRegistry = ProgressRegistry(prompt_id="", dynprompt=DynamicPrompt({})) +global_progress_registry: ProgressRegistry = ProgressRegistry( + prompt_id="", dynprompt=DynamicPrompt({}) +) + def reset_progress_state(prompt_id: str, dynprompt: DynamicPrompt) -> None: global global_progress_registry @@ -280,9 +331,11 @@ def reset_progress_state(prompt_id: str, dynprompt: DynamicPrompt) -> None: # Create new registry global_progress_registry = ProgressRegistry(prompt_id, dynprompt) + def add_progress_handler(handler: ProgressHandler) -> None: handler.set_registry(global_progress_registry) global_progress_registry.register_handler(handler) + def get_progress_state() -> ProgressRegistry: return global_progress_registry diff --git a/main.py b/main.py index be2cc6cf0..80416f055 100644 --- a/main.py +++ b/main.py @@ -13,6 +13,7 @@ import logging import sys from comfy_execution.progress import get_progress_state from comfy_execution.utils import get_executing_context +from comfy_api import feature_flags if __name__ == "__main__": #NOTE: These do not do anything on core ComfyUI, they are for custom nodes. @@ -246,9 +247,17 @@ def hijack_progress(server_instance): server_instance.send_sync("progress", progress, server_instance.client_id) if preview_image is not None: - # Also send old method for backward compatibility - # TODO - Remove after this repo is updated to frontend with metadata support - server_instance.send_sync(BinaryEventTypes.UNENCODED_PREVIEW_IMAGE, preview_image, server_instance.client_id) + # Only send old method if client doesn't support preview metadata + if not feature_flags.supports_feature( + server_instance.sockets_metadata, + server_instance.client_id, + "supports_preview_metadata", + ): + server_instance.send_sync( + BinaryEventTypes.UNENCODED_PREVIEW_IMAGE, + preview_image, + server_instance.client_id, + ) comfy.utils.set_progress_bar_global_hook(hook) diff --git a/server.py b/server.py index b86f122ba..e8bad9f4e 100644 --- a/server.py +++ b/server.py @@ -26,6 +26,7 @@ import mimetypes from comfy.cli_args import args import comfy.utils import comfy.model_management +from comfy_api import feature_flags import node_helpers from comfyui_version import __version__ from app.frontend_management import FrontendManager @@ -174,6 +175,7 @@ class PromptServer(): max_upload_size = round(args.max_upload_size * 1024 * 1024) self.app = web.Application(client_max_size=max_upload_size, middlewares=middlewares) self.sockets = dict() + self.sockets_metadata = dict() self.web_root = ( FrontendManager.init_frontend(args.front_end_version) if args.front_end_root is None @@ -198,20 +200,53 @@ class PromptServer(): else: sid = uuid.uuid4().hex + # Store WebSocket for backward compatibility self.sockets[sid] = ws + # Store metadata separately + self.sockets_metadata[sid] = {"feature_flags": {}} try: # Send initial state to the new client - await self.send("status", { "status": self.get_queue_info(), 'sid': sid }, sid) + await self.send("status", {"status": self.get_queue_info(), "sid": sid}, sid) # On reconnect if we are the currently executing client send the current node if self.client_id == sid and self.last_node_id is not None: await self.send("executing", { "node": self.last_node_id }, sid) + # Flag to track if we've received the first message + first_message = True + async for msg in ws: if msg.type == aiohttp.WSMsgType.ERROR: logging.warning('ws connection closed with exception %s' % ws.exception()) + elif msg.type == aiohttp.WSMsgType.TEXT: + try: + data = json.loads(msg.data) + # Check if first message is feature flags + if first_message and data.get("type") == "feature_flags": + # Store client feature flags + client_flags = data.get("data", {}) + self.sockets_metadata[sid]["feature_flags"] = client_flags + + # Send server feature flags in response + await self.send( + "feature_flags", + feature_flags.get_server_features(), + sid, + ) + + logging.info( + f"Feature flags negotiated for client {sid}: {client_flags}" + ) + first_message = False + except json.JSONDecodeError: + logging.warning( + f"Invalid JSON received from client {sid}: {msg.data}" + ) + except Exception as e: + logging.error(f"Error processing WebSocket message: {e}") finally: self.sockets.pop(sid, None) + self.sockets_metadata.pop(sid, None) return ws @routes.get("/") @@ -544,6 +579,10 @@ class PromptServer(): } return web.json_response(system_stats) + @routes.get("/features") + async def get_features(request): + return web.json_response(feature_flags.get_server_features()) + @routes.get("/prompt") async def get_prompt(request): return web.json_response(self.get_queue_info()) @@ -882,10 +921,10 @@ class PromptServer(): ssl_ctx = None scheme = "http" if args.tls_keyfile and args.tls_certfile: - ssl_ctx = ssl.SSLContext(protocol=ssl.PROTOCOL_TLS_SERVER, verify_mode=ssl.CERT_NONE) - ssl_ctx.load_cert_chain(certfile=args.tls_certfile, + ssl_ctx = ssl.SSLContext(protocol=ssl.PROTOCOL_TLS_SERVER, verify_mode=ssl.CERT_NONE) + ssl_ctx.load_cert_chain(certfile=args.tls_certfile, keyfile=args.tls_keyfile) - scheme = "https" + scheme = "https" if verbose: logging.info("Starting server\n") diff --git a/tests-unit/feature_flags_test.py b/tests-unit/feature_flags_test.py new file mode 100644 index 000000000..f2702cfc8 --- /dev/null +++ b/tests-unit/feature_flags_test.py @@ -0,0 +1,98 @@ +"""Tests for feature flags functionality.""" + +from comfy_api.feature_flags import ( + get_connection_feature, + supports_feature, + get_server_features, + SERVER_FEATURE_FLAGS, +) + + +class TestFeatureFlags: + """Test suite for feature flags functions.""" + + def test_get_server_features_returns_copy(self): + """Test that get_server_features returns a copy of the server flags.""" + features = get_server_features() + # Verify it's a copy by modifying it + features["test_flag"] = True + # Original should be unchanged + assert "test_flag" not in SERVER_FEATURE_FLAGS + + def test_get_server_features_contains_expected_flags(self): + """Test that server features contain expected flags.""" + features = get_server_features() + assert "supports_preview_metadata" in features + assert features["supports_preview_metadata"] is True + assert "max_upload_size" in features + assert isinstance(features["max_upload_size"], (int, float)) + + def test_get_connection_feature_with_missing_sid(self): + """Test getting feature for non-existent session ID.""" + sockets_metadata = {} + result = get_connection_feature(sockets_metadata, "missing_sid", "some_feature") + assert result is False # Default value + + def test_get_connection_feature_with_custom_default(self): + """Test getting feature with custom default value.""" + sockets_metadata = {} + result = get_connection_feature( + sockets_metadata, "missing_sid", "some_feature", default="custom_default" + ) + assert result == "custom_default" + + def test_get_connection_feature_with_feature_flags(self): + """Test getting feature from connection with feature flags.""" + sockets_metadata = { + "sid1": { + "feature_flags": { + "supports_preview_metadata": True, + "custom_feature": "value", + }, + } + } + result = get_connection_feature(sockets_metadata, "sid1", "supports_preview_metadata") + assert result is True + + result = get_connection_feature(sockets_metadata, "sid1", "custom_feature") + assert result == "value" + + def test_get_connection_feature_missing_feature(self): + """Test getting non-existent feature from connection.""" + sockets_metadata = { + "sid1": {"feature_flags": {"existing_feature": True}} + } + result = get_connection_feature(sockets_metadata, "sid1", "missing_feature") + assert result is False + + def test_supports_feature_returns_boolean(self): + """Test that supports_feature always returns boolean.""" + sockets_metadata = { + "sid1": { + "feature_flags": { + "bool_feature": True, + "string_feature": "value", + "none_feature": None, + }, + } + } + + # True boolean feature + assert supports_feature(sockets_metadata, "sid1", "bool_feature") is True + + # Non-boolean values should return False + assert supports_feature(sockets_metadata, "sid1", "string_feature") is False + assert supports_feature(sockets_metadata, "sid1", "none_feature") is False + assert supports_feature(sockets_metadata, "sid1", "missing_feature") is False + + def test_supports_feature_with_missing_connection(self): + """Test supports_feature with missing connection.""" + sockets_metadata = {} + assert supports_feature(sockets_metadata, "missing_sid", "any_feature") is False + + def test_empty_feature_flags_dict(self): + """Test connection with empty feature flags dictionary.""" + sockets_metadata = {"sid1": {"feature_flags": {}}} + result = get_connection_feature(sockets_metadata, "sid1", "any_feature") + assert result is False + assert supports_feature(sockets_metadata, "sid1", "any_feature") is False diff --git a/tests-unit/websocket_feature_flags_test.py b/tests-unit/websocket_feature_flags_test.py new file mode 100644 index 000000000..b8d7ee258 --- /dev/null +++ b/tests-unit/websocket_feature_flags_test.py @@ -0,0 +1,128 @@ +"""Simplified tests for WebSocket feature flags functionality.""" +from unittest.mock import Mock, patch +from comfy_api import feature_flags + + +class TestWebSocketFeatureFlags: + """Test suite for WebSocket feature flags integration.""" + + @patch('main.server') + def test_preview_message_with_feature_support(self, mock_server): + """Test that UNENCODED_PREVIEW_IMAGE is not sent when client supports metadata.""" + # Setup mock server with client that supports preview metadata + mock_server.client_id = "test_client" + mock_server.sockets_metadata = { + "test_client": { + "feature_flags": {"supports_preview_metadata": True} + } + } + mock_server.send_sync = Mock() + + # Import the function we're testing + from main import hijack_progress + + # Call hijack_progress + hijack_progress(mock_server) + + # Verify feature flags can check support + assert feature_flags.supports_feature( + mock_server.sockets_metadata, + "test_client", + "supports_preview_metadata" + ) is True + + @patch('main.server') + def test_preview_message_without_feature_support(self, mock_server): + """Test that UNENCODED_PREVIEW_IMAGE is sent when client doesn't support metadata.""" + # Setup mock server with legacy client + mock_server.client_id = "legacy_client" + mock_server.sockets_metadata = { + "legacy_client": { + "feature_flags": {} # No features + } + } + mock_server.send_sync = Mock() + + # Import the function we're testing + from main import hijack_progress + + # Call hijack_progress + hijack_progress(mock_server) + + # Verify feature flags check returns False + assert feature_flags.supports_feature( + mock_server.sockets_metadata, + "legacy_client", + "supports_preview_metadata" + ) is False + + def test_server_feature_flags_response(self): + """Test server feature flags are properly formatted.""" + features = feature_flags.get_server_features() + + # Check expected server features + assert "supports_preview_metadata" in features + assert features["supports_preview_metadata"] is True + assert "max_upload_size" in features + assert isinstance(features["max_upload_size"], (int, float)) + + def test_progress_py_checks_feature_flags(self): + """Test that progress.py checks feature flags before sending metadata.""" + # This simulates the check in progress.py + client_id = "test_client" + sockets_metadata = {"test_client": {"feature_flags": {}}} + + # The actual check would be in progress.py + supports_metadata = feature_flags.supports_feature( + sockets_metadata, client_id, "supports_preview_metadata" + ) + + assert supports_metadata is False + + def test_multiple_clients_different_features(self): + """Test handling multiple clients with different feature support.""" + sockets_metadata = { + "modern_client": { + "feature_flags": {"supports_preview_metadata": True} + }, + "legacy_client": { + "feature_flags": {} + } + } + + # Check modern client + assert feature_flags.supports_feature( + sockets_metadata, "modern_client", "supports_preview_metadata" + ) is True + + # Check legacy client + assert feature_flags.supports_feature( + sockets_metadata, "legacy_client", "supports_preview_metadata" + ) is False + + def test_feature_negotiation_message_format(self): + """Test the format of feature negotiation messages.""" + # Client message format + client_message = { + "type": "feature_flags", + "data": { + "supports_preview_metadata": True, + "api_version": "1.0.0" + } + } + + # Verify structure + assert client_message["type"] == "feature_flags" + assert "supports_preview_metadata" in client_message["data"] + + # Server response format (what would be sent) + server_features = feature_flags.get_server_features() + server_message = { + "type": "feature_flags", + "data": server_features + } + + # Verify structure + assert server_message["type"] == "feature_flags" + assert "supports_preview_metadata" in server_message["data"] + assert server_message["data"]["supports_preview_metadata"] is True