Compare commits

...
Author SHA1 Message Date
Robert Wojciechowski bbd8936c65 refactor: clarify load-balanced dispatch reservations 2026-04-26 22:51:01 +00:00
Robert Wojciechowski a91f9fb081 Update pyproject.toml 2026-04-08 10:01:49 +10:00
Robert Wojciechowski 6874d2735f Merge branch 'fix/issue-78-queue-status' 2026-04-07 23:41:58 +00:00
Robert Wojciechowski 0e9ee7fbac test: make async helper regression portable 2026-04-07 23:39:18 +00:00
Robert Wojciechowski 5386de10e3 fix: include queue metadata timestamps 2026-04-07 23:34:15 +00:00
Robert Wojciechowski 41f9e44945 fix: return native prompt queue metadata 2026-04-07 23:19:37 +00:00
Robert Wojciechowski 7792d613bb Update FUNDING.yml 2026-04-04 17:42:09 +11:00
Robert Wojciechowski 6e698512b7 Update pyproject.toml 2026-04-04 17:36:24 +11:00
Robert Wojciechowski 79cc9f5ad5 Update registry publish action ref 2026-04-04 06:35:44 +00:00
Robert Wojciechowski 4eba87cba3 Update pyproject.toml 2026-04-04 17:09:55 +11:00
Robert Wojciechowski 27e08de94c Document ComfyUI Desktop support 2026-04-04 06:06:59 +00:00
Robert Wojciechowski c8453a139b Use loopback callback URLs for local workers 2026-04-04 00:15:00 +00:00
Robert Wojciechowski aae831e1e5 Fix stale master callback port selection 2026-04-04 00:05:49 +00:00
Robert Wojciechowski e7ab67733b Add ComfyUI Desktop worker support 2026-04-03 21:54:51 +00:00
Robert Wojciechowski dd55ff740e fix: forward all queuePrompt args through interceptor
The interceptor only captured (number, prompt) and dropped the third
options argument. This caused partialExecutionTargets to be lost,
making ComfyUI execute all output nodes instead of just the selected
one when using Execute Selected Output.

Fixes #76
2026-03-26 22:00:53 +00:00
Robert Wojciechowski a6d0b82d35 Update README.md 2026-03-02 16:10:42 +11:00
Robert Wojciechowski a694272b40 Harden worker probe response validation and add tests 2026-03-01 22:36:57 +00:00
Robert Wojciechowski 16ea22a643 Update pyproject.toml 2026-02-28 15:58:46 +11:00
21 changed files with 1198 additions and 248 deletions
-1
View File
@@ -1,4 +1,3 @@
# These are supported funding model platforms
github: robertvoy
buy_me_a_coffee: robertvoy
+1 -1
View File
@@ -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 }}
+46 -56
View File
@@ -15,7 +15,7 @@
---
## Key Features
## Key Features
#### Parallel Workflow Processing
- Run your workflow on multiple GPUs simultaneously with varied seeds, collect results on the master
@@ -27,22 +27,12 @@
- Intelligent distribution
- Handles single images and videos
#### Ease of Use
- Auto-setup local workers; easily add remote/cloud ones
- Convert any workflow to distributed with 2 nodes
- JSON configuration with UI controls
---
## Current Architecture
- Workflow-level load balancing is controlled by **Distributed Collector** via the `load_balance` toggle.
- There is **no Distributed Queue node** anymore.
- With `load_balance=true`, orchestration selects one least-busy execution participant:
- If master participation is enabled, master is included as a candidate.
- If master is in orchestrator-only mode, only workers are considered.
---
#### Ease of Use
- Auto-setup local workers; easily add remote/cloud ones
- Convert any workflow to distributed with 2 nodes
- JSON configuration with UI controls
---
## Worker Types
@@ -61,7 +51,6 @@ ComfyUI Distributed supports three types of workers:
## Requirements
- ComfyUI
> Note: Desktop app not currently supported
- Multiple NVIDIA GPUs
> No additional GPUs? Use [Cloud Workers](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/worker-setup-guides.md#cloud-workers)
- That's it
@@ -92,19 +81,19 @@ Join Runpod with [this link](https://get.runpod.io/0bw29uf3ug0p) and unlock a sp
## Workflow Examples
### Basic Parallel Generation
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
### Basic Parallel Generation
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
![Clipboard Image (6)](https://github.com/user-attachments/assets/9598c94c-d9b4-4ccf-ab16-a21398220aeb)
> [Download workflow](/workflows/distributed-txt2img.json)
1. Open your ComfyUI workflow
2. Add **Distributed Seed** → connect to sampler's seed
3. Add **Distributed Collector** → after VAE Decode
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
5. Enable workers in the UI
6. Run the workflow!
2. Add **Distributed Seed** → connect to sampler's seed
3. Add **Distributed Collector** → after VAE Decode
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
5. Enable workers in the UI
6. Run the workflow!
### Parallel WAN Generation
Generate multiple videos in the time it takes to generate one. Each worker uses a different seed.
@@ -122,12 +111,12 @@ Generate multiple videos in the time it takes to generate one. Each worker uses
7. Enable workers in the UI
8. Run the workflow!
### Distributed Image Upscaling
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
### Distributed Image Upscaling
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
![Clipboard Image (3)](https://github.com/user-attachments/assets/ffb57a0d-7b75-4497-96d2-875d60865a1a)
> [Download workflow](/workflows/distributed-upscale.json)
> [Download workflow](/workflows/distributed-upscale.json)
1. Load your image
2. Upscale with ESRGAN or similar
@@ -153,7 +142,7 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
---
## Developer API
## Developer API
Control your distributed cluster programmatically without opening the browser.
@@ -161,32 +150,32 @@ Control your distributed cluster programmatically without opening the browser.
* **Functionality:** Accepts a standard ComfyUI workflow JSON, automatically distributes it to available workers, and returns the execution 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).
---
## Distributed Value
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
- The node shows only currently enabled workers.
- If worker enablement changes, worker fields update automatically.
- When disconnected, it resets to default string mode and clears per-worker overrides.
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
---
## Nodes
| Node | Description |
|------|-------------|
| **Distributed Seed** | Generates unique seeds for each worker |
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
| **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 |
> **⚠️ 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).
---
## Distributed Value
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
- The node shows only currently enabled workers.
- If worker enablement changes, worker fields update automatically.
- When disconnected, it resets to default string mode and clears per-worker overrides.
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
---
## Nodes
| Node | Description |
|------|-------------|
| **Distributed Seed** | Generates unique seeds for each worker |
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
| **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 |
| **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 |
@@ -206,7 +195,7 @@ No, it does not speed up the generation of a single image or video. Instead, it
<details>
<summary>Does it work with the ComfyUI desktop app?</summary>
Currently, it is not compatible with the ComfyUI desktop app.
Yes, it does now.
</details>
<details>
@@ -249,3 +238,4 @@ Buy me a coffee at: https://buymeacoffee.com/robertvoy
+3 -1
View File
@@ -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,
})
+66 -34
View File
@@ -1,5 +1,6 @@
import asyncio
import json
import time
import uuid
import aiohttp
@@ -207,9 +208,13 @@ async def _probe_worker_queue(worker, semaphore, probe_timeout):
payload = await probe_worker(worker_url, timeout=probe_timeout)
if payload is None:
return None
try:
reserved_slots = int(worker.get("reserved_slots", 0) or 0)
except (TypeError, ValueError):
reserved_slots = 0
return {
"worker": worker,
"queue_remaining": _extract_queue_remaining(payload),
"queue_remaining": _extract_queue_remaining(payload) + max(reserved_slots, 0),
}
@@ -222,17 +227,7 @@ def _select_idle_round_robin(statuses):
return statuses[index]
async def select_least_busy_worker(
workers,
trace_execution_id=None,
probe_concurrency=8,
probe_timeout=3.0,
):
"""Select one worker by queue depth, round-robin among idle workers."""
if not workers:
return None
probe_limit = parse_positive_int(probe_concurrency, 8)
async def _probe_worker_queues(workers, probe_limit, probe_timeout):
probe_semaphore = asyncio.Semaphore(probe_limit)
statuses = await asyncio.gather(
*[
@@ -240,29 +235,66 @@ async def select_least_busy_worker(
for worker in workers
]
)
statuses = [status for status in statuses if status is not None]
if not statuses:
if trace_execution_id:
trace_info(trace_execution_id, "Least-busy selection failed: no worker queue probes succeeded.")
else:
log("[Distributed] Least-busy selection failed: no worker queue probes succeeded.")
return None
return [status for status in statuses if status is not None]
def _select_worker_status(statuses, require_idle=False):
idle_statuses = [status for status in statuses if status["queue_remaining"] == 0]
if idle_statuses:
selected = _select_idle_round_robin(idle_statuses)
else:
selected = min(statuses, key=lambda status: status["queue_remaining"])
return _select_idle_round_robin(idle_statuses)
if require_idle:
return None
return min(statuses, key=lambda status: status["queue_remaining"])
worker = selected["worker"]
queue_remaining = selected["queue_remaining"]
if trace_execution_id:
trace_debug(
trace_execution_id,
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}",
)
else:
debug_log(
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}"
)
return worker
async def select_least_busy_worker(
workers,
trace_execution_id=None,
probe_concurrency=8,
probe_timeout=3.0,
require_idle=False,
idle_poll_interval=1.0,
idle_wait_timeout=None,
):
"""Select a worker by queue depth, optionally waiting until one is idle."""
if not workers:
return None
probe_limit = parse_positive_int(probe_concurrency, 8)
poll_interval = max(float(idle_poll_interval or 0), 0.0)
deadline = None
if idle_wait_timeout is not None:
deadline = time.monotonic() + max(float(idle_wait_timeout), 0.0)
while True:
statuses = await _probe_worker_queues(workers, probe_limit, probe_timeout)
if not statuses:
if trace_execution_id:
trace_info(trace_execution_id, "Least-busy selection failed: no worker queue probes succeeded.")
else:
log("[Distributed] Least-busy selection failed: no worker queue probes succeeded.")
return None
selected = _select_worker_status(statuses, require_idle=require_idle)
if selected is not None:
worker = selected["worker"]
queue_remaining = selected["queue_remaining"]
if trace_execution_id:
trace_debug(
trace_execution_id,
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}",
)
else:
debug_log(
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}"
)
return worker
if deadline is not None and time.monotonic() >= deadline:
if trace_execution_id:
trace_info(trace_execution_id, "No idle load-balance worker became available before timeout.")
else:
debug_log("No idle load-balance worker became available before timeout.")
return None
await asyncio.sleep(poll_interval)
+264 -129
View File
@@ -13,7 +13,7 @@ from ..utils.constants import (
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
)
from ..utils.logging import debug_log, log
from ..utils.network import build_master_url
from ..utils.network import build_master_url, build_master_callback_url
from ..utils.trace_logger import trace_debug
from .schemas import parse_positive_float, parse_positive_int
from .orchestration.dispatch import (
@@ -46,12 +46,159 @@ def ensure_distributed_state(server_instance=None):
ps.distributed_pending_jobs = {}
if not hasattr(ps, "distributed_jobs_lock"):
ps.distributed_jobs_lock = asyncio.Lock()
if not hasattr(ps, "distributed_worker_reservations"):
ps.distributed_worker_reservations = {}
# Initialize top-level distributed queue state at module import time.
ensure_distributed_state()
def _worker_id(worker):
return str(worker.get("id"))
def _non_negative_int(value, default=0):
try:
parsed = int(value if value is not None else default)
except (TypeError, ValueError):
parsed = default
return max(parsed, 0)
async def _release_worker_reservation(worker_id):
if not worker_id:
return
ensure_distributed_state()
async with prompt_server.distributed_jobs_lock:
reservations = prompt_server.distributed_worker_reservations
current = _non_negative_int(reservations.get(worker_id, 0))
if current <= 1:
reservations.pop(worker_id, None)
else:
reservations[worker_id] = current - 1
async def _candidate_workers_with_reservations(candidate_workers):
ensure_distributed_state()
async with prompt_server.distributed_jobs_lock:
reservations = dict(prompt_server.distributed_worker_reservations)
return [
{
**worker,
"reserved_slots": reservations.get(_worker_id(worker), 0),
}
for worker in candidate_workers
]
async def _reserve_selected_worker(selected_worker):
selected_worker_id = _worker_id(selected_worker)
selected_reserved_slots = _non_negative_int(selected_worker.get("reserved_slots", 0))
async with prompt_server.distributed_jobs_lock:
reservations = prompt_server.distributed_worker_reservations
current_reserved_slots = _non_negative_int(reservations.get(selected_worker_id, 0))
if current_reserved_slots != selected_reserved_slots:
return None
reservations[selected_worker_id] = current_reserved_slots + 1
return selected_worker_id
def _normalized_idle_poll_interval(value):
try:
poll_interval = float(value or 0)
except (TypeError, ValueError):
poll_interval = 1.0
return max(poll_interval, 0.05)
async def _select_and_reserve_load_balance_worker(
candidate_workers,
execution_trace_id,
worker_probe_concurrency,
idle_poll_interval,
):
poll_interval = _normalized_idle_poll_interval(idle_poll_interval)
while candidate_workers:
reserved_candidates = await _candidate_workers_with_reservations(candidate_workers)
selected_worker = await select_least_busy_worker(
reserved_candidates,
trace_execution_id=execution_trace_id,
probe_concurrency=worker_probe_concurrency,
require_idle=True,
idle_poll_interval=0,
idle_wait_timeout=0,
)
if selected_worker is None:
await asyncio.sleep(poll_interval)
continue
reserved_worker_id = await _reserve_selected_worker(selected_worker)
if reserved_worker_id is not None:
return selected_worker
await asyncio.sleep(0)
return None
def _load_balance_candidate_workers(active_workers, delegate_master, master_url):
candidate_workers = list(active_workers)
if not delegate_master:
candidate_workers.append(
{
"id": "master",
"name": "Master",
"host": master_url,
"type": "local",
}
)
return candidate_workers
async def _resolve_load_balance_selection(
active_workers,
delegate_master,
master_url,
config,
execution_trace_id,
worker_probe_concurrency,
):
candidate_workers = _load_balance_candidate_workers(active_workers, delegate_master, master_url)
selected_worker = None
if candidate_workers:
settings = (config.get("settings", {}) or {})
idle_poll_interval = parse_positive_float(
settings.get("load_balance_idle_poll_interval_seconds"),
1.0,
)
selected_worker = await _select_and_reserve_load_balance_worker(
candidate_workers,
execution_trace_id,
worker_probe_concurrency,
idle_poll_interval,
)
if selected_worker is not None and _worker_id(selected_worker) == "master":
trace_debug(
execution_trace_id,
"Load-balance selected master for execution (workers skipped).",
)
return [], False, "master"
if selected_worker is not None:
trace_debug(
execution_trace_id,
f"Load-balance selected worker {selected_worker.get('id')} (master set to delegate-only).",
)
return [selected_worker], True, _worker_id(selected_worker)
trace_debug(
execution_trace_id,
"Load-balance requested but no idle/probeable execution candidates were available.",
)
return [], False, None
async def _ensure_distributed_queue(job_id):
"""Ensure a queue exists for the given distributed job ID."""
ensure_distributed_state()
@@ -144,6 +291,7 @@ async def _prepare_worker_payload(
enabled_ids,
job_id_map,
master_url,
config,
delegate_master,
trace_execution_id,
worker_prep_semaphore,
@@ -153,6 +301,11 @@ async def _prepare_worker_payload(
"""Prepare one worker prompt payload with bounded concurrency and media-sync timeout."""
async with worker_prep_semaphore:
worker_prompt = prompt_index.copy_prompt()
worker_master_url = build_master_callback_url(
worker,
config=config,
prompt_server_instance=prompt_server,
)
worker_type = str(worker.get("type") or "local").strip().lower()
is_remote_like = bool(worker.get("host")) and worker_type != "local"
@@ -167,7 +320,7 @@ async def _prepare_worker_payload(
worker["id"],
enabled_ids,
job_id_map,
master_url,
worker_master_url,
delegate_master,
prompt_index,
)
@@ -202,7 +355,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()
@@ -253,141 +406,123 @@ async def orchestrate_distributed_execution(
probe_concurrency=worker_probe_concurrency,
)
reserved_worker_id = None
if load_balance_requested:
candidate_workers = list(active_workers)
if not delegate_master:
# Include master in load balancing only when master participation is enabled.
candidate_workers.append(
{
"id": "master",
"name": "Master",
"host": master_url,
"type": "local",
}
active_workers, delegate_master, reserved_worker_id = await _resolve_load_balance_selection(
active_workers,
delegate_master,
master_url,
config,
execution_trace_id,
worker_probe_concurrency,
)
try:
enabled_ids = [worker["id"] for worker in active_workers]
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
if not job_id_map:
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
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", {}),
)
selected_worker = None
if candidate_workers:
selected_worker = await select_least_busy_worker(
candidate_workers,
trace_execution_id=execution_trace_id,
probe_concurrency=worker_probe_concurrency,
)
if selected_worker is None and candidate_workers:
for job_id in job_id_map.values():
await _ensure_distributed_queue(job_id)
master_prompt = prompt_index.copy_prompt()
master_prompt = apply_participant_overrides(
master_prompt,
"master",
enabled_ids,
job_id_map,
master_url,
delegate_master,
prompt_index,
)
if delegate_master:
collector_ids = find_nodes_by_class(master_prompt, "DistributedCollector")
upscale_nodes = find_nodes_by_class(master_prompt, "UltimateSDUpscaleDistributed")
if upscale_nodes:
debug_log(
"Delegate-only master mode currently does not support UltimateSDUpscaleDistributed nodes; running full prompt on master."
)
elif not collector_ids:
debug_log(
"Delegate-only master mode requested but no collectors found in master prompt. Running full prompt on master."
)
else:
master_prompt = prepare_delegate_master_prompt(master_prompt, collector_ids)
if active_workers:
trace_debug(
execution_trace_id,
"Load-balance selection probe failed; using first available candidate.",
"Active distributed workers: "
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers),
)
selected_worker = candidate_workers[0]
if selected_worker is not None and str(selected_worker.get("id")) == "master":
# Master selected as least busy; run master workload only.
active_workers = []
delegate_master = False
trace_debug(
execution_trace_id,
"Load-balance selected master for execution (workers skipped).",
worker_payloads = []
if active_workers:
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
worker_payloads = await asyncio.gather(
*[
_prepare_worker_payload(
worker,
prompt_index,
enabled_ids,
job_id_map,
master_url,
config,
delegate_master,
execution_trace_id,
worker_prep_semaphore,
media_sync_semaphore,
media_sync_timeout_seconds,
)
for worker in active_workers
]
)
elif selected_worker is not None:
active_workers = [selected_worker]
# Worker selected as least busy; keep master orchestrator-only for this run.
delegate_master = True
trace_debug(
execution_trace_id,
f"Load-balance selected worker {selected_worker.get('id')} (master set to delegate-only).",
if worker_payloads:
await asyncio.gather(
*[
dispatch_worker_prompt(
worker,
wprompt,
workflow_meta,
client_id,
use_websocket=use_websocket,
trace_execution_id=execution_trace_id,
)
for worker, wprompt in worker_payloads
]
)
else:
trace_debug(
execution_trace_id,
"Load-balance requested but no execution candidates were available.",
)
active_workers = []
delegate_master = False
enabled_ids = [worker["id"] for worker in active_workers]
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
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
for job_id in job_id_map.values():
await _ensure_distributed_queue(job_id)
master_prompt = prompt_index.copy_prompt()
master_prompt = apply_participant_overrides(
master_prompt,
"master",
enabled_ids,
job_id_map,
master_url,
delegate_master,
prompt_index,
)
if delegate_master:
collector_ids = find_nodes_by_class(master_prompt, "DistributedCollector")
upscale_nodes = find_nodes_by_class(master_prompt, "UltimateSDUpscaleDistributed")
if upscale_nodes:
debug_log(
"Delegate-only master mode currently does not support UltimateSDUpscaleDistributed nodes; running full prompt on master."
)
elif not collector_ids:
debug_log(
"Delegate-only master mode requested but no collectors found in master prompt. Running full prompt on master."
)
else:
master_prompt = prepare_delegate_master_prompt(master_prompt, collector_ids)
if active_workers:
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,
"Active distributed workers: "
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers),
f"Orchestration complete: prompt_id={prompt_id}, dispatched_workers={len(worker_payloads)}, delegate_master={delegate_master}",
)
worker_payloads = []
if active_workers:
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
worker_payloads = await asyncio.gather(
*[
_prepare_worker_payload(
worker,
prompt_index,
enabled_ids,
job_id_map,
master_url,
delegate_master,
execution_trace_id,
worker_prep_semaphore,
media_sync_semaphore,
media_sync_timeout_seconds,
)
for worker in active_workers
]
)
if worker_payloads:
await asyncio.gather(
*[
dispatch_worker_prompt(
worker,
wprompt,
workflow_meta,
client_id,
use_websocket=use_websocket,
trace_execution_id=execution_trace_id,
)
for worker, wprompt in worker_payloads
]
)
prompt_id = await queue_prompt_payload(master_prompt, workflow_meta, client_id)
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
finally:
await _release_worker_reservation(reserved_worker_id)
+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.0"
version = "1.4.4"
license = {file = "LICENSE"}
dependencies = []
+5 -3
View 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):
+91
View File
@@ -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()
+58
View File
@@ -250,5 +250,63 @@ class DispatchSelectionTests(unittest.IsolatedAsyncioTestCase):
self.assertIsNone(selected)
async def test_select_least_busy_worker_waits_for_idle_when_required(self):
workers = [
{"id": "w1", "name": "Worker 1"},
{"id": "w2", "name": "Worker 2"},
]
queue_sequences = {
"w1": [1, 1],
"w2": [1, 0],
}
async def fake_probe(worker_url, timeout=3.0):
worker_id = worker_url.rsplit("/", 1)[-1]
sequence = queue_sequences[worker_id]
value = sequence.pop(0) if sequence else 0
return {"exec_info": {"queue_remaining": value}}
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
dispatch,
"probe_worker",
side_effect=fake_probe,
):
dispatch._least_busy_rr_index = 0
selected = await dispatch.select_least_busy_worker(
workers,
probe_concurrency=2,
require_idle=True,
idle_poll_interval=0,
)
self.assertEqual(selected["id"], "w2")
async def test_select_least_busy_worker_counts_reserved_slots_as_busy(self):
workers = [
{"id": "w1", "name": "Worker 1", "reserved_slots": 1},
{"id": "w2", "name": "Worker 2"},
]
async def fake_probe(worker_url, timeout=3.0):
return {"exec_info": {"queue_remaining": 0}}
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
dispatch,
"probe_worker",
side_effect=fake_probe,
):
dispatch._least_busy_rr_index = 0
selected = await dispatch.select_least_busy_worker(
workers,
probe_concurrency=2,
require_idle=True,
idle_poll_interval=0,
idle_wait_timeout=0,
)
self.assertEqual(selected["id"], "w2")
if __name__ == "__main__":
unittest.main()
+221
View File
@@ -0,0 +1,221 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
from unittest.mock import AsyncMock, patch
def _load_queue_orchestration_module():
module_path = Path(__file__).resolve().parents[1] / "api" / "queue_orchestration.py"
package_name = "dist_queue_orchestration_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
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
orchestration_pkg = types.ModuleType(f"{package_name}.api.orchestration")
orchestration_pkg.__path__ = []
sys.modules[f"{package_name}.api.orchestration"] = orchestration_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
prompt_server_instance = types.SimpleNamespace(
distributed_pending_jobs={},
distributed_worker_reservations={},
)
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
sys.modules["server"] = server_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.queue_prompt_payload = AsyncMock(
return_value={"prompt_id": "prompt-master", "number": 1, "node_errors": {}}
)
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {
"workers": [{"id": "w1", "name": "Worker 1", "enabled": True, "host": "worker-1", "port": 8188}],
"settings": {"load_balance_idle_poll_interval_seconds": 0},
}
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.ORCHESTRATION_MEDIA_SYNC_CONCURRENCY = 4
constants_module.ORCHESTRATION_MEDIA_SYNC_TIMEOUT = 5.0
constants_module.ORCHESTRATION_WORKER_PROBE_CONCURRENCY = 8
constants_module.ORCHESTRATION_WORKER_PREP_CONCURRENCY = 4
sys.modules[f"{package_name}.utils.constants"] = constants_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
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.build_master_url = lambda *_args, **_kwargs: "http://master:8188"
network_module.build_master_callback_url = lambda *_args, **_kwargs: "http://master:8188"
sys.modules[f"{package_name}.utils.network"] = network_module
trace_module = types.ModuleType(f"{package_name}.utils.trace_logger")
trace_module.trace_debug = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.trace_logger"] = trace_module
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
schemas_module.parse_positive_int = lambda value, default: int(value if value is not None else default)
schemas_module.parse_positive_float = lambda value, default: float(value if value is not None else default)
sys.modules[f"{package_name}.api.schemas"] = schemas_module
dispatch_module = types.ModuleType(f"{package_name}.api.orchestration.dispatch")
dispatch_module.dispatch_worker_prompt = AsyncMock()
dispatch_module.select_active_workers = AsyncMock(return_value=([{"id": "w1", "name": "Worker 1"}], False))
dispatch_module.select_least_busy_worker = AsyncMock()
sys.modules[f"{package_name}.api.orchestration.dispatch"] = dispatch_module
media_sync_module = types.ModuleType(f"{package_name}.api.orchestration.media_sync")
media_sync_module.convert_paths_for_platform = lambda prompt, _sep: prompt
media_sync_module.fetch_worker_path_separator = AsyncMock(return_value=None)
media_sync_module.sync_worker_media = AsyncMock()
sys.modules[f"{package_name}.api.orchestration.media_sync"] = media_sync_module
prompt_transform_module = types.ModuleType(f"{package_name}.api.orchestration.prompt_transform")
class _PromptIndex:
def __init__(self, prompt):
self.prompt = prompt
self.inputs_by_node = {node_id: node.get("inputs", {}) for node_id, node in prompt.items()}
def nodes_for_class(self, class_type):
return [
node_id
for node_id, node in self.prompt.items()
if node.get("class_type") == class_type
]
def copy_prompt(self):
return {node_id: dict(node) for node_id, node in self.prompt.items()}
prompt_transform_module.PromptIndex = _PromptIndex
prompt_transform_module.apply_participant_overrides = lambda prompt, *_args, **_kwargs: prompt
prompt_transform_module.find_nodes_by_class = lambda prompt, class_type: [
node_id for node_id, node in prompt.items() if node.get("class_type") == class_type
]
prompt_transform_module.generate_job_id_map = lambda _prompt_index, _prefix: {"10": "job-10"}
prompt_transform_module.prepare_delegate_master_prompt = lambda prompt, _collector_ids: prompt
prompt_transform_module.prune_prompt_for_worker = lambda prompt: prompt
sys.modules[f"{package_name}.api.orchestration.prompt_transform"] = prompt_transform_module
spec = importlib.util.spec_from_file_location(
f"{package_name}.api.queue_orchestration",
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
queue_orchestration = _load_queue_orchestration_module()
class LoadBalanceQueueingTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
queue_orchestration.ensure_distributed_state()
queue_orchestration.prompt_server.distributed_worker_reservations.clear()
queue_orchestration.prompt_server.distributed_pending_jobs.clear()
async def test_load_balance_reservation_selection_refreshes_reservations_between_polls(self):
queue_orchestration.prompt_server.distributed_worker_reservations["w1"] = 1
calls = []
async def fake_select(candidates, **_kwargs):
reserved_slots = candidates[0].get("reserved_slots")
calls.append(reserved_slots)
if reserved_slots:
await queue_orchestration._release_worker_reservation("w1")
return None
return candidates[0]
with patch.object(queue_orchestration, "select_least_busy_worker", side_effect=fake_select):
selected = await queue_orchestration._select_and_reserve_load_balance_worker(
[{"id": "w1", "name": "Worker 1"}],
"exec-test",
worker_probe_concurrency=1,
idle_poll_interval=0,
)
self.assertEqual(selected["id"], "w1")
self.assertEqual(calls, [1, 0])
self.assertEqual(queue_orchestration.prompt_server.distributed_worker_reservations, {"w1": 1})
async def test_load_balance_does_not_fallback_to_first_candidate_when_idle_selection_fails(self):
prompt = {
"10": {
"class_type": "DistributedCollector",
"inputs": {"load_balance": True},
}
}
with patch.object(
queue_orchestration,
"_select_and_reserve_load_balance_worker",
new=AsyncMock(return_value=None),
), patch.object(
queue_orchestration,
"dispatch_worker_prompt",
new=AsyncMock(),
) as dispatch_mock:
prompt_id, prompt_number, worker_count, node_errors = await queue_orchestration.orchestrate_distributed_execution(
prompt,
workflow_meta={},
client_id="client-1",
trace_execution_id="exec-test",
)
self.assertEqual((prompt_id, prompt_number, worker_count, node_errors), ("prompt-master", 1, 0, {}))
dispatch_mock.assert_not_awaited()
async def test_load_balance_releases_worker_reservation_when_dispatch_fails(self):
prompt = {
"10": {
"class_type": "DistributedCollector",
"inputs": {"load_balance": True},
}
}
async def select_first_candidate(candidates, **_kwargs):
return candidates[0]
with patch.object(
queue_orchestration,
"select_least_busy_worker",
side_effect=select_first_candidate,
), patch.object(
queue_orchestration,
"dispatch_worker_prompt",
new=AsyncMock(side_effect=RuntimeError("dispatch failed")),
):
with self.assertRaisesRegex(RuntimeError, "dispatch failed"):
await queue_orchestration.orchestrate_distributed_execution(
prompt,
workflow_meta={},
client_id="client-1",
trace_execution_id="exec-test",
)
self.assertEqual(queue_orchestration.prompt_server.distributed_worker_reservations, {})
if __name__ == "__main__":
unittest.main()
+35 -1
View File
@@ -90,14 +90,48 @@ class NetworkHelpersTests(unittest.TestCase):
"https://master.example.com",
)
def test_build_master_url_ignores_stale_saved_port_and_uses_runtime_port(self):
cfg = {"master": {"host": "192.168.68.56", "port": 8001}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8188",
)
def test_build_master_url_keeps_explicit_port_in_host(self):
cfg = {"master": {"host": "192.168.68.56:8001"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8001",
)
def test_build_master_url_falls_back_to_server_address(self):
cfg = {"master": {"host": ""}}
cfg = {"master": {"host": "", "port": 8001}}
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190)
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://127.0.0.1:8190",
)
def test_build_master_callback_url_uses_loopback_for_local_worker(self):
cfg = {"master": {"host": "192.168.68.56"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
worker = {"id": "w1", "type": "local", "host": "localhost", "port": 8189}
self.assertEqual(
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
"http://127.0.0.1:8001",
)
def test_build_master_callback_url_keeps_public_master_url_for_remote_worker(self):
cfg = {"master": {"host": "192.168.68.56"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
worker = {"id": "w2", "type": "remote", "host": "192.168.68.99", "port": 8189}
self.assertEqual(
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8001",
)
if __name__ == "__main__":
unittest.main()
+134
View File
@@ -0,0 +1,134 @@
import importlib.util
import sys
import types
import unittest
from argparse import Namespace
from pathlib import Path
from unittest.mock import patch
def _load_process_module(module_filename: str):
module_path = Path(__file__).resolve().parents[1] / "workers" / "process" / module_filename
package_name = "dist_proc_testpkg"
module_name = module_filename[:-3]
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
workers_pkg = types.ModuleType(f"{package_name}.workers")
workers_pkg.__path__ = []
sys.modules[f"{package_name}.workers"] = workers_pkg
process_pkg = types.ModuleType(f"{package_name}.workers.process")
process_pkg.__path__ = []
sys.modules[f"{package_name}.workers.process"] = process_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
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
process_module = types.ModuleType(f"{package_name}.utils.process")
process_module.get_python_executable = lambda: "/usr/bin/test-python"
sys.modules[f"{package_name}.utils.process"] = process_module
spec = importlib.util.spec_from_file_location(
f"{package_name}.workers.process.{module_name}",
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
root_discovery_module = _load_process_module("root_discovery.py")
launch_builder_module = _load_process_module("launch_builder.py")
class ComfyRootDiscoveryTests(unittest.TestCase):
def test_prefers_loaded_comfyui_module_path(self):
discovery = root_discovery_module.ComfyRootDiscovery()
server_module = types.SimpleNamespace(__file__="/opt/ComfyUI/server.py")
def fake_exists(path):
return path == "/opt/ComfyUI/main.py"
with patch.dict(sys.modules, {"server": server_module}, clear=False), \
patch.object(root_discovery_module.os.path, "exists", side_effect=fake_exists), \
patch.dict(root_discovery_module.os.environ, {}, clear=True):
self.assertEqual(discovery.find_comfy_root(), "/opt/ComfyUI")
class LaunchCommandBuilderTests(unittest.TestCase):
def test_inherits_runtime_layout_args_for_desktop(self):
builder = launch_builder_module.LaunchCommandBuilder()
runtime_args = Namespace(
listen="127.0.0.1",
base_directory="C:/Users/test/ComfyUI",
temp_directory=None,
input_directory="C:/Users/test/ComfyUI/input",
output_directory="C:/Users/test/ComfyUI/output",
user_directory="C:/Users/test/ComfyUI/user",
front_end_root="C:/Program Files/ComfyUI/web_custom_versions/desktop_app",
extra_model_paths_config=[["C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml"]],
enable_manager=True,
disable_manager_ui=False,
enable_manager_legacy_ui=False,
windows_standalone_build=True,
log_stdout=True,
verbose="INFO",
enable_cors_header="*",
)
comfy_module = types.ModuleType("comfy")
comfy_cli_args = types.ModuleType("comfy.cli_args")
comfy_cli_args.args = runtime_args
worker_config = {
"port": 9001,
"extra_args": "--preview-method auto",
}
def fake_exists(path):
return path == "/desktop/ComfyUI/main.py"
with patch.dict(
sys.modules,
{"comfy": comfy_module, "comfy.cli_args": comfy_cli_args},
clear=False,
), patch.object(launch_builder_module.os.path, "exists", side_effect=fake_exists):
cmd = builder.build_launch_command(worker_config, "/desktop/ComfyUI")
self.assertEqual(cmd[:2], ["/usr/bin/test-python", "/desktop/ComfyUI/main.py"])
self.assertIn("--listen", cmd)
self.assertIn("127.0.0.1", cmd)
self.assertIn("--base-directory", cmd)
self.assertIn("C:/Users/test/ComfyUI", cmd)
self.assertIn("--input-directory", cmd)
self.assertIn("--output-directory", cmd)
self.assertIn("--user-directory", cmd)
self.assertIn("--front-end-root", cmd)
self.assertIn("--extra-model-paths-config", cmd)
self.assertIn("C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml", cmd)
self.assertIn("--enable-manager", cmd)
self.assertIn("--windows-standalone-build", cmd)
self.assertIn("--log-stdout", cmd)
self.assertIn("--disable-auto-launch", cmd)
self.assertIn("--enable-cors-header", cmd)
self.assertIn("*", cmd)
self.assertIn("--port", cmd)
self.assertIn("9001", cmd)
self.assertNotIn("--auto-launch", cmd)
if __name__ == "__main__":
unittest.main()
+16 -2
View File
@@ -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
+43 -8
View File
@@ -66,6 +66,25 @@ def normalize_host(value):
return host.split("/")[0]
def _split_host_and_port(host):
if not host:
return host, None
if host.startswith("["):
match = re.match(r"^(\[[^\]]+\])(?::(\d+))?$", host)
if match:
parsed_port = int(match.group(2)) if match.group(2) else None
return match.group(1), parsed_port
return host, None
if host.count(":") == 1:
candidate_host, candidate_port = host.rsplit(":", 1)
if candidate_port.isdigit():
return candidate_host, int(candidate_port)
return host, None
def build_worker_url(worker, endpoint=""):
"""Construct the worker base URL with optional endpoint."""
host = (worker.get("host") or "").strip()
@@ -126,12 +145,7 @@ def build_master_url(config=None, prompt_server_instance=None):
prompt_server_instance = prompt_server_instance or server.PromptServer.instance
master_cfg = (config or {}).get("master", {}) or {}
configured_host = (master_cfg.get("host") or "").strip()
configured_port = master_cfg.get("port")
default_port = getattr(prompt_server_instance, "port", 8188) or 8188
try:
port = int(configured_port or default_port)
except (TypeError, ValueError):
port = int(default_port)
runtime_port = getattr(prompt_server_instance, "port", 8188) or 8188
def _needs_https(hostname):
hostname = hostname.lower()
@@ -149,10 +163,11 @@ def build_master_url(config=None, prompt_server_instance=None):
if configured_host.startswith(("http://", "https://")):
return configured_host.rstrip("/")
host = configured_host
host, explicit_port = _split_host_and_port(configured_host)
port = explicit_port if explicit_port is not None else int(runtime_port)
scheme = "https" if _needs_https(host) or port == 443 else "http"
default_port_for_scheme = 443 if scheme == "https" else 80
if configured_port is None and scheme == "https" and _needs_https(host):
if explicit_port is None and scheme == "https" and _needs_https(host):
port = default_port_for_scheme
port_part = "" if port == default_port_for_scheme else f":{port}"
return f"{scheme}://{host}{port_part}"
@@ -160,7 +175,27 @@ def build_master_url(config=None, prompt_server_instance=None):
address = getattr(prompt_server_instance, "address", "127.0.0.1") or "127.0.0.1"
if address in ("0.0.0.0", "::"):
address = "127.0.0.1"
port = int(runtime_port)
scheme = "https" if port == 443 else "http"
default_port_for_scheme = 443 if scheme == "https" else 80
port_part = "" if port == default_port_for_scheme else f":{port}"
return f"{scheme}://{address}{port_part}"
def build_master_callback_url(worker, config=None, prompt_server_instance=None):
"""Build the callback URL a specific worker should use to reach the master."""
prompt_server_instance = prompt_server_instance or server.PromptServer.instance
worker_type = str((worker or {}).get("type") or "").strip().lower()
worker_host = normalize_host((worker or {}).get("host"))
local_hosts = {"", "localhost", "127.0.0.1", "::1", "[::1]", "0.0.0.0"}
is_local_worker = worker_type == "local" or worker_host in local_hosts
if is_local_worker:
port = int(getattr(prompt_server_instance, "port", 8188) or 8188)
scheme = "https" if port == 443 else "http"
default_port_for_scheme = 443 if scheme == "https" else 80
port_part = "" if port == default_port_for_scheme else f":{port}"
return f"{scheme}://127.0.0.1{port_part}"
return build_master_url(config=config, prompt_server_instance=prompt_server_instance)
+23 -2
View File
@@ -178,11 +178,32 @@ export function createApiClient(baseUrl) {
return { ok: false, status: response.status, queueRemaining: null };
}
const data = await response.json().catch(() => ({}));
let data;
try {
data = await response.json();
} catch {
return { ok: false, status: response.status, queueRemaining: null };
}
if (!data || typeof data !== "object" || Array.isArray(data)) {
return { ok: false, status: response.status, queueRemaining: null };
}
const execInfo = data.exec_info;
if (!execInfo || typeof execInfo !== "object" || Array.isArray(execInfo)) {
return { ok: false, status: response.status, queueRemaining: null };
}
const rawQueueRemaining = execInfo.queue_remaining;
const queueRemaining = Number(rawQueueRemaining);
if (!Number.isFinite(queueRemaining)) {
return { ok: false, status: response.status, queueRemaining: null };
}
return {
ok: true,
status: response.status,
queueRemaining: data.exec_info?.queue_remaining || 0,
queueRemaining: Math.max(0, queueRemaining),
};
} finally {
clearTimeout(timeoutId);
+2 -2
View File
@@ -4,7 +4,7 @@ import { TIMEOUTS, NODE_CLASSES, generateUUID } from './constants.js';
import { checkAllWorkerStatuses, getWorkerUrl } from './workerLifecycle.js';
export function setupInterceptor(extension) {
api.queuePrompt = async (number, prompt) => {
api.queuePrompt = async (number, prompt, ...rest) => {
if (extension.isEnabled) {
const hasCollector = findNodesByClass(prompt.output, NODE_CLASSES.DISTRIBUTED_COLLECTOR).length > 0;
const hasDistUpscale = findNodesByClass(prompt.output, NODE_CLASSES.UPSCALE_DISTRIBUTED).length > 0;
@@ -18,7 +18,7 @@ export function setupInterceptor(extension) {
return result;
}
}
return extension.originalQueuePrompt(number, prompt);
return extension.originalQueuePrompt(number, prompt, ...rest);
};
}
+95
View File
@@ -0,0 +1,95 @@
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { createApiClient } from "../apiClient.js";
describe("apiClient probeWorker", () => {
let originalFetch;
beforeEach(() => {
originalFetch = globalThis.fetch;
globalThis.fetch = vi.fn();
});
afterEach(() => {
globalThis.fetch = originalFetch;
vi.restoreAllMocks();
});
it("returns ok=true when /prompt returns valid exec_info payload", async () => {
globalThis.fetch.mockResolvedValue({
ok: true,
status: 200,
json: vi.fn().mockResolvedValue({ exec_info: { queue_remaining: 2 } }),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: true, status: 200, queueRemaining: 2 });
});
it("returns ok=false on non-200 responses", async () => {
globalThis.fetch.mockResolvedValue({
ok: false,
status: 503,
json: vi.fn(),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: false, status: 503, queueRemaining: null });
});
it("returns ok=false when response JSON is invalid", async () => {
globalThis.fetch.mockResolvedValue({
ok: true,
status: 200,
json: vi.fn().mockRejectedValue(new Error("invalid json")),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: false, status: 200, queueRemaining: null });
});
it("returns ok=false when exec_info is missing", async () => {
globalThis.fetch.mockResolvedValue({
ok: true,
status: 200,
json: vi.fn().mockResolvedValue({}),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: false, status: 200, queueRemaining: null });
});
it("returns ok=false when queue_remaining is not numeric", async () => {
globalThis.fetch.mockResolvedValue({
ok: true,
status: 200,
json: vi.fn().mockResolvedValue({ exec_info: { queue_remaining: "n/a" } }),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: false, status: 200, queueRemaining: null });
});
it("clamps negative queue_remaining to zero", async () => {
globalThis.fetch.mockResolvedValue({
ok: true,
status: 200,
json: vi.fn().mockResolvedValue({ exec_info: { queue_remaining: -5 } }),
});
const client = createApiClient("http://127.0.0.1:8188");
const result = await client.probeWorker("http://worker.local:8190", 1000);
expect(result).toEqual({ ok: true, status: 200, queueRemaining: 0 });
});
});
+68 -3
View File
@@ -10,6 +10,62 @@ from ...utils.process import get_python_executable
class LaunchCommandBuilder:
"""Build command-lines for launching worker ComfyUI processes."""
def _extend_arg(self, cmd, flag, value):
if value in (None, "", [], ()):
return
cmd.extend([flag, str(value)])
def _extend_grouped_args(self, cmd, flag, values):
for group in values or []:
flattened = [str(item) for item in group if item]
if flattened:
cmd.append(flag)
cmd.extend(flattened)
def _get_runtime_args(self):
try:
from comfy.cli_args import args
return args
except Exception as exc:
debug_log(f"Could not read current ComfyUI CLI args for worker launch: {exc}")
return None
def _build_runtime_launch_args(self):
args = self._get_runtime_args()
if args is None:
return []
inherited = []
self._extend_arg(inherited, "--listen", getattr(args, "listen", None))
self._extend_arg(inherited, "--base-directory", getattr(args, "base_directory", None))
self._extend_arg(inherited, "--temp-directory", getattr(args, "temp_directory", None))
self._extend_arg(inherited, "--input-directory", getattr(args, "input_directory", None))
self._extend_arg(inherited, "--output-directory", getattr(args, "output_directory", None))
self._extend_arg(inherited, "--user-directory", getattr(args, "user_directory", None))
self._extend_arg(inherited, "--front-end-root", getattr(args, "front_end_root", None))
self._extend_grouped_args(
inherited,
"--extra-model-paths-config",
getattr(args, "extra_model_paths_config", None),
)
if getattr(args, "enable_manager", False):
inherited.append("--enable-manager")
if getattr(args, "disable_manager_ui", False):
inherited.append("--disable-manager-ui")
if getattr(args, "enable_manager_legacy_ui", False):
inherited.append("--enable-manager-legacy-ui")
if getattr(args, "windows_standalone_build", False):
inherited.append("--windows-standalone-build")
if getattr(args, "log_stdout", False):
inherited.append("--log-stdout")
verbose = getattr(args, "verbose", None)
if verbose and verbose != "INFO":
inherited.extend(["--verbose", str(verbose)])
return inherited
def _find_windows_terminal(self):
"""Find Windows Terminal executable."""
possible_paths = [
@@ -39,10 +95,19 @@ class LaunchCommandBuilder:
cmd = [
get_python_executable(),
main_py,
"--port",
str(worker_config["port"]),
"--enable-cors-header",
]
cmd.extend(self._build_runtime_launch_args())
cmd.extend(["--port", str(worker_config["port"])])
current_args = self._get_runtime_args()
current_cors = getattr(current_args, "enable_cors_header", None) if current_args else None
cmd.append("--enable-cors-header")
if current_cors is not None:
cmd.append(str(current_cors))
if "--disable-auto-launch" not in cmd:
cmd.append("--disable-auto-launch")
debug_log(f"Using main.py: {main_py}")
else:
error_msg = f"Could not find main.py in {comfy_root}\n"
+1
View File
@@ -33,6 +33,7 @@ class ProcessLifecycle:
env["CUDA_VISIBLE_DEVICES"] = str(worker_config.get("cuda_device", 0))
env["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
env["COMFYUI_MASTER_PID"] = str(os.getpid())
env["COMFYUI_IS_WORKER"] = "1"
cmd = self._manager.build_launch_command(worker_config, comfy_root)
cwd = comfy_root
+25 -4
View File
@@ -1,4 +1,5 @@
import os
import sys
from ...utils.logging import debug_log, log
@@ -6,6 +7,21 @@ from ...utils.logging import debug_log, log
class ComfyRootDiscovery:
"""Resolve the ComfyUI root directory across local and container layouts."""
def _find_root_from_loaded_modules(self):
"""Use already-imported ComfyUI modules to locate the runtime root."""
for module_name in ("server", "folder_paths", "main"):
module = sys.modules.get(module_name)
module_file = getattr(module, "__file__", None)
if not module_file:
continue
candidate = os.path.dirname(os.path.abspath(module_file))
if os.path.exists(os.path.join(candidate, "main.py")):
debug_log(f"Found ComfyUI root via loaded module {module_name}: {candidate}")
return candidate
return None
def find_comfy_root(self):
# Start from current file location.
current_dir = os.path.dirname(os.path.abspath(__file__))
@@ -17,12 +33,17 @@ class ComfyRootDiscovery:
debug_log(f"Found ComfyUI root via COMFYUI_ROOT environment variable: {env_root}")
return env_root
# Method 2: Try going up from custom_nodes directory.
# Method 2: Inspect the already-loaded ComfyUI runtime modules.
runtime_root = self._find_root_from_loaded_modules()
if runtime_root:
return runtime_root
# Method 3: Try going up from custom_nodes directory.
if os.path.exists(os.path.join(potential_root, "main.py")):
debug_log(f"Found ComfyUI root via directory traversal: {potential_root}")
return potential_root
# Method 3: Look for common Docker paths.
# Method 4: Look for common Docker paths.
docker_paths = [
"/basedir",
"/ComfyUI",
@@ -37,7 +58,7 @@ class ComfyRootDiscovery:
debug_log(f"Found ComfyUI root in Docker path: {path}")
return path
# Method 4: Search upwards for main.py.
# Method 5: Search upwards for main.py.
search_dir = current_dir
for _ in range(5):
if os.path.exists(os.path.join(search_dir, "main.py")):
@@ -48,7 +69,7 @@ class ComfyRootDiscovery:
break
search_dir = parent
# Method 5: Try to import and use folder_paths.
# Method 6: Try to import and use folder_paths.
try:
import folder_paths