Compare commits

...
Author SHA1 Message Date
Robert Wojciechowski 539f94ba39 fix: harden local worker startup preflight and database paths
Handle auto-launch before the master binds, allow TIME_WAIT restarts without accepting active listeners, resolve relative paths in the worker cwd, and round-trip SQLite URLs through SQLAlchemy. Add regressions and remove the README additions.
2026-10-10 22:13:13 +00:00
Robert Wojciechowski fbfb411402 fix: isolate local worker databases and avoid port conflicts 2026-10-10 21:59:18 +00:00
Robert Wojciechowski 32ac027e0e docs: align selected copy with current behavior 2026-07-12 11:13:22 +00:00
Robert Wojciechowski fec959e1ac chore: bump version to 1.4.7 2026-07-12 10:40:27 +00:00
Robert Wojciechowski cedc8d4591 feat: support audio-only distributed workflows (#87)
Support image, audio, or combined distributed collection with guarded audio-only worker, master, endpoint, preview, and delegate-only paths.
2026-07-12 20:36:43 +10:00
Robert Wojciechowski db01824d6e Update pyproject.toml 2026-06-06 09:10:43 +10:00
Robert Wojciechowski 6e7818bee6 fix: preserve delegate-only master config inputs (#86)
Fixes #85
2026-06-05 17:12:25 +10:00
Robert Wojciechowski c246d965e7 Update pyproject.toml 2026-05-28 09:00:17 +10:00
Robert Wojciechowski d6ce26f143 Merge pull request #84 from robertvoy/fix/issue-83
Fix collector handling of ComfyUI list inputs
2026-05-28 08:59:02 +10:00
24 changed files with 1746 additions and 91 deletions
+2 -2
View File
@@ -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 |
+11 -4
View File
@@ -295,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()
+237 -10
View File
@@ -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",
+14 -1
View File
@@ -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)
+18 -8
View File
@@ -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
+2 -2
View File
@@ -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 -1
View File
@@ -26,6 +26,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DistributedModelName": "Distributed Model Name",
"DistributedValue": "Distributed Value",
"ImageBatchDivider": "Image Batch Divider",
"AudioBatchDivider": "Audio Batch Divider",
"AudioBatchDivider": "Audio Segment Divider",
"DistributedEmptyImage": "Distributed Empty Image",
}
+73 -45
View File
@@ -29,7 +29,6 @@ class DistributedCollectorNode:
def INPUT_TYPES(s):
return {
"required": {
"images": ("IMAGE",),
"load_balance": (
"BOOLEAN",
{
@@ -38,7 +37,10 @@ class DistributedCollectorNode:
},
),
},
"optional": { "audio": ("AUDIO",) },
"optional": {
"images": ("IMAGE",),
"audio": ("AUDIO",),
},
"hidden": {
"multi_job_id": ("STRING", {"default": ""}),
"is_worker": ("BOOLEAN", {"default": False}),
@@ -102,8 +104,9 @@ class DistributedCollectorNode:
return None
return {"waveform": torch.cat(waveforms, dim=-1), "sample_rate": sample_rate}
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):
images = self._normalize_images_input(images)
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)
@@ -115,6 +118,14 @@ class DistributedCollectorNode:
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}
@@ -141,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.
@@ -257,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):
@@ -289,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:
@@ -332,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
@@ -511,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)
+6 -6
View File
@@ -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):
+2 -2
View File
@@ -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"
}),
}
}
+1 -1
View File
@@ -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.4"
version = "1.4.7"
license = {file = "LICENSE"}
dependencies = []
+71
View File
@@ -302,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,
+34 -1
View File
@@ -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}]}
+159
View File
@@ -124,6 +124,49 @@ def test_collector_opts_into_comfyui_list_inputs():
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)
@@ -203,3 +246,119 @@ def test_worker_list_input_sends_one_completion_sequence_with_last_only_on_final
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"])
+546 -3
View File
@@ -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)
@@ -285,17 +323,522 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
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_does_not_preserve_non_primitive_upstream_for_collector(self):
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
self.assertNotIn("2", result)
self.assertNotEqual(result["3"]["inputs"]["images"], ["2", 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"])
+113
View File
@@ -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()
+219 -1
View File
@@ -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()
+10 -1
View File
@@ -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" : "",
+2 -2
View File
@@ -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,
+65
View File
@@ -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);
});
});
+97
View File
@@ -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")
+49
View File
@@ -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
+13
View File
@@ -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))
+1 -1
View File
@@ -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"