Compare commits
21
Commits
@@ -1,4 +1,3 @@
|
||||
# These are supported funding model platforms
|
||||
|
||||
github: robertvoy
|
||||
buy_me_a_coffee: robertvoy
|
||||
|
||||
@@ -19,6 +19,6 @@ jobs:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
with:
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -147,7 +147,7 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
|
||||
Control your distributed cluster programmatically without opening the browser.
|
||||
|
||||
* **Endpoint:** `POST /distributed/queue`
|
||||
* **Functionality:** Accepts a standard ComfyUI workflow JSON, automatically distributes it to available workers, and returns the execution ID.
|
||||
* **Functionality:** Accepts a ComfyUI API-format prompt, dispatches it to the requested reachable workers, and returns the master `prompt_id`.
|
||||
* **Documentation:** [See API Examples & Scripts](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/comfyui-distributed-api.md)
|
||||
|
||||
> **⚠️ Security Warning:** Do not expose your ComfyUI port to the public internet. If you need remote access, run ComfyUI behind a secure proxy (like Cloudflare or a VPN).
|
||||
@@ -175,7 +175,7 @@ Use **Distributed Value** when you want per-worker overrides (for example, diffe
|
||||
| **Distributed Value** | Outputs per-worker override values with fallback to default |
|
||||
| **Ultimate SD Upscale Distributed** | Distributes upscale tiles across workers |
|
||||
| **Image Batch Divider** | Splits image batches for multi-GPU output |
|
||||
| **Audio Batch Divider** | Splits audio batches for multi-GPU output |
|
||||
| **Audio Segment Divider** | Splits an audio waveform into up to ten sequential time segments |
|
||||
| **Distributed Model Name** | Passes model paths to workers, enabling workflows to use models not present on the master in orchestrator-only mode |
|
||||
| **Distributed Empty Image** | Produces an empty IMAGE batch used when the master delegates all work |
|
||||
|
||||
|
||||
+15
-22
@@ -1,29 +1,22 @@
|
||||
# Import everything needed from the main module
|
||||
from .distributed import (
|
||||
NODE_CLASS_MAPPINGS as DISTRIBUTED_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as DISTRIBUTED_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
"""ComfyUI-Distributed's native V3 extension entrypoint."""
|
||||
from comfy_api.v0_0_2 import ComfyExtension, io
|
||||
|
||||
# Import utilities
|
||||
from .utils.config import ensure_config_exists, CONFIG_FILE
|
||||
from .utils.logging import debug_log
|
||||
from .nodes.v3 import NODES
|
||||
from .runtime.bootstrap import initialize
|
||||
|
||||
# Import distributed upscale nodes
|
||||
from .nodes.distributed_upscale import (
|
||||
NODE_CLASS_MAPPINGS as UPSCALE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as UPSCALE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
WEB_DIRECTORY = './web'
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
ensure_config_exists()
|
||||
class DistributedExtension(ComfyExtension):
|
||||
async def on_load(self) -> None:
|
||||
initialize()
|
||||
|
||||
# Merge node mappings
|
||||
NODE_CLASS_MAPPINGS = {**DISTRIBUTED_CLASS_MAPPINGS, **UPSCALE_CLASS_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**DISTRIBUTED_DISPLAY_NAME_MAPPINGS, **UPSCALE_DISPLAY_NAME_MAPPINGS}
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return list(NODES)
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
debug_log("Loaded Distributed nodes.")
|
||||
debug_log(f"Config file: {CONFIG_FILE}")
|
||||
debug_log(f"Available nodes: {list(NODE_CLASS_MAPPINGS.keys())}")
|
||||
async def comfy_entrypoint() -> DistributedExtension:
|
||||
return DistributedExtension()
|
||||
|
||||
|
||||
__all__ = ['comfy_entrypoint', 'WEB_DIRECTORY']
|
||||
|
||||
+14
-5
@@ -217,7 +217,7 @@ async def distributed_queue_endpoint(request):
|
||||
return await handle_api_error(request, exc, 400)
|
||||
|
||||
try:
|
||||
prompt_id, worker_count = await orchestrate_distributed_execution(
|
||||
prompt_id, prompt_number, worker_count, node_errors = await orchestrate_distributed_execution(
|
||||
payload.prompt,
|
||||
payload.workflow_meta,
|
||||
payload.client_id,
|
||||
@@ -227,6 +227,8 @@ async def distributed_queue_endpoint(request):
|
||||
)
|
||||
return web.json_response({
|
||||
"prompt_id": prompt_id,
|
||||
"number": prompt_number,
|
||||
"node_errors": node_errors,
|
||||
"worker_count": worker_count,
|
||||
"auto_prepare_supported": True,
|
||||
})
|
||||
@@ -293,17 +295,24 @@ async def job_complete_endpoint(request):
|
||||
errors.append("worker_id: expected non-empty string")
|
||||
if not isinstance(batch_idx, int) or batch_idx < 0:
|
||||
errors.append("batch_idx: expected non-negative integer")
|
||||
if not isinstance(image_payload, str) or not image_payload.strip():
|
||||
image_provided = image_payload is not None
|
||||
audio_provided = audio_payload is not None
|
||||
if not image_provided and not audio_provided:
|
||||
errors.append("expected at least one of image or audio")
|
||||
if image_provided and (not isinstance(image_payload, str) or not image_payload.strip()):
|
||||
errors.append("image: expected non-empty base64 PNG string")
|
||||
if audio_payload is not None and not isinstance(audio_payload, dict):
|
||||
if audio_provided and not isinstance(audio_payload, dict):
|
||||
errors.append("audio: expected object when provided")
|
||||
if not isinstance(is_last, bool):
|
||||
errors.append("is_last: expected boolean")
|
||||
if errors:
|
||||
return await handle_api_error(request, errors, 400)
|
||||
|
||||
tensor = _decode_canonical_png_tensor(image_payload)
|
||||
decoded_audio = _decode_audio_payload(audio_payload) if audio_payload is not None else None
|
||||
try:
|
||||
tensor = _decode_canonical_png_tensor(image_payload) if image_provided else None
|
||||
decoded_audio = _decode_audio_payload(audio_payload) if audio_provided else None
|
||||
except ValueError as exc:
|
||||
return await handle_api_error(request, exc, 400)
|
||||
multi_job_id = job_id.strip()
|
||||
worker_id = worker_id.strip()
|
||||
|
||||
|
||||
@@ -125,6 +125,209 @@ def _find_upstream_nodes(prompt_obj, start_ids):
|
||||
return connected
|
||||
|
||||
|
||||
_DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES = {
|
||||
"PrimitiveBoolean",
|
||||
"PrimitiveFloat",
|
||||
"PrimitiveInt",
|
||||
"PrimitiveNode",
|
||||
"PrimitiveString",
|
||||
}
|
||||
|
||||
_DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES = {
|
||||
"LoadImage",
|
||||
}
|
||||
|
||||
_DELEGATE_MASTER_SAFE_SCALAR_TYPES = {"BOOLEAN", "FLOAT", "INT", "STRING"}
|
||||
_DELEGATE_MASTER_SAFE_LIST_TYPES = {"LIST"}
|
||||
|
||||
# ComfyUI 0.23 exposes CreateList via the newer schema API rather than the
|
||||
# legacy RETURN_TYPES/INPUT_TYPES attributes. Treat it as a safe config utility
|
||||
# only after its connected inputs recursively prove safe.
|
||||
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES = {"CreateList"}
|
||||
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES = {
|
||||
"CreateList": ("inputs.",),
|
||||
}
|
||||
|
||||
# Test hook. At runtime this stays None and the ComfyUI node registry is loaded lazily.
|
||||
_DELEGATE_MASTER_NODE_CLASS_MAPPINGS = None
|
||||
|
||||
|
||||
def _get_delegate_master_node_class_mappings():
|
||||
"""Return ComfyUI node-class mappings when available."""
|
||||
if _DELEGATE_MASTER_NODE_CLASS_MAPPINGS is not None:
|
||||
return _DELEGATE_MASTER_NODE_CLASS_MAPPINGS
|
||||
try:
|
||||
import nodes as comfy_nodes # type: ignore
|
||||
except Exception: # pragma: no cover - depends on ComfyUI runtime imports
|
||||
return {}
|
||||
return getattr(comfy_nodes, "NODE_CLASS_MAPPINGS", {}) or {}
|
||||
|
||||
|
||||
def _get_delegate_master_node_class(class_type):
|
||||
mappings = _get_delegate_master_node_class_mappings()
|
||||
return mappings.get(class_type) if isinstance(mappings, dict) else None
|
||||
|
||||
|
||||
def _normalize_delegate_master_return_type(return_type):
|
||||
if return_type is None:
|
||||
return ""
|
||||
return str(return_type).strip().upper()
|
||||
|
||||
|
||||
def _delegate_master_type_is_safe_scalar(type_name):
|
||||
return type_name in _DELEGATE_MASTER_SAFE_SCALAR_TYPES
|
||||
|
||||
|
||||
def _delegate_master_type_is_safe_config(type_name):
|
||||
return _delegate_master_type_is_safe_scalar(type_name) or type_name in _DELEGATE_MASTER_SAFE_LIST_TYPES
|
||||
|
||||
|
||||
def _delegate_master_output_is_safe_scalar(class_type, output_index):
|
||||
"""Return True when a registered node output is lightweight config data."""
|
||||
if class_type in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES:
|
||||
return True
|
||||
node_class = _get_delegate_master_node_class(class_type)
|
||||
return_types = getattr(node_class, "RETURN_TYPES", ()) if node_class is not None else ()
|
||||
try:
|
||||
output_type = return_types[int(output_index)]
|
||||
except (IndexError, TypeError, ValueError):
|
||||
return False
|
||||
return _delegate_master_type_is_safe_config(_normalize_delegate_master_return_type(output_type))
|
||||
|
||||
|
||||
def _get_delegate_master_input_types(class_type):
|
||||
node_class = _get_delegate_master_node_class(class_type)
|
||||
input_types = getattr(node_class, "INPUT_TYPES", None) if node_class is not None else None
|
||||
if callable(input_types):
|
||||
try:
|
||||
input_types = input_types()
|
||||
except TypeError:
|
||||
return {}
|
||||
return input_types if isinstance(input_types, dict) else {}
|
||||
|
||||
|
||||
def _normalize_delegate_master_input_type(input_spec):
|
||||
if isinstance(input_spec, (list, tuple)) and input_spec:
|
||||
return _normalize_delegate_master_return_type(input_spec[0])
|
||||
return _normalize_delegate_master_return_type(input_spec)
|
||||
|
||||
|
||||
def _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
"""Return True when a registered downstream input expects config data."""
|
||||
for prefix in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES.get(class_type, ()):
|
||||
if input_name.startswith(prefix):
|
||||
return True
|
||||
input_types = _get_delegate_master_input_types(class_type)
|
||||
for section_name in ("required", "optional"):
|
||||
section = input_types.get(section_name, {})
|
||||
if isinstance(section, dict) and input_name in section:
|
||||
input_type = _normalize_delegate_master_input_type(section[input_name])
|
||||
return _delegate_master_type_is_safe_config(input_type)
|
||||
return False
|
||||
|
||||
|
||||
def _is_delegate_master_always_retained_upstream_node(node):
|
||||
if not isinstance(node, dict):
|
||||
return False
|
||||
class_type = node.get("class_type")
|
||||
return isinstance(class_type, str) and class_type in _DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES
|
||||
|
||||
|
||||
def _is_delegate_master_retained_upstream_node(node, output_index=0):
|
||||
"""Return True for lightweight upstream nodes safe to keep on the master."""
|
||||
if not isinstance(node, dict):
|
||||
return False
|
||||
class_type = node.get("class_type")
|
||||
if not isinstance(class_type, str):
|
||||
return False
|
||||
return (
|
||||
class_type in _DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES
|
||||
or class_type.startswith("Primitive")
|
||||
or _delegate_master_output_is_safe_scalar(class_type, output_index)
|
||||
)
|
||||
|
||||
|
||||
def _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
node_id,
|
||||
output_index,
|
||||
memo,
|
||||
visiting,
|
||||
):
|
||||
"""Return safe retained branch nodes, or None when the branch is not safe."""
|
||||
node_id = str(node_id)
|
||||
cache_key = (node_id, output_index)
|
||||
if cache_key in memo:
|
||||
cached = memo[cache_key]
|
||||
return None if cached is None else set(cached)
|
||||
if cache_key in visiting:
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
|
||||
node = prompt_obj.get(node_id)
|
||||
if not _is_delegate_master_retained_upstream_node(node, output_index):
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
|
||||
visiting.add(cache_key)
|
||||
retained = {node_id}
|
||||
inputs = node.get("inputs", {}) if isinstance(node, dict) else {}
|
||||
class_type = node.get("class_type") if isinstance(node, dict) else None
|
||||
for input_name, value in inputs.items():
|
||||
if not (isinstance(value, list) and len(value) == 2):
|
||||
continue
|
||||
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
source_id = str(value[0])
|
||||
branch = _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
source_id,
|
||||
value[1],
|
||||
memo,
|
||||
visiting,
|
||||
)
|
||||
if branch is None:
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = None
|
||||
return None
|
||||
retained.update(branch)
|
||||
|
||||
visiting.remove(cache_key)
|
||||
memo[cache_key] = frozenset(retained)
|
||||
return retained
|
||||
|
||||
|
||||
def _find_delegate_master_retained_upstream_nodes(prompt_obj, start_ids):
|
||||
"""Return lightweight upstream nodes needed by kept delegate-master nodes."""
|
||||
connected = set()
|
||||
memo = {}
|
||||
for node_id in start_ids:
|
||||
node = prompt_obj.get(str(node_id)) or {}
|
||||
inputs = node.get("inputs", {})
|
||||
class_type = node.get("class_type") if isinstance(node, dict) else None
|
||||
for input_name, value in inputs.items():
|
||||
if not (isinstance(value, list) and len(value) == 2):
|
||||
continue
|
||||
source_node = prompt_obj.get(str(value[0]))
|
||||
if _is_delegate_master_always_retained_upstream_node(source_node):
|
||||
connected.add(str(value[0]))
|
||||
continue
|
||||
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
|
||||
continue
|
||||
branch = _collect_delegate_master_retained_upstream_branch(
|
||||
prompt_obj,
|
||||
value[0],
|
||||
value[1],
|
||||
memo,
|
||||
set(),
|
||||
)
|
||||
if branch is not None:
|
||||
connected.update(branch)
|
||||
return connected
|
||||
|
||||
|
||||
def prune_prompt_for_worker(prompt_obj):
|
||||
"""Prune worker prompt to distributed nodes and their upstream dependencies."""
|
||||
collector_ids = find_nodes_by_class(prompt_obj, "DistributedCollector")
|
||||
@@ -148,16 +351,33 @@ def prune_prompt_for_worker(prompt_obj):
|
||||
downstream = _find_downstream_nodes(prompt_obj, [dist_id])
|
||||
has_removed_downstream = any(node_id != dist_id for node_id in downstream)
|
||||
if has_removed_downstream:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {
|
||||
"images": [dist_id, 0],
|
||||
},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {
|
||||
"title": "Preview Image (auto-added)",
|
||||
},
|
||||
}
|
||||
original_node = prompt_obj.get(str(dist_id), {})
|
||||
class_type = original_node.get("class_type")
|
||||
inputs = original_node.get("inputs", {})
|
||||
image_connected = class_type != "DistributedCollector" or (
|
||||
isinstance(inputs.get("images"), list)
|
||||
and len(inputs["images"]) == 2
|
||||
)
|
||||
audio_connected = (
|
||||
class_type == "DistributedCollector"
|
||||
and isinstance(inputs.get("audio"), list)
|
||||
and len(inputs["audio"]) == 2
|
||||
)
|
||||
|
||||
if image_connected:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {"images": [dist_id, 0]},
|
||||
"class_type": "PreviewImage",
|
||||
"_meta": {"title": "Preview Image (auto-added)"},
|
||||
}
|
||||
elif audio_connected:
|
||||
preview_id = next_id()
|
||||
pruned_prompt[preview_id] = {
|
||||
"inputs": {"audio": [dist_id, 1]},
|
||||
"class_type": "PreviewAudio",
|
||||
"_meta": {"title": "Preview Audio (auto-added)"},
|
||||
}
|
||||
|
||||
return pruned_prompt
|
||||
|
||||
@@ -167,6 +387,9 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
downstream = _find_downstream_nodes(prompt_obj, collector_ids)
|
||||
nodes_to_keep = set(collector_ids)
|
||||
nodes_to_keep.update(downstream)
|
||||
nodes_to_keep.update(
|
||||
_find_delegate_master_retained_upstream_nodes(prompt_obj, nodes_to_keep)
|
||||
)
|
||||
|
||||
pruned_prompt = {}
|
||||
for node_id in nodes_to_keep:
|
||||
@@ -194,6 +417,10 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
|
||||
collector_entry = pruned_prompt.get(collector_id)
|
||||
if not collector_entry:
|
||||
continue
|
||||
original_inputs = (prompt_obj.get(collector_id) or {}).get("inputs", {})
|
||||
original_images = original_inputs.get("images")
|
||||
if not (isinstance(original_images, list) and len(original_images) == 2):
|
||||
continue
|
||||
placeholder_id = next_id()
|
||||
pruned_prompt[placeholder_id] = {
|
||||
"class_type": "DistributedEmptyImage",
|
||||
|
||||
@@ -208,7 +208,7 @@ async def orchestrate_distributed_execution(
|
||||
"""Core orchestration logic for the /distributed/queue endpoint.
|
||||
|
||||
Returns:
|
||||
tuple[str, int]: (prompt_id, worker_count)
|
||||
tuple[str, int, int, dict]: (prompt_id, number, worker_count, node_errors)
|
||||
"""
|
||||
ensure_distributed_state()
|
||||
execution_trace_id = trace_execution_id or _generate_execution_trace_id()
|
||||
@@ -317,8 +317,18 @@ async def orchestrate_distributed_execution(
|
||||
|
||||
if not job_id_map:
|
||||
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
|
||||
prompt_id = await queue_prompt_payload(prompt_obj, workflow_meta, client_id)
|
||||
return prompt_id, 0
|
||||
queue_result = await queue_prompt_payload(
|
||||
prompt_obj,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
include_queue_metadata=True,
|
||||
)
|
||||
return (
|
||||
queue_result["prompt_id"],
|
||||
queue_result["number"],
|
||||
0,
|
||||
queue_result.get("node_errors", {}),
|
||||
)
|
||||
|
||||
for job_id in job_id_map.values():
|
||||
await _ensure_distributed_queue(job_id)
|
||||
@@ -392,9 +402,17 @@ async def orchestrate_distributed_execution(
|
||||
]
|
||||
)
|
||||
|
||||
prompt_id = await queue_prompt_payload(master_prompt, workflow_meta, client_id)
|
||||
queue_result = await queue_prompt_payload(
|
||||
master_prompt,
|
||||
workflow_meta,
|
||||
client_id,
|
||||
include_queue_metadata=True,
|
||||
)
|
||||
prompt_id = queue_result["prompt_id"]
|
||||
prompt_number = queue_result["number"]
|
||||
node_errors = queue_result.get("node_errors", {})
|
||||
trace_debug(
|
||||
execution_trace_id,
|
||||
f"Orchestration complete: prompt_id={prompt_id}, dispatched_workers={len(worker_payloads)}, delegate_master={delegate_master}",
|
||||
)
|
||||
return prompt_id, len(worker_payloads)
|
||||
return prompt_id, prompt_number, len(worker_payloads), node_errors
|
||||
|
||||
+14
-1
@@ -16,12 +16,14 @@ from ..utils.logging import debug_log, log
|
||||
from ..utils.network import (
|
||||
build_worker_url,
|
||||
get_client_session,
|
||||
get_server_port,
|
||||
handle_api_error,
|
||||
normalize_host,
|
||||
probe_worker,
|
||||
)
|
||||
from ..utils.constants import CHUNK_SIZE
|
||||
from ..workers import get_worker_manager
|
||||
from ..workers.ports import allocate_worker_ports
|
||||
from .schemas import require_fields, validate_worker_id
|
||||
from ..workers.detection import (
|
||||
get_machine_id,
|
||||
@@ -277,6 +279,13 @@ def _get_cuda_info():
|
||||
def _collect_network_info_sync():
|
||||
"""Collect network/cuda info in a worker thread to avoid blocking route handlers."""
|
||||
cuda_device, cuda_device_count, physical_device_count = _get_cuda_info()
|
||||
device_count = physical_device_count if physical_device_count > 0 else cuda_device_count
|
||||
master_port = get_server_port()
|
||||
config = load_config()
|
||||
worker_ports = []
|
||||
if not config.get("settings", {}).get("has_auto_populated_workers") and not config.get("workers"):
|
||||
worker_count = device_count - (1 if cuda_device is not None and 0 <= cuda_device < device_count else 0)
|
||||
worker_ports = allocate_worker_ports(master_port, config.get("workers", []), worker_count)
|
||||
hostname = socket.gethostname()
|
||||
all_ips = get_network_ips()
|
||||
recommended_ip = get_recommended_ip(all_ips)
|
||||
@@ -285,7 +294,9 @@ def _collect_network_info_sync():
|
||||
"all_ips": all_ips,
|
||||
"recommended_ip": recommended_ip,
|
||||
"cuda_device": cuda_device,
|
||||
"cuda_device_count": physical_device_count if physical_device_count > 0 else cuda_device_count,
|
||||
"cuda_device_count": device_count,
|
||||
"master_port": master_port,
|
||||
"local_worker_ports": worker_ports,
|
||||
}
|
||||
|
||||
|
||||
@@ -483,6 +494,8 @@ async def launch_worker_endpoint(request):
|
||||
"message": f"Worker {worker['name']} launched",
|
||||
"log_file": log_file
|
||||
})
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 400)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 500)
|
||||
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# conftest.py — project-level pytest configuration.
|
||||
#
|
||||
# Problem: custom_nodes/ComfyUI-Distributed/__init__.py uses relative imports
|
||||
# (from .distributed import ...) that fail when pytest tries to import it as a
|
||||
# (from .nodes.v3 import ...) that fail when pytest tries to import it as a
|
||||
# standalone module during Package.setup() for the root package node.
|
||||
#
|
||||
# Fix: patch Package.setup() to skip the root-package's __init__.py import.
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
"""
|
||||
ComfyUI-Distributed: thin entry point.
|
||||
All implementation lives in workers/, nodes/, api/.
|
||||
"""
|
||||
import atexit
|
||||
import os
|
||||
|
||||
import server
|
||||
|
||||
from .utils.config import ensure_config_exists
|
||||
from .utils.logging import debug_log
|
||||
from .utils.network import cleanup_client_session
|
||||
from .workers import get_worker_manager
|
||||
from .workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
from .upscale.job_store import ensure_tile_jobs_initialized
|
||||
from .nodes import (
|
||||
NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS,
|
||||
ImageBatchDivider,
|
||||
DistributedCollectorNode,
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from . import api # noqa: F401 - triggers all @routes.* registrations
|
||||
from .api.queue_orchestration import ensure_distributed_state
|
||||
|
||||
ensure_config_exists()
|
||||
|
||||
# Aiohttp session cleanup
|
||||
async def _cleanup_session():
|
||||
await cleanup_client_session()
|
||||
|
||||
|
||||
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
|
||||
|
||||
# Initialize distributed job state on prompt_server
|
||||
prompt_server = server.PromptServer.instance
|
||||
ensure_distributed_state(prompt_server)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# Worker startup
|
||||
if not os.environ.get('COMFYUI_IS_WORKER'):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
@@ -19,8 +19,8 @@ This document describes the **public HTTP API** added to ComfyUI-Distributed to
|
||||
- `POST /distributed/queue` — queues a workflow using the same distributed orchestration rules as the UI:
|
||||
- Detects distributed nodes in the prompt (`DistributedCollector`, `UltimateSDUpscaleDistributed`).
|
||||
- Resolves enabled/selected workers.
|
||||
- Pings workers (`GET /prompt`) to include only reachable ones.
|
||||
- Dispatches the workflow to workers (`POST /prompt`).
|
||||
- Probes and dispatches workers through `/distributed/worker_ws` by default.
|
||||
- If `settings.websocket_orchestration=false`, probes with `GET /prompt` and dispatches with `POST /prompt` instead.
|
||||
- Queues the master workflow in ComfyUI’s prompt queue.
|
||||
- If any `DistributedCollector` has `load_balance=true`, selects one least-busy participant for this run.
|
||||
|
||||
@@ -61,7 +61,8 @@ Queue a workflow for distributed execution.
|
||||
#### Fields
|
||||
|
||||
- `prompt` (required unless `workflow.prompt` is present, object)
|
||||
- The ComfyUI prompt/workflow graph, same shape as used by `POST /prompt`.
|
||||
- A complete ComfyUI API-format prompt graph, using the same shape as `POST /prompt`.
|
||||
- This is not the normal visual workflow export from the ComfyUI editor.
|
||||
- `workflow` (optional, object)
|
||||
- Workflow metadata that ComfyUI normally stores in `extra_pnginfo.workflow`.
|
||||
- If you don’t care about UI metadata, you can omit it.
|
||||
@@ -123,10 +124,14 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
|
||||
|
||||
## Worker requirements (important)
|
||||
|
||||
For a worker to participate, it must be reachable from the master:
|
||||
For a worker to participate, it must be reachable from the master. By default:
|
||||
|
||||
- WebSocket probe and dispatch: `<worker-base>/distributed/worker_ws` must accept the connection.
|
||||
|
||||
If `settings.websocket_orchestration=false`:
|
||||
|
||||
- Health check: `GET <worker-base>/prompt` must return HTTP 200.
|
||||
- Dispatch: `POST <worker-base>/prompt` must accept the workflow.
|
||||
- Dispatch: `POST <worker-base>/prompt` must accept the prompt.
|
||||
|
||||
Also, for collector-based flows:
|
||||
|
||||
@@ -201,7 +206,7 @@ Proxy endpoint on master that fetches logs from a configured remote/cloud worker
|
||||
|
||||
## Examples
|
||||
|
||||
### 1) Minimal `curl`
|
||||
### 1) Minimal `curl` request envelope
|
||||
|
||||
```bash
|
||||
curl -X POST "http://127.0.0.1:8188/distributed/queue" \
|
||||
@@ -209,18 +214,23 @@ curl -X POST "http://127.0.0.1:8188/distributed/queue" \
|
||||
-d @payload.json
|
||||
```
|
||||
|
||||
Where `payload.json` contains at least:
|
||||
`payload.json` must contain a complete ComfyUI API-format prompt. The abbreviated envelope below illustrates the request shape but is not directly executable:
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": {
|
||||
"1": {"class_type": "KSampler", "inputs": {} }
|
||||
"<node_id>": {
|
||||
"class_type": "<node class>",
|
||||
"inputs": {"<required_input>": "<value or connection>"}
|
||||
}
|
||||
},
|
||||
"enabled_worker_ids": [],
|
||||
"client_id": "external-client"
|
||||
}
|
||||
```
|
||||
|
||||
Export or construct a valid API-format prompt with all required node inputs and at least one output node before submitting it.
|
||||
|
||||
### 2) Python (`requests`)
|
||||
|
||||
```python
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
# Native V3 node API
|
||||
|
||||
The package registers one `ComfyExtension` through `comfy_entrypoint()` and
|
||||
imports the versioned `comfy_api.v0_0_2` API. Use a ComfyUI version that provides
|
||||
that API; there is no V1 registration fallback.
|
||||
|
||||
## Compatibility
|
||||
|
||||
- All eight node IDs, display names, categories, visible input order, defaults
|
||||
and output order are retained. Existing execution algorithms remain in the
|
||||
private collector, utilities and upscale modules.
|
||||
- The image/audio dividers explicitly declare ten typed outputs. Their existing
|
||||
frontend extensions still show only the selected number of outputs. This
|
||||
replaces the V1 `ByPassTypeTuple` indexing workaround without changing saved
|
||||
IMAGE/AUDIO socket indices or the ten returned values.
|
||||
- Standard hidden context uses `cls.hidden`. Worker/orchestration metadata keeps
|
||||
its existing prompt-input names and Python defaults via `accept_all_inputs`;
|
||||
it does not become visible widgets.
|
||||
- The collector retains list-input handling. Upscale retains its always-changing
|
||||
fingerprint and creates a private runtime helper for each V3 execution, so
|
||||
mutable helper state is not attached to sanitized V3 class clones.
|
||||
- Routes, distributed state and the existing worker startup/shutdown hooks are
|
||||
initialized by `runtime/bootstrap.py` during `ComfyExtension.on_load()`.
|
||||
Worker mode still suppresses automatic worker launch. `distributed.py` is no
|
||||
longer an entrypoint.
|
||||
|
||||
## Verification
|
||||
|
||||
Run ordinary unit tests from this repository:
|
||||
|
||||
```bash
|
||||
python -m pytest tests -q -o addopts=
|
||||
```
|
||||
|
||||
Opt into real-framework acceptance using a ComfyUI checkout and its interpreter:
|
||||
|
||||
```bash
|
||||
COMFYUI_SOURCE_ROOT=/path/to/ComfyUI \
|
||||
/path/to/ComfyUI/.venv/bin/python -m pytest tests -q -o addopts=
|
||||
```
|
||||
|
||||
The acceptance subprocess uses CPU mode and the real ComfyUI loader, input
|
||||
parser, V3 class preparation, prompt validation and `PromptExecutor`. It compares
|
||||
all node schemas with the V1 fixture, checks all five bundled workflows' node
|
||||
IDs and socket/link contracts, exercises injected worker metadata, image/audio
|
||||
lists, collector aggregation and divider outputs, rejects invalid upscale enums,
|
||||
and decodes an actual preview PNG referenced by executor history.
|
||||
|
||||
The upscale GPU/model boundary is mocked to check argument forwarding and
|
||||
per-execution helper isolation. Full checkpoint inference, browser canvas
|
||||
acceptance and multi-host HTTP transport are not exercised. No HTTP listener,
|
||||
workers or model downloads are started; preview files use a temporary scratch
|
||||
directory. The test does not install this branch into a live custom-node folder.
|
||||
@@ -79,7 +79,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
|
||||
📺 [Watch Tutorial](https://www.youtube.com/watch?v=wxKKWMQhYTk)
|
||||
|
||||
**On Runpod:**
|
||||
> If using your own template, make sure you launch ComfyUI with the `--enable-cors-header` argument and you `git clone ComfyUI-Distributed` into custom_nodes. ⚠️ **Required!**
|
||||
> If using your own template, launch ComfyUI with `--listen --enable-cors-header` and clone `ComfyUI-Distributed` into `custom_nodes`. ⚠️ **Required!**
|
||||
|
||||
1. Register a [Runpod](https://get.runpod.io/0bw29uf3ug0p) account.
|
||||
2. On Runpod, go to Storage > New Network Volume and create a volume that will store the models you need. Start with 40 GB, you can always add more later. Learn more [about Network Volumes](https://docs.runpod.io/pods/storage/create-network-volumes).
|
||||
@@ -92,7 +92,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
|
||||
- SAGE_ATTENTION: optional optimisation (set to true/false)
|
||||
5. Deploy your pod.
|
||||
6. Connect to your pod using JupyterLabs. This gives us access to the pod's file system.
|
||||
7. Download models into /workspaces/ComfyUI/models/ (these will remain on your network drive even after you terminate the pod). Example commands below:
|
||||
7. Download models into `/workspace/ComfyUI/models/` (these will remain on your network drive even after you terminate the pod). Example commands below:
|
||||
```
|
||||
# Download from CivitAI
|
||||
comfy model download --url https://civitai.com/api/download/models/1759168 --relative-path /workspace/ComfyUI/models/checkpoints --set-civitai-api-token $CIVITAI_API_TOKEN
|
||||
|
||||
+1
-31
@@ -1,31 +1 @@
|
||||
from .utilities import (
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
ImageBatchDivider,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from .collector import DistributedCollectorNode
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DistributedCollector": DistributedCollectorNode,
|
||||
"DistributedSeed": DistributedSeed,
|
||||
"DistributedModelName": DistributedModelName,
|
||||
"DistributedValue": DistributedValue,
|
||||
"ImageBatchDivider": ImageBatchDivider,
|
||||
"AudioBatchDivider": AudioBatchDivider,
|
||||
"DistributedEmptyImage": DistributedEmptyImage,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DistributedCollector": "Distributed Collector",
|
||||
"DistributedSeed": "Distributed Seed",
|
||||
"DistributedModelName": "Distributed Model Name",
|
||||
"DistributedValue": "Distributed Value",
|
||||
"ImageBatchDivider": "Image Batch Divider",
|
||||
"AudioBatchDivider": "Audio Batch Divider",
|
||||
"DistributedEmptyImage": "Distributed Empty Image",
|
||||
}
|
||||
"""Private execution helpers; public registration lives in nodes.v3."""
|
||||
|
||||
+131
-44
@@ -22,13 +22,13 @@ prompt_server = _server.PromptServer.instance
|
||||
|
||||
|
||||
class DistributedCollectorNode:
|
||||
INPUT_IS_LIST = True
|
||||
EMPTY_AUDIO = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"load_balance": (
|
||||
"BOOLEAN",
|
||||
{
|
||||
@@ -37,7 +37,10 @@ class DistributedCollectorNode:
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": { "audio": ("AUDIO",) },
|
||||
"optional": {
|
||||
"images": ("IMAGE",),
|
||||
"audio": ("AUDIO",),
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": ("STRING", {"default": ""}),
|
||||
"is_worker": ("BOOLEAN", {"default": False}),
|
||||
@@ -55,7 +58,74 @@ class DistributedCollectorNode:
|
||||
FUNCTION = "run"
|
||||
CATEGORY = "image"
|
||||
|
||||
def run(self, images, load_balance=False, audio=None, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False, delegate_only=False):
|
||||
@staticmethod
|
||||
def _unwrap_list_input(value):
|
||||
"""Unwrap scalar inputs when ComfyUI passes them via INPUT_IS_LIST."""
|
||||
if isinstance(value, (list, tuple)) and len(value) == 1:
|
||||
return value[0]
|
||||
return value
|
||||
|
||||
def _normalize_images_input(self, images):
|
||||
"""Collapse ComfyUI list IMAGE inputs into a normal batched IMAGE tensor."""
|
||||
if isinstance(images, (list, tuple)):
|
||||
if not images:
|
||||
raise ValueError("Collector received an empty image list")
|
||||
if not all(isinstance(image, torch.Tensor) for image in images):
|
||||
raise TypeError("Collector expected IMAGE list items to be torch.Tensor instances")
|
||||
if len(images) == 1:
|
||||
return ensure_contiguous(images[0])
|
||||
return ensure_contiguous(torch.cat([ensure_contiguous(image) for image in images], dim=0))
|
||||
return ensure_contiguous(images)
|
||||
|
||||
def _normalize_audio_input(self, audio):
|
||||
"""Collapse ComfyUI list AUDIO inputs into a single AUDIO payload when present."""
|
||||
if not isinstance(audio, (list, tuple)):
|
||||
return audio
|
||||
|
||||
audio_items = [item for item in audio if item is not None]
|
||||
if not audio_items:
|
||||
return None
|
||||
if len(audio_items) == 1:
|
||||
return audio_items[0]
|
||||
|
||||
waveforms = []
|
||||
sample_rate = 44100
|
||||
for item in audio_items:
|
||||
if not isinstance(item, dict):
|
||||
raise TypeError("Collector expected AUDIO list items to be dictionaries")
|
||||
waveform = item.get("waveform")
|
||||
if waveform is None or waveform.numel() == 0:
|
||||
continue
|
||||
waveforms.append(waveform)
|
||||
if sample_rate == 44100:
|
||||
sample_rate = item.get("sample_rate", 44100)
|
||||
|
||||
if not waveforms:
|
||||
return None
|
||||
return {"waveform": torch.cat(waveforms, dim=-1), "sample_rate": sample_rate}
|
||||
|
||||
def run(self, images=None, load_balance=False, audio=None, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False, delegate_only=False):
|
||||
if images is not None:
|
||||
images = self._normalize_images_input(images)
|
||||
audio = self._normalize_audio_input(audio)
|
||||
load_balance = self._unwrap_list_input(load_balance)
|
||||
multi_job_id = self._unwrap_list_input(multi_job_id)
|
||||
is_worker = self._unwrap_list_input(is_worker)
|
||||
master_url = self._unwrap_list_input(master_url)
|
||||
enabled_worker_ids = self._unwrap_list_input(enabled_worker_ids)
|
||||
worker_batch_size = self._unwrap_list_input(worker_batch_size)
|
||||
worker_id = self._unwrap_list_input(worker_id)
|
||||
pass_through = self._unwrap_list_input(pass_through)
|
||||
delegate_only = self._unwrap_list_input(delegate_only)
|
||||
|
||||
remote_only_master = (
|
||||
bool(multi_job_id)
|
||||
and not is_worker
|
||||
and (delegate_only or is_master_delegate_only())
|
||||
)
|
||||
if images is None and audio is None and not remote_only_master:
|
||||
raise ValueError("DistributedCollector requires at least one image or audio input")
|
||||
|
||||
# Create empty audio if not provided
|
||||
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
|
||||
|
||||
@@ -82,41 +152,56 @@ class DistributedCollectorNode:
|
||||
return result
|
||||
|
||||
async def send_batch_to_master(self, image_batch, audio, multi_job_id, master_url, worker_id):
|
||||
"""Send image batch to master via canonical JSON envelopes."""
|
||||
batch_size = image_batch.shape[0]
|
||||
if batch_size == 0:
|
||||
return
|
||||
|
||||
"""Send an image batch, optionally with audio, or an audio-only completion."""
|
||||
encoded_audio = encode_audio_payload(audio)
|
||||
|
||||
session = await get_client_session()
|
||||
url = f"{master_url}/distributed/job_complete"
|
||||
for batch_idx in range(batch_size):
|
||||
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
|
||||
byte_io = io.BytesIO()
|
||||
img.save(byte_io, format='PNG', compress_level=0)
|
||||
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
|
||||
payload = {
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": int(batch_idx),
|
||||
"image": f"data:image/png;base64,{encoded_image}",
|
||||
"is_last": bool(batch_idx == batch_size - 1),
|
||||
}
|
||||
if payload["is_last"] and encoded_audio is not None:
|
||||
payload["audio"] = encoded_audio
|
||||
|
||||
payloads = []
|
||||
batch_size = 0 if image_batch is None else image_batch.shape[0]
|
||||
if batch_size == 0:
|
||||
if encoded_audio is None:
|
||||
raise ValueError("Worker completion requires image or audio data")
|
||||
payloads.append(
|
||||
{
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": 0,
|
||||
"audio": encoded_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
else:
|
||||
for batch_idx in range(batch_size):
|
||||
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
|
||||
byte_io = io.BytesIO()
|
||||
img.save(byte_io, format='PNG', compress_level=0)
|
||||
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
|
||||
payload = {
|
||||
"job_id": str(multi_job_id),
|
||||
"worker_id": str(worker_id),
|
||||
"batch_idx": int(batch_idx),
|
||||
"image": f"data:image/png;base64,{encoded_image}",
|
||||
"is_last": bool(batch_idx == batch_size - 1),
|
||||
}
|
||||
if payload["is_last"] and encoded_audio is not None:
|
||||
payload["audio"] = encoded_audio
|
||||
payloads.append(payload)
|
||||
|
||||
for payload in payloads:
|
||||
timeout_seconds = 60 if "image" in payload else 600
|
||||
try:
|
||||
async with session.post(
|
||||
url,
|
||||
json=payload,
|
||||
timeout=aiohttp.ClientTimeout(total=60),
|
||||
timeout=aiohttp.ClientTimeout(total=timeout_seconds),
|
||||
) as response:
|
||||
response.raise_for_status()
|
||||
except Exception as e:
|
||||
log(f"Worker - Failed to send canonical image envelope to master: {e}")
|
||||
media_type = "image/audio" if "image" in payload else "audio-only"
|
||||
log(f"Worker - Failed to send canonical {media_type} envelope to master: {e}")
|
||||
debug_log(f"Worker - Full error details: URL={url}")
|
||||
raise # Re-raise to handle at caller level
|
||||
raise
|
||||
|
||||
def _combine_audio(self, master_audio, worker_audio, empty_audio, worker_order=None):
|
||||
"""Combine audio from master and workers into a single audio output.
|
||||
@@ -198,8 +283,8 @@ class DistributedCollectorNode:
|
||||
images_on_cpu,
|
||||
delegate_mode: bool,
|
||||
fallback_images,
|
||||
) -> torch.Tensor:
|
||||
"""Assemble final tensor: master first, then workers in enabled order."""
|
||||
):
|
||||
"""Assemble final tensor, or return None when the job contains only audio."""
|
||||
ordered_tensors = []
|
||||
if not delegate_mode and images_on_cpu is not None:
|
||||
for i in range(master_batch_size):
|
||||
@@ -230,15 +315,15 @@ class DistributedCollectorNode:
|
||||
|
||||
if cpu_tensors:
|
||||
return ensure_contiguous(torch.cat(cpu_tensors, dim=0))
|
||||
elif fallback_images is not None:
|
||||
if fallback_images is not None:
|
||||
return ensure_contiguous(fallback_images)
|
||||
else:
|
||||
raise ValueError("No image data collected from master or workers")
|
||||
return None
|
||||
|
||||
async def execute(self, images, audio, load_balance=False, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", delegate_only=False):
|
||||
if is_worker:
|
||||
# Worker mode: send images and audio to master in a single batch
|
||||
debug_log(f"Worker - Job {multi_job_id} complete. Sending {images.shape[0]} image(s) to master")
|
||||
image_count = 0 if images is None else images.shape[0]
|
||||
debug_log(f"Worker - Job {multi_job_id} complete. Sending {image_count} image(s) to master")
|
||||
await self.send_batch_to_master(images, audio, multi_job_id, master_url, worker_id)
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
else:
|
||||
@@ -273,14 +358,15 @@ class DistributedCollectorNode:
|
||||
master_audio = None
|
||||
debug_log(f"Master - Job {multi_job_id}: Delegate-only mode enabled, collecting exclusively from {num_workers} workers")
|
||||
else:
|
||||
images_on_cpu = images.cpu()
|
||||
master_batch_size = images.shape[0]
|
||||
if images is None:
|
||||
images_on_cpu = None
|
||||
master_batch_size = 0
|
||||
else:
|
||||
images_on_cpu = ensure_contiguous(images.cpu())
|
||||
master_batch_size = images.shape[0]
|
||||
master_audio = audio # Keep master's audio for later
|
||||
debug_log(f"Master - Job {multi_job_id}: Master has {master_batch_size} images, collecting from {num_workers} workers...")
|
||||
|
||||
# Ensure master images are contiguous
|
||||
images_on_cpu = ensure_contiguous(images_on_cpu)
|
||||
|
||||
|
||||
# Initialize storage for collected images and audio
|
||||
worker_images = {} # Dict to store images by worker_id and index
|
||||
@@ -452,18 +538,19 @@ class DistributedCollectorNode:
|
||||
if multi_job_id in prompt_server.distributed_pending_jobs:
|
||||
del prompt_server.distributed_pending_jobs[multi_job_id]
|
||||
|
||||
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
|
||||
try:
|
||||
combined = self._reorder_and_combine_tensors(
|
||||
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
|
||||
)
|
||||
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
|
||||
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
|
||||
|
||||
# Combine audio from master and workers
|
||||
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
|
||||
if combined is None:
|
||||
debug_log(f"Master - Job {multi_job_id} complete with audio only")
|
||||
else:
|
||||
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
|
||||
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
|
||||
|
||||
return (combined, combined_audio)
|
||||
except Exception as e:
|
||||
log(f"Master - Error combining images: {e}")
|
||||
# Return just the master images as fallback
|
||||
return (images, audio if audio is not None else self.EMPTY_AUDIO)
|
||||
# Preserve collected audio even when image assembly fails.
|
||||
return (images, combined_audio)
|
||||
|
||||
@@ -56,22 +56,22 @@ class UltimateSDUpscaleDistributed(
|
||||
"""
|
||||
Distributed version of Ultimate SD Upscale (No Upscale).
|
||||
|
||||
Supports three processing modes:
|
||||
Supports two currently selected processing modes:
|
||||
1. Single GPU: No workers available, process everything locally
|
||||
2. Static Mode: Small batches, distributes tiles across workers (flattened)
|
||||
3. Dynamic Mode: Large batches, assigns whole images to workers dynamically
|
||||
2. Distributed tile queue: Workers pull tiles from a shared queue
|
||||
|
||||
Features:
|
||||
- Multi-mode batch handling for efficient video/image upscaling
|
||||
- Tile-based batch handling for video/image upscaling
|
||||
- Tiled VAE support for memory efficiency
|
||||
- Dynamic load balancing for large batches
|
||||
- Shared work queue so faster workers can process more tiles
|
||||
- Backward compatible with single-image workflows
|
||||
|
||||
Environment Variables:
|
||||
- COMFYUI_MAX_BATCH: Chunk size for tile sending (default 20)
|
||||
- COMFYUI_MAX_PAYLOAD_SIZE: Max API payload bytes (default 50MB)
|
||||
|
||||
Threshold: dynamic_threshold input controls mode switch (default 8)
|
||||
The hidden dynamic_threshold input is retained for workflow compatibility but
|
||||
does not affect the current mode-selection policy.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
@@ -268,12 +268,3 @@ class UltimateSDUpscaleDistributed(
|
||||
|
||||
# Ensure initialization before registering routes
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# Node registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": UltimateSDUpscaleDistributed,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": "Ultimate SD Upscale Distributed (No Upscale)",
|
||||
}
|
||||
|
||||
+2
-2
@@ -269,7 +269,7 @@ class ImageBatchDivider:
|
||||
|
||||
|
||||
class AudioBatchDivider:
|
||||
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
|
||||
"""Divides an audio waveform into sequential segments along the time dimension."""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -282,7 +282,7 @@ class AudioBatchDivider:
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of parts to divide the audio into"
|
||||
"tooltip": "Number of sequential time segments to create"
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
"""Native V3 schemas with the existing execution algorithms kept intact.
|
||||
|
||||
Runtime objects are private, per-execution helpers, not registered V1 nodes.
|
||||
This avoids sharing mutable instance state through V3's sanitized class clones.
|
||||
"""
|
||||
import comfy.samplers
|
||||
from comfy_api.v0_0_2 import io
|
||||
|
||||
from . import utilities as _utilities
|
||||
from .collector import DistributedCollectorNode as _CollectorRuntime
|
||||
from .distributed_upscale import UltimateSDUpscaleDistributed as _UpscaleRuntime
|
||||
|
||||
|
||||
class DistributedSeed(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedSeed', display_name='Distributed Seed', category='utils',
|
||||
inputs=[io.Int.Input('seed', default=1125899906842, min=0,
|
||||
max=1125899906842624, force_input=False)],
|
||||
outputs=[io.Int.Output(display_name='seed')],
|
||||
accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, seed, is_worker=False, worker_id=''):
|
||||
return io.NodeOutput(*_utilities.DistributedSeed().distribute(seed, is_worker, worker_id))
|
||||
|
||||
|
||||
class DistributedValue(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedValue', display_name='Distributed Value', category='utils',
|
||||
inputs=[io.String.Input('default_value', default=''),
|
||||
io.String.Input('worker_values', default='{}')],
|
||||
outputs=[io.AnyType.Output(display_name='value')],
|
||||
accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, default_value, worker_values='{}', is_worker=False, worker_id=''):
|
||||
return io.NodeOutput(*_utilities.DistributedValue().distribute(
|
||||
default_value, worker_values, is_worker, worker_id))
|
||||
|
||||
|
||||
class DistributedModelName(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedModelName', display_name='Distributed Model Name', category='utils',
|
||||
inputs=[io.String.Input('text', default='')],
|
||||
outputs=[io.AnyType.Output(display_name='output')],
|
||||
hidden=[io.Hidden.unique_id, io.Hidden.extra_pnginfo], is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, text):
|
||||
result = _utilities.DistributedModelName().log_input(
|
||||
text, unique_id=cls.hidden.unique_id, extra_pnginfo=cls.hidden.extra_pnginfo)
|
||||
return io.NodeOutput(*result['result'], ui=result['ui'])
|
||||
|
||||
|
||||
class ImageBatchDivider(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='ImageBatchDivider', display_name='Image Batch Divider', category='image',
|
||||
inputs=[io.Image.Input('images'),
|
||||
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
|
||||
display_mode=io.NumberDisplay.number,
|
||||
tooltip='Number of parts to divide the batch into')],
|
||||
# The existing frontend still displays only divide_by sockets.
|
||||
outputs=[io.Image.Output(display_name=f'batch_{index + 1}') for index in range(10)],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images, divide_by):
|
||||
return io.NodeOutput(*_utilities.ImageBatchDivider().divide_batch(images, divide_by))
|
||||
|
||||
|
||||
class AudioBatchDivider(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='AudioBatchDivider', display_name='Audio Segment Divider', category='audio',
|
||||
inputs=[io.Audio.Input('audio'),
|
||||
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
|
||||
display_mode=io.NumberDisplay.number,
|
||||
tooltip='Number of sequential time segments to create')],
|
||||
outputs=[io.Audio.Output(display_name=f'audio_{index + 1}') for index in range(10)],
|
||||
is_output_node=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, audio, divide_by):
|
||||
return io.NodeOutput(*_utilities.AudioBatchDivider().divide_audio(audio, divide_by))
|
||||
|
||||
|
||||
class DistributedEmptyImage(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedEmptyImage', display_name='Distributed Empty Image', category='image',
|
||||
inputs=[io.Int.Input('height', default=64, min=1, max=4096, step=1),
|
||||
io.Int.Input('width', default=64, min=1, max=4096, step=1),
|
||||
io.Int.Input('channels', default=3, min=1, max=4, step=1)],
|
||||
outputs=[io.Image.Output()],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, height, width, channels):
|
||||
return io.NodeOutput(*_utilities.DistributedEmptyImage().create(height, width, channels))
|
||||
|
||||
|
||||
class DistributedCollector(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='DistributedCollector', display_name='Distributed Collector', category='image',
|
||||
inputs=[io.Boolean.Input('load_balance', default=False,
|
||||
tooltip='Run this workflow on one least-busy participant (master included when participating).'),
|
||||
io.Image.Input('images', optional=True), io.Audio.Input('audio', optional=True)],
|
||||
outputs=[io.Image.Output(display_name='images'), io.Audio.Output(display_name='audio')],
|
||||
is_input_list=True, accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, images=None, load_balance=False, audio=None, multi_job_id='',
|
||||
is_worker=False, master_url='', enabled_worker_ids='[]', worker_batch_size=1,
|
||||
worker_id='', pass_through=False, delegate_only=False):
|
||||
return io.NodeOutput(*_CollectorRuntime().run(
|
||||
images=images, load_balance=load_balance, audio=audio, multi_job_id=multi_job_id,
|
||||
is_worker=is_worker, master_url=master_url, enabled_worker_ids=enabled_worker_ids,
|
||||
worker_batch_size=worker_batch_size, worker_id=worker_id,
|
||||
pass_through=pass_through, delegate_only=delegate_only))
|
||||
|
||||
|
||||
class UltimateSDUpscaleDistributed(io.ComfyNode):
|
||||
@classmethod
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id='UltimateSDUpscaleDistributed',
|
||||
display_name='Ultimate SD Upscale Distributed (No Upscale)', category='image/upscaling',
|
||||
inputs=[
|
||||
io.Image.Input('upscaled_image'), io.Model.Input('model'),
|
||||
io.Conditioning.Input('positive'), io.Conditioning.Input('negative'), io.Vae.Input('vae'),
|
||||
io.Int.Input('seed', default=0, min=0, max=0xffffffffffffffff),
|
||||
io.Int.Input('steps', default=20, min=1, max=10000),
|
||||
io.Float.Input('cfg', default=8.0, min=0.0, max=100.0),
|
||||
io.Combo.Input('sampler_name', options=comfy.samplers.KSampler.SAMPLERS),
|
||||
io.Combo.Input('scheduler', options=comfy.samplers.KSampler.SCHEDULERS),
|
||||
io.Float.Input('denoise', default=0.5, min=0.0, max=1.0, step=0.01),
|
||||
io.Int.Input('tile_width', default=512, min=64, max=2048, step=8),
|
||||
io.Int.Input('tile_height', default=512, min=64, max=2048, step=8),
|
||||
io.Int.Input('padding', default=32, min=0, max=256, step=8),
|
||||
io.Int.Input('mask_blur', default=8, min=0, max=256),
|
||||
io.Boolean.Input('force_uniform_tiles', default=True),
|
||||
io.Boolean.Input('tiled_decode', default=False),
|
||||
], outputs=[io.Image.Output()], accept_all_inputs=True,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def fingerprint_inputs(cls, **kwargs):
|
||||
return _UpscaleRuntime.IS_CHANGED(**kwargs)
|
||||
|
||||
@classmethod
|
||||
def execute(cls, upscaled_image, model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
|
||||
mask_blur, force_uniform_tiles, tiled_decode, multi_job_id='', is_worker=False,
|
||||
master_url='', enabled_worker_ids='[]', worker_id='', tile_indices='', dynamic_threshold=8):
|
||||
return io.NodeOutput(*_UpscaleRuntime().run(
|
||||
upscaled_image, model, positive, negative, vae, seed, steps, cfg,
|
||||
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
|
||||
mask_blur, force_uniform_tiles, tiled_decode, multi_job_id, is_worker,
|
||||
master_url, enabled_worker_ids, worker_id, tile_indices, dynamic_threshold))
|
||||
|
||||
|
||||
NODES = [DistributedCollector, DistributedSeed, DistributedModelName, DistributedValue,
|
||||
ImageBatchDivider, AudioBatchDivider, DistributedEmptyImage, UltimateSDUpscaleDistributed]
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "ComfyUI-Distributed"
|
||||
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
|
||||
version = "1.4.1"
|
||||
version = "1.5.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = []
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Internal extension lifecycle support (not a node provider)."""
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Initialize routes, distributed state and the existing worker lifecycle."""
|
||||
import atexit
|
||||
import os
|
||||
|
||||
import server
|
||||
|
||||
from ..utils.config import CONFIG_FILE, ensure_config_exists
|
||||
from ..utils.logging import debug_log
|
||||
from ..workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
from ..upscale.job_store import ensure_tile_jobs_initialized
|
||||
|
||||
_initialized = False
|
||||
|
||||
|
||||
def initialize():
|
||||
"""Called by ComfyExtension.on_load; initialize once per loaded package."""
|
||||
global _initialized
|
||||
if _initialized:
|
||||
return
|
||||
|
||||
ensure_config_exists()
|
||||
from .. import api # noqa: F401 - registers the existing @routes.* handlers
|
||||
from ..api.queue_orchestration import ensure_distributed_state
|
||||
|
||||
ensure_distributed_state(server.PromptServer.instance)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
if not os.environ.get('COMFYUI_IS_WORKER'):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
|
||||
_initialized = True
|
||||
debug_log('Loaded Distributed nodes.')
|
||||
debug_log(f'Config file: {CONFIG_FILE}')
|
||||
@@ -162,7 +162,7 @@ def _load_job_routes_module():
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
queue_orchestration_module = types.ModuleType(f"{package_name}.api.queue_orchestration")
|
||||
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 1))
|
||||
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 7, 1, {}))
|
||||
sys.modules[f"{package_name}.api.queue_orchestration"] = queue_orchestration_module
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -221,7 +221,7 @@ job_routes = _load_job_routes_module()
|
||||
|
||||
|
||||
class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_id(self):
|
||||
async def test_distributed_queue_happy_path_returns_prompt_metadata(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"prompt": {"1": {"class_type": "Node"}},
|
||||
@@ -233,12 +233,14 @@ class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
|
||||
with patch.object(
|
||||
job_routes,
|
||||
"orchestrate_distributed_execution",
|
||||
new=AsyncMock(return_value=("prompt_123", 2)),
|
||||
new=AsyncMock(return_value=("prompt_123", 42, 2, {})),
|
||||
):
|
||||
response = await job_routes.distributed_queue_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
self.assertEqual(response.payload.get("prompt_id"), "prompt_123")
|
||||
self.assertEqual(response.payload.get("number"), 42)
|
||||
self.assertEqual(response.payload.get("node_errors"), {})
|
||||
self.assertTrue(response.payload.get("auto_prepare_supported"))
|
||||
|
||||
async def test_distributed_queue_missing_prompt_returns_400(self):
|
||||
@@ -300,6 +302,77 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
self.assertEqual(queued["audio"]["sample_rate"], 44100)
|
||||
self.assertEqual(tuple(queued["audio"]["waveform"].shape), (1, 2, 4))
|
||||
|
||||
async def test_job_complete_accepts_audio_without_image(self):
|
||||
queue = asyncio.Queue()
|
||||
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
job_routes.prompt_server.distributed_pending_jobs = {"audio-only-job": queue}
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "audio-only-job",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"audio": self._encoded_audio_payload(),
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
with patch.object(job_routes, "_decode_canonical_png_tensor") as decode_image:
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 200)
|
||||
decode_image.assert_not_called()
|
||||
queued = await queue.get()
|
||||
self.assertIsNone(queued["tensor"])
|
||||
self.assertEqual(queued["audio"]["sample_rate"], 44100)
|
||||
|
||||
async def test_job_complete_rejects_payload_without_image_or_audio(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("image or audio", response.payload.get("message", "").lower())
|
||||
|
||||
async def test_job_complete_rejects_invalid_image_even_when_audio_is_present(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"image": 123,
|
||||
"audio": self._encoded_audio_payload(),
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("image", response.payload.get("message", "").lower())
|
||||
|
||||
async def test_job_complete_returns_400_for_malformed_audio(self):
|
||||
request = _FakeRequest(
|
||||
{
|
||||
"job_id": "job-1",
|
||||
"worker_id": "worker-1",
|
||||
"batch_idx": 0,
|
||||
"audio": {"data": "AAAA", "shape": [1, 2], "dtype": "float32"},
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
|
||||
response = await job_routes.job_complete_endpoint(request)
|
||||
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("audio.shape", response.payload.get("message", "").lower())
|
||||
|
||||
def test_decode_audio_payload_rejects_bad_shape(self):
|
||||
bad = {
|
||||
"sample_rate": 44100,
|
||||
|
||||
@@ -164,26 +164,11 @@ class ConvertPathsForPlatformTests(unittest.TestCase):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class FindMediaReferencesTests(unittest.TestCase):
|
||||
def test_finds_image_input(self):
|
||||
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.png"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("photo.png", refs)
|
||||
|
||||
def test_finds_video_input(self):
|
||||
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"video": "clip.mp4"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("clip.mp4", refs)
|
||||
|
||||
def test_finds_file_input_for_load_video(self):
|
||||
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"file": "1 - Copy.mp4"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("1 - Copy.mp4", refs)
|
||||
|
||||
def test_finds_audio_input(self):
|
||||
prompt = {"1": {"class_type": "LoadAudio", "inputs": {"audio": "track.wav"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
self.assertIn("track.wav", refs)
|
||||
|
||||
def test_strips_annotation_suffix(self):
|
||||
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.jpg [abc123]"}}}
|
||||
refs = ms._find_media_references(prompt)
|
||||
|
||||
@@ -106,7 +106,7 @@ def _load_worker_routes_module():
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
workers_pkg.__path__ = [str(module_path.parents[1] / "workers")]
|
||||
workers_pkg.get_worker_manager = lambda: _DummyWorkerManager()
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
@@ -213,6 +213,7 @@ def _load_worker_routes_module():
|
||||
raise RuntimeError("not used in these tests")
|
||||
|
||||
network_module.get_client_session = _get_client_session
|
||||
network_module.get_server_port = lambda: 8189
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
@@ -269,6 +270,38 @@ worker_routes = _load_worker_routes_module()
|
||||
|
||||
|
||||
class WorkerRoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_network_info_reports_actual_master_and_available_ports(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 1, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=["127.0.0.1"]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": []}), \
|
||||
patch.object(worker_routes, "get_server_port", return_value=8189), \
|
||||
patch.object(worker_routes, "allocate_worker_ports", return_value=[8191, 8192, 8193]) as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["master_port"], 8189)
|
||||
self.assertEqual(response.payload["local_worker_ports"], [8191, 8192, 8193])
|
||||
self.assertEqual(response.payload["cuda_device_count"], 4)
|
||||
allocate.assert_called_once_with(8189, [], 3)
|
||||
|
||||
async def test_network_info_does_not_reallocate_existing_configuration(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 4, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=[]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": [{"id": "existing", "port": 8189}]}), \
|
||||
patch.object(worker_routes, "allocate_worker_ports") as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["local_worker_ports"], [])
|
||||
allocate.assert_not_called()
|
||||
|
||||
async def test_launch_conflict_is_a_clear_client_error(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8189}]}
|
||||
with patch.object(worker_routes, "get_worker_manager", return_value=manager), \
|
||||
patch.object(worker_routes, "load_config", return_value=config), \
|
||||
patch.object(manager, "launch_worker", side_effect=ValueError("Worker port 8189 conflicts with the master port 8189")):
|
||||
response = await worker_routes.launch_worker_endpoint(_FakeRequest({"worker_id": "worker-a"}))
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("master port 8189", response.payload["message"])
|
||||
self.assertEqual(manager.processes, {})
|
||||
|
||||
async def test_launch_worker_valid_id_returns_200(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8188}]}
|
||||
|
||||
+691
@@ -0,0 +1,691 @@
|
||||
{
|
||||
"DistributedCollector": {
|
||||
"input": {
|
||||
"required": {
|
||||
"load_balance": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false,
|
||||
"tooltip": "Run this workflow on one least-busy participant (master included when participating)."
|
||||
}
|
||||
]
|
||||
},
|
||||
"optional": {
|
||||
"images": [
|
||||
"IMAGE"
|
||||
],
|
||||
"audio": [
|
||||
"AUDIO"
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"master_url": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"enabled_worker_ids": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "[]"
|
||||
}
|
||||
],
|
||||
"worker_batch_size": [
|
||||
"INT",
|
||||
{
|
||||
"default": 1,
|
||||
"min": 1,
|
||||
"max": 1024
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"pass_through": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"delegate_only": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"load_balance"
|
||||
],
|
||||
"optional": [
|
||||
"images",
|
||||
"audio"
|
||||
],
|
||||
"hidden": [
|
||||
"multi_job_id",
|
||||
"is_worker",
|
||||
"master_url",
|
||||
"enabled_worker_ids",
|
||||
"worker_batch_size",
|
||||
"worker_id",
|
||||
"pass_through",
|
||||
"delegate_only"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE",
|
||||
"AUDIO"
|
||||
],
|
||||
"output_name": [
|
||||
"images",
|
||||
"audio"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": true,
|
||||
"output_node": false,
|
||||
"category": "image",
|
||||
"display_name": "Distributed Collector"
|
||||
},
|
||||
"DistributedSeed": {
|
||||
"input": {
|
||||
"required": {
|
||||
"seed": [
|
||||
"INT",
|
||||
{
|
||||
"default": 1125899906842,
|
||||
"min": 0,
|
||||
"max": 1125899906842624,
|
||||
"forceInput": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"seed"
|
||||
],
|
||||
"hidden": [
|
||||
"is_worker",
|
||||
"worker_id"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"INT"
|
||||
],
|
||||
"output_name": [
|
||||
"seed"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Seed"
|
||||
},
|
||||
"DistributedModelName": {
|
||||
"input": {
|
||||
"required": {
|
||||
"text": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"unique_id": "UNIQUE_ID",
|
||||
"extra_pnginfo": "EXTRA_PNGINFO"
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"text"
|
||||
],
|
||||
"hidden": [
|
||||
"unique_id",
|
||||
"extra_pnginfo"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"output"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Model Name"
|
||||
},
|
||||
"DistributedValue": {
|
||||
"input": {
|
||||
"required": {
|
||||
"default_value": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"worker_values": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "{}"
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"default_value",
|
||||
"worker_values"
|
||||
],
|
||||
"hidden": [
|
||||
"is_worker",
|
||||
"worker_id"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"value"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "utils",
|
||||
"display_name": "Distributed Value"
|
||||
},
|
||||
"ImageBatchDivider": {
|
||||
"input": {
|
||||
"required": {
|
||||
"images": [
|
||||
"IMAGE"
|
||||
],
|
||||
"divide_by": [
|
||||
"INT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 1,
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of parts to divide the batch into"
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"images",
|
||||
"divide_by"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"batch_1",
|
||||
"batch_2",
|
||||
"batch_3",
|
||||
"batch_4",
|
||||
"batch_5",
|
||||
"batch_6",
|
||||
"batch_7",
|
||||
"batch_8",
|
||||
"batch_9",
|
||||
"batch_10"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "image",
|
||||
"display_name": "Image Batch Divider"
|
||||
},
|
||||
"AudioBatchDivider": {
|
||||
"input": {
|
||||
"required": {
|
||||
"audio": [
|
||||
"AUDIO"
|
||||
],
|
||||
"divide_by": [
|
||||
"INT",
|
||||
{
|
||||
"default": 2,
|
||||
"min": 1,
|
||||
"max": 10,
|
||||
"step": 1,
|
||||
"display": "number",
|
||||
"tooltip": "Number of sequential time segments to create"
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"audio",
|
||||
"divide_by"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*",
|
||||
"*"
|
||||
],
|
||||
"output_name": [
|
||||
"audio_1",
|
||||
"audio_2",
|
||||
"audio_3",
|
||||
"audio_4",
|
||||
"audio_5",
|
||||
"audio_6",
|
||||
"audio_7",
|
||||
"audio_8",
|
||||
"audio_9",
|
||||
"audio_10"
|
||||
],
|
||||
"output_is_list": [
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false,
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": true,
|
||||
"category": "audio",
|
||||
"display_name": "Audio Segment Divider"
|
||||
},
|
||||
"DistributedEmptyImage": {
|
||||
"input": {
|
||||
"required": {
|
||||
"height": [
|
||||
"INT",
|
||||
{
|
||||
"default": 64,
|
||||
"min": 1,
|
||||
"max": 4096,
|
||||
"step": 1
|
||||
}
|
||||
],
|
||||
"width": [
|
||||
"INT",
|
||||
{
|
||||
"default": 64,
|
||||
"min": 1,
|
||||
"max": 4096,
|
||||
"step": 1
|
||||
}
|
||||
],
|
||||
"channels": [
|
||||
"INT",
|
||||
{
|
||||
"default": 3,
|
||||
"min": 1,
|
||||
"max": 4,
|
||||
"step": 1
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"height",
|
||||
"width",
|
||||
"channels"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_name": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "image",
|
||||
"display_name": "Distributed Empty Image"
|
||||
},
|
||||
"UltimateSDUpscaleDistributed": {
|
||||
"input": {
|
||||
"required": {
|
||||
"upscaled_image": [
|
||||
"IMAGE"
|
||||
],
|
||||
"model": [
|
||||
"MODEL"
|
||||
],
|
||||
"positive": [
|
||||
"CONDITIONING"
|
||||
],
|
||||
"negative": [
|
||||
"CONDITIONING"
|
||||
],
|
||||
"vae": [
|
||||
"VAE"
|
||||
],
|
||||
"seed": [
|
||||
"INT",
|
||||
{
|
||||
"default": 0,
|
||||
"min": 0,
|
||||
"max": 18446744073709551615
|
||||
}
|
||||
],
|
||||
"steps": [
|
||||
"INT",
|
||||
{
|
||||
"default": 20,
|
||||
"min": 1,
|
||||
"max": 10000
|
||||
}
|
||||
],
|
||||
"cfg": [
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 8.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0
|
||||
}
|
||||
],
|
||||
"sampler_name": [
|
||||
[
|
||||
"euler",
|
||||
"euler_cfg_pp",
|
||||
"euler_ancestral",
|
||||
"euler_ancestral_cfg_pp",
|
||||
"heun",
|
||||
"heunpp2",
|
||||
"exp_heun_2_x0",
|
||||
"exp_heun_2_x0_sde",
|
||||
"dpm_2",
|
||||
"dpm_2_ancestral",
|
||||
"lms",
|
||||
"dpm_fast",
|
||||
"dpm_adaptive",
|
||||
"dpmpp_2s_ancestral",
|
||||
"dpmpp_2s_ancestral_cfg_pp",
|
||||
"dpmpp_sde",
|
||||
"dpmpp_sde_gpu",
|
||||
"dpmpp_2m",
|
||||
"dpmpp_2m_cfg_pp",
|
||||
"dpmpp_2m_sde",
|
||||
"dpmpp_2m_sde_gpu",
|
||||
"dpmpp_2m_sde_heun",
|
||||
"dpmpp_2m_sde_heun_gpu",
|
||||
"dpmpp_3m_sde",
|
||||
"dpmpp_3m_sde_gpu",
|
||||
"ddpm",
|
||||
"lcm",
|
||||
"ipndm",
|
||||
"ipndm_v",
|
||||
"deis",
|
||||
"cfgpp_ud10_ab",
|
||||
"res_multistep",
|
||||
"res_multistep_cfg_pp",
|
||||
"res_multistep_ancestral",
|
||||
"res_multistep_ancestral_cfg_pp",
|
||||
"gradient_estimation",
|
||||
"gradient_estimation_cfg_pp",
|
||||
"er_sde",
|
||||
"seeds_2",
|
||||
"seeds_3",
|
||||
"sa_solver",
|
||||
"sa_solver_pece",
|
||||
"ddim",
|
||||
"uni_pc",
|
||||
"uni_pc_bh2"
|
||||
]
|
||||
],
|
||||
"scheduler": [
|
||||
[
|
||||
"simple",
|
||||
"sgm_uniform",
|
||||
"karras",
|
||||
"exponential",
|
||||
"ddim_uniform",
|
||||
"beta",
|
||||
"normal",
|
||||
"linear_quadratic",
|
||||
"kl_optimal"
|
||||
]
|
||||
],
|
||||
"denoise": [
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 0.5,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"step": 0.01
|
||||
}
|
||||
],
|
||||
"tile_width": [
|
||||
"INT",
|
||||
{
|
||||
"default": 512,
|
||||
"min": 64,
|
||||
"max": 2048,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"tile_height": [
|
||||
"INT",
|
||||
{
|
||||
"default": 512,
|
||||
"min": 64,
|
||||
"max": 2048,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"padding": [
|
||||
"INT",
|
||||
{
|
||||
"default": 32,
|
||||
"min": 0,
|
||||
"max": 256,
|
||||
"step": 8
|
||||
}
|
||||
],
|
||||
"mask_blur": [
|
||||
"INT",
|
||||
{
|
||||
"default": 8,
|
||||
"min": 0,
|
||||
"max": 256
|
||||
}
|
||||
],
|
||||
"force_uniform_tiles": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": true
|
||||
}
|
||||
],
|
||||
"tiled_decode": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
]
|
||||
},
|
||||
"hidden": {
|
||||
"multi_job_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"is_worker": [
|
||||
"BOOLEAN",
|
||||
{
|
||||
"default": false
|
||||
}
|
||||
],
|
||||
"master_url": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"enabled_worker_ids": [
|
||||
"STRING",
|
||||
{
|
||||
"default": "[]"
|
||||
}
|
||||
],
|
||||
"worker_id": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"tile_indices": [
|
||||
"STRING",
|
||||
{
|
||||
"default": ""
|
||||
}
|
||||
],
|
||||
"dynamic_threshold": [
|
||||
"INT",
|
||||
{
|
||||
"default": 8,
|
||||
"min": 1,
|
||||
"max": 64
|
||||
}
|
||||
]
|
||||
}
|
||||
},
|
||||
"input_order": {
|
||||
"required": [
|
||||
"upscaled_image",
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
"vae",
|
||||
"seed",
|
||||
"steps",
|
||||
"cfg",
|
||||
"sampler_name",
|
||||
"scheduler",
|
||||
"denoise",
|
||||
"tile_width",
|
||||
"tile_height",
|
||||
"padding",
|
||||
"mask_blur",
|
||||
"force_uniform_tiles",
|
||||
"tiled_decode"
|
||||
],
|
||||
"hidden": [
|
||||
"multi_job_id",
|
||||
"is_worker",
|
||||
"master_url",
|
||||
"enabled_worker_ids",
|
||||
"worker_id",
|
||||
"tile_indices",
|
||||
"dynamic_threshold"
|
||||
]
|
||||
},
|
||||
"output": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_name": [
|
||||
"IMAGE"
|
||||
],
|
||||
"output_is_list": [
|
||||
false
|
||||
],
|
||||
"is_input_list": false,
|
||||
"output_node": false,
|
||||
"category": "image/upscaling",
|
||||
"display_name": "Ultimate SD Upscale Distributed (No Upscale)"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,252 @@
|
||||
"""Exercise the actual loader, schemas and executor in a fresh CPU process.
|
||||
|
||||
No HTTP listener, workers, model downloads or live installation changes.
|
||||
The original contracts were captured from 32ac027 using the same core.
|
||||
"""
|
||||
import asyncio
|
||||
import inspect
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import tempfile
|
||||
from unittest.mock import patch
|
||||
|
||||
from PIL import Image
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
COMFY_ROOT = Path(sys.argv[1]).resolve()
|
||||
sys.path.insert(0, str(COMFY_ROOT))
|
||||
os.environ['COMFYUI_IS_WORKER'] = '1'
|
||||
import comfy.cli_args
|
||||
comfy.cli_args.args.cpu = True
|
||||
comfy.cli_args.args.disable_assets = True
|
||||
from app.assets.manager import default_asset_manager
|
||||
from comfy_api.v0_0_2 import io
|
||||
import execution
|
||||
import nodes
|
||||
import server
|
||||
import torch
|
||||
|
||||
BASELINE = json.loads((ROOT / 'tests/fixtures/v1_node_contracts.json').read_text())
|
||||
|
||||
|
||||
def normalized_input(value):
|
||||
kind = value[0]
|
||||
opts = dict(value[1]) if len(value) > 1 else {}
|
||||
if kind == 'STRING':
|
||||
# V3 serializes the same default single-line widget explicitly.
|
||||
opts.setdefault('multiline', False)
|
||||
if kind == 'COMBO':
|
||||
kind = opts['options']
|
||||
opts = {key: val for key, val in opts.items() if key != 'options'}
|
||||
if isinstance(kind, list):
|
||||
# Single-selection is the V1 dropdown default as well.
|
||||
opts.setdefault('multiselect', False)
|
||||
return [kind, opts]
|
||||
|
||||
|
||||
async def map_node(cls, values, extra=None):
|
||||
"""Use the actual input parser and sanitized V3 executor clones."""
|
||||
prepared, missing, hidden = execution.get_input_data(values, cls, unique_id='probe', extra_data=extra or {})
|
||||
assert not missing, missing
|
||||
# Use a worker thread for synchronous nodes: collector/upscale bridge back
|
||||
# to PromptServer's running loop, just as ComfyUI's prompt worker does.
|
||||
def run():
|
||||
return asyncio.run(execution._async_map_node_over_list(
|
||||
'v3-acceptance', 'probe', cls, prepared, cls.FUNCTION, v3_data=hidden))
|
||||
return await asyncio.to_thread(run)
|
||||
|
||||
|
||||
def check_saved_workflows(mapping):
|
||||
checked = 0
|
||||
occurrences = 0
|
||||
for path in sorted((ROOT / 'workflows').glob('*.json')):
|
||||
workflow = json.loads(path.read_text())
|
||||
local_nodes = {node['id']: node for node in workflow['nodes']}
|
||||
for saved in workflow['nodes']:
|
||||
if saved['type'] not in mapping:
|
||||
continue
|
||||
occurrences += 1
|
||||
info = mapping[saved['type']].GET_NODE_INFO_V1()
|
||||
inputs = {**info['input'].get('required', {}), **info['input'].get('optional', {})}
|
||||
for socket in saved.get('inputs', []):
|
||||
assert socket['name'] in inputs, (path.name, saved['id'], socket)
|
||||
assert socket['type'] == inputs[socket['name']][0], (path.name, socket)
|
||||
for index, output in enumerate(saved.get('outputs', [])):
|
||||
assert output['type'] == info['output'][index], (path.name, index)
|
||||
for link in workflow['links']:
|
||||
if link[1] == saved['id']:
|
||||
assert link[2] < len(info['output']), (path.name, link)
|
||||
target = local_nodes[link[3]]['inputs'][link[4]]
|
||||
assert target['type'] == info['output'][link[2]], (path.name, link)
|
||||
checked += 1
|
||||
assert checked == 5 and occurrences == 10, (checked, occurrences)
|
||||
print('SAVED_WORKFLOW_CONTRACTS_OK', checked, occurrences)
|
||||
|
||||
|
||||
async def check_execution(mapping, module, prompt_server, asset_manager):
|
||||
result = await map_node(mapping['DistributedSeed'],
|
||||
{'seed': 123, 'is_worker': True, 'worker_id': 'worker_2'})
|
||||
assert result[0].result == (126,)
|
||||
values = json.dumps({'_type': 'INT', '3': '17'})
|
||||
result = await map_node(mapping['DistributedValue'],
|
||||
{'default_value': '4', 'worker_values': values,
|
||||
'is_worker': True, 'worker_id': 'worker_2'})
|
||||
assert result[0].result == (17,)
|
||||
metadata = {'workflow': {'nodes': [{'id': 'probe', 'widgets_values': []}]}}
|
||||
result = await map_node(mapping['DistributedModelName'], {'text': 'model.ckpt'},
|
||||
{'extra_pnginfo': metadata})
|
||||
assert result[0].result == ('model.ckpt',)
|
||||
assert result[0].ui == {'text': ['model.ckpt']}
|
||||
assert metadata['workflow']['nodes'][0]['widgets_values'] == [['model.ckpt']]
|
||||
result = await map_node(mapping['DistributedEmptyImage'],
|
||||
{'height': 8, 'width': 8, 'channels': 3})
|
||||
empty = result[0].result[0]
|
||||
assert empty.shape == (0, 8, 8, 3) and empty.numel() == 0
|
||||
images = torch.arange(10 * 8 * 8 * 3, dtype=torch.float32).reshape(10, 8, 8, 3)
|
||||
for count in (1, 3, 10):
|
||||
result = await map_node(mapping['ImageBatchDivider'], {'images': images, 'divide_by': count})
|
||||
assert len(result[0].result) == 10
|
||||
assert torch.equal(torch.cat(result[0].result[:count]), images)
|
||||
audio = {'waveform': torch.arange(33, dtype=torch.float32).reshape(1, 1, 33), 'sample_rate': 24000}
|
||||
result = await map_node(mapping['AudioBatchDivider'], {'audio': audio, 'divide_by': 3})
|
||||
assert len(result[0].result) == 10
|
||||
assert torch.equal(torch.cat([item['waveform'] for item in result[0].result[:3]], dim=-1), audio['waveform'])
|
||||
# INPUT_IS_LIST must preserve all images/audio but unwrap transport scalars.
|
||||
collector = mapping['DistributedCollector']
|
||||
prepared, missing, hidden = execution.get_input_data(
|
||||
{'load_balance': False, 'multi_job_id': '', 'pass_through': True}, collector, 'collector-list')
|
||||
assert not missing
|
||||
prepared['images'] = [images[:2], images[2:]]
|
||||
prepared['audio'] = [audio, audio]
|
||||
result = await execution._async_map_node_over_list(
|
||||
'v3-acceptance', 'collector-list', collector, prepared, collector.FUNCTION, v3_data=hidden)
|
||||
assert torch.equal(result[0].result[0], images)
|
||||
assert result[0].result[1]['waveform'].shape[-1] == 66
|
||||
# Exercise real async aggregation without an HTTP endpoint or worker.
|
||||
queue = asyncio.Queue()
|
||||
await queue.put({'worker_id': 'worker_2', 'tensor': images[2:3], 'image_index': 0, 'is_last': True})
|
||||
prompt_server.distributed_pending_jobs['v3-aggregate'] = queue
|
||||
result = await map_node(collector,
|
||||
{'images': images[:2], 'load_balance': False,
|
||||
'multi_job_id': 'v3-aggregate', 'enabled_worker_ids': '["worker_2"]'})
|
||||
assert torch.equal(result[0].result[0], images[:3])
|
||||
assert 'v3-aggregate' not in prompt_server.distributed_pending_jobs
|
||||
|
||||
# Prove V3's per-execution class clones do not leak helper instance state;
|
||||
# replace only the GPU/model boundary, not the input parser or executor.
|
||||
upscale_runtime = sys.modules[module.__name__ + '.nodes.distributed_upscale'].UltimateSDUpscaleDistributed
|
||||
seen = []
|
||||
def fake_upscale(self, *args):
|
||||
assert not hasattr(self, 'acceptance_marker')
|
||||
self.acceptance_marker = True
|
||||
seen.append(args)
|
||||
return (args[0],)
|
||||
inputs = {'upscaled_image': images, 'model': object(), 'positive': [], 'negative': [],
|
||||
'vae': object(), 'seed': 1, 'steps': 1, 'cfg': 1.0,
|
||||
'sampler_name': 'euler', 'scheduler': 'normal', 'denoise': 0.5,
|
||||
'tile_width': 64, 'tile_height': 64, 'padding': 0, 'mask_blur': 0,
|
||||
'force_uniform_tiles': True, 'tiled_decode': False,
|
||||
'multi_job_id': 'tile-job', 'is_worker': True, 'master_url': 'http://master.invalid',
|
||||
'enabled_worker_ids': '["worker_2"]', 'worker_id': 'worker_2',
|
||||
'tile_indices': '[2,3]', 'dynamic_threshold': 9}
|
||||
with patch.object(upscale_runtime, 'run', fake_upscale):
|
||||
for _ in range(2):
|
||||
result = await map_node(mapping['UltimateSDUpscaleDistributed'], inputs)
|
||||
assert result[0].result[0] is images
|
||||
assert len(seen) == 2 and seen[0][-7:] == tuple(inputs[name] for name in (
|
||||
'multi_job_id', 'is_worker', 'master_url', 'enabled_worker_ids', 'worker_id',
|
||||
'tile_indices', 'dynamic_threshold')), seen[0]
|
||||
import math
|
||||
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id='tile-job'))
|
||||
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id=''))
|
||||
|
||||
graph = {
|
||||
'1': {'class_type': 'EmptyImage', 'inputs': {'height': 8, 'width': 8, 'batch_size': 10, 'color': 0}},
|
||||
'2': {'class_type': 'DistributedCollector', 'inputs': {'images': ['1', 0], 'load_balance': False,
|
||||
'pass_through': True}},
|
||||
'3': {'class_type': 'ImageBatchDivider', 'inputs': {'images': ['2', 0], 'divide_by': 10}},
|
||||
'4': {'class_type': 'PreviewImage', 'inputs': {'images': ['3', 9]}},
|
||||
}
|
||||
valid = await execution.validate_prompt('v3-graph', graph, None)
|
||||
assert valid[0], valid
|
||||
invalid = {'1': {'class_type': 'UltimateSDUpscaleDistributed',
|
||||
'inputs': {**inputs, 'sampler_name': 'INVALID_ENUM', 'scheduler': 'INVALID_ENUM'}},
|
||||
'2': {'class_type': 'PreviewImage', 'inputs': {'images': ['1', 0]}}}
|
||||
rejected = await execution.validate_prompt('v3-invalid', invalid, None)
|
||||
assert not rejected[0]
|
||||
errors = [error for entry in rejected[3].values() for error in entry['errors']]
|
||||
enum_errors = [error['extra_info']['input_name'] for error in errors if error['type'] == 'value_not_in_list']
|
||||
assert {'sampler_name', 'scheduler'} <= set(enum_errors), errors
|
||||
import folder_paths
|
||||
scratch = Path(os.environ.get('TMPDIR', Path.home() / '.hermes/cache/scratch'))
|
||||
with tempfile.TemporaryDirectory(prefix='v3-preview-', dir=scratch) as temp:
|
||||
with patch.object(folder_paths, 'temp_directory', temp):
|
||||
executor = execution.PromptExecutor(
|
||||
prompt_server, cache_args={'ram': 0, 'ram_inactive': 0}, asset_manager=asset_manager)
|
||||
await asyncio.to_thread(executor.execute, graph, 'v3-graph', {}, valid[2])
|
||||
assert executor.success, executor.status_messages
|
||||
history = executor.history_result
|
||||
record = history['outputs']['4']['images'][0]
|
||||
preview = Path(temp) / record.get('subfolder', '') / record['filename']
|
||||
with Image.open(preview) as image:
|
||||
assert image.size == (8, 8) and image.mode == 'RGB'
|
||||
assert image.getextrema() == ((0, 0), (0, 0), (0, 0))
|
||||
print('V3_EXECUTION_OK eight nodes; upscale GPU boundary mocked; preview artifact verified')
|
||||
|
||||
|
||||
async def main():
|
||||
asset_manager = default_asset_manager()
|
||||
prompt_server = server.PromptServer(asyncio.get_running_loop(), asset_manager)
|
||||
assert not (ROOT / 'distributed.py').exists(), 'obsolete root bootstrap remains'
|
||||
assert await nodes.load_custom_node(str(ROOT)), 'ComfyUI loader rejected the pack'
|
||||
module = sys.modules[str(ROOT).replace('.', '_x_')]
|
||||
assert not hasattr(module, 'NODE_CLASS_MAPPINGS'), 'V1 map shadows V3 entrypoint'
|
||||
extension = await module.comfy_entrypoint()
|
||||
classes = await extension.get_node_list()
|
||||
mapping = {cls.GET_SCHEMA().node_id: cls for cls in classes}
|
||||
assert len(classes) == len(mapping) == len(BASELINE) == 8
|
||||
assert set(mapping) == set(BASELINE)
|
||||
for node_id, cls in mapping.items():
|
||||
assert issubclass(cls, io.ComfyNode)
|
||||
assert nodes.NODE_CLASS_MAPPINGS[node_id] is cls
|
||||
old = dict(BASELINE[node_id])
|
||||
if node_id in ('ImageBatchDivider', 'AudioBatchDivider'):
|
||||
# V1's ByPassTypeTuple advertises '*' when indexed, while its
|
||||
# underlying tuple and existing frontend declare IMAGE/AUDIO.
|
||||
# Native V3 declares all ten existing typed sockets explicitly.
|
||||
assert old['output'] == ['*'] * 10
|
||||
old['output'] = ['IMAGE' if node_id == 'ImageBatchDivider' else 'AUDIO'] * 10
|
||||
new = json.loads(json.dumps(cls.GET_NODE_INFO_V1()))
|
||||
for group in ('required', 'optional'):
|
||||
old_inputs = old['input'].get(group, {})
|
||||
new_inputs = new['input'].get(group, {})
|
||||
assert list(old_inputs) == list(new_inputs), (node_id, group, 'input order')
|
||||
for name, original in old_inputs.items():
|
||||
assert normalized_input(original) == normalized_input(new_inputs[name]), (node_id, name, original, new_inputs[name])
|
||||
for key in ('output', 'output_name', 'output_is_list', 'is_input_list', 'output_node', 'category', 'display_name'):
|
||||
assert old[key] == new[key], (node_id, key, old[key], new[key])
|
||||
# Standard context lives in cls.hidden; orchestrator metadata remains
|
||||
# accepted by its original kwarg name, without creating new widgets.
|
||||
signature = inspect.signature(cls.execute)
|
||||
for name, field in old['input'].get('hidden', {}).items():
|
||||
if isinstance(field, list):
|
||||
assert cls.GET_SCHEMA().accept_all_inputs, node_id
|
||||
assert name in signature.parameters, (node_id, name)
|
||||
assert signature.parameters[name].default == field[1]['default'], (node_id, name)
|
||||
else:
|
||||
assert name in new['input']['hidden'], (node_id, name)
|
||||
expected_hidden = []
|
||||
if node_id == 'DistributedModelName':
|
||||
expected_hidden.extend(['unique_id', 'extra_pnginfo'])
|
||||
if old['output_node']:
|
||||
expected_hidden.extend(name for name in ['prompt', 'extra_pnginfo'] if name not in expected_hidden)
|
||||
assert list(new['input'].get('hidden', {})) == expected_hidden, (node_id, new['input'].get('hidden'))
|
||||
print('SCHEMA_PARITY_OK', len(mapping))
|
||||
check_saved_workflows(mapping)
|
||||
await check_execution(mapping, module, prompt_server, asset_manager)
|
||||
print('V3_ACCEPTANCE_OK')
|
||||
|
||||
|
||||
asyncio.run(main())
|
||||
@@ -0,0 +1,91 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
class _PromptQueue:
|
||||
def __init__(self):
|
||||
self.items = []
|
||||
|
||||
def put(self, item):
|
||||
self.items.append(item)
|
||||
|
||||
|
||||
def _load_async_helpers_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "utils" / "async_helpers.py"
|
||||
package_name = "dist_async_helpers_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
execution_module = types.ModuleType("execution")
|
||||
|
||||
async def _validate_prompt(prompt_id, prompt, partial_execution_targets):
|
||||
return (True, None, ["9"], {})
|
||||
|
||||
execution_module.validate_prompt = _validate_prompt
|
||||
execution_module.SENSITIVE_EXTRA_DATA_KEYS = []
|
||||
sys.modules["execution"] = execution_module
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
trigger_on_prompt=lambda payload: payload,
|
||||
number=12,
|
||||
prompt_queue=_PromptQueue(),
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.get_server_loop = lambda: None
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.utils.async_helpers", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module, prompt_server
|
||||
|
||||
|
||||
async_helpers, prompt_server = _load_async_helpers_module()
|
||||
|
||||
|
||||
class QueuePromptPayloadTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_queue_prompt_payload_includes_create_time_and_client_metadata(self):
|
||||
result = await async_helpers.queue_prompt_payload(
|
||||
{"1": {"class_type": "Node"}},
|
||||
workflow_meta={"id": "workflow-1"},
|
||||
client_id="client-1",
|
||||
include_queue_metadata=True,
|
||||
)
|
||||
|
||||
self.assertIsInstance(result["prompt_id"], str)
|
||||
self.assertTrue(result["prompt_id"])
|
||||
self.assertEqual(result["number"], 12)
|
||||
self.assertEqual(result["node_errors"], {})
|
||||
|
||||
self.assertEqual(prompt_server.number, 13)
|
||||
self.assertEqual(len(prompt_server.prompt_queue.items), 1)
|
||||
queued_item = prompt_server.prompt_queue.items[0]
|
||||
self.assertEqual(queued_item[0], 12)
|
||||
extra_data = queued_item[3]
|
||||
self.assertEqual(extra_data["client_id"], "client-1")
|
||||
self.assertIn("create_time", extra_data)
|
||||
self.assertIsInstance(extra_data["create_time"], int)
|
||||
self.assertGreater(extra_data["create_time"], 0)
|
||||
self.assertEqual(extra_data["extra_pnginfo"]["workflow"], {"id": "workflow-1"})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,364 @@
|
||||
import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
def _load_collector_module():
|
||||
module_path = Path(__file__).resolve().parents[1] / "nodes" / "collector.py"
|
||||
package_name = "dist_collector_list_testpkg"
|
||||
|
||||
for mod_name in list(sys.modules):
|
||||
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
|
||||
del sys.modules[mod_name]
|
||||
|
||||
root_pkg = types.ModuleType(package_name)
|
||||
root_pkg.__path__ = []
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
|
||||
nodes_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.nodes"] = nodes_pkg
|
||||
|
||||
utils_pkg = types.ModuleType(f"{package_name}.utils")
|
||||
utils_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
class _Routes:
|
||||
def post(self, _path):
|
||||
return lambda fn: fn
|
||||
|
||||
def get(self, _path):
|
||||
return lambda fn: fn
|
||||
|
||||
prompt_server = types.SimpleNamespace(
|
||||
routes=_Routes(),
|
||||
distributed_jobs_lock=None,
|
||||
distributed_pending_jobs={},
|
||||
)
|
||||
server_module = types.ModuleType("server")
|
||||
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
|
||||
sys.modules["server"] = server_module
|
||||
|
||||
comfy_module = types.ModuleType("comfy")
|
||||
model_management = types.ModuleType("comfy.model_management")
|
||||
|
||||
class InterruptProcessingException(Exception):
|
||||
pass
|
||||
|
||||
model_management.InterruptProcessingException = InterruptProcessingException
|
||||
model_management.throw_exception_if_processing_interrupted = lambda: None
|
||||
comfy_module.model_management = model_management
|
||||
|
||||
comfy_utils = types.ModuleType("comfy.utils")
|
||||
|
||||
class ProgressBar:
|
||||
def __init__(self, _total):
|
||||
self.total = _total
|
||||
self.updates = []
|
||||
|
||||
def update(self, value):
|
||||
self.updates.append(value)
|
||||
|
||||
comfy_utils.ProgressBar = ProgressBar
|
||||
comfy_module.utils = comfy_utils
|
||||
sys.modules["comfy"] = comfy_module
|
||||
sys.modules["comfy.model_management"] = model_management
|
||||
sys.modules["comfy.utils"] = comfy_utils
|
||||
|
||||
aiohttp_module = types.ModuleType("aiohttp")
|
||||
aiohttp_module.ClientTimeout = lambda total: types.SimpleNamespace(total=total)
|
||||
sys.modules["aiohttp"] = aiohttp_module
|
||||
|
||||
logging_module = types.ModuleType(f"{package_name}.utils.logging")
|
||||
logging_module.debug_log = lambda *_args, **_kwargs: None
|
||||
logging_module.log = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.logging"] = logging_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.get_worker_timeout_seconds = lambda: 0.1
|
||||
config_module.load_config = lambda: {"workers": []}
|
||||
config_module.is_master_delegate_only = lambda: False
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.HEARTBEAT_INTERVAL = 1.0
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
|
||||
image_module = types.ModuleType(f"{package_name}.utils.image")
|
||||
def _ensure_contiguous(tensor):
|
||||
return tensor.contiguous() if hasattr(tensor, "contiguous") else tensor
|
||||
|
||||
image_module.ensure_contiguous = _ensure_contiguous
|
||||
image_module.tensor_to_pil = lambda *_args, **_kwargs: None
|
||||
image_module.pil_to_tensor = lambda value: value
|
||||
sys.modules[f"{package_name}.utils.image"] = image_module
|
||||
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.build_worker_url = lambda worker: "http://worker"
|
||||
network_module.get_client_session = lambda: None
|
||||
network_module.probe_worker = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
audio_payload_module = types.ModuleType(f"{package_name}.utils.audio_payload")
|
||||
audio_payload_module.encode_audio_payload = lambda audio: audio
|
||||
sys.modules[f"{package_name}.utils.audio_payload"] = audio_payload_module
|
||||
|
||||
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
|
||||
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
|
||||
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(f"{package_name}.nodes.collector", module_path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
assert spec is not None and spec.loader is not None
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def test_collector_opts_into_comfyui_list_inputs():
|
||||
collector = _load_collector_module().DistributedCollectorNode
|
||||
|
||||
assert collector.INPUT_IS_LIST is True
|
||||
|
||||
|
||||
def test_collector_exposes_images_as_optional_input():
|
||||
input_types = _load_collector_module().DistributedCollectorNode.INPUT_TYPES()
|
||||
|
||||
assert "images" not in input_types["required"]
|
||||
assert input_types["optional"]["images"] == ("IMAGE",)
|
||||
|
||||
|
||||
def test_audio_only_pass_through_returns_no_images_and_preserves_audio():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
|
||||
|
||||
images, returned_audio = collector.run(images=None, audio=[audio])
|
||||
|
||||
assert images is None
|
||||
assert returned_audio is audio
|
||||
|
||||
|
||||
def test_collector_rejects_missing_images_and_audio():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
|
||||
try:
|
||||
collector.run(images=None, audio=None)
|
||||
except ValueError as exc:
|
||||
assert "image or audio" in str(exc).lower()
|
||||
else:
|
||||
raise AssertionError("Expected collector to reject a run with no media input")
|
||||
|
||||
|
||||
def test_delegate_only_master_allows_no_local_media_input():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
|
||||
images, audio = collector.run(
|
||||
images=None,
|
||||
audio=None,
|
||||
multi_job_id=["delegate-audio-job"],
|
||||
delegate_only=[True],
|
||||
enabled_worker_ids=["[]"],
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
|
||||
|
||||
def test_pass_through_collapses_comfyui_image_list_to_batch_and_unwraps_hidden_inputs():
|
||||
collector = _load_collector_module().DistributedCollectorNode()
|
||||
first = torch.zeros(1, 2, 2, 3)
|
||||
second = torch.ones(1, 2, 2, 3)
|
||||
|
||||
images, audio = collector.run(
|
||||
images=[first, second],
|
||||
load_balance=[False],
|
||||
audio=[None],
|
||||
multi_job_id=[""],
|
||||
is_worker=[False],
|
||||
master_url=[""],
|
||||
enabled_worker_ids=["[]"],
|
||||
worker_batch_size=[1],
|
||||
worker_id=[""],
|
||||
pass_through=[False],
|
||||
delegate_only=[False],
|
||||
)
|
||||
|
||||
assert tuple(images.shape) == (2, 2, 2, 3)
|
||||
assert torch.equal(images[0:1], first)
|
||||
assert torch.equal(images[1:2], second)
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
|
||||
|
||||
def test_worker_list_input_sends_one_completion_sequence_with_last_only_on_final_item():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
first = torch.zeros(1, 2, 2, 3)
|
||||
second = torch.ones(1, 2, 2, 3)
|
||||
posted_payloads = []
|
||||
|
||||
class _FakeImage:
|
||||
def save(self, fp, format=None, compress_level=None):
|
||||
fp.write(b"png-bytes")
|
||||
|
||||
class _FakeResponse:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
class _FakeSession:
|
||||
def post(self, url, json, timeout):
|
||||
posted_payloads.append(json)
|
||||
return _FakeResponse()
|
||||
|
||||
async def _fake_get_client_session():
|
||||
return _FakeSession()
|
||||
|
||||
module.tensor_to_pil = lambda *_args, **_kwargs: _FakeImage()
|
||||
module.get_client_session = _fake_get_client_session
|
||||
module.encode_audio_payload = lambda _audio: None
|
||||
|
||||
images, audio = collector.run(
|
||||
images=[first, second],
|
||||
load_balance=[False],
|
||||
audio=[None],
|
||||
multi_job_id=["job-list-1"],
|
||||
is_worker=[True],
|
||||
master_url=["http://master"],
|
||||
enabled_worker_ids=["[]"],
|
||||
worker_batch_size=[1],
|
||||
worker_id=["worker-a"],
|
||||
pass_through=[False],
|
||||
delegate_only=[False],
|
||||
)
|
||||
|
||||
assert tuple(images.shape) == (2, 2, 2, 3)
|
||||
assert tuple(audio["waveform"].shape) == (1, 2, 1)
|
||||
assert len(posted_payloads) == 2
|
||||
assert [payload["batch_idx"] for payload in posted_payloads] == [0, 1]
|
||||
assert [payload["is_last"] for payload in posted_payloads] == [False, True]
|
||||
assert {payload["job_id"] for payload in posted_payloads} == {"job-list-1"}
|
||||
assert {payload["worker_id"] for payload in posted_payloads} == {"worker-a"}
|
||||
|
||||
|
||||
def test_audio_only_worker_sends_one_completion_without_image():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
|
||||
posted = []
|
||||
|
||||
class _FakeResponse:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
class _FakeSession:
|
||||
def post(self, url, json, timeout):
|
||||
posted.append((url, json, timeout.total))
|
||||
return _FakeResponse()
|
||||
|
||||
async def _fake_get_client_session():
|
||||
return _FakeSession()
|
||||
|
||||
module.get_client_session = _fake_get_client_session
|
||||
module.encode_audio_payload = lambda value: {"encoded": value is audio}
|
||||
|
||||
images, returned_audio = collector.run(
|
||||
images=None,
|
||||
audio=[audio],
|
||||
multi_job_id=["audio-job"],
|
||||
is_worker=[True],
|
||||
master_url=["http://master"],
|
||||
worker_id=["worker-a"],
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert returned_audio is audio
|
||||
assert len(posted) == 1
|
||||
assert posted[0][0] == "http://master/distributed/job_complete"
|
||||
assert posted[0][2] == 600
|
||||
assert posted[0][1] == {
|
||||
"job_id": "audio-job",
|
||||
"worker_id": "worker-a",
|
||||
"batch_idx": 0,
|
||||
"audio": {"encoded": True},
|
||||
"is_last": True,
|
||||
}
|
||||
|
||||
|
||||
def test_audio_only_master_combines_local_and_worker_audio():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
master_audio = {"waveform": torch.ones(1, 2, 2), "sample_rate": 48000}
|
||||
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
|
||||
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
queue = asyncio.Queue()
|
||||
queue.put_nowait(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"image_index": 0,
|
||||
"tensor": None,
|
||||
"audio": worker_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
module.prompt_server.distributed_pending_jobs = {"audio-job": queue}
|
||||
|
||||
images, combined_audio = asyncio.run(
|
||||
collector.execute(
|
||||
images=None,
|
||||
audio=master_audio,
|
||||
multi_job_id="audio-job",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
)
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert combined_audio["sample_rate"] == 48000
|
||||
assert tuple(combined_audio["waveform"].shape) == (1, 2, 5)
|
||||
assert torch.equal(combined_audio["waveform"][..., :2], master_audio["waveform"])
|
||||
assert torch.equal(combined_audio["waveform"][..., 2:], worker_audio["waveform"])
|
||||
|
||||
|
||||
def test_delegate_only_audio_collects_worker_audio_without_placeholder_image():
|
||||
module = _load_collector_module()
|
||||
collector = module.DistributedCollectorNode()
|
||||
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
|
||||
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
|
||||
queue = asyncio.Queue()
|
||||
queue.put_nowait(
|
||||
{
|
||||
"worker_id": "worker-a",
|
||||
"image_index": 0,
|
||||
"tensor": None,
|
||||
"audio": worker_audio,
|
||||
"is_last": True,
|
||||
}
|
||||
)
|
||||
module.prompt_server.distributed_pending_jobs = {"delegate-audio-job": queue}
|
||||
|
||||
images, combined_audio = asyncio.run(
|
||||
collector.execute(
|
||||
images=None,
|
||||
audio=None,
|
||||
multi_job_id="delegate-audio-job",
|
||||
enabled_worker_ids='["worker-a"]',
|
||||
delegate_only=True,
|
||||
)
|
||||
)
|
||||
|
||||
assert images is None
|
||||
assert combined_audio["sample_rate"] == 48000
|
||||
assert torch.equal(combined_audio["waveform"], worker_audio["waveform"])
|
||||
@@ -86,27 +86,25 @@ class ParseTilesFromFormTests(unittest.TestCase):
|
||||
|
||||
# --- happy paths ---
|
||||
|
||||
def test_single_tile_returns_one_entry(self):
|
||||
def test_single_tile_returns_image_and_metadata(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
self.assertEqual(len(tiles), 1)
|
||||
|
||||
def test_multiple_tiles_all_returned(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(len(tiles), 3)
|
||||
|
||||
def test_tile_image_is_pil_image(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
self.assertIsInstance(tiles[0]["image"], PILImage.Image)
|
||||
|
||||
def test_tile_metadata_fields_are_parsed(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1))
|
||||
tile = tiles[0]
|
||||
self.assertIsInstance(tile["image"], PILImage.Image)
|
||||
self.assertEqual(tile["tile_idx"], 0)
|
||||
self.assertEqual(tile["x"], 0)
|
||||
self.assertEqual(tile["y"], 0)
|
||||
self.assertEqual(tile["extracted_width"], 64)
|
||||
self.assertEqual(tile["extracted_height"], 64)
|
||||
|
||||
def test_multiple_tiles_preserve_count_order_and_coordinates(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(len(tiles), 3)
|
||||
for i, tile in enumerate(tiles):
|
||||
self.assertEqual(tile["tile_idx"], i)
|
||||
self.assertEqual(tiles[1]["x"], 64)
|
||||
self.assertEqual(tiles[2]["x"], 128)
|
||||
|
||||
def test_padding_is_parsed_from_form(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(1, padding=16))
|
||||
self.assertEqual(tiles[0]["padding"], 16)
|
||||
@@ -138,16 +136,6 @@ class ParseTilesFromFormTests(unittest.TestCase):
|
||||
self.assertNotIn("batch_idx", tiles[0])
|
||||
self.assertNotIn("global_idx", tiles[0])
|
||||
|
||||
def test_tile_indices_match_metadata_order(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
for i, tile in enumerate(tiles):
|
||||
self.assertEqual(tile["tile_idx"], i)
|
||||
|
||||
def test_x_coordinates_reflect_metadata(self):
|
||||
tiles = pp._parse_tiles_from_form(_make_form(3))
|
||||
self.assertEqual(tiles[1]["x"], 64)
|
||||
self.assertEqual(tiles[2]["x"], 128)
|
||||
|
||||
# --- error cases ---
|
||||
|
||||
def test_missing_tiles_metadata_raises_value_error(self):
|
||||
|
||||
@@ -71,6 +71,15 @@ def _collector_only_prompt():
|
||||
}
|
||||
|
||||
|
||||
def _audio_only_prompt():
|
||||
"""1(Audio source) → 2(DistributedCollector) → 3(Audio sink)."""
|
||||
return {
|
||||
"1": {"class_type": "LoadAudio", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {"audio": ["1", 0]}},
|
||||
"3": {"class_type": "SaveAudio", "inputs": {"audio": ["2", 1]}},
|
||||
}
|
||||
|
||||
|
||||
def _delegate_prompt():
|
||||
"""1 → 2 → 3(DistributedCollector) → 4(SaveImage)"""
|
||||
return {
|
||||
@@ -217,6 +226,35 @@ class PrunePromptForWorkerTests(unittest.TestCase):
|
||||
self.assertEqual(len(preview_nodes), 1)
|
||||
self.assertEqual(preview_nodes[0]["inputs"]["images"], ["4", 0])
|
||||
|
||||
def test_injects_preview_audio_for_audio_only_collector(self):
|
||||
result = pt.prune_prompt_for_worker(_audio_only_prompt())
|
||||
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewAudio"]
|
||||
self.assertEqual(len(preview_nodes), 1)
|
||||
self.assertEqual(preview_nodes[0]["inputs"]["audio"], ["2", 1])
|
||||
self.assertFalse(any(n.get("class_type") == "PreviewImage" for n in result.values()))
|
||||
|
||||
def test_prefers_image_preview_when_collector_has_images_and_audio(self):
|
||||
prompt = _linear_prompt()
|
||||
prompt["6"] = {"class_type": "LoadAudio", "inputs": {}}
|
||||
prompt["4"]["inputs"]["audio"] = ["6", 0]
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
self.assertEqual(
|
||||
len([n for n in result.values() if n.get("class_type") == "PreviewImage"]),
|
||||
1,
|
||||
)
|
||||
self.assertFalse(any(n.get("class_type") == "PreviewAudio" for n in result.values()))
|
||||
|
||||
def test_preserves_image_preview_for_distributed_upscale(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "LoadImage", "inputs": {}},
|
||||
"2": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"upscaled_image": ["1", 0]}},
|
||||
"3": {"class_type": "SaveImage", "inputs": {"images": ["2", 0]}},
|
||||
}
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
|
||||
self.assertEqual(len(preview_nodes), 1)
|
||||
self.assertEqual(preview_nodes[0]["inputs"]["images"], ["2", 0])
|
||||
|
||||
def test_no_preview_image_when_no_downstream(self):
|
||||
result = pt.prune_prompt_for_worker(_collector_only_prompt())
|
||||
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
|
||||
@@ -242,7 +280,7 @@ class PrunePromptForWorkerTests(unittest.TestCase):
|
||||
def test_upscale_node_is_treated_as_distributed(self):
|
||||
prompt = {
|
||||
"1": {"class_type": "KSampler", "inputs": {}},
|
||||
"2": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"image": ["1", 0]}},
|
||||
"2": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"upscaled_image": ["1", 0]}},
|
||||
"3": {"class_type": "SaveImage", "inputs": {"images": ["2", 0]}},
|
||||
}
|
||||
result = pt.prune_prompt_for_worker(prompt)
|
||||
@@ -264,8 +302,8 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
|
||||
self.assertNotIn("1", result)
|
||||
self.assertNotIn("2", result)
|
||||
|
||||
def test_removes_dangling_upstream_refs(self):
|
||||
"""Collector must not retain dangling refs to pruned upstream nodes."""
|
||||
def test_replaces_dangling_upstream_ref_with_one_empty_image_placeholder(self):
|
||||
"""Collector must point to exactly one valid placeholder after pruning."""
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
collector_inputs = result["3"].get("inputs", {})
|
||||
@@ -276,26 +314,520 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
|
||||
self.assertNotEqual(source_id, "2")
|
||||
self.assertIn(source_id, result)
|
||||
self.assertEqual(result[source_id].get("class_type"), "DistributedEmptyImage")
|
||||
|
||||
def test_injects_empty_image_placeholder(self):
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
empty_nodes = [(nid, n) for nid, n in result.items() if n.get("class_type") == "DistributedEmptyImage"]
|
||||
self.assertEqual(len(empty_nodes), 1)
|
||||
placeholder_id = empty_nodes[0][0]
|
||||
self.assertEqual(result["3"]["inputs"]["images"], [placeholder_id, 0])
|
||||
|
||||
def test_audio_only_collector_does_not_get_image_placeholder(self):
|
||||
prompt = _audio_only_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["2"])
|
||||
empty_nodes = [n for n in result.values() if n.get("class_type") == "DistributedEmptyImage"]
|
||||
self.assertEqual(empty_nodes, [])
|
||||
self.assertNotIn("images", result["2"].get("inputs", {}))
|
||||
self.assertNotIn("audio", result["2"].get("inputs", {}))
|
||||
|
||||
def test_one_placeholder_per_collector(self):
|
||||
"""Two collectors → two placeholders."""
|
||||
prompt = {
|
||||
"1": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {}},
|
||||
"1": {"class_type": "DistributedCollector", "inputs": {"images": ["10", 0]}},
|
||||
"2": {"class_type": "DistributedCollector", "inputs": {"images": ["11", 0]}},
|
||||
"3": {"class_type": "SaveImage", "inputs": {"images": ["1", 0]}},
|
||||
"10": {"class_type": "LoadImage", "inputs": {}},
|
||||
"11": {"class_type": "LoadImage", "inputs": {}},
|
||||
}
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["1", "2"])
|
||||
empty_nodes = [n for n in result.values() if n.get("class_type") == "DistributedEmptyImage"]
|
||||
self.assertEqual(len(empty_nodes), 2)
|
||||
|
||||
def test_preserves_primitive_string_for_downstream_required_input(self):
|
||||
"""Delegate-only master keeps primitive inputs needed by SaveImage."""
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {"SaveImage": save_image}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "test/input_bug"}},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["15", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertIn("15", result)
|
||||
self.assertEqual(result["15"], prompt["15"])
|
||||
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["15", 0])
|
||||
|
||||
def test_preserves_load_image_for_switch_alternate_required_input(self):
|
||||
"""Delegate-only master keeps LoadImage inputs needed by switches."""
|
||||
prompt = {
|
||||
"848": {"class_type": "LoadImage", "inputs": {"image": "input_bug2_00003_.png"}},
|
||||
"862": {"class_type": "VAEDecode", "inputs": {}},
|
||||
"854": {"class_type": "DistributedCollector", "inputs": {"images": ["862", 0]}},
|
||||
"850": {
|
||||
"class_type": "ComfySwitchNode",
|
||||
"inputs": {
|
||||
"on_false": ["848", 0],
|
||||
"on_true": ["854", 0],
|
||||
},
|
||||
},
|
||||
"851": {"class_type": "PreviewImage", "inputs": {"images": ["850", 0]}},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["854"])
|
||||
|
||||
self.assertIn("848", result)
|
||||
self.assertEqual(result["850"]["inputs"]["on_false"], ["848", 0])
|
||||
self.assertEqual(result["850"]["inputs"]["on_true"], ["854", 0])
|
||||
self.assertNotIn("862", result)
|
||||
|
||||
def test_preserves_registered_string_utility_subgraph_for_downstream_required_input(self):
|
||||
"""Delegate-only master keeps scalar utility chains used by SaveImage."""
|
||||
string_concat = type(
|
||||
"StringConcatenate",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("STRING",),
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"string_a": ("STRING",),
|
||||
"string_b": ("STRING",),
|
||||
}
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"SaveImage": save_image,
|
||||
"StringConcatenate": string_concat,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "test"}},
|
||||
"16": {"class_type": "PrimitiveString", "inputs": {"value": "input_bug3"}},
|
||||
"17": {
|
||||
"class_type": "StringConcatenate",
|
||||
"inputs": {
|
||||
"string_a": ["15", 0],
|
||||
"string_b": ["16", 0],
|
||||
},
|
||||
},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["17", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertIn("15", result)
|
||||
self.assertIn("16", result)
|
||||
self.assertIn("17", result)
|
||||
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["17", 0])
|
||||
self.assertEqual(result["17"]["inputs"]["string_a"], ["15", 0])
|
||||
self.assertEqual(result["17"]["inputs"]["string_b"], ["16", 0])
|
||||
|
||||
def test_preserves_registered_multi_string_join_subgraph_for_downstream_required_input(self):
|
||||
"""Delegate-only master keeps multi-input scalar utility chains."""
|
||||
join_string_multi = type(
|
||||
"JoinStringMulti",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("STRING",),
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {"string_1": ("STRING",)},
|
||||
"optional": {
|
||||
"string_2": ("STRING",),
|
||||
"string_3": ("STRING",),
|
||||
},
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"JoinStringMulti": join_string_multi,
|
||||
"SaveImage": save_image,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "test"}},
|
||||
"16": {"class_type": "PrimitiveString", "inputs": {"value": "input"}},
|
||||
"17": {"class_type": "PrimitiveString", "inputs": {"value": "bug4"}},
|
||||
"18": {
|
||||
"class_type": "JoinStringMulti",
|
||||
"inputs": {
|
||||
"string_1": ["15", 0],
|
||||
"string_2": ["16", 0],
|
||||
"string_3": ["17", 0],
|
||||
},
|
||||
},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["18", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertIn("15", result)
|
||||
self.assertIn("16", result)
|
||||
self.assertIn("17", result)
|
||||
self.assertIn("18", result)
|
||||
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["18", 0])
|
||||
|
||||
def test_does_not_preserve_scalar_utility_with_heavy_upstream_dependency(self):
|
||||
"""Scalar utility nodes are retained only when their full input branch is safe."""
|
||||
string_concat = type(
|
||||
"StringConcatenate",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("STRING",),
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"string_a": ("STRING",),
|
||||
"string_b": ("STRING",),
|
||||
}
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"SaveImage": save_image,
|
||||
"StringConcatenate": string_concat,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "test"}},
|
||||
"17": {
|
||||
"class_type": "StringConcatenate",
|
||||
"inputs": {
|
||||
"string_a": ["15", 0],
|
||||
"string_b": ["8", 0],
|
||||
},
|
||||
},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["17", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertNotIn("8", result)
|
||||
self.assertNotIn("17", result)
|
||||
self.assertNotIn("filename_prefix", result["9"]["inputs"])
|
||||
|
||||
def test_does_not_preserve_scalar_output_for_non_scalar_downstream_input(self):
|
||||
"""Scalar outputs are retained only for scalar/config downstream inputs."""
|
||||
string_provider = type("StringProvider", (), {"RETURN_TYPES": ("STRING",)})
|
||||
image_consumer = type(
|
||||
"ImageConsumer",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"mask": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"ImageConsumer": image_consumer,
|
||||
"StringProvider": string_provider,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"17": {"class_type": "StringProvider", "inputs": {}},
|
||||
"9": {
|
||||
"class_type": "ImageConsumer",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"mask": ["17", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertNotIn("17", result)
|
||||
self.assertNotIn("mask", result["9"]["inputs"])
|
||||
|
||||
def test_preserves_scalar_list_join_subgraph_for_downstream_required_input(self):
|
||||
"""Delegate-only master keeps list-shaped scalar config chains."""
|
||||
create_list = type(
|
||||
"CreateList",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("LIST",),
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {"inputs.input0": ("STRING",)},
|
||||
"optional": {
|
||||
"inputs.input1": ("STRING",),
|
||||
"inputs.input2": ("STRING",),
|
||||
},
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
string_data_list_join = type(
|
||||
"StringDataListJoin",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("STRING",),
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"strings": ("LIST",),
|
||||
"delimiter": ("STRING",),
|
||||
}
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"CreateList": create_list,
|
||||
"SaveImage": save_image,
|
||||
"StringDataListJoin": string_data_list_join,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "test"}},
|
||||
"17": {"class_type": "PrimitiveString", "inputs": {"value": "input_bug"}},
|
||||
"29": {"class_type": "PrimitiveString", "inputs": {"value": "new"}},
|
||||
"28": {
|
||||
"class_type": "CreateList",
|
||||
"inputs": {
|
||||
"inputs.input0": ["15", 0],
|
||||
"inputs.input1": ["17", 0],
|
||||
"inputs.input2": ["29", 0],
|
||||
},
|
||||
},
|
||||
"32": {
|
||||
"class_type": "StringDataListJoin",
|
||||
"inputs": {
|
||||
"strings": ["28", 0],
|
||||
"delimiter": "/",
|
||||
},
|
||||
},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["32", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
self.assertIn("15", result)
|
||||
self.assertIn("17", result)
|
||||
self.assertIn("28", result)
|
||||
self.assertIn("29", result)
|
||||
self.assertIn("32", result)
|
||||
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["32", 0])
|
||||
self.assertEqual(result["32"]["inputs"]["strings"], ["28", 0])
|
||||
self.assertEqual(result["28"]["inputs"]["inputs.input0"], ["15", 0])
|
||||
self.assertNotIn("8", result)
|
||||
|
||||
def test_preserves_builtin_create_list_string_data_list_join_subgraph(self):
|
||||
"""Delegate-only master handles ComfyUI 0.23 CreateList data-list output."""
|
||||
string_data_list_join = type(
|
||||
"StringDataListJoin",
|
||||
(),
|
||||
{
|
||||
"RETURN_TYPES": ("STRING",),
|
||||
"INPUT_IS_LIST": True,
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"strings": ("STRING", {"forceInput": True}),
|
||||
"sep": ("STRING", {"default": " "}),
|
||||
}
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
save_image = type(
|
||||
"SaveImage",
|
||||
(),
|
||||
{
|
||||
"INPUT_TYPES": classmethod(
|
||||
lambda cls: {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
"filename_prefix": ("STRING",),
|
||||
}
|
||||
}
|
||||
)
|
||||
},
|
||||
)
|
||||
previous_mappings = getattr(pt, "_DELEGATE_MASTER_NODE_CLASS_MAPPINGS", None)
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = {
|
||||
"Basic data handling: StringDataListJoin": string_data_list_join,
|
||||
"SaveImage": save_image,
|
||||
}
|
||||
try:
|
||||
prompt = {
|
||||
"8": {"class_type": "KSampler", "inputs": {}},
|
||||
"11": {"class_type": "DistributedCollector", "inputs": {"images": ["8", 0]}},
|
||||
"15": {"class_type": "PrimitiveString", "inputs": {"value": "input_bug"}},
|
||||
"17": {"class_type": "PrimitiveString", "inputs": {"value": "test"}},
|
||||
"29": {"class_type": "PrimitiveString", "inputs": {"value": "new"}},
|
||||
"28": {
|
||||
"class_type": "CreateList",
|
||||
"inputs": {
|
||||
"inputs.input0": ["17", 0],
|
||||
"inputs.input1": ["15", 0],
|
||||
"inputs.input2": ["29", 0],
|
||||
},
|
||||
},
|
||||
"32": {
|
||||
"class_type": "Basic data handling: StringDataListJoin",
|
||||
"inputs": {
|
||||
"strings": ["28", 0],
|
||||
"sep": "/",
|
||||
},
|
||||
},
|
||||
"9": {
|
||||
"class_type": "SaveImage",
|
||||
"inputs": {
|
||||
"images": ["11", 0],
|
||||
"filename_prefix": ["32", 0],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["11"])
|
||||
finally:
|
||||
pt._DELEGATE_MASTER_NODE_CLASS_MAPPINGS = previous_mappings
|
||||
|
||||
for node_id in ("15", "17", "28", "29", "32"):
|
||||
self.assertIn(node_id, result)
|
||||
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["32", 0])
|
||||
self.assertEqual(result["32"]["inputs"]["strings"], ["28", 0])
|
||||
self.assertEqual(result["28"]["inputs"]["inputs.input0"], ["17", 0])
|
||||
self.assertNotIn("8", result)
|
||||
|
||||
def test_result_is_independent_copy(self):
|
||||
prompt = _delegate_prompt()
|
||||
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
"""Real-framework acceptance; opt in with COMFYUI_SOURCE_ROOT."""
|
||||
import os
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def test_native_v3_runtime():
|
||||
comfy_root = os.environ.get('COMFYUI_SOURCE_ROOT')
|
||||
if not comfy_root:
|
||||
pytest.skip('Set COMFYUI_SOURCE_ROOT to run native ComfyUI V3 acceptance')
|
||||
helper = Path(__file__).parent / 'helpers' / 'v3_runtime_check.py'
|
||||
result = subprocess.run(
|
||||
[sys.executable, str(helper), str(Path(comfy_root).resolve())],
|
||||
capture_output=True, text=True, timeout=120,
|
||||
env={**os.environ, 'COMFYUI_IS_WORKER': '1'},
|
||||
)
|
||||
assert result.returncode == 0, result.stdout + '\n' + result.stderr
|
||||
assert 'V3_ACCEPTANCE_OK' in result.stdout
|
||||
@@ -0,0 +1,113 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import socket
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def load_ports():
|
||||
path = Path(__file__).resolve().parents[1] / "workers" / "ports.py"
|
||||
spec = importlib.util.spec_from_file_location("worker_ports_test", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
class WorkerPortTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.ports = load_ports()
|
||||
|
||||
def test_allocation_starts_above_actual_master_and_skips_assigned_and_occupied(self):
|
||||
workers = [
|
||||
{"id": "existing", "host": "localhost", "port": 8190, "enabled": False},
|
||||
{"id": "remote", "host": "remote.example", "port": 8192},
|
||||
]
|
||||
with patch.object(self.ports, "is_port_available", side_effect=lambda port: port != 8191):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, workers, 3), [8192, 8193, 8194])
|
||||
|
||||
def test_local_host_forms_reserve_ports(self):
|
||||
for host in ["::1", "[::1]", "http://localhost/", "HTTPS://LOCALHOST", "0.0.0.0", None]:
|
||||
with self.subTest(host=host), patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, [{"id": "local", "host": host, "port": 8190}], 1), [8191])
|
||||
|
||||
def test_exhaustion_is_explicit_not_a_partial_allocation(self):
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
with self.assertRaisesRegex(ValueError, "available.*ports"):
|
||||
self.ports.allocate_worker_ports(65534, [], 2)
|
||||
|
||||
def test_launch_rejects_master_port_without_reassigning_configuration(self):
|
||||
worker = {"id": "manual", "port": 8189}
|
||||
with self.assertRaisesRegex(ValueError, "master.*8189"):
|
||||
self.ports.validate_worker_port(worker, 8189, [])
|
||||
self.assertEqual(worker["port"], 8189)
|
||||
|
||||
def test_launch_rejects_another_local_workers_reserved_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "other", "host": None, "port": 8190, "enabled": False}
|
||||
with self.assertRaisesRegex(ValueError, "other"):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_ignores_itself_and_remote_workers_with_same_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "remote", "host": "example.com", "port": 8190}
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_rejects_an_occupied_port(self):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener:
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen()
|
||||
port = listener.getsockname()[1]
|
||||
self.assertFalse(self.ports.is_port_available(port))
|
||||
with self.assertRaisesRegex(ValueError, "already in use"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
@unittest.skipIf(os.name == "nt", "asyncio does not reuse addresses on Windows")
|
||||
def test_recently_closed_connection_does_not_block_worker_restart(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
port = listener.getsockname()[1]
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=2) as client:
|
||||
accepted, _ = listener.accept()
|
||||
with accepted:
|
||||
accepted.settimeout(2)
|
||||
accepted.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(client.recv(1), b"")
|
||||
client.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(accepted.recv(1), b"")
|
||||
self.assertTrue(self.ports.is_port_available(port))
|
||||
self.ports.validate_worker_port({"id": "restart", "port": port}, 1, [])
|
||||
with socket.socket() as restarted:
|
||||
restarted.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
restarted.bind(("127.0.0.1", port))
|
||||
restarted.listen(1)
|
||||
|
||||
def test_active_listener_with_reuseaddr_is_still_a_conflict(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
@unittest.skipUnless(socket.has_ipv6, "IPv6 unavailable")
|
||||
def test_ipv6_only_listener_is_still_a_conflict(self):
|
||||
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as listener:
|
||||
listener.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
|
||||
try:
|
||||
listener.bind(("::1", 0))
|
||||
except OSError as exc:
|
||||
self.skipTest(f"IPv6 loopback unavailable: {exc}")
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
def test_invalid_ports_fail_clearly(self):
|
||||
for port in [0, 65536, "not-a-port", None, True, 8190.5]:
|
||||
with self.subTest(port=port), self.assertRaisesRegex(ValueError, "port"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,11 +1,16 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import types
|
||||
import unittest
|
||||
from argparse import Namespace
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
|
||||
|
||||
def _load_process_module(module_filename: str):
|
||||
module_path = Path(__file__).resolve().parents[1] / "workers" / "process" / module_filename
|
||||
@@ -21,7 +26,7 @@ def _load_process_module(module_filename: str):
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
workers_pkg.__path__ = [str(module_path.parents[1])]
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
process_pkg = types.ModuleType(f"{package_name}.workers.process")
|
||||
@@ -39,8 +44,23 @@ def _load_process_module(module_filename: str):
|
||||
|
||||
process_module = types.ModuleType(f"{package_name}.utils.process")
|
||||
process_module.get_python_executable = lambda: "/usr/bin/test-python"
|
||||
process_module.is_process_alive = lambda _pid: False
|
||||
process_module.terminate_process = lambda *_args, **_kwargs: None
|
||||
sys.modules[f"{package_name}.utils.process"] = process_module
|
||||
|
||||
config_module = types.ModuleType(f"{package_name}.utils.config")
|
||||
config_module.load_config = lambda: {"workers": [], "settings": {"stop_workers_on_master_exit": False}}
|
||||
config_module.save_config = lambda _config: None
|
||||
sys.modules[f"{package_name}.utils.config"] = config_module
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
constants_module.PROCESS_TERMINATION_TIMEOUT = 1
|
||||
constants_module.PROCESS_WAIT_TIMEOUT = 1
|
||||
constants_module.WORKER_CHECK_INTERVAL = 0.01
|
||||
sys.modules[f"{package_name}.utils.constants"] = constants_module
|
||||
network_module = types.ModuleType(f"{package_name}.utils.network")
|
||||
network_module.get_server_port = lambda: 8189
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
spec = importlib.util.spec_from_file_location(
|
||||
f"{package_name}.workers.process.{module_name}",
|
||||
module_path,
|
||||
@@ -70,6 +90,145 @@ class ComfyRootDiscoveryTests(unittest.TestCase):
|
||||
|
||||
|
||||
class LaunchCommandBuilderTests(unittest.TestCase):
|
||||
def build_command(self, root, worker=None, runtime=None):
|
||||
builder = launch_builder_module.LaunchCommandBuilder()
|
||||
worker = worker or {"id": "worker-a", "port": 8190}
|
||||
runtime = runtime if runtime is not None else Namespace(database_url=None)
|
||||
with patch.object(builder, "_get_runtime_args", return_value=runtime):
|
||||
return builder.build_launch_command(worker, str(root))
|
||||
|
||||
def test_database_is_unique_and_stable_by_worker_id(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
first = self.build_command(root, {"id": "worker-a", "port": 8190})
|
||||
renamed = self.build_command(root, {
|
||||
"id": "worker-a", "port": 9001, "name": "Renamed", "cuda_device": 3,
|
||||
})
|
||||
second = self.build_command(root, {"id": "worker-b", "port": 8191})
|
||||
database = first[first.index("--database-url") + 1]
|
||||
self.assertEqual(database, renamed[renamed.index("--database-url") + 1])
|
||||
self.assertNotEqual(database, second[second.index("--database-url") + 1])
|
||||
self.assertTrue(database.startswith("sqlite:///" + root.as_posix() + "/"))
|
||||
self.assertTrue(Path(database.removeprefix("sqlite:///")).parent.is_dir())
|
||||
self.assertNotIn("--disable-assets", first)
|
||||
|
||||
def test_database_uses_effective_user_or_base_directory(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra_args, expected in [
|
||||
(f'--user-directory "{root / "custom user"}"', root / "custom user"),
|
||||
(f'--base-directory "{root / "custom base"}"', root / "custom base" / "user"),
|
||||
]:
|
||||
with self.subTest(extra_args=extra_args):
|
||||
cmd = self.build_command(root, {
|
||||
"id": "worker-a", "port": 8190, "extra_args": extra_args,
|
||||
})
|
||||
database = cmd[cmd.index("--database-url") + 1]
|
||||
self.assertTrue(database.startswith("sqlite:///" + expected.as_posix() + "/"))
|
||||
cmd = self.build_command(root, runtime=Namespace(
|
||||
database_url="sqlite:///master.db", user_directory=str(root / "runtime user"),
|
||||
))
|
||||
database = cmd[cmd.index("--database-url") + 1]
|
||||
self.assertTrue(database.startswith("sqlite:///" + (root / "runtime user").as_posix() + "/"))
|
||||
self.assertNotEqual(database, "sqlite:///master.db")
|
||||
|
||||
def test_relative_database_directories_use_worker_cwd(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory) / "ComfyUI"
|
||||
root.mkdir()
|
||||
(root / "main.py").touch()
|
||||
launcher = Path(directory) / "launcher"
|
||||
launcher.mkdir()
|
||||
cases = [
|
||||
(Namespace(database_url=None, user_directory="profiles"), "", root / "profiles"),
|
||||
(Namespace(database_url=None, base_directory="data"), "", root / "data/user"),
|
||||
(Namespace(database_url=None), "--user-directory=profiles", root / "profiles"),
|
||||
(Namespace(database_url=None), "--base-directory data", root / "data/user"),
|
||||
]
|
||||
for runtime, extra, expected in cases:
|
||||
with self.subTest(runtime=runtime, extra=extra), \
|
||||
patch.object(os, "getcwd", return_value=str(launcher)):
|
||||
cmd = self.build_command(root, {
|
||||
"id": "worker-a", "port": 8190, "extra_args": extra,
|
||||
}, runtime)
|
||||
database = make_url(cmd[cmd.index("--database-url") + 1]).database
|
||||
self.assertEqual(Path(database).parent, expected / "distributed/workers")
|
||||
self.assertFalse((launcher / "profiles").exists())
|
||||
self.assertFalse((launcher / "data").exists())
|
||||
|
||||
def test_database_url_round_trips_special_characters_without_collisions(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for name in ["user?profile", "user%3Fprofile"]:
|
||||
with self.subTest(name=name):
|
||||
user = root / name
|
||||
runtime = Namespace(database_url=None, user_directory=str(user))
|
||||
# SQLAlchemy 2.0 cannot round-trip '?' in a filename. Fail
|
||||
# clearly on those versions rather than silently sharing a DB.
|
||||
sample = (user / "test.db").as_posix()
|
||||
serialized = URL.create("sqlite", database=sample).render_as_string()
|
||||
if make_url(serialized).database != sample:
|
||||
with self.assertRaisesRegex(ValueError, "SQLAlchemy.*database path"):
|
||||
self.build_command(root, runtime=runtime)
|
||||
self.assertFalse(user.exists())
|
||||
continue
|
||||
databases = []
|
||||
for worker_id in ["worker-a", "worker-b"]:
|
||||
cmd = self.build_command(root, {"id": worker_id, "port": 8190}, runtime)
|
||||
url = cmd[cmd.index("--database-url") + 1]
|
||||
database = make_url(url).database
|
||||
self.assertEqual(Path(database).parent, user / "distributed/workers")
|
||||
engine = create_engine(url)
|
||||
try:
|
||||
with engine.connect() as connection:
|
||||
actual = connection.exec_driver_sql("PRAGMA database_list").one()[2]
|
||||
self.assertEqual(Path(actual), Path(database))
|
||||
finally:
|
||||
engine.dispose()
|
||||
databases.append(database)
|
||||
self.assertNotEqual(*databases)
|
||||
|
||||
def test_explicit_database_and_disabled_assets_are_preserved(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra in ["--database-url sqlite:///explicit.db", "--database-url=sqlite:///explicit.db", "--disable-assets"]:
|
||||
with self.subTest(extra=extra):
|
||||
cmd = self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
|
||||
self.assertEqual(sum(arg.split("=")[0] == "--database-url" for arg in cmd), 0 if extra == "--disable-assets" else 1)
|
||||
self.assertFalse((root / "user").exists())
|
||||
|
||||
def test_older_comfyui_without_database_flag_keeps_existing_command(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
cmd = self.build_command(root, runtime=Namespace())
|
||||
self.assertNotIn("--database-url", cmd)
|
||||
self.assertFalse((root / "user").exists())
|
||||
|
||||
def test_worker_id_cannot_escape_database_directory(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
databases = []
|
||||
for worker_id in ["../outside", "a/b", "a_b"]:
|
||||
cmd = self.build_command(root, {"id": worker_id, "port": 8190})
|
||||
path = Path(cmd[cmd.index("--database-url") + 1].removeprefix("sqlite:///"))
|
||||
self.assertTrue(path.is_relative_to(root / "user"))
|
||||
databases.append(path)
|
||||
self.assertEqual(len(set(databases)), 3)
|
||||
|
||||
def test_extra_args_cannot_silently_override_the_configured_port(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
(root / "main.py").touch()
|
||||
for extra in ["--port 8189", "--port=8189"]:
|
||||
with self.subTest(extra=extra), self.assertRaisesRegex(ValueError, "configured worker port"):
|
||||
self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
|
||||
|
||||
def test_inherits_runtime_layout_args_for_desktop(self):
|
||||
builder = launch_builder_module.LaunchCommandBuilder()
|
||||
runtime_args = Namespace(
|
||||
@@ -130,5 +289,64 @@ class LaunchCommandBuilderTests(unittest.TestCase):
|
||||
self.assertNotIn("--auto-launch", cmd)
|
||||
|
||||
|
||||
class ProcessLaunchPortTests(unittest.TestCase):
|
||||
def test_launch_before_master_listens_uses_configured_master_port(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
cli_args = types.ModuleType("comfy.cli_args")
|
||||
cli_args.args = Namespace(port=8189)
|
||||
for missing in [AttributeError("PromptServer has no port"), None]:
|
||||
with self.subTest(missing=missing), tempfile.TemporaryDirectory() as directory:
|
||||
manager = types.SimpleNamespace(
|
||||
find_comfy_root=lambda: directory,
|
||||
build_launch_command=lambda _worker, _root: ["python", "main.py"],
|
||||
processes={}, save_processes=lambda: None,
|
||||
)
|
||||
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
|
||||
patch.object(lifecycle_module, "get_server_port", side_effect=missing if isinstance(missing, Exception) else None, return_value=None), \
|
||||
patch.object(lifecycle_module, "validate_worker_port") as validate, \
|
||||
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
|
||||
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
|
||||
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
|
||||
spawn.assert_called_once()
|
||||
self.assertEqual(pid, 1234)
|
||||
|
||||
def test_pre_listen_launch_still_rejects_the_configured_master_port(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
cli_args = types.ModuleType("comfy.cli_args")
|
||||
cli_args.args = Namespace(port=8189)
|
||||
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
|
||||
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
|
||||
patch.object(lifecycle_module, "get_server_port", side_effect=AttributeError("port")), \
|
||||
patch.object(lifecycle_module.subprocess, "Popen") as spawn:
|
||||
with self.assertRaisesRegex(ValueError, "master port 8189"):
|
||||
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
|
||||
spawn.assert_not_called()
|
||||
|
||||
def test_conflict_is_rejected_before_building_or_spawning(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
|
||||
with patch.object(lifecycle_module.subprocess, "Popen") as spawn:
|
||||
with self.assertRaisesRegex(ValueError, "master port 8189"):
|
||||
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
|
||||
spawn.assert_not_called()
|
||||
self.assertEqual(manager.processes, {})
|
||||
|
||||
def test_available_port_reaches_existing_launch_path(self):
|
||||
lifecycle_module = _load_process_module("lifecycle.py")
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
manager = types.SimpleNamespace(
|
||||
find_comfy_root=lambda: directory,
|
||||
build_launch_command=lambda _worker, _root: ["python", "main.py", "--port", "8190"],
|
||||
processes={}, save_processes=lambda: None,
|
||||
)
|
||||
with patch.object(lifecycle_module, "validate_worker_port") as validate, \
|
||||
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
|
||||
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
|
||||
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
|
||||
spawn.assert_called_once()
|
||||
self.assertEqual(pid, 1234)
|
||||
self.assertEqual(manager.processes["manual"]["pid"], 1234)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
+16
-2
@@ -3,6 +3,7 @@ Async helper utilities for ComfyUI-Distributed.
|
||||
"""
|
||||
import asyncio
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
import execution
|
||||
import server
|
||||
@@ -104,7 +105,12 @@ class PromptValidationError(RuntimeError):
|
||||
super().__init__(f"Invalid prompt: {merged}")
|
||||
|
||||
|
||||
async def queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
|
||||
async def queue_prompt_payload(
|
||||
prompt_obj,
|
||||
workflow_meta=None,
|
||||
client_id=None,
|
||||
include_queue_metadata=False,
|
||||
):
|
||||
"""Validate and queue a prompt via ComfyUI's prompt queue."""
|
||||
payload = {"prompt": prompt_obj}
|
||||
payload = prompt_server.trigger_on_prompt(payload)
|
||||
@@ -117,7 +123,7 @@ async def queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
|
||||
node_errors = valid[3] if len(valid) > 3 else {}
|
||||
raise PromptValidationError(error_payload, node_errors)
|
||||
|
||||
extra_data = {}
|
||||
extra_data = {"create_time": int(time.time() * 1000)}
|
||||
if workflow_meta:
|
||||
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
|
||||
if client_id:
|
||||
@@ -132,4 +138,12 @@ async def queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
|
||||
prompt_server.number = number + 1
|
||||
prompt_queue_item = (number, prompt_id, prompt, extra_data, valid[2], sensitive)
|
||||
prompt_server.prompt_queue.put(prompt_queue_item)
|
||||
|
||||
if include_queue_metadata:
|
||||
return {
|
||||
"prompt_id": prompt_id,
|
||||
"number": number,
|
||||
"node_errors": {},
|
||||
}
|
||||
|
||||
return prompt_id
|
||||
|
||||
+9
-12
@@ -2,7 +2,6 @@ import { app } from "/scripts/app.js";
|
||||
import { ENDPOINTS } from "./constants.js";
|
||||
|
||||
const NODE_CLASS = "DistributedValue";
|
||||
const CONVERTED_WIDGET = "converted-widget";
|
||||
const DYNAMIC_DEFAULT_WIDGET = "_dv_default";
|
||||
const DYNAMIC_WORKER_WIDGET_PREFIX = "_dv_worker_";
|
||||
const WORKERS_CHANGED_EVENT = "distributed:workers-changed";
|
||||
@@ -42,15 +41,13 @@ function getDynamicWorkerWidgets(node) {
|
||||
return (node.widgets || []).filter((w) => w.name.startsWith(DYNAMIC_WORKER_WIDGET_PREFIX));
|
||||
}
|
||||
|
||||
function hideWidgetForGood(node, widget, suffix = "") {
|
||||
function hideWidgetForGood(widget) {
|
||||
if (!widget) return;
|
||||
if (typeof widget.type === "string" && widget.type.startsWith(CONVERTED_WIDGET)) return;
|
||||
|
||||
widget.origType = widget.type;
|
||||
widget.origComputeSize = widget.computeSize;
|
||||
widget.origSerializeValue = widget.serializeValue;
|
||||
widget.computeSize = () => [0, -4];
|
||||
widget.type = `${CONVERTED_WIDGET}${suffix}`;
|
||||
// These are serialized backing fields, not widgets converted into sockets.
|
||||
// Keep their native type/serializer and mutate the shared options in place
|
||||
// so both the classic canvas and Nodes 2.0 suppress the whole widget row.
|
||||
widget.options ??= {};
|
||||
widget.options.hidden = true;
|
||||
|
||||
// Hide any attached DOM element (multiline widgets).
|
||||
if (widget.element) widget.element.style.display = "none";
|
||||
@@ -58,14 +55,14 @@ function hideWidgetForGood(node, widget, suffix = "") {
|
||||
|
||||
if (widget.linkedWidgets) {
|
||||
for (const linked of widget.linkedWidgets) {
|
||||
hideWidgetForGood(node, linked, `:${widget.name}`);
|
||||
hideWidgetForGood(linked);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function hideRawWidgets(node) {
|
||||
hideWidgetForGood(node, getRawDefaultWidget(node), ":default_value");
|
||||
hideWidgetForGood(node, getRawWorkerValuesWidget(node), ":worker_values");
|
||||
hideWidgetForGood(getRawDefaultWidget(node));
|
||||
hideWidgetForGood(getRawWorkerValuesWidget(node));
|
||||
}
|
||||
|
||||
function removeDynamicDefaultWidget(node) {
|
||||
|
||||
+10
-1
@@ -42,6 +42,15 @@ export async function detectMasterIP(extension) {
|
||||
if (shouldAutoPopulate) {
|
||||
extension.log(`Auto-populating workers based on ${extension.cudaDeviceCount} CUDA devices (excluding master on CUDA ${extension.masterCudaDevice})`, "info");
|
||||
|
||||
const workerCount = extension.cudaDeviceCount - (
|
||||
Number.isInteger(extension.masterCudaDevice) &&
|
||||
extension.masterCudaDevice >= 0 &&
|
||||
extension.masterCudaDevice < extension.cudaDeviceCount ? 1 : 0
|
||||
);
|
||||
const workerPorts = data.local_worker_ports;
|
||||
if (!Array.isArray(workerPorts) || workerPorts.length !== workerCount) {
|
||||
throw new Error("Could not allocate enough available local worker ports; check the server log");
|
||||
}
|
||||
const newWorkers = [];
|
||||
let workerNum = 1;
|
||||
let portOffset = 0;
|
||||
@@ -56,7 +65,7 @@ export async function detectMasterIP(extension) {
|
||||
id: generateUUID(),
|
||||
name: `Worker ${workerNum}`,
|
||||
host: isRunpod ? null : "localhost",
|
||||
port: 8189 + portOffset,
|
||||
port: workerPorts[portOffset],
|
||||
cuda_device: i,
|
||||
enabled: true,
|
||||
extra_args: isRunpod ? "--listen" : "",
|
||||
|
||||
@@ -72,7 +72,7 @@ export function renderSettingsSection(extension) {
|
||||
createCheckboxSetting(
|
||||
"setting-debug",
|
||||
"Debug Mode",
|
||||
"Enable verbose logging in the browser console.",
|
||||
"Enable verbose logging in the browser console and ComfyUI server output.",
|
||||
extension.config?.settings?.debug || false,
|
||||
(event) => extension._updateSetting("debug", event.target.checked)
|
||||
)
|
||||
@@ -100,7 +100,7 @@ export function renderSettingsSection(extension) {
|
||||
createNumberSetting(
|
||||
"setting-worker-timeout",
|
||||
"Worker Timeout",
|
||||
"Seconds without a heartbeat before a worker is considered timed out. Default 60.",
|
||||
"Maximum result-wait and heartbeat inactivity period before recovery begins. Busy workers may receive additional grace. Default: 60 seconds.",
|
||||
extension.config?.settings?.worker_timeout_seconds ?? 60,
|
||||
10,
|
||||
1,
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
import { readFileSync } from "node:fs";
|
||||
import vm from "node:vm";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
// Evaluate this repository's browser entrypoint without needing a ComfyUI server.
|
||||
// The real-renderer acceptance additionally checks the input rows in Chromium.
|
||||
const source = readFileSync(new URL("../distributedValue.js", import.meta.url), "utf8")
|
||||
.replace(/^import .*;\r?\n/gm, "");
|
||||
|
||||
function setup({ connected = false, targetType = "STRING" } = {}) {
|
||||
let extension;
|
||||
const listeners = new Map();
|
||||
const tasks = [];
|
||||
const workers = [
|
||||
{ id: "gpu-a", name: "GPU A", enabled: true },
|
||||
{ id: "gpu-b", name: "GPU B", enabled: true },
|
||||
{ id: "off", name: "Disabled", enabled: false },
|
||||
];
|
||||
const target = {
|
||||
inputs: [{ name: "value" }],
|
||||
widgets: [{ name: "value", type: targetType === "INT" ? "number" : "string", options: { step: 1, precision: 0 } }],
|
||||
};
|
||||
const graph = { links: { 1: { target_id: 2, target_slot: 0 } }, getNodeById: () => target };
|
||||
const raw = ["default_value", "worker_values"].map((name, index) => ({
|
||||
name,
|
||||
type: "text",
|
||||
value: index === 0 ? "saved default" : '{"1":"saved worker"}',
|
||||
options: { multiline: false },
|
||||
computeSize: vi.fn(() => [200, 20]),
|
||||
serializeValue: vi.fn(function () { return this.value; }),
|
||||
}));
|
||||
const node = {
|
||||
comfyClass: "DistributedValue", graph, widgets: [...raw], size: [200, 100],
|
||||
inputs: raw.map(w => ({ name: w.name, widget: { name: w.name }, link: null })),
|
||||
outputs: [{ links: connected ? [1] : [] }],
|
||||
computeSize: () => [200, 100], setSize: vi.fn(), setDirtyCanvas: vi.fn(),
|
||||
addWidget(type, name, value, callback, options) {
|
||||
const widget = { type, name, value, callback, options };
|
||||
this.widgets.push(widget);
|
||||
return widget;
|
||||
},
|
||||
configure(data) {
|
||||
data.widgets_values.forEach((value, i) => { this.widgets[i].value = value; });
|
||||
return "configured";
|
||||
},
|
||||
};
|
||||
vm.runInNewContext(source, {
|
||||
app: { graph, registerExtension: ext => { extension = ext; } },
|
||||
ENDPOINTS: { CONFIG: "/distributed/config" },
|
||||
fetch: async () => ({ ok: true, json: async () => ({ workers }) }),
|
||||
window: { addEventListener: (name, handler) => listeners.set(name, handler) },
|
||||
setTimeout: callback => { tasks.push(callback); }, console,
|
||||
}, { filename: "distributedValue.js" });
|
||||
return { node, raw, extension, listeners, tasks, flush: () => { while (tasks.length) tasks.shift()(); } };
|
||||
}
|
||||
|
||||
describe("Distributed Value internal widget visibility", () => {
|
||||
it("hides raw fields natively without converting their types or replacing their options", async () => {
|
||||
const h = setup();
|
||||
const options = h.raw.map(w => w.options);
|
||||
const sizes = h.raw.map(w => w.computeSize);
|
||||
const serializers = h.raw.map(w => w.serializeValue);
|
||||
const slots = h.node.inputs;
|
||||
await h.extension.nodeCreated(h.node);
|
||||
|
||||
for (const [i, widget] of h.raw.entries()) {
|
||||
expect(widget.type).toBe("text");
|
||||
expect(widget.options).toBe(options[i]);
|
||||
expect(widget.options.hidden).toBe(true);
|
||||
expect(widget.computeSize).toBe(sizes[i]);
|
||||
expect(widget.serializeValue).toBe(serializers[i]);
|
||||
expect(h.node.widgets[i]).toBe(widget);
|
||||
}
|
||||
expect(h.node.inputs).toBe(slots);
|
||||
expect(h.node.widgets.slice(2).map(w => w.label)).toEqual(["default_value", "GPU A", "GPU B"]);
|
||||
expect(h.node.widgets.slice(2).every(w => !w.options.hidden)).toBe(true);
|
||||
});
|
||||
|
||||
it("keeps typed edits in the original serialized fields across rebuild/configure", async () => {
|
||||
const h = setup({ connected: true, targetType: "INT" });
|
||||
await h.extension.nodeCreated(h.node);
|
||||
h.node.widgets.find(w => w.name === "_dv_default").callback(42);
|
||||
h.node.widgets.find(w => w.name === "_dv_worker_1").callback(7);
|
||||
h.node.widgets.find(w => w.name === "_dv_worker_2").callback(9);
|
||||
const values = h.raw.map(w => w.serializeValue());
|
||||
expect(values[0]).toBe(42);
|
||||
expect(JSON.parse(values[1])).toMatchObject({
|
||||
_type: "INT", 1: "7", 2: "9", _by_worker_id: { "gpu-a": "7", "gpu-b": "9" },
|
||||
});
|
||||
expect(h.node.configure({ widgets_values: values })).toBe("configured");
|
||||
h.flush();
|
||||
expect(h.raw.map(w => w.serializeValue())).toEqual(values);
|
||||
expect(h.raw.every(w => w.options.hidden && w.type === "text")).toBe(true);
|
||||
expect(h.node.widgets.slice(2).map(w => w.value)).toEqual([42, 7, 9]);
|
||||
});
|
||||
|
||||
it("reapplies visibility on worker refresh without duplicating raw or dynamic fields", async () => {
|
||||
const h = setup();
|
||||
await h.extension.nodeCreated(h.node);
|
||||
h.raw[0].options.hidden = false;
|
||||
await h.listeners.get("distributed:workers-changed")({ detail: { workers: [
|
||||
{ id: "gpu-b", name: "Renamed B", enabled: true },
|
||||
] } });
|
||||
expect(h.raw.every(w => w.options.hidden)).toBe(true);
|
||||
expect(h.node.widgets.map(w => w.name)).toEqual(["default_value", "worker_values", "_dv_default", "_dv_worker_1"]);
|
||||
expect(h.node.widgets[3].label).toBe("Renamed B");
|
||||
});
|
||||
|
||||
it("hides linked backing widgets and handles absent options without altering serializers", async () => {
|
||||
const h = setup();
|
||||
delete h.raw[1].options;
|
||||
const linked = { type: "text", options: {}, serializeValue: vi.fn(() => "linked") };
|
||||
h.raw[0].linkedWidgets = [linked];
|
||||
await h.extension.nodeCreated(h.node);
|
||||
expect(h.raw[1].options.hidden).toBe(true);
|
||||
expect(linked.options.hidden).toBe(true);
|
||||
expect(linked.type).toBe("text");
|
||||
expect(linked.serializeValue()).toBe("linked");
|
||||
});
|
||||
});
|
||||
@@ -1,15 +0,0 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import { buildWorkerWebSocketUrl } from "../urlUtils.js";
|
||||
|
||||
|
||||
describe("execution decision helpers", () => {
|
||||
it("buildWorkerWebSocketUrl converts http/https to ws/wss", () => {
|
||||
expect(buildWorkerWebSocketUrl("http://worker.local:8188")).toBe(
|
||||
"ws://worker.local:8188/distributed/worker_ws"
|
||||
);
|
||||
expect(buildWorkerWebSocketUrl("https://worker.example.com")).toBe(
|
||||
"wss://worker.example.com/distributed/worker_ws"
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,65 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { detectMasterIP } from "../masterDetection.js";
|
||||
|
||||
describe("automatic local worker ports", () => {
|
||||
let originalWindow;
|
||||
beforeEach(() => {
|
||||
originalWindow = globalThis.window;
|
||||
globalThis.window = { location: { hostname: "localhost", port: "443" } };
|
||||
});
|
||||
afterEach(() => { globalThis.window = originalWindow; });
|
||||
|
||||
function makeExtension(overrides = {}) {
|
||||
return {
|
||||
config: { workers: [], settings: {}, master: {} },
|
||||
log: vi.fn(),
|
||||
api: {
|
||||
getNetworkInfo: vi.fn().mockResolvedValue({
|
||||
cuda_device: 0, cuda_device_count: 3,
|
||||
master_port: 8189, local_worker_ports: [8191, 8193],
|
||||
...overrides,
|
||||
}),
|
||||
updateMaster: vi.fn().mockResolvedValue({}),
|
||||
updateWorker: vi.fn().mockResolvedValue({}),
|
||||
updateSetting: vi.fn().mockResolvedValue({}),
|
||||
},
|
||||
ui: { updateMasterDisplay: vi.fn() },
|
||||
app: {},
|
||||
loadConfig: vi.fn().mockResolvedValue(),
|
||||
};
|
||||
}
|
||||
|
||||
it("uses server-allocated ports, not hardcoded ports or the browser proxy port", async () => {
|
||||
const extension = makeExtension();
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers.map(worker => worker.port)).toEqual([8191, 8193]);
|
||||
expect(extension.config.workers.map(worker => worker.cuda_device)).toEqual([1, 2]);
|
||||
expect(extension.api.updateWorker).toHaveBeenCalledTimes(2);
|
||||
});
|
||||
|
||||
it("does not overwrite an existing manually configured worker", async () => {
|
||||
const extension = makeExtension();
|
||||
const worker = { id: "manual", host: "localhost", port: 8189 };
|
||||
extension.config.workers = [worker];
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers).toEqual([worker]);
|
||||
expect(extension.api.updateWorker).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("fails without partially saving workers if not enough ports are available", async () => {
|
||||
const extension = makeExtension({ local_worker_ports: [8191] });
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers).toEqual([]);
|
||||
expect(extension.api.updateWorker).not.toHaveBeenCalled();
|
||||
expect(extension.api.updateSetting).not.toHaveBeenCalled();
|
||||
expect(extension.log).toHaveBeenCalledWith(expect.stringContaining("port"), "error");
|
||||
});
|
||||
|
||||
it("preserves Runpod worker host and launch settings", async () => {
|
||||
globalThis.window.location.hostname = "pod-8189.proxy.runpod.net";
|
||||
const extension = makeExtension();
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers.map(worker => worker.port)).toEqual([8191, 8193]);
|
||||
expect(extension.config.workers.every(worker => worker.host === null && worker.extra_args === "--listen")).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -106,11 +106,6 @@ describe("buildWorkerWebSocketUrl", () => {
|
||||
"wss://worker.example.com/distributed/worker_ws"
|
||||
);
|
||||
});
|
||||
|
||||
it("always appends /distributed/worker_ws", () => {
|
||||
const url = buildWorkerWebSocketUrl("http://worker.local:8188");
|
||||
expect(url.endsWith("/distributed/worker_ws")).toBe(true);
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest";
|
||||
import { getWorkerUrl } from "../workerLifecycle.js";
|
||||
|
||||
|
||||
describe("workerLifecycle URL construction", () => {
|
||||
describe("workerLifecycle URL wiring", () => {
|
||||
let originalWindow;
|
||||
|
||||
beforeEach(() => {
|
||||
@@ -22,31 +22,9 @@ describe("workerLifecycle URL construction", () => {
|
||||
globalThis.window = originalWindow;
|
||||
});
|
||||
|
||||
it("builds local worker URL with explicit local port", () => {
|
||||
// URL variants are covered in urlUtils.test.js; keep the wrapper wiring here.
|
||||
it("forwards the worker and endpoint using window.location", () => {
|
||||
const worker = { id: "w1", port: 8189, type: "local" };
|
||||
expect(getWorkerUrl({}, worker, "/prompt")).toBe("http://127.0.0.1:8189/prompt");
|
||||
});
|
||||
|
||||
it("builds remote worker URL with host:port", () => {
|
||||
const worker = { id: "w2", host: "worker.example.com", port: 9000, type: "remote" };
|
||||
expect(getWorkerUrl({}, worker, "/prompt")).toBe("http://worker.example.com:9000/prompt");
|
||||
});
|
||||
|
||||
it("builds cloud worker URL as https", () => {
|
||||
const worker = { id: "w3", host: "cloud.example.com", port: 443, type: "cloud" };
|
||||
expect(getWorkerUrl({}, worker, "/prompt")).toBe("https://cloud.example.com/prompt");
|
||||
});
|
||||
|
||||
it("rewrites runpod proxy hostname for local worker ports", () => {
|
||||
globalThis.window = {
|
||||
location: {
|
||||
hostname: "podabc.proxy.runpod.net",
|
||||
protocol: "https:",
|
||||
port: "",
|
||||
origin: "https://podabc.proxy.runpod.net",
|
||||
},
|
||||
};
|
||||
const worker = { id: "w4", port: 8189, type: "local" };
|
||||
expect(getWorkerUrl({}, worker, "/prompt")).toBe("https://podabc-8189.proxy.runpod.net/prompt");
|
||||
});
|
||||
});
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Host-local worker port selection and launch preflight (no ComfyUI imports)."""
|
||||
import errno
|
||||
import os
|
||||
import socket
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
|
||||
def _port_number(value):
|
||||
try:
|
||||
port = int(value)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError("Worker port must be an integer between 1 and 65535") from None
|
||||
if isinstance(value, (bool, float)) or not 1 <= port <= 65535:
|
||||
raise ValueError("Worker port must be an integer between 1 and 65535")
|
||||
return port
|
||||
|
||||
|
||||
def _local_config(worker):
|
||||
local_hosts = ("localhost", "127.0.0.1", "0.0.0.0", "::1", "::")
|
||||
host = (worker.get("host") or "localhost").strip().lower()
|
||||
if worker.get("type") == "local" or host in local_hosts:
|
||||
return True
|
||||
try:
|
||||
host = urlsplit(host if "://" in host else "//" + host).hostname
|
||||
except ValueError:
|
||||
return False
|
||||
return host in local_hosts
|
||||
|
||||
|
||||
def _reserved_ports(workers, exclude_id=None):
|
||||
reserved = {}
|
||||
for worker in workers:
|
||||
if exclude_id is not None and str(worker.get("id")) == str(exclude_id):
|
||||
continue
|
||||
if not _local_config(worker):
|
||||
continue
|
||||
try:
|
||||
port = _port_number(worker.get("port"))
|
||||
except ValueError:
|
||||
continue
|
||||
reserved[port] = worker.get("name") or worker.get("id") or "another local worker"
|
||||
return reserved
|
||||
|
||||
|
||||
def is_port_available(port):
|
||||
"""Match asyncio's address reuse while rejecting active IPv4/IPv6 listeners."""
|
||||
addresses = [(socket.AF_INET, "0.0.0.0")]
|
||||
if socket.has_ipv6:
|
||||
addresses.append((socket.AF_INET6, "::"))
|
||||
for family, address in addresses:
|
||||
try:
|
||||
with socket.socket(family, socket.SOCK_STREAM) as probe:
|
||||
if os.name != "nt":
|
||||
# Like asyncio, allow a restart while old connections are
|
||||
# in TIME_WAIT. Do not enable SO_REUSEPORT.
|
||||
probe.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
elif hasattr(socket, "SO_EXCLUSIVEADDRUSE"):
|
||||
# Windows SO_REUSEADDR can steal an active listener's port.
|
||||
probe.setsockopt(socket.SOL_SOCKET, socket.SO_EXCLUSIVEADDRUSE, 1)
|
||||
probe.bind((address, port))
|
||||
probe.listen(1)
|
||||
except OSError as exc:
|
||||
if family == socket.AF_INET6 and exc.errno in (errno.EAFNOSUPPORT, errno.EADDRNOTAVAIL):
|
||||
continue
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def allocate_worker_ports(master_port, workers, count):
|
||||
"""Select ports above the actual master, skipping reservations/listeners.
|
||||
|
||||
This is a snapshot, not a reservation. Launch must validate again because
|
||||
another process can claim a suggested port before the user starts a worker.
|
||||
"""
|
||||
master_port = _port_number(master_port)
|
||||
reserved = _reserved_ports(workers)
|
||||
selected = []
|
||||
if count == 0:
|
||||
return selected
|
||||
for port in range(master_port + 1, 65536):
|
||||
if port not in reserved and is_port_available(port):
|
||||
selected.append(port)
|
||||
if len(selected) == count:
|
||||
return selected
|
||||
raise ValueError(f"Not enough available worker ports above master port {master_port}; need {count}")
|
||||
|
||||
|
||||
def validate_worker_port(worker, master_port, workers):
|
||||
"""Reject conflicts without silently changing a configured worker port."""
|
||||
port = _port_number(worker.get("port"))
|
||||
if port == _port_number(master_port):
|
||||
raise ValueError(f"Worker port {port} conflicts with the master port {master_port}; choose another worker port")
|
||||
reserved = _reserved_ports(workers, exclude_id=worker.get("id"))
|
||||
if port in reserved:
|
||||
raise ValueError(f"Worker port {port} is assigned to local worker {reserved[port]}; choose another worker port")
|
||||
if not is_port_available(port):
|
||||
raise ValueError(f"Worker port {port} is already in use on this host; choose another worker port")
|
||||
@@ -1,4 +1,5 @@
|
||||
import glob
|
||||
import hashlib
|
||||
import os
|
||||
import shlex
|
||||
import shutil
|
||||
@@ -87,6 +88,51 @@ class LaunchCommandBuilder:
|
||||
return wt_path
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _arg_value(cmd, flag):
|
||||
"""Return the last value, matching argparse's override semantics."""
|
||||
value = None
|
||||
for index, arg in enumerate(cmd):
|
||||
if arg == flag and index + 1 < len(cmd):
|
||||
value = cmd[index + 1]
|
||||
elif arg.startswith(flag + "="):
|
||||
value = arg.split("=", 1)[1]
|
||||
return value
|
||||
|
||||
def _add_worker_database(self, cmd, worker_config, comfy_root):
|
||||
args = self._get_runtime_args()
|
||||
# The parsed namespace advertises support without launching ComfyUI
|
||||
# (or importing another checkout's argument parser).
|
||||
if args is None or not hasattr(args, "database_url"):
|
||||
return
|
||||
if any(arg.split("=", 1)[0] in ("--database-url", "--disable-assets") for arg in cmd):
|
||||
return
|
||||
|
||||
user_directory = self._arg_value(cmd, "--user-directory")
|
||||
base_directory = self._arg_value(cmd, "--base-directory")
|
||||
if not user_directory:
|
||||
user_directory = os.path.join(base_directory or comfy_root, "user")
|
||||
if not os.path.isabs(user_directory):
|
||||
# Popen uses comfy_root as cwd, which can differ from the master's.
|
||||
user_directory = os.path.join(comfy_root, user_directory)
|
||||
database_directory = os.path.abspath(os.path.join(user_directory, "distributed", "workers"))
|
||||
# Hash the immutable ID: names/GPU/ports may change, and IDs must not
|
||||
# become filesystem paths or collide after filename sanitization.
|
||||
worker_id = str(worker_config["id"])
|
||||
filename = hashlib.sha256(worker_id.encode("utf-8")).hexdigest() + ".db"
|
||||
database_path = os.path.join(database_directory, filename).replace(os.sep, "/")
|
||||
# Use the same URL codec as ComfyUI; raw '?' and '%xx' can otherwise
|
||||
# truncate/decode the path and make distinct workers share a database.
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
database_url = URL.create("sqlite", database=database_path).render_as_string()
|
||||
if make_url(database_url).database != database_path:
|
||||
raise ValueError(
|
||||
"Installed SQLAlchemy cannot safely encode the worker database path; "
|
||||
"use a user directory without '?' or an explicit --database-url with a safe path"
|
||||
)
|
||||
os.makedirs(database_directory, exist_ok=True)
|
||||
cmd.extend(["--database-url", database_url])
|
||||
|
||||
def build_launch_command(self, worker_config, comfy_root):
|
||||
"""Build the command to launch a worker."""
|
||||
main_py = os.path.join(comfy_root, "main.py")
|
||||
@@ -141,4 +187,7 @@ class LaunchCommandBuilder:
|
||||
raise ValueError(f"Invalid characters in extra_args: {arg}. Forbidden: {forbidden}")
|
||||
cmd.extend(extra_args_list)
|
||||
|
||||
if int(self._arg_value(cmd, "--port")) != int(worker_config["port"]):
|
||||
raise ValueError("extra_args must not override the configured worker port; set the port in the UI/config instead")
|
||||
self._add_worker_database(cmd, worker_config, comfy_root)
|
||||
return cmd
|
||||
|
||||
@@ -7,7 +7,9 @@ import time
|
||||
from ...utils.config import load_config, save_config
|
||||
from ...utils.constants import PROCESS_TERMINATION_TIMEOUT, PROCESS_WAIT_TIMEOUT, WORKER_CHECK_INTERVAL
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.network import get_server_port
|
||||
from ...utils.process import get_python_executable, is_process_alive, terminate_process
|
||||
from ..ports import validate_worker_port
|
||||
|
||||
try:
|
||||
import psutil
|
||||
@@ -28,6 +30,17 @@ class ProcessLifecycle:
|
||||
"""Launch a worker process with logging."""
|
||||
_ = show_window # Kept for API compatibility.
|
||||
comfy_root = self._manager.find_comfy_root()
|
||||
config = load_config()
|
||||
try:
|
||||
master_port = get_server_port()
|
||||
except AttributeError:
|
||||
master_port = None
|
||||
if master_port is None:
|
||||
# The auto-launch timer can fire during custom-node imports, before
|
||||
# PromptServer assigns its port. Still reserve the configured port.
|
||||
from comfy.cli_args import args
|
||||
master_port = args.port
|
||||
validate_worker_port(worker_config, master_port, config.get("workers", []))
|
||||
|
||||
env = os.environ.copy()
|
||||
env["CUDA_VISIBLE_DEVICES"] = str(worker_config.get("cuda_device", 0))
|
||||
|
||||
@@ -814,7 +814,7 @@
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"If all your GPUs are the same/similar, set static_distribution to true\n"
|
||||
"Workers pull tiles from a shared queue, so faster workers can process more tiles automatically.\n"
|
||||
],
|
||||
"color": "#432",
|
||||
"bgcolor": "#653"
|
||||
|
||||
Reference in New Issue
Block a user