mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-03 01:07:05 +08:00
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.
This commit is contained in:
parent
1f7ddb3499
commit
bd78a3455f
69
comfy_api/feature_flags.py
Normal file
69
comfy_api/feature_flags.py
Normal file
@ -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()
|
||||||
@ -6,6 +6,8 @@ from abc import ABC
|
|||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
from comfy_execution.graph import DynamicPrompt
|
from comfy_execution.graph import DynamicPrompt
|
||||||
from protocol import BinaryEventTypes
|
from protocol import BinaryEventTypes
|
||||||
|
from comfy_api import feature_flags
|
||||||
|
|
||||||
|
|
||||||
class NodeState(Enum):
|
class NodeState(Enum):
|
||||||
Pending = "pending"
|
Pending = "pending"
|
||||||
@ -13,19 +15,23 @@ class NodeState(Enum):
|
|||||||
Finished = "finished"
|
Finished = "finished"
|
||||||
Error = "error"
|
Error = "error"
|
||||||
|
|
||||||
|
|
||||||
class NodeProgressState(TypedDict):
|
class NodeProgressState(TypedDict):
|
||||||
"""
|
"""
|
||||||
A class to represent the state of a node's progress.
|
A class to represent the state of a node's progress.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
state: NodeState
|
state: NodeState
|
||||||
value: float
|
value: float
|
||||||
max: float
|
max: float
|
||||||
|
|
||||||
|
|
||||||
class ProgressHandler(ABC):
|
class ProgressHandler(ABC):
|
||||||
"""
|
"""
|
||||||
Abstract base class for progress handlers.
|
Abstract base class for progress handlers.
|
||||||
Progress handlers receive progress updates and display them in various ways.
|
Progress handlers receive progress updates and display them in various ways.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, name: str):
|
def __init__(self, name: str):
|
||||||
self.name = name
|
self.name = name
|
||||||
self.enabled = True
|
self.enabled = True
|
||||||
@ -37,8 +43,15 @@ class ProgressHandler(ABC):
|
|||||||
"""Called when a node starts processing"""
|
"""Called when a node starts processing"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def update_handler(self, node_id: str, value: float, max_value: float,
|
def update_handler(
|
||||||
state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None):
|
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"""
|
"""Called when a node's progress is updated"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@ -58,10 +71,12 @@ class ProgressHandler(ABC):
|
|||||||
"""Disable this handler"""
|
"""Disable this handler"""
|
||||||
self.enabled = False
|
self.enabled = False
|
||||||
|
|
||||||
|
|
||||||
class CLIProgressHandler(ProgressHandler):
|
class CLIProgressHandler(ProgressHandler):
|
||||||
"""
|
"""
|
||||||
Handler that displays progress using tqdm progress bars in the CLI.
|
Handler that displays progress using tqdm progress bars in the CLI.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__("cli")
|
super().__init__("cli")
|
||||||
self.progress_bars: Dict[str, tqdm] = {}
|
self.progress_bars: Dict[str, tqdm] = {}
|
||||||
@ -75,12 +90,19 @@ class CLIProgressHandler(ProgressHandler):
|
|||||||
desc=f"Node {node_id}",
|
desc=f"Node {node_id}",
|
||||||
unit="steps",
|
unit="steps",
|
||||||
leave=True,
|
leave=True,
|
||||||
position=len(self.progress_bars)
|
position=len(self.progress_bars),
|
||||||
)
|
)
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def update_handler(self, node_id: str, value: float, max_value: float,
|
def update_handler(
|
||||||
state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None):
|
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
|
# Handle case where start_handler wasn't called
|
||||||
if node_id not in self.progress_bars:
|
if node_id not in self.progress_bars:
|
||||||
self.progress_bars[node_id] = tqdm(
|
self.progress_bars[node_id] = tqdm(
|
||||||
@ -88,7 +110,7 @@ class CLIProgressHandler(ProgressHandler):
|
|||||||
desc=f"Node {node_id}",
|
desc=f"Node {node_id}",
|
||||||
unit="steps",
|
unit="steps",
|
||||||
leave=True,
|
leave=True,
|
||||||
position=len(self.progress_bars)
|
position=len(self.progress_bars),
|
||||||
)
|
)
|
||||||
self.progress_bars[node_id].update(value)
|
self.progress_bars[node_id].update(value)
|
||||||
else:
|
else:
|
||||||
@ -119,10 +141,12 @@ class CLIProgressHandler(ProgressHandler):
|
|||||||
bar.close()
|
bar.close()
|
||||||
self.progress_bars.clear()
|
self.progress_bars.clear()
|
||||||
|
|
||||||
|
|
||||||
class WebUIProgressHandler(ProgressHandler):
|
class WebUIProgressHandler(ProgressHandler):
|
||||||
"""
|
"""
|
||||||
Handler that sends progress updates to the WebUI via WebSockets.
|
Handler that sends progress updates to the WebUI via WebSockets.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, server_instance):
|
def __init__(self, server_instance):
|
||||||
super().__init__("webui")
|
super().__init__("webui")
|
||||||
self.server_instance = server_instance
|
self.server_instance = server_instance
|
||||||
@ -145,17 +169,16 @@ class WebUIProgressHandler(ProgressHandler):
|
|||||||
"prompt_id": prompt_id,
|
"prompt_id": prompt_id,
|
||||||
"display_node_id": self.registry.dynprompt.get_display_node_id(node_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),
|
"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()
|
for node_id, state in nodes.items()
|
||||||
if state["state"] != NodeState.Pending
|
if state["state"] != NodeState.Pending
|
||||||
}
|
}
|
||||||
|
|
||||||
# Send a combined progress_state message with all node states
|
# Send a combined progress_state message with all node states
|
||||||
self.server_instance.send_sync("progress_state", {
|
self.server_instance.send_sync(
|
||||||
"prompt_id": prompt_id,
|
"progress_state", {"prompt_id": prompt_id, "nodes": active_nodes}
|
||||||
"nodes": active_nodes
|
)
|
||||||
})
|
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def start_handler(self, node_id: str, state: NodeProgressState, prompt_id: str):
|
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)
|
self._send_progress_state(prompt_id, self.registry.nodes)
|
||||||
|
|
||||||
@override
|
@override
|
||||||
def update_handler(self, node_id: str, value: float, max_value: float,
|
def update_handler(
|
||||||
state: NodeProgressState, prompt_id: str, image: Optional[Image.Image] = None):
|
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
|
# Send progress state of all nodes
|
||||||
if self.registry:
|
if self.registry:
|
||||||
self._send_progress_state(prompt_id, self.registry.nodes)
|
self._send_progress_state(prompt_id, self.registry.nodes)
|
||||||
if image:
|
if image:
|
||||||
metadata = {
|
# Only send new format if client supports it
|
||||||
"node_id": node_id,
|
if feature_flags.supports_feature(
|
||||||
"prompt_id": prompt_id,
|
self.server_instance.sockets_metadata,
|
||||||
"display_node_id": self.registry.dynprompt.get_display_node_id(node_id),
|
self.server_instance.client_id,
|
||||||
"parent_node_id": self.registry.dynprompt.get_parent_node_id(node_id),
|
"supports_preview_metadata",
|
||||||
"real_node_id": self.registry.dynprompt.get_real_node_id(node_id)
|
):
|
||||||
}
|
metadata = {
|
||||||
self.server_instance.send_sync(BinaryEventTypes.PREVIEW_IMAGE_WITH_METADATA, (image, metadata), self.server_instance.client_id)
|
"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
|
@override
|
||||||
def finish_handler(self, node_id: str, state: NodeProgressState, prompt_id: str):
|
def finish_handler(self, node_id: str, state: NodeProgressState, prompt_id: str):
|
||||||
@ -186,10 +229,12 @@ class WebUIProgressHandler(ProgressHandler):
|
|||||||
if self.registry:
|
if self.registry:
|
||||||
self._send_progress_state(prompt_id, self.registry.nodes)
|
self._send_progress_state(prompt_id, self.registry.nodes)
|
||||||
|
|
||||||
|
|
||||||
class ProgressRegistry:
|
class ProgressRegistry:
|
||||||
"""
|
"""
|
||||||
Registry that maintains node progress state and notifies registered handlers.
|
Registry that maintains node progress state and notifies registered handlers.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, prompt_id: str, dynprompt: DynamicPrompt):
|
def __init__(self, prompt_id: str, dynprompt: DynamicPrompt):
|
||||||
self.prompt_id = prompt_id
|
self.prompt_id = prompt_id
|
||||||
self.dynprompt = dynprompt
|
self.dynprompt = dynprompt
|
||||||
@ -221,9 +266,7 @@ class ProgressRegistry:
|
|||||||
"""Ensure a node entry exists"""
|
"""Ensure a node entry exists"""
|
||||||
if node_id not in self.nodes:
|
if node_id not in self.nodes:
|
||||||
self.nodes[node_id] = NodeProgressState(
|
self.nodes[node_id] = NodeProgressState(
|
||||||
state = NodeState.Pending,
|
state=NodeState.Pending, value=0, max=1
|
||||||
value = 0,
|
|
||||||
max = 1
|
|
||||||
)
|
)
|
||||||
return self.nodes[node_id]
|
return self.nodes[node_id]
|
||||||
|
|
||||||
@ -239,7 +282,9 @@ class ProgressRegistry:
|
|||||||
if handler.enabled:
|
if handler.enabled:
|
||||||
handler.start_handler(node_id, entry, self.prompt_id)
|
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"""
|
"""Update progress for a node"""
|
||||||
entry = self.ensure_entry(node_id)
|
entry = self.ensure_entry(node_id)
|
||||||
entry["state"] = NodeState.Running
|
entry["state"] = NodeState.Running
|
||||||
@ -249,7 +294,9 @@ class ProgressRegistry:
|
|||||||
# Notify all enabled handlers
|
# Notify all enabled handlers
|
||||||
for handler in self.handlers.values():
|
for handler in self.handlers.values():
|
||||||
if handler.enabled:
|
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:
|
def finish_progress(self, node_id: str) -> None:
|
||||||
"""Finish progress tracking for a node"""
|
"""Finish progress tracking for a node"""
|
||||||
@ -267,8 +314,12 @@ class ProgressRegistry:
|
|||||||
for handler in self.handlers.values():
|
for handler in self.handlers.values():
|
||||||
handler.reset()
|
handler.reset()
|
||||||
|
|
||||||
|
|
||||||
# Global registry instance
|
# 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:
|
def reset_progress_state(prompt_id: str, dynprompt: DynamicPrompt) -> None:
|
||||||
global global_progress_registry
|
global global_progress_registry
|
||||||
@ -280,9 +331,11 @@ def reset_progress_state(prompt_id: str, dynprompt: DynamicPrompt) -> None:
|
|||||||
# Create new registry
|
# Create new registry
|
||||||
global_progress_registry = ProgressRegistry(prompt_id, dynprompt)
|
global_progress_registry = ProgressRegistry(prompt_id, dynprompt)
|
||||||
|
|
||||||
|
|
||||||
def add_progress_handler(handler: ProgressHandler) -> None:
|
def add_progress_handler(handler: ProgressHandler) -> None:
|
||||||
handler.set_registry(global_progress_registry)
|
handler.set_registry(global_progress_registry)
|
||||||
global_progress_registry.register_handler(handler)
|
global_progress_registry.register_handler(handler)
|
||||||
|
|
||||||
|
|
||||||
def get_progress_state() -> ProgressRegistry:
|
def get_progress_state() -> ProgressRegistry:
|
||||||
return global_progress_registry
|
return global_progress_registry
|
||||||
|
|||||||
15
main.py
15
main.py
@ -13,6 +13,7 @@ import logging
|
|||||||
import sys
|
import sys
|
||||||
from comfy_execution.progress import get_progress_state
|
from comfy_execution.progress import get_progress_state
|
||||||
from comfy_execution.utils import get_executing_context
|
from comfy_execution.utils import get_executing_context
|
||||||
|
from comfy_api import feature_flags
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
#NOTE: These do not do anything on core ComfyUI, they are for custom nodes.
|
#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)
|
server_instance.send_sync("progress", progress, server_instance.client_id)
|
||||||
if preview_image is not None:
|
if preview_image is not None:
|
||||||
# Also send old method for backward compatibility
|
# Only send old method if client doesn't support preview metadata
|
||||||
# TODO - Remove after this repo is updated to frontend with metadata support
|
if not feature_flags.supports_feature(
|
||||||
server_instance.send_sync(BinaryEventTypes.UNENCODED_PREVIEW_IMAGE, preview_image, server_instance.client_id)
|
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)
|
comfy.utils.set_progress_bar_global_hook(hook)
|
||||||
|
|
||||||
|
|||||||
47
server.py
47
server.py
@ -26,6 +26,7 @@ import mimetypes
|
|||||||
from comfy.cli_args import args
|
from comfy.cli_args import args
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
from comfy_api import feature_flags
|
||||||
import node_helpers
|
import node_helpers
|
||||||
from comfyui_version import __version__
|
from comfyui_version import __version__
|
||||||
from app.frontend_management import FrontendManager
|
from app.frontend_management import FrontendManager
|
||||||
@ -174,6 +175,7 @@ class PromptServer():
|
|||||||
max_upload_size = round(args.max_upload_size * 1024 * 1024)
|
max_upload_size = round(args.max_upload_size * 1024 * 1024)
|
||||||
self.app = web.Application(client_max_size=max_upload_size, middlewares=middlewares)
|
self.app = web.Application(client_max_size=max_upload_size, middlewares=middlewares)
|
||||||
self.sockets = dict()
|
self.sockets = dict()
|
||||||
|
self.sockets_metadata = dict()
|
||||||
self.web_root = (
|
self.web_root = (
|
||||||
FrontendManager.init_frontend(args.front_end_version)
|
FrontendManager.init_frontend(args.front_end_version)
|
||||||
if args.front_end_root is None
|
if args.front_end_root is None
|
||||||
@ -198,20 +200,53 @@ class PromptServer():
|
|||||||
else:
|
else:
|
||||||
sid = uuid.uuid4().hex
|
sid = uuid.uuid4().hex
|
||||||
|
|
||||||
|
# Store WebSocket for backward compatibility
|
||||||
self.sockets[sid] = ws
|
self.sockets[sid] = ws
|
||||||
|
# Store metadata separately
|
||||||
|
self.sockets_metadata[sid] = {"feature_flags": {}}
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Send initial state to the new client
|
# 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
|
# 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:
|
if self.client_id == sid and self.last_node_id is not None:
|
||||||
await self.send("executing", { "node": self.last_node_id }, sid)
|
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:
|
async for msg in ws:
|
||||||
if msg.type == aiohttp.WSMsgType.ERROR:
|
if msg.type == aiohttp.WSMsgType.ERROR:
|
||||||
logging.warning('ws connection closed with exception %s' % ws.exception())
|
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:
|
finally:
|
||||||
self.sockets.pop(sid, None)
|
self.sockets.pop(sid, None)
|
||||||
|
self.sockets_metadata.pop(sid, None)
|
||||||
return ws
|
return ws
|
||||||
|
|
||||||
@routes.get("/")
|
@routes.get("/")
|
||||||
@ -544,6 +579,10 @@ class PromptServer():
|
|||||||
}
|
}
|
||||||
return web.json_response(system_stats)
|
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")
|
@routes.get("/prompt")
|
||||||
async def get_prompt(request):
|
async def get_prompt(request):
|
||||||
return web.json_response(self.get_queue_info())
|
return web.json_response(self.get_queue_info())
|
||||||
@ -882,10 +921,10 @@ class PromptServer():
|
|||||||
ssl_ctx = None
|
ssl_ctx = None
|
||||||
scheme = "http"
|
scheme = "http"
|
||||||
if args.tls_keyfile and args.tls_certfile:
|
if args.tls_keyfile and args.tls_certfile:
|
||||||
ssl_ctx = ssl.SSLContext(protocol=ssl.PROTOCOL_TLS_SERVER, verify_mode=ssl.CERT_NONE)
|
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.load_cert_chain(certfile=args.tls_certfile,
|
||||||
keyfile=args.tls_keyfile)
|
keyfile=args.tls_keyfile)
|
||||||
scheme = "https"
|
scheme = "https"
|
||||||
|
|
||||||
if verbose:
|
if verbose:
|
||||||
logging.info("Starting server\n")
|
logging.info("Starting server\n")
|
||||||
|
|||||||
98
tests-unit/feature_flags_test.py
Normal file
98
tests-unit/feature_flags_test.py
Normal file
@ -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
|
||||||
128
tests-unit/websocket_feature_flags_test.py
Normal file
128
tests-unit/websocket_feature_flags_test.py
Normal file
@ -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
|
||||||
Loading…
x
Reference in New Issue
Block a user