Compare commits

..
Author SHA1 Message Date
robertvoy 1d4dcb985e Feat: Audio Batch Divider 2026-01-25 17:32:24 +11:00
robertvoy b6d4756392 collector supports audio 2026-01-24 18:32:24 +11:00
robertvoy 044c54b780 merged fork 2026-01-24 14:19:22 +11:00
robertvoy b768e7fc8a https handling in host field 2025-12-30 12:11:12 +11:00
robertvoy 9805049268 DistributedModelName 2025-12-30 11:17:44 +11:00
robertvoy 399ed7c7d1 feat: initial setup for cloudflare integration 2025-12-29 09:38:03 +11:00
Robert Wojciechowski 64ba4b75db Update pyproject.toml 2025-12-28 10:55:48 +11:00
Robert Wojciechowski d353d118d7 Update README.md 2025-12-28 10:54:16 +11:00
robertvoy 5dd10cef39 organise docs 2025-12-28 10:48:54 +11:00
Robert Wojciechowski 8eba225f5e Merge pull request #61 from umanets/feature/distributed-queue-api
Add public POST /distributed/queue API endpoint + docs
2025-12-28 10:36:59 +11:00
Igor Umanets 9bc2cfc43e Update README-API.md with links to video walkthrough and examples repository 2025-12-26 16:42:18 +01:00
Igor Umanets 370b48438b Add instructions for retrieving enabled worker IDs in API documentation 2025-12-26 15:12:05 +01:00
Igor Umanets 71c7a99e1e Add API documentation for distributed queue 2025-12-26 14:05:15 +01:00
Igor Umanets def7c0df1d Add distributed queue API orchestration 2025-12-26 10:47:34 +01:00
Robert Wojciechowski 5f5320b34e Update pyproject.toml 2025-10-28 08:01:06 +11:00
Robert Wojciechowski 5a33da828c Merge pull request #50 from IT-BillDeng/feature/disable-master
Feature/disable master
2025-10-28 08:00:30 +11:00
IT-BillDeng ef8e0542c8 Document master participation, orchestrator-only, and fallback modes in the worker setup guide 2025-10-27 03:03:17 +08:00
IT-BillDeng 89b7812fb4 Improve master participation logic with fallback handling and UI messaging 2025-10-27 03:00:57 +08:00
IT-BillDeng f13b953a82 Enable UI toggle for master participation with live status feedback and delegate-mode orchestration. 2025-10-27 02:29:51 +08:00
Robert Wojciechowski 1e3e094cdb Update README.md 2025-10-25 09:40:53 +11:00
Robert Wojciechowski 7bcf9290ba Update video-upscaler-runpod-preset.md 2025-09-15 12:18:44 +10:00
Robert Wojciechowski 6aec567071 Update README.md 2025-09-15 12:01:27 +10:00
20 changed files with 2751 additions and 198 deletions
+4
View File
@@ -0,0 +1,4 @@
bin/
logs/
gpu_config.json
__pycache__/
+14 -9
View File
@@ -5,7 +5,7 @@
<a href="/docs/worker-setup-guides.md"><img src="https://img.shields.io/badge/Setup_Guides-grey?style=flat&logo=gitbook&logoColor=white" alt="Setup Guides"></a>
<a href="/workflows"><img src="https://img.shields.io/badge/Workflows-grey?style=flat&logo=json&logoColor=white" alt="Workflows"></a>
<a href="https://buymeacoffee.com/robertvoy"><img src="https://img.shields.io/badge/Donation-grey?style=flat&logo=buymeacoffee&logoColor=white" alt="Donation"></a>
<a href="https://x.com/robertvoy3D"><img src="https://img.shields.io/twitter/follow/robertvoy3D" alt="Twitter"></a>
<a href="https://x.com/rbw_ai"><img src="https://img.shields.io/twitter/follow/rbw_ai" alt="Twitter"></a>
<br><br>
</div>
@@ -25,7 +25,7 @@
#### Distributed Upscaling
- Accelerate Ultimate SD Upscale by distributing tiles across GPUs
- Intelligent distribution
- Handles single images and batches
- Handles single images and videos
#### Ease of Use
- Auto-setup local workers; easily add remote/cloud ones
@@ -142,6 +142,18 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
---
## Developer API
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.
* **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).
---
## FAQ
<details>
@@ -192,10 +204,3 @@ Your support helps keep this project thriving.
Buy me a coffee at: https://buymeacoffee.com/robertvoy
+684 -58
View File
@@ -4,11 +4,13 @@ from PIL import Image
import folder_paths
import os
import json
import re
import asyncio
import aiohttp
from aiohttp import web
import io
import server
import execution
import comfy.model_management
import subprocess
import platform
@@ -23,11 +25,20 @@ from comfy.utils import ProgressBar
# Import shared utilities
from .utils.logging import debug_log, log
from .utils.config import CONFIG_FILE, get_default_config, load_config, save_config, ensure_config_exists, get_worker_timeout_seconds
from .utils.config import (
CONFIG_FILE,
get_default_config,
load_config,
save_config,
ensure_config_exists,
get_worker_timeout_seconds,
is_master_delegate_only,
)
from .utils.image import tensor_to_pil, pil_to_tensor, ensure_contiguous
from .utils.process import is_process_alive, terminate_process, get_python_executable
from .utils.network import handle_api_error, get_server_port, get_server_loop, get_client_session, cleanup_client_session
from .utils.async_helpers import run_async_in_server_loop
from .utils.cloudflare import cloudflare_tunnel_manager
from .utils.constants import (
WORKER_JOB_TIMEOUT, PROCESS_TERMINATION_TIMEOUT, WORKER_CHECK_INTERVAL,
STATUS_CHECK_INTERVAL, CHUNK_SIZE, LOG_TAIL_BYTES, WORKER_LOG_PATTERN,
@@ -51,6 +62,17 @@ def cleanup():
atexit.register(cleanup)
def normalize_host(value):
if value is None:
return None
if not isinstance(value, str):
return value
host = value.strip()
if not host:
return host
host = re.sub(r"^https?://", "", host, flags=re.IGNORECASE)
return host.split("/")[0]
# --- API Endpoints ---
@server.PromptServer.instance.routes.get("/distributed/config")
async def get_config_endpoint(request):
@@ -75,6 +97,102 @@ async def queue_status_endpoint(request):
except Exception as e:
return await handle_api_error(request, e, 500)
prompt_server = server.PromptServer.instance
async def _queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
"""Validate and queue a prompt via ComfyUI's prompt queue."""
payload = {"prompt": prompt_obj}
payload = prompt_server.trigger_on_prompt(payload)
prompt = payload["prompt"]
prompt_id = str(uuid.uuid4())
valid = await execution.validate_prompt(prompt_id, prompt, None)
if not valid[0]:
raise RuntimeError(f"Invalid prompt: {valid[1]}")
extra_data = {}
if workflow_meta:
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
if client_id:
extra_data["client_id"] = client_id
sensitive = {}
for key in getattr(execution, "SENSITIVE_EXTRA_DATA_KEYS", []):
if key in extra_data:
sensitive[key] = extra_data.pop(key)
number = getattr(prompt_server, "number", 0)
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)
return prompt_id
@server.PromptServer.instance.routes.get("/distributed/worker_ws")
async def worker_ws_endpoint(request):
"""WebSocket endpoint for worker prompt dispatch."""
ws = web.WebSocketResponse(heartbeat=30)
await ws.prepare(request)
async for msg in ws:
if msg.type == aiohttp.WSMsgType.TEXT:
try:
data = json.loads(msg.data or "{}")
except json.JSONDecodeError:
await ws.send_json({
"type": "dispatch_ack",
"request_id": None,
"ok": False,
"error": "Invalid JSON payload.",
})
continue
if data.get("type") != "dispatch_prompt":
await ws.send_json({
"type": "dispatch_ack",
"request_id": data.get("request_id"),
"ok": False,
"error": "Unsupported websocket message type.",
})
continue
prompt = data.get("prompt")
if not isinstance(prompt, dict):
await ws.send_json({
"type": "dispatch_ack",
"request_id": data.get("request_id"),
"ok": False,
"error": "Field 'prompt' must be an object.",
})
continue
try:
prompt_id = await _queue_prompt_payload(
prompt,
workflow_meta=data.get("workflow"),
client_id=data.get("client_id"),
)
await ws.send_json({
"type": "dispatch_ack",
"request_id": data.get("request_id"),
"ok": True,
"prompt_id": prompt_id,
})
except Exception as exc:
await ws.send_json({
"type": "dispatch_ack",
"request_id": data.get("request_id"),
"ok": False,
"error": str(exc),
})
elif msg.type == aiohttp.WSMsgType.ERROR:
log(f"[Distributed] Worker websocket error: {ws.exception()}")
return ws
@server.PromptServer.instance.routes.post("/distributed/worker/clear_launching")
async def clear_launching_state(request):
"""Clear the launching flag when worker is confirmed running."""
@@ -261,6 +379,49 @@ async def get_network_info_endpoint(request):
"recommended_ip": None
})
@server.PromptServer.instance.routes.get("/distributed/tunnel/status")
async def tunnel_status_endpoint(request):
"""Return Cloudflare tunnel status and last known details."""
try:
status = cloudflare_tunnel_manager.get_status()
config = load_config()
master_host = (config.get("master") or {}).get("host")
return web.json_response({
"status": "success",
"tunnel": status,
"master_host": master_host
})
except Exception as e:
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.post("/distributed/tunnel/start")
async def tunnel_start_endpoint(request):
"""Start a Cloudflare tunnel pointing at the current ComfyUI server."""
try:
result = await cloudflare_tunnel_manager.start_tunnel()
config = load_config()
return web.json_response({
"status": "success",
"tunnel": result,
"master_host": (config.get("master") or {}).get("host")
})
except Exception as e:
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.post("/distributed/tunnel/stop")
async def tunnel_stop_endpoint(request):
"""Stop the managed Cloudflare tunnel if running."""
try:
result = await cloudflare_tunnel_manager.stop_tunnel()
config = load_config()
return web.json_response({
"status": "success",
"tunnel": result,
"master_host": (config.get("master") or {}).get("host")
})
except Exception as e:
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.get("/distributed/system_info")
async def get_system_info_endpoint(request):
"""Get system information including machine ID for local worker detection."""
@@ -315,7 +476,7 @@ async def update_worker_endpoint(request):
if data["host"] is None:
worker.pop("host", None)
else:
worker["host"] = data["host"]
worker["host"] = normalize_host(data["host"])
# Handle cuda_device field - remove it if None
if "cuda_device" in data:
@@ -344,7 +505,7 @@ async def update_worker_endpoint(request):
new_worker = {
"id": worker_id,
"name": data["name"],
"host": data.get("host", "localhost"),
"host": normalize_host(data.get("host", "localhost")),
"port": data["port"],
"cuda_device": data["cuda_device"],
"enabled": data.get("enabled", False),
@@ -621,7 +782,8 @@ async def get_local_worker_status_endpoint(request):
for worker in config.get("workers", []):
# Only check local workers
if not worker.get("host") or worker.get("host") in ["localhost", "127.0.0.1"]:
host = normalize_host(worker.get("host")) or ""
if not host or host in ["localhost", "127.0.0.1"]:
worker_id = worker["id"]
port = worker["port"]
@@ -1272,7 +1434,7 @@ async def get_client_session():
# Local worker detection functions
async def is_local_worker(worker_config):
"""Check if a worker is running on the same machine as the master."""
host = worker_config.get('host', 'localhost')
host = normalize_host(worker_config.get('host', 'localhost')) or 'localhost'
if host in ['localhost', '127.0.0.1', '0.0.0.0', ''] or worker_config.get('type') == 'local':
return True
@@ -1289,7 +1451,7 @@ async def is_same_physical_host(worker_config):
master_machine_id = get_machine_id()
# Fetch worker's machine ID via API
host = worker_config.get('host', 'localhost')
host = normalize_host(worker_config.get('host', 'localhost')) or 'localhost'
port = worker_config.get('port', 8188)
session = await get_client_session()
@@ -1352,11 +1514,11 @@ async def get_comms_channel(worker_id, worker_config):
return f"http://127.0.0.1:{worker_config['port']}"
elif worker_config.get('type') == 'cloud' and is_runpod_environment():
# Runpod same-host optimization (if detected)
host = worker_config.get('host', 'localhost')
host = normalize_host(worker_config.get('host', 'localhost')) or 'localhost'
return f"http://{host}:{worker_config['port']}"
else:
# Remote worker: use configured endpoint
host = worker_config.get('host', 'localhost')
host = normalize_host(worker_config.get('host', 'localhost')) or 'localhost'
return f"http://{host}:{worker_config['port']}"
# Auto-launch workers if enabled
@@ -1383,7 +1545,7 @@ def auto_launch_workers():
worker_name = worker.get('name', f'Worker {worker_id}')
# Skip remote workers
host = worker.get('host', 'localhost').lower()
host = (normalize_host(worker.get('host', 'localhost')) or 'localhost').lower()
if host not in ['localhost', '127.0.0.1', '', None]:
debug_log(f"Skipping remote worker {worker_name} (host: {host})")
continue
@@ -1437,6 +1599,10 @@ async def async_cleanup_and_exit(signum=None):
else:
print("\n[Distributed] Master shutting down, workers will continue running")
worker_manager.save_processes()
try:
await cloudflare_tunnel_manager.stop_tunnel()
except Exception as tunnel_error:
debug_log(f"Error stopping Cloudflare tunnel during shutdown: {tunnel_error}")
except Exception as e:
print(f"[Distributed] Error during cleanup: {e}")
@@ -1495,6 +1661,17 @@ def sync_cleanup():
else:
print("\n[Distributed] Master shutting down, workers will continue running")
worker_manager.save_processes()
try:
loop = asyncio.get_event_loop()
if loop.is_running():
loop.create_task(cloudflare_tunnel_manager.stop_tunnel())
else:
loop.run_until_complete(cloudflare_tunnel_manager.stop_tunnel())
except RuntimeError:
# No running loop; create a temporary one
asyncio.run(cloudflare_tunnel_manager.stop_tunnel())
except Exception as tunnel_error:
debug_log(f"Error stopping Cloudflare tunnel during sync cleanup: {tunnel_error}")
except Exception as e:
print(f"[Distributed] Error during cleanup: {e}")
@@ -1517,6 +1694,10 @@ if not os.environ.get('COMFYUI_MASTER_PID'):
worker_manager.cleanup_all()
else:
worker_manager.save_processes()
try:
asyncio.run(cloudflare_tunnel_manager.stop_tunnel())
except Exception as tunnel_error:
print(f"[Distributed] Error stopping Cloudflare tunnel: {tunnel_error}")
except Exception as cleanup_error:
print(f"[Distributed] Error during cleanup: {cleanup_error}")
sys.exit(0)
@@ -1536,6 +1717,45 @@ if not hasattr(prompt_server, 'distributed_pending_jobs'):
prompt_server.distributed_pending_jobs = {}
prompt_server.distributed_jobs_lock = asyncio.Lock()
from .distributed_queue_api import orchestrate_distributed_execution
@server.PromptServer.instance.routes.post("/distributed/queue")
async def distributed_queue_endpoint(request):
"""Queue a distributed workflow, mirroring the UI orchestration pipeline."""
try:
data = await request.json()
except Exception as exc:
return await handle_api_error(request, f"Invalid JSON payload: {exc}", 400)
prompt = data.get("prompt")
if not isinstance(prompt, dict):
return await handle_api_error(request, "Field 'prompt' must be an object", 400)
workflow_meta = data.get("workflow")
client_id = data.get("client_id")
delegate_master = data.get("delegate_master")
enabled_ids = data.get("enabled_worker_ids")
if enabled_ids is not None:
if not isinstance(enabled_ids, list):
return await handle_api_error(request, "enabled_worker_ids must be a list of worker IDs", 400)
enabled_ids = [str(worker_id) for worker_id in enabled_ids]
try:
prompt_id, worker_count = await orchestrate_distributed_execution(
prompt,
workflow_meta,
client_id,
enabled_worker_ids=enabled_ids,
delegate_master=delegate_master,
)
return web.json_response({
"prompt_id": prompt_id,
"worker_count": worker_count,
})
except Exception as exc:
return await handle_api_error(request, exc, 500)
@server.PromptServer.instance.routes.post("/distributed/load_image")
async def load_image_endpoint(request):
"""Load an image or video file and return it as base64 data with hash."""
@@ -1767,6 +1987,21 @@ async def job_complete_endpoint(request):
log(f"Error processing image from worker {worker_id}: {e}")
return await handle_api_error(request, f"Image processing error: {e}", 400)
# Parse audio data if present
audio_data = None
audio_waveform_field = data.get('audio_waveform')
audio_sample_rate = data.get('audio_sample_rate')
if audio_waveform_field is not None:
try:
waveform_bytes = audio_waveform_field.file.read()
waveform_buffer = io.BytesIO(waveform_bytes)
waveform_tensor = torch.load(waveform_buffer, weights_only=True)
sample_rate = int(audio_sample_rate) if audio_sample_rate else 44100
audio_data = {"waveform": waveform_tensor, "sample_rate": sample_rate}
debug_log(f"Received audio from worker {worker_id}: shape={waveform_tensor.shape}, sample_rate={sample_rate}")
except Exception as e:
log(f"Error parsing audio from worker {worker_id}: {e}")
# Put batch into queue
async with prompt_server.distributed_jobs_lock:
debug_log(f"Current pending jobs: {list(prompt_server.distributed_pending_jobs.keys())}")
@@ -1777,7 +2012,8 @@ async def job_complete_endpoint(request):
'worker_id': worker_id,
'tensors': tensors,
'indices': indices,
'is_last': is_last
'is_last': is_last,
'audio': audio_data
})
debug_log(f"Received batch result for job {multi_job_id} from worker {worker_id}, size={len(tensors)}")
else:
@@ -1786,7 +2022,8 @@ async def job_complete_endpoint(request):
'tensor': tensors[0],
'worker_id': worker_id,
'image_index': int(image_index) if image_index else 0,
'is_last': is_last
'is_last': is_last,
'audio': audio_data
})
debug_log(f"Received single result for job {multi_job_id} from worker {worker_id}")
@@ -1805,6 +2042,7 @@ class DistributedCollectorNode:
def INPUT_TYPES(s):
return {
"required": { "images": ("IMAGE",) },
"optional": { "audio": ("AUDIO",) },
"hidden": {
"multi_job_id": ("STRING", {"default": ""}),
"is_worker": ("BOOLEAN", {"default": False}),
@@ -1813,48 +2051,51 @@ class DistributedCollectorNode:
"worker_batch_size": ("INT", {"default": 1, "min": 1, "max": 1024}),
"worker_id": ("STRING", {"default": ""}),
"pass_through": ("BOOLEAN", {"default": False}),
"delegate_only": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_TYPES = ("IMAGE", "AUDIO")
RETURN_NAMES = ("images", "audio")
FUNCTION = "run"
CATEGORY = "image"
def run(self, images, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False):
def run(self, images, 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):
# Create empty audio if not provided
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
if not multi_job_id or pass_through:
if pass_through:
print(f"[Distributed Collector] Pass-through mode enabled, returning images unchanged")
return (images,)
return (images, audio if audio is not None else empty_audio)
# Use async helper to run in server loop
result = run_async_in_server_loop(
self.execute(images, multi_job_id, is_worker, master_url, enabled_worker_ids, worker_batch_size, worker_id)
self.execute(images, audio, multi_job_id, is_worker, master_url, enabled_worker_ids, worker_batch_size, worker_id, delegate_only)
)
return result
async def send_batch_to_master(self, image_batch, multi_job_id, master_url, worker_id):
"""Send image batch to master, chunked if large."""
async def send_batch_to_master(self, image_batch, audio, multi_job_id, master_url, worker_id):
"""Send image batch and optional audio to master, chunked if large."""
batch_size = image_batch.shape[0]
if batch_size == 0:
if batch_size == 0 and audio is None:
return
for start in range(0, batch_size, MAX_BATCH):
chunk = image_batch[start:start + MAX_BATCH]
chunk_size = chunk.shape[0]
is_chunk_last = (start + chunk_size == batch_size) # True only for final chunk
data = aiohttp.FormData()
data.add_field('multi_job_id', multi_job_id)
data.add_field('worker_id', str(worker_id))
data.add_field('is_last', str(is_chunk_last))
data.add_field('batch_size', str(chunk_size))
# Chunk metadata: Absolute index from full batch
metadata = [{'index': start + j} for j in range(chunk_size)]
data.add_field('images_metadata', json.dumps(metadata), content_type='application/json')
# Add chunk images
for j in range(chunk_size):
# Convert tensor slice to PIL
@@ -1863,9 +2104,21 @@ class DistributedCollectorNode:
img.save(byte_io, format='PNG', compress_level=0)
byte_io.seek(0)
data.add_field(f'image_{j}', byte_io, filename=f'image_{j}.png', content_type='image/png')
# Add audio data only on the final chunk to avoid duplication
if is_chunk_last and audio is not None:
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate", 44100)
if waveform is not None and waveform.numel() > 0:
# Serialize waveform tensor to bytes
audio_bytes = io.BytesIO()
torch.save(waveform, audio_bytes)
audio_bytes.seek(0)
data.add_field('audio_waveform', audio_bytes, filename='audio.pt', content_type='application/octet-stream')
data.add_field('audio_sample_rate', str(sample_rate))
debug_log(f"Worker - Including audio: shape={waveform.shape}, sample_rate={sample_rate}")
try:
session = await get_client_session()
url = f"{master_url}/distributed/job_complete"
async with session.post(url, data=data) as response:
@@ -1875,30 +2128,76 @@ class DistributedCollectorNode:
debug_log(f"Worker - Full error details: URL={url}")
raise # Re-raise to handle at caller level
def _combine_audio(self, master_audio, worker_audio, empty_audio):
"""Combine audio from master and workers into a single audio output."""
audio_pieces = []
sample_rate = 44100
# Add master audio first if present
if master_audio is not None:
waveform = master_audio.get("waveform")
if waveform is not None and waveform.numel() > 0:
audio_pieces.append(waveform)
sample_rate = master_audio.get("sample_rate", 44100)
# Add worker audio in sorted order
for worker_id_str in sorted(worker_audio.keys()):
w_audio = worker_audio[worker_id_str]
if w_audio is not None:
waveform = w_audio.get("waveform")
if waveform is not None and waveform.numel() > 0:
audio_pieces.append(waveform)
# Use first available sample rate
if sample_rate == 44100:
sample_rate = w_audio.get("sample_rate", 44100)
if not audio_pieces:
return empty_audio
try:
# Concatenate along the samples dimension (dim=-1)
# Ensure all pieces have same batch and channel dimensions
combined_waveform = torch.cat(audio_pieces, dim=-1)
debug_log(f"Master - Combined audio: {len(audio_pieces)} pieces, final shape={combined_waveform.shape}")
return {"waveform": combined_waveform, "sample_rate": sample_rate}
except Exception as e:
log(f"Master - Error combining audio: {e}")
return empty_audio
async def execute(self, images, audio, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", delegate_only=False):
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
async def execute(self, images, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id=""):
if is_worker:
# Worker mode: send images to master in a single batch
# 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")
await self.send_batch_to_master(images, multi_job_id, master_url, worker_id)
return (images,)
await self.send_batch_to_master(images, audio, multi_job_id, master_url, worker_id)
return (images, audio if audio is not None else empty_audio)
else:
# Master mode: collect images from workers
delegate_mode = delegate_only or is_master_delegate_only()
# Master mode: collect images and audio from workers
enabled_workers = json.loads(enabled_worker_ids)
num_workers = len(enabled_workers)
if num_workers == 0:
return (images,)
images_on_cpu = images.cpu()
master_batch_size = images.shape[0]
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
return (images, audio if audio is not None else empty_audio)
if delegate_mode:
master_batch_size = 0
images_on_cpu = None
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]
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
worker_audio = {} # Dict to store audio by worker_id
# Get the existing queue - it should already exist from prepare_job
async with prompt_server.distributed_jobs_lock:
@@ -1966,19 +2265,25 @@ class DistributedCollectorNode:
# Single image mode (backward compat)
image_index = result['image_index']
tensor = result['tensor']
debug_log(f"Master - Got single result from worker {worker_id}, image {image_index}, is_last={is_last}")
if worker_id not in worker_images:
worker_images[worker_id] = {}
worker_images[worker_id][image_index] = tensor
collected_count += 1
# Collect audio data if present
result_audio = result.get('audio')
if result_audio is not None:
worker_audio[worker_id] = result_audio
debug_log(f"Master - Got audio from worker {worker_id}")
# Record activity and refresh timeout baseline
last_activity = time.time()
base_timeout = float(get_worker_timeout_seconds())
if is_last:
workers_done.add(worker_id)
p.update(1) # +1 per completed worker
@@ -2004,7 +2309,7 @@ class DistributedCollectorNode:
if not wrec:
debug_log(f"Collector probe: worker {wid} not found in config")
continue
host = wrec.get('host') or 'localhost'
host = normalize_host(wrec.get('host') or 'localhost') or 'localhost'
port = int(wrec.get('port', 8188))
url = f"http://{host}:{port}/prompt"
try:
@@ -2101,9 +2406,10 @@ class DistributedCollectorNode:
# Pattern: master img 1, master img 2, worker 1 img 1, worker 1 img 2, worker 2 img 1, worker 2 img 2, etc.
ordered_tensors = []
# Add master images first
for i in range(master_batch_size):
ordered_tensors.append(images_on_cpu[i:i+1])
# Add master images first (if any)
if not delegate_mode and images_on_cpu is not None:
for i in range(master_batch_size):
ordered_tensors.append(images_on_cpu[i:i+1])
# Add worker images in order
# The worker IDs in worker_images are already strings (e.g., "1", "2")
@@ -2123,15 +2429,25 @@ class DistributedCollectorNode:
cpu_tensors.append(t)
try:
combined = torch.cat(cpu_tensors, dim=0)
if cpu_tensors:
combined = torch.cat(cpu_tensors, dim=0)
else:
# No tensors collected (likely delegate mode with no worker output)
combined = ensure_contiguous(images) if images is not None else images
if combined is None:
raise ValueError("No image data collected from master or workers")
# Ensure the combined tensor is contiguous and properly formatted
combined = ensure_contiguous(combined)
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total (master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
return (combined,)
# Combine audio from master and workers
combined_audio = self._combine_audio(master_audio, worker_audio, empty_audio)
return (combined, combined_audio)
except Exception as e:
log(f"Master - Error combining images: {e}")
# Return just the master images as fallback
return (images,)
return (images, audio if audio is not None else empty_audio)
# --- Distributor Node ---
class DistributedSeed:
@@ -2195,6 +2511,68 @@ class AnyType(str):
any_type = AnyType("*")
class DistributedModelName:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"default": ""}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
RETURN_TYPES = (any_type,)
RETURN_NAMES = ("output",)
FUNCTION = "log_input"
OUTPUT_NODE = True
CATEGORY = "utils"
def _stringify(self, value):
if isinstance(value, str):
return value
if isinstance(value, (int, float, bool)):
return str(value)
try:
return json.dumps(value, indent=4)
except Exception:
return str(value)
def _update_workflow(self, extra_pnginfo, unique_id, values):
if not extra_pnginfo:
return
info = extra_pnginfo[0] if isinstance(extra_pnginfo, list) else extra_pnginfo
if not isinstance(info, dict) or "workflow" not in info:
return
node_id = None
if isinstance(unique_id, list) and unique_id:
node_id = str(unique_id[0])
elif unique_id is not None:
node_id = str(unique_id)
if not node_id:
return
workflow = info["workflow"]
node = next((x for x in workflow["nodes"] if str(x.get("id")) == node_id), None)
if node:
node["widgets_values"] = [values]
def log_input(self, text, unique_id=None, extra_pnginfo=None):
values = []
if isinstance(text, list):
for val in text:
values.append(self._stringify(val))
else:
values.append(self._stringify(text))
# Keep widget display in workflow metadata if available.
self._update_workflow(extra_pnginfo, unique_id, values)
if isinstance(values, list) and len(values) == 1:
return {"ui": {"text": values}, "result": (values[0],)}
return {"ui": {"text": values}, "result": (values,)}
class ByPassTypeTuple(tuple):
def __getitem__(self, index):
if index > 0:
@@ -2262,13 +2640,261 @@ class ImageBatchDivider:
return tuple(outputs)
NODE_CLASS_MAPPINGS = {
class AudioBatchDivider:
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"audio": ("AUDIO",),
"divide_by": ("INT", {
"default": 2,
"min": 1,
"max": 10,
"step": 1,
"display": "number",
"tooltip": "Number of parts to divide the audio into"
}),
}
}
RETURN_TYPES = ByPassTypeTuple(("AUDIO",)) # Flexible for variable outputs
RETURN_NAMES = ByPassTypeTuple(tuple([f"audio_{i+1}" for i in range(10)]))
FUNCTION = "divide_audio"
OUTPUT_NODE = True
CATEGORY = "audio"
def divide_audio(self, audio, divide_by):
import torch
waveform = audio.get("waveform")
sample_rate = audio.get("sample_rate", 44100)
if waveform is None or waveform.numel() == 0:
# Return empty audio for all outputs
empty_audio = {"waveform": torch.zeros(1, 2, 1), "sample_rate": sample_rate}
return tuple([empty_audio] * 10)
total_splits = min(divide_by, 10) # Cap to max 10
# Waveform shape: [batch, channels, samples]
total_samples = waveform.shape[-1]
samples_per_split = total_samples // total_splits
remainder = total_samples % total_splits
outputs = []
start_idx = 0
for i in range(total_splits):
current_samples = samples_per_split + (1 if i < remainder else 0)
end_idx = start_idx + current_samples
split_waveform = waveform[..., start_idx:end_idx]
outputs.append({
"waveform": split_waveform,
"sample_rate": sample_rate
})
start_idx = end_idx
# Pad with empty audio up to max (10) to match RETURN_TYPES length
empty_audio = {
"waveform": torch.zeros(waveform.shape[0], waveform.shape[1], 1,
dtype=waveform.dtype, device=waveform.device),
"sample_rate": sample_rate
}
while len(outputs) < 10:
outputs.append(empty_audio)
return tuple(outputs)
class DistributedEmptyImage:
"""Produces an empty IMAGE batch used when the master delegates all work."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"height": ("INT", {"default": 64, "min": 1, "max": 4096, "step": 1}),
"width": ("INT", {"default": 64, "min": 1, "max": 4096, "step": 1}),
"channels": ("INT", {"default": 3, "min": 1, "max": 4, "step": 1}),
}
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "create"
CATEGORY = "image"
def create(self, height, width, channels):
import torch
shape = (0, height, width, channels)
tensor = torch.zeros(shape, dtype=torch.float32)
return (tensor,)
# --- Distributed Queue Node ---
class DistributedQueueNode:
"""Routes entire workflows to the least-busy worker for job scheduling/load balancing."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {},
"hidden": {
"prompt": "PROMPT",
"extra_pnginfo": "EXTRA_PNGINFO",
"client_id": "CLIENT_ID",
"skip_dispatch": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("worker_id", "prompt_id")
FUNCTION = "queue"
OUTPUT_NODE = True
CATEGORY = "utils"
def queue(self, prompt=None, extra_pnginfo=None, client_id=None, skip_dispatch=False):
if skip_dispatch:
debug_log("DistributedQueue: skip_dispatch enabled; not dispatching.")
return ("", "")
if not prompt:
log("DistributedQueue: No prompt supplied; skipping dispatch.")
return ("", "")
return run_async_in_server_loop(
self._queue_on_worker(prompt, extra_pnginfo, client_id)
)
async def _queue_on_worker(self, prompt_obj, workflow_meta, client_id):
config = load_config()
enabled_workers = [
worker for worker in config.get("workers", [])
if worker.get("enabled", False)
]
if not enabled_workers:
log("DistributedQueue: No enabled workers found.")
return ("", "")
statuses = await self._fetch_worker_statuses(enabled_workers)
if not statuses:
log("DistributedQueue: No reachable workers found.")
return ("", "")
selected = self._select_worker(statuses)
worker = selected["worker"]
queue_remaining = selected["queue_remaining"]
prompt_copy = json.loads(json.dumps(prompt_obj))
self._mark_skip_dispatch(prompt_copy)
payload = {"prompt": prompt_copy}
extra_data = {}
if workflow_meta:
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
if client_id:
extra_data["client_id"] = client_id
if extra_data:
payload["extra_data"] = extra_data
url = self._build_worker_url(worker, "/prompt")
session = await get_client_session()
async with session.post(
url,
json=payload,
timeout=aiohttp.ClientTimeout(total=60),
) as resp:
resp.raise_for_status()
response_payload = await resp.json()
prompt_id = str(response_payload.get("prompt_id", ""))
worker_id = str(worker.get("id", ""))
log(
f"DistributedQueue: queued prompt {prompt_id} on worker {worker_id} (queue_remaining={queue_remaining})."
)
return (worker_id, prompt_id)
async def _fetch_worker_statuses(self, workers):
session = await get_client_session()
statuses = []
for worker in workers:
url = self._build_worker_url(worker, "/prompt")
try:
async with session.get(
url,
timeout=aiohttp.ClientTimeout(total=3),
) as resp:
if resp.status != 200:
continue
data = await resp.json()
queue_remaining = int(data.get("exec_info", {}).get("queue_remaining", 0))
statuses.append({
"worker": worker,
"queue_remaining": queue_remaining,
})
except Exception as e:
debug_log(f"DistributedQueue: Worker status check failed ({url}): {e}")
return statuses
def _select_worker(self, statuses):
idle_workers = [status for status in statuses if status["queue_remaining"] == 0]
if idle_workers:
return self._select_round_robin(idle_workers)
return min(statuses, key=lambda status: status["queue_remaining"])
def _select_round_robin(self, statuses):
prompt_server = server.PromptServer.instance
if not hasattr(prompt_server, "distributed_queue_rr_index"):
prompt_server.distributed_queue_rr_index = 0
index = prompt_server.distributed_queue_rr_index % len(statuses)
prompt_server.distributed_queue_rr_index += 1
return statuses[index]
def _mark_skip_dispatch(self, prompt_obj):
for node in prompt_obj.values():
if isinstance(node, dict) and node.get("class_type") == "DistributedQueue":
inputs = node.setdefault("inputs", {})
inputs["skip_dispatch"] = True
def _build_worker_url(self, worker, endpoint=""):
host = (worker.get("host") or "").strip()
port = int(worker.get("port", worker.get("listen_port", 8188)) or 8188)
if not host:
host = getattr(server.PromptServer.instance, "address", "127.0.0.1") or "127.0.0.1"
if host.startswith(("http://", "https://")):
base = host.rstrip("/")
else:
is_cloud = worker.get("type") == "cloud" or host.endswith(".proxy.runpod.net") or port == 443
scheme = "https" if is_cloud else "http"
default_port = 443 if scheme == "https" else 80
port_part = "" if port == default_port else f":{port}"
base = f"{scheme}://{host}{port_part}"
return f"{base}{endpoint}"
NODE_CLASS_MAPPINGS = {
"DistributedQueue": DistributedQueueNode,
"DistributedCollector": DistributedCollectorNode,
"DistributedSeed": DistributedSeed,
"ImageBatchDivider": ImageBatchDivider
"DistributedModelName": DistributedModelName,
"ImageBatchDivider": ImageBatchDivider,
"AudioBatchDivider": AudioBatchDivider,
"DistributedEmptyImage": DistributedEmptyImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"DistributedQueue": "Distributed Queue",
"DistributedCollector": "Distributed Collector",
"DistributedSeed": "Distributed Seed",
"ImageBatchDivider": "Image Batch Divider"
"DistributedModelName": "Distributed Model Name",
"ImageBatchDivider": "Image Batch Divider",
"AudioBatchDivider": "Audio Batch Divider",
"DistributedEmptyImage": "Distributed Empty Image",
}
+605
View File
@@ -0,0 +1,605 @@
import asyncio
import json
import time
import uuid
from collections import deque
import aiohttp
import execution
import server
from .utils.config import load_config
from .utils.logging import debug_log, log
prompt_server = server.PromptServer.instance
def ensure_distributed_state():
"""Ensure prompt_server has the state used by distributed queue orchestration."""
if not hasattr(prompt_server, "distributed_pending_jobs"):
prompt_server.distributed_pending_jobs = {}
if not hasattr(prompt_server, "distributed_jobs_lock"):
prompt_server.distributed_jobs_lock = asyncio.Lock()
if not hasattr(prompt_server, "distributed_worker_ws"):
prompt_server.distributed_worker_ws = {}
async def _get_client_session():
"""Get or create aiohttp client session (shared with distributed.py)."""
if not hasattr(prompt_server, "_distributed_session"):
prompt_server._distributed_session = aiohttp.ClientSession()
return prompt_server._distributed_session
class PromptIndex:
"""Cache prompt metadata for faster worker/master prompt preparation."""
def __init__(self, prompt_obj):
self._prompt_json = json.dumps(prompt_obj)
self.nodes_by_class = {}
self.class_by_node = {}
self.inputs_by_node = {}
for node_id, node in _iter_prompt_nodes(prompt_obj):
class_type = node.get("class_type")
node_id_str = str(node_id)
if class_type:
self.nodes_by_class.setdefault(class_type, []).append(node_id_str)
self.class_by_node[node_id_str] = class_type
self.inputs_by_node[node_id_str] = node.get("inputs", {})
self._upstream_cache = {}
def copy_prompt(self):
return json.loads(self._prompt_json)
def nodes_for_class(self, class_name):
return self.nodes_by_class.get(class_name, [])
def has_upstream(self, start_node_id, target_class):
cache_key = (str(start_node_id), target_class)
if cache_key in self._upstream_cache:
return self._upstream_cache[cache_key]
visited = set()
stack = [str(start_node_id)]
while stack:
node_id = stack.pop()
if node_id in visited:
continue
visited.add(node_id)
inputs = self.inputs_by_node.get(node_id, {})
for value in inputs.values():
if isinstance(value, list) and len(value) == 2:
upstream_id = str(value[0])
if self.class_by_node.get(upstream_id) == target_class:
self._upstream_cache[cache_key] = True
return True
if upstream_id in self.inputs_by_node:
stack.append(upstream_id)
self._upstream_cache[cache_key] = False
return False
def _iter_prompt_nodes(prompt_obj):
for node_id, node in prompt_obj.items():
if isinstance(node, dict):
yield str(node_id), node
def _find_nodes_by_class(prompt_obj, class_name):
nodes = []
for node_id, node in _iter_prompt_nodes(prompt_obj):
if node.get("class_type") == class_name:
nodes.append(node_id)
return nodes
def _find_downstream_nodes(prompt_obj, start_ids):
"""Return all nodes reachable downstream from the provided IDs."""
adjacency = {}
for node_id, node in _iter_prompt_nodes(prompt_obj):
inputs = node.get("inputs", {})
for value in inputs.values():
if isinstance(value, list) and len(value) == 2:
source_id = str(value[0])
adjacency.setdefault(source_id, set()).add(str(node_id))
connected = set(start_ids)
queue = deque(start_ids)
while queue:
current = queue.popleft()
for dependent in adjacency.get(current, ()): # pragma: no branch - simple iteration
if dependent not in connected:
connected.add(dependent)
queue.append(dependent)
return connected
def _create_numeric_id_generator(prompt_obj):
"""Return a closure that yields new numeric string IDs."""
max_id = 0
for node_id in prompt_obj.keys():
try:
numeric = int(node_id)
except (TypeError, ValueError):
continue
max_id = max(max_id, numeric)
counter = max_id
def _next_id():
nonlocal counter
counter += 1
return str(counter)
return _next_id
def _prepare_delegate_master_prompt(prompt_obj, collector_ids):
"""Prune master prompt so it only executes post-collector nodes in delegate mode."""
downstream = _find_downstream_nodes(prompt_obj, collector_ids)
nodes_to_keep = set(collector_ids)
nodes_to_keep.update(downstream)
pruned_prompt = {}
for node_id in nodes_to_keep:
node = prompt_obj.get(node_id)
if node is not None:
pruned_prompt[node_id] = json.loads(json.dumps(node))
pruned_ids = set(pruned_prompt.keys())
for node_id, node in pruned_prompt.items():
inputs = node.get("inputs")
if not inputs:
continue
for input_name, input_value in list(inputs.items()):
if isinstance(input_value, list) and len(input_value) == 2:
source_id = str(input_value[0])
if source_id not in pruned_ids:
inputs.pop(input_name, None)
debug_log(
f"Removed upstream reference '{input_name}' from node {node_id} for delegate-only master prompt."
)
next_id = _create_numeric_id_generator(pruned_prompt)
for collector_id in collector_ids:
collector_entry = pruned_prompt.get(collector_id)
if not collector_entry:
continue
placeholder_id = next_id()
pruned_prompt[placeholder_id] = {
"class_type": "DistributedEmptyImage",
"inputs": {
"height": 64,
"width": 64,
"channels": 3,
},
"_meta": {
"title": "Distributed Empty Image (auto-added)",
},
}
collector_entry.setdefault("inputs", {})["images"] = [placeholder_id, 0]
debug_log(
f"Inserted placeholder node {placeholder_id} for collector {collector_id} in delegate-only master prompt."
)
return pruned_prompt
async def _ensure_distributed_queue(job_id):
"""Ensure a queue exists for the given distributed job ID."""
ensure_distributed_state()
async with prompt_server.distributed_jobs_lock:
if job_id not in prompt_server.distributed_pending_jobs:
prompt_server.distributed_pending_jobs[job_id] = asyncio.Queue()
def _generate_job_id_map(prompt_index, prefix):
"""Create stable per-node job IDs for distributed nodes."""
job_map = {}
distributed_nodes = prompt_index.nodes_for_class("DistributedCollector") + prompt_index.nodes_for_class(
"UltimateSDUpscaleDistributed"
)
for node_id in distributed_nodes:
job_map[node_id] = f"{prefix}_{node_id}"
return job_map
def _resolve_enabled_workers(config, requested_ids=None):
"""Return a list of worker configs that should participate."""
workers = []
for worker in config.get("workers", []):
worker_id = str(worker.get("id") or "").strip()
if not worker_id:
continue
if requested_ids is not None:
if worker_id not in requested_ids:
continue
elif not worker.get("enabled", False):
continue
raw_port = worker.get("port", worker.get("listen_port", 8188))
try:
port = int(raw_port or 8188)
except (TypeError, ValueError):
log(f"[Distributed] Invalid port '{raw_port}' for worker {worker_id}; defaulting to 8188.")
port = 8188
workers.append(
{
"id": worker_id,
"name": worker.get("name", worker_id),
"host": worker.get("host"),
"port": port,
"type": worker.get("type", "local"),
}
)
return workers
def _resolve_master_url():
"""Best-effort reconstruction of the master's public URL."""
cfg = load_config()
master_cfg = cfg.get("master", {}) or {}
configured_host = (master_cfg.get("host") or "").strip()
configured_port = master_cfg.get("port")
default_port = getattr(prompt_server, "port", 8188) or 8188
port = int(configured_port or default_port)
def _needs_https(hostname):
hostname = hostname.lower()
https_domains = (
".proxy.runpod.net",
".ngrok-free.app",
".ngrok-free.dev",
".ngrok.io",
".trycloudflare.com",
".cloudflare.dev",
)
return any(hostname.endswith(suffix) for suffix in https_domains)
if configured_host:
if configured_host.startswith(("http://", "https://")):
return configured_host.rstrip("/")
host = configured_host
scheme = "https" if _needs_https(host) or port == 443 else "http"
default_port_for_scheme = 443 if scheme == "https" else 80
# For ngrok/cloud domains without explicit port, default to the scheme's default
if configured_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}"
address = getattr(prompt_server, "address", "127.0.0.1") or "127.0.0.1"
if address in ("0.0.0.0", "::"):
address = "127.0.0.1"
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_worker_url(worker, endpoint=""):
"""Construct the worker base URL with optional endpoint."""
host = (worker.get("host") or "").strip()
port = int(worker.get("port", 8188) or 8188)
if not host:
host = getattr(prompt_server, "address", "127.0.0.1") or "127.0.0.1"
if host.startswith(("http://", "https://")):
base = host.rstrip("/")
else:
is_cloud = worker.get("type") == "cloud" or host.endswith(".proxy.runpod.net") or port == 443
scheme = "https" if is_cloud else "http"
default_port = 443 if scheme == "https" else 80
port_part = "" if port == default_port else f":{port}"
base = f"{scheme}://{host}{port_part}"
return f"{base}{endpoint}"
async def _worker_is_active(worker):
"""Ping worker's /prompt endpoint to confirm it's reachable."""
url = _build_worker_url(worker, "/prompt")
session = await _get_client_session()
try:
async with session.get(url, timeout=aiohttp.ClientTimeout(total=3)) as resp:
return resp.status == 200
except Exception:
return False
async def _worker_ws_is_active(worker):
"""Ping worker's websocket endpoint to confirm it's reachable."""
session = await _get_client_session()
url = _build_worker_url(worker, "/distributed/worker_ws")
try:
ws = await session.ws_connect(url, heartbeat=20, timeout=3)
await ws.close()
return True
except Exception:
return False
async def _get_worker_ws_connection(worker):
"""Return a cached websocket connection to a worker (creating if needed)."""
ensure_distributed_state()
worker_id = worker.get("id")
existing = prompt_server.distributed_worker_ws.get(worker_id)
if existing and not existing.closed:
return existing
session = await _get_client_session()
url = _build_worker_url(worker, "/distributed/worker_ws")
ws = await session.ws_connect(url, heartbeat=20, timeout=10)
prompt_server.distributed_worker_ws[worker_id] = ws
return ws
async def _close_worker_ws(worker_id):
ws = prompt_server.distributed_worker_ws.pop(worker_id, None)
if ws and not ws.closed:
await ws.close()
async def _dispatch_worker_prompt(worker, prompt_obj, workflow_meta, client_id=None, use_websocket=False):
"""Send the prepared prompt to a worker ComfyUI instance."""
url = _build_worker_url(worker, "/prompt")
payload = {"prompt": prompt_obj}
extra_data = {}
if workflow_meta:
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
if extra_data:
payload["extra_data"] = extra_data
if use_websocket:
request_id = uuid.uuid4().hex
ws_payload = {
"type": "dispatch_prompt",
"request_id": request_id,
"prompt": prompt_obj,
"workflow": workflow_meta,
"client_id": client_id,
}
try:
ws = await _get_worker_ws_connection(worker)
await ws.send_json(ws_payload)
response = await asyncio.wait_for(ws.receive(), timeout=60)
if response.type == aiohttp.WSMsgType.TEXT:
data = json.loads(response.data or "{}")
if (
data.get("type") == "dispatch_ack"
and data.get("request_id") == request_id
and data.get("ok")
):
return
error_msg = data.get("error") or "Worker rejected websocket dispatch."
raise RuntimeError(error_msg)
raise RuntimeError("Unexpected websocket response from worker.")
except Exception as exc:
worker_id = worker.get("id")
await _close_worker_ws(worker_id)
log(f"[Distributed] Websocket dispatch failed for worker {worker_id}: {exc}")
session = await _get_client_session()
async with session.post(
url,
json=payload,
timeout=aiohttp.ClientTimeout(total=60),
) as resp:
resp.raise_for_status()
async def _queue_master_prompt(prompt_obj, workflow_meta, client_id):
"""Queue the master prompt through ComfyUI's prompt queue."""
payload = {"prompt": prompt_obj}
payload = prompt_server.trigger_on_prompt(payload)
prompt = payload["prompt"]
prompt_id = str(uuid.uuid4())
valid = await execution.validate_prompt(prompt_id, prompt, None)
if not valid[0]:
raise RuntimeError(f"Invalid prompt: {valid[1]}")
extra_data = {}
if workflow_meta:
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
if client_id:
extra_data["client_id"] = client_id
sensitive = {}
for key in getattr(execution, "SENSITIVE_EXTRA_DATA_KEYS", []):
if key in extra_data:
sensitive[key] = extra_data.pop(key)
number = getattr(prompt_server, "number", 0)
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)
return prompt_id
def _apply_participant_overrides(
prompt_copy,
participant_id,
enabled_worker_ids,
job_id_map,
master_url,
delegate_master,
prompt_index,
):
"""Return a prompt copy with hidden inputs configured for master/worker."""
is_master = participant_id == "master"
worker_index_map = {wid: idx for idx, wid in enumerate(enabled_worker_ids)}
enabled_json = json.dumps(enabled_worker_ids)
# Distributed seed nodes
for node_id in prompt_index.nodes_for_class("DistributedSeed"):
node = prompt_copy.get(node_id, {})
inputs = node.setdefault("inputs", {})
inputs["is_worker"] = not is_master
if not is_master:
idx = worker_index_map.get(participant_id, 0)
inputs["worker_id"] = f"worker_{idx}"
else:
inputs["worker_id"] = ""
# Distributed collectors
for node_id in prompt_index.nodes_for_class("DistributedCollector"):
node = prompt_copy.get(node_id, {})
inputs = node.setdefault("inputs", {})
if prompt_index.has_upstream(node_id, "UltimateSDUpscaleDistributed"):
inputs["pass_through"] = True
continue
unique_id = job_id_map.get(node_id, node_id)
inputs["multi_job_id"] = unique_id
inputs["is_worker"] = not is_master
if is_master:
inputs["enabled_worker_ids"] = enabled_json
inputs["delegate_only"] = bool(delegate_master)
inputs.pop("master_url", None)
inputs.pop("worker_id", None)
else:
inputs["master_url"] = master_url
inputs["worker_id"] = participant_id
inputs["enabled_worker_ids"] = enabled_json
inputs["delegate_only"] = False
# Distributed upscaler nodes
for node_id in prompt_index.nodes_for_class("UltimateSDUpscaleDistributed"):
node = prompt_copy.get(node_id, {})
inputs = node.setdefault("inputs", {})
unique_id = job_id_map.get(node_id, node_id)
inputs["multi_job_id"] = unique_id
inputs["is_worker"] = not is_master
if is_master:
inputs["enabled_worker_ids"] = enabled_json
inputs.pop("master_url", None)
inputs.pop("worker_id", None)
else:
inputs["master_url"] = master_url
inputs["worker_id"] = participant_id
inputs["enabled_worker_ids"] = enabled_json
return prompt_copy
async def orchestrate_distributed_execution(
prompt_obj,
workflow_meta,
client_id,
enabled_worker_ids=None,
delegate_master=None,
):
"""Core orchestration logic for the /distributed/queue endpoint.
Returns:
tuple[str, int]: (prompt_id, worker_count)
"""
ensure_distributed_state()
config = load_config()
use_websocket = bool(config.get("settings", {}).get("websocket_orchestration", False))
requested_ids = enabled_worker_ids if enabled_worker_ids is not None else None
workers = _resolve_enabled_workers(config, requested_ids)
prompt_index = PromptIndex(prompt_obj)
# Respect master delegate-only configuration
if delegate_master is None:
delegate_master = bool(config.get("settings", {}).get("master_delegate_only", False))
if not workers and delegate_master:
debug_log("Delegate-only requested but no workers are enabled. Falling back to master execution.")
delegate_master = False
# Filter to active workers
active_workers = []
for worker in workers:
if use_websocket:
is_active = await _worker_ws_is_active(worker)
else:
is_active = await _worker_is_active(worker)
if is_active:
active_workers.append(worker)
else:
log(f"[Distributed] Worker {worker['name']} ({worker['id']}) is offline, skipping.")
if not active_workers and delegate_master:
debug_log("All workers offline while delegate-only requested; enabling master participation.")
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:
prompt_id = await _queue_master_prompt(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_url = _resolve_master_url()
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:
debug_log(
"Active distributed workers: "
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers)
)
worker_payloads = []
for worker in active_workers:
worker_prompt = prompt_index.copy_prompt()
worker_prompt = _apply_participant_overrides(
worker_prompt,
worker["id"],
enabled_ids,
job_id_map,
master_url,
delegate_master,
prompt_index,
)
worker_payloads.append((worker, worker_prompt))
if worker_payloads:
await asyncio.gather(
*[
_dispatch_worker_prompt(worker, wprompt, workflow_meta, client_id, use_websocket=use_websocket)
for worker, wprompt in worker_payloads
]
)
prompt_id = await _queue_master_prompt(master_prompt, workflow_meta, client_id)
return prompt_id, len(worker_payloads)
+20 -2
View File
@@ -63,6 +63,24 @@ def sync_wrapper(async_func):
# Note: tensor_to_pil and pil_to_tensor are imported from utils.image
def _parse_enabled_worker_ids(enabled_worker_ids):
"""Parse enabled worker IDs from either JSON or list input."""
if isinstance(enabled_worker_ids, list):
return [str(worker_id) for worker_id in enabled_worker_ids]
if not enabled_worker_ids:
return []
if isinstance(enabled_worker_ids, str):
try:
parsed = json.loads(enabled_worker_ids)
except json.JSONDecodeError:
log("USDU Dist: Invalid enabled_worker_ids JSON; defaulting to no workers.")
return []
if isinstance(parsed, list):
return [str(wid) for wid in parsed]
return []
class UltimateSDUpscaleDistributed:
"""
Distributed version of Ultimate SD Upscale (No Upscale).
@@ -248,8 +266,8 @@ class UltimateSDUpscaleDistributed:
if num_workers > 0 and num_tiles_per_image > 1:
mode = "static"
log(f"USDU Dist: Workers {num_workers}")
log(f"USDU Dist: Workers {num_workers} | Mode {mode} | Threshold {dynamic_threshold}")
if mode == "single_gpu":
# No workers, process all tiles locally
return self.process_single_gpu(upscaled_image, model, positive, negative, vae,
+201
View File
@@ -0,0 +1,201 @@
# ComfyUI-Distributed API (Experimental)
This document describes the **public HTTP API** added to ComfyUI-Distributed to allow queueing *distributed* workflows from external tools (scripts, services, CI jobs, render farms, etc.) without using the ComfyUI web UI.
## Demo
- Video walkthrough: https://youtu.be/yiQlPd0MzLk
## Examples Repository
- Examples repo: https://github.com/umanets/ComfyUI-Distributed-API-examples.git
---
## Overview
### What this adds
- `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`).
- Queues the master workflow in ComfyUI’s prompt queue.
### What it does *not* add
- Authentication/authorization.
- A separate “job status” API for distributed results (you still use ComfyUI’s normal prompt history / websocket flow, and the existing `/distributed/queue_status/{job_id}` behavior for collector queues).
---
## Endpoint: `POST /distributed/queue`
Queue a workflow for distributed execution.
### URL
- `http://<master-host>:<master-port>/distributed/queue`
### Headers
- `Content-Type: application/json`
### Request Body
```json
{
"prompt": { "<node_id>": { "class_type": "...", "inputs": { } } },
"workflow": { },
"client_id": "optional",
"delegate_master": false,
"enabled_worker_ids": ["1", "2"]
}
```
#### Fields
- `prompt` (required, object)
- The ComfyUI prompt/workflow graph, same shape as used by `POST /prompt`.
- `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.
- `client_id` (optional, string)
- Passed through as `extra_data.client_id` (useful if you consume ComfyUI websocket events).
- `delegate_master` (optional, boolean)
- If `true`, attempts “workers-only” execution for workflows based on `DistributedCollector`.
- Current limitation: delegate-only mode **does not support** `UltimateSDUpscaleDistributed` and will fall back to running the full prompt on master.
- `enabled_worker_ids` (optional, array of strings)
- If provided, only these worker IDs will be considered.
- If omitted, the plugin uses workers marked as enabled in the UI config.
##### How to get `enabled_worker_ids`
Worker IDs come from the plugin config (`GET /distributed/config`) under `workers[].id`.
Example (bash + `jq`):
```bash
curl -s "http://127.0.0.1:8188/distributed/config" \
| jq -r '.workers[] | "id=\(.id)\tname=\(.name)\tenabled=\(.enabled)\thost=\(.host)\tport=\(.port)\ttype=\(.type)"'
```
Example (PowerShell):
```powershell
$cfg = Invoke-RestMethod "http://127.0.0.1:8188/distributed/config"
$cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -AutoSize
```
### Response Body
```json
{
"prompt_id": "<uuid>",
"worker_count": 2
}
```
- `prompt_id` — the master prompt id queued into ComfyUI.
- `worker_count` — number of workers that received a dispatched prompt (only those that passed the health check).
### Status Codes
- `200` — queued.
- `400` — invalid JSON or invalid body.
- `500` — orchestration/dispatch failure (see server logs for details).
---
## Worker requirements (important)
For a worker to participate, it must be reachable from the master:
- Health check: `GET <worker-base>/prompt` must return HTTP 200.
- Dispatch: `POST <worker-base>/prompt` must accept the workflow.
Also, for collector-based flows:
- Workers will send results back to the master via `POST /distributed/job_complete` (that route must be reachable from workers).
### CORS note
If you call the API from a browser (not from a backend), ensure the master ComfyUI is started with `--enable-cors-header`.
---
## Examples
### 1) Minimal `curl`
```bash
curl -X POST "http://127.0.0.1:8188/distributed/queue" \
-H "Content-Type: application/json" \
-d @payload.json
```
Where `payload.json` contains at least:
```json
{
"prompt": {
"1": {"class_type": "KSampler", "inputs": {} }
}
}
```
### 2) Python (`requests`)
```python
import requests
url = "http://127.0.0.1:8188/distributed/queue"
payload = {
"prompt": {...},
"workflow": {...},
"delegate_master": False,
"enabled_worker_ids": ["1", "2"],
}
r = requests.post(url, json=payload, timeout=60)
r.raise_for_status()
print(r.json())
```
### 3) JavaScript (`fetch`)
```js
const url = "http://127.0.0.1:8188/distributed/queue";
const payload = {
prompt: {/* ... */},
workflow: {/* ... */},
delegate_master: false,
enabled_worker_ids: ["1", "2"],
};
const resp = await fetch(url, {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify(payload),
});
if (!resp.ok) throw new Error(await resp.text());
console.log(await resp.json());
```
---
## Operational notes / gotchas
- If the workflow contains **no distributed nodes**, the endpoint falls back to normal master queueing and returns `worker_count: 0`.
- Worker selection is “best-effort”: offline workers are skipped.
- For public URLs/tunnels: prefer configuring `master.host` with an explicit scheme (`https://...`) to avoid ambiguity.
---
## Changelog (this feature)
- Added `POST /distributed/queue` endpoint.
- Added orchestration module used by the endpoint.
+3 -2
View File
@@ -15,5 +15,6 @@
7. In ComfyUI, open the GPU panel on the left.
> If you set SAGE_ATTENTION to true, add "--use-sage-attention" to Extra Args on the workers.
8. Launch the workers.
9. Upload video, add prompt and run workflow.
10. Right-click the Video Combine node and click Save Preview to save the video.
9. [Load the workflow.](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/workflows/distributed-upscale-video.json)
10. Upload video, add prompt and run workflow.
11. Right-click the Video Combine node and click Save Preview to save the video.
+8
View File
@@ -6,6 +6,14 @@
<img width="600" src="https://github.com/user-attachments/assets/609c42aa-8a1c-4a3f-939e-f3552fa1d54f" />
### Master participation modes
The master can either contribute GPU work or stay in **orchestrator-only** mode:
- **Participating**: Master renders alongside workers, useful when you want every available GPU.
- **Orchestrator-only**: Master sends jobs to selected workers but skips local rendering. Enable this by opening the Distributed panel and unchecking the master toggle. The master card will display *“Master disabled: running as orchestrator only.”*
- **Fallback**: If orchestrator-only is enabled but no workers remain selected, the master automatically re-enables execution to guarantee the workflow still runs. The UI shows a green *“Master fallback execution active”* badge so you know work is executing locally again.
### Types of Workers
- **Local workers**: Additional GPUs on the same machine as the master
+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.1.0"
version = "1.2.1"
license = {file = "LICENSE"}
dependencies = []
+397
View File
@@ -0,0 +1,397 @@
"""
Lightweight Cloudflare Tunnel manager.
Responsibilities:
- Locate or download the cloudflared binary for the current platform.
- Start/stop a quick tunnel pointing to the local ComfyUI server.
- Track status and persist minimal state in the shared config.
"""
import asyncio
import os
import platform
import re
import shutil
import signal
import stat
import subprocess
import threading
import time
from urllib import request, error as urlerror
from .config import load_config, save_config
from .logging import debug_log, log
from .network import get_server_port
from .process import is_process_alive, terminate_process
# Regex to capture the generated public URL from cloudflared output
PUBLIC_URL_PATTERN = re.compile(r"(https?://[\w.-]+\.(?:trycloudflare\.com|cloudflare\.dev))", re.IGNORECASE)
# Default timeout (seconds) to wait for a tunnel URL before giving up
TUNNEL_START_TIMEOUT = 25
def _normalize_host(value):
if not value or not isinstance(value, str):
return ""
value = value.strip()
value = re.sub(r"^https?://", "", value, flags=re.IGNORECASE)
return value.split("/")[0]
class CloudflareTunnelManager:
def __init__(self):
self.process = None
self.pid = None
self.public_url = None
self.last_error = None
self.log_file = None
self.status = "stopped"
self.previous_master_host = None
self._lock = asyncio.Lock()
self._url_event = None
self._reader_thread = None
self._loop = None
self._recent_logs = []
self.binary_path = None
self._restore_state()
# --- Binary discovery & download ---
@property
def base_dir(self):
return os.path.dirname(os.path.dirname(__file__))
@property
def bin_dir(self):
return os.path.join(self.base_dir, "bin")
def _get_asset_name(self):
system = platform.system().lower()
machine = platform.machine().lower()
if system == "windows":
if "arm" in machine:
return "cloudflared-windows-arm64.exe"
return "cloudflared-windows-amd64.exe"
if system == "darwin":
if machine in ("arm64", "aarch64"):
return "cloudflared-darwin-arm64"
return "cloudflared-darwin-amd64"
if system == "linux":
if machine in ("arm64", "aarch64"):
return "cloudflared-linux-arm64"
return "cloudflared-linux-amd64"
raise RuntimeError(f"Unsupported platform for cloudflared: {system}/{machine}")
def _download_cloudflared(self):
asset = self._get_asset_name()
url = f"https://github.com/cloudflare/cloudflared/releases/latest/download/{asset}"
os.makedirs(self.bin_dir, exist_ok=True)
target_path = os.path.join(self.bin_dir, "cloudflared.exe" if asset.endswith(".exe") else "cloudflared")
debug_log(f"Downloading cloudflared from {url}")
try:
with request.urlopen(url, timeout=30) as resp:
with open(target_path, "wb") as f:
shutil.copyfileobj(resp, f)
except urlerror.URLError as exc:
raise RuntimeError(f"Failed to download cloudflared: {exc}") from exc
# Make executable
st = os.stat(target_path)
os.chmod(target_path, st.st_mode | stat.S_IEXEC)
debug_log(f"Downloaded cloudflared to {target_path}")
return target_path
def ensure_binary(self):
# Allow explicit override
env_path = os.environ.get("CLOUDFLARED_PATH")
if env_path and os.path.exists(env_path):
self.binary_path = env_path
return env_path
# Check cached path
cached = self.binary_path
if cached and os.path.exists(cached):
return cached
# Local bin directory
candidate = os.path.join(self.bin_dir, "cloudflared.exe" if platform.system().lower() == "windows" else "cloudflared")
if os.path.exists(candidate):
self.binary_path = candidate
return candidate
# System PATH
path_binary = shutil.which("cloudflared")
if path_binary:
self.binary_path = path_binary
return path_binary
# Download on demand
self.binary_path = self._download_cloudflared()
return self.binary_path
# --- State helpers ---
def _persist_state(self, status=None, public_url=None, pid=None, log_file=None, previous_host=None, master_host=None):
cfg = load_config()
tunnel_cfg = cfg.get("tunnel", {}) if isinstance(cfg.get("tunnel", {}), dict) else {}
if status is not None:
tunnel_cfg["status"] = status
if public_url is not None:
tunnel_cfg["public_url"] = public_url
if pid is not None:
tunnel_cfg["pid"] = pid
if log_file is not None:
tunnel_cfg["log_file"] = log_file
if previous_host is not None:
tunnel_cfg["previous_master_host"] = previous_host
if master_host is not None:
cfg.setdefault("master", {})["host"] = master_host
cfg["tunnel"] = tunnel_cfg
save_config(cfg)
def _restore_state(self):
cfg = load_config()
tunnel_cfg = cfg.get("tunnel", {}) if isinstance(cfg.get("tunnel", {}), dict) else {}
self.public_url = tunnel_cfg.get("public_url") or None
self.previous_master_host = tunnel_cfg.get("previous_master_host")
self.log_file = tunnel_cfg.get("log_file")
pid = tunnel_cfg.get("pid")
if pid and is_process_alive(pid):
self.pid = pid
self.status = tunnel_cfg.get("status", "running")
debug_log(f"Detected existing cloudflared process (pid={pid})")
else:
# Clear stale info
self._persist_state(status="stopped", public_url="", pid=None, log_file=None)
self.status = "stopped"
self.pid = None
# --- Process output handling ---
def _append_log(self, line):
if self.log_file:
try:
with open(self.log_file, "a", encoding="utf-8", errors="replace") as f:
f.write(line + "\n")
except Exception as exc: # pragma: no cover - best effort logging
debug_log(f"Failed to write tunnel log: {exc}")
self._recent_logs.append(line)
# Keep last ~200 lines to avoid unbounded growth
if len(self._recent_logs) > 200:
self._recent_logs = self._recent_logs[-200:]
def _reader(self):
assert self.process is not None
loop = self._loop
for raw_line in iter(self.process.stdout.readline, ""):
line = raw_line.strip()
if not line:
continue
self._append_log(line)
match = PUBLIC_URL_PATTERN.search(line)
if match and not self.public_url:
self.public_url = match.group(1).rstrip("/")
self.status = "running"
if self._url_event and loop:
loop.call_soon_threadsafe(self._url_event.set)
# Capture obvious error strings
if "error" in line.lower() and not self.last_error:
self.last_error = line
# Process exited
if self.status == "starting" and not self.public_url:
self.status = "error"
if not self.last_error:
self.last_error = "Cloudflare tunnel exited before becoming ready"
if self._url_event and loop:
loop.call_soon_threadsafe(self._url_event.set)
elif self.status == "running":
# Treat unexpected exit as stopped so UI can retry
self.status = "stopped"
elif self.status == "stopped":
# Normal stop path, nothing to do
pass
# --- Public API ---
async def start_tunnel(self):
async with self._lock:
if self.process and self.process.poll() is None:
return {
"status": self.status,
"public_url": self.public_url,
"pid": self.process.pid,
"log_file": self.log_file,
}
# If we had a stale pid, ensure it's not running
if self.pid and is_process_alive(self.pid):
debug_log(f"Stopping stale cloudflared pid {self.pid} before starting a new one")
await self.stop_tunnel()
binary = await asyncio.to_thread(self.ensure_binary)
port = get_server_port()
self.status = "starting"
self.last_error = None
self.public_url = None
self._recent_logs = []
# Remember current master host so we can restore it if the tunnel stops
config = load_config()
master_host = (config.get("master") or {}).get("host") or ""
tunnel_cfg = config.get("tunnel") or {}
if tunnel_cfg.get("previous_master_host"):
self.previous_master_host = tunnel_cfg.get("previous_master_host")
else:
self.previous_master_host = master_host
os.makedirs(os.path.join(self.base_dir, "logs"), exist_ok=True)
timestamp = time.strftime("%Y%m%d-%H%M%S")
self.log_file = os.path.join(self.base_dir, "logs", f"cloudflare-{timestamp}.log")
cmd = [
binary,
"tunnel",
"--no-autoupdate",
"--url",
f"http://127.0.0.1:{port}",
]
debug_log(f"Starting cloudflared: {' '.join(cmd)}")
try:
self.process = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
)
except FileNotFoundError:
self.status = "error"
raise RuntimeError("cloudflared binary not found")
except Exception as exc:
self.status = "error"
raise RuntimeError(f"Failed to start cloudflared: {exc}") from exc
self.pid = self.process.pid
self._persist_state(
status="starting",
pid=self.pid,
log_file=self.log_file,
previous_host=self.previous_master_host
)
# Kick off background reader
self._loop = asyncio.get_running_loop()
self._url_event = asyncio.Event()
self._reader_thread = threading.Thread(target=self._reader, daemon=True)
self._reader_thread.start()
try:
await asyncio.wait_for(self._url_event.wait(), timeout=TUNNEL_START_TIMEOUT)
except asyncio.TimeoutError:
self.last_error = "Timed out waiting for Cloudflare to assign a URL"
await self.stop_tunnel()
raise RuntimeError(self.last_error)
if self.status != "running" or not self.public_url:
error_msg = self.last_error or "Cloudflare tunnel failed to start"
await self.stop_tunnel()
raise RuntimeError(error_msg)
debug_log(f"Cloudflare tunnel ready at {self.public_url}")
master_host = _normalize_host(self.public_url)
self._persist_state(
status="running",
public_url=self.public_url,
pid=self.pid,
log_file=self.log_file,
previous_host=self.previous_master_host or "",
master_host=master_host
)
return {
"status": self.status,
"public_url": self.public_url,
"pid": self.pid,
"log_file": self.log_file,
}
async def stop_tunnel(self):
async with self._lock:
pid = self.process.pid if self.process else self.pid
if not pid:
self._persist_state(status="stopped", public_url=None, pid=None, log_file=None)
self.status = "stopped"
return {"status": "stopped"}
debug_log(f"Stopping cloudflared (pid={pid})")
if self.process:
terminate_process(self.process, timeout=5)
else:
try:
os.kill(pid, signal.SIGTERM)
time.sleep(0.5)
except Exception as exc: # pragma: no cover - best effort
debug_log(f"Error stopping cloudflared pid {pid}: {exc}")
config = load_config()
tunnel_cfg = config.get("tunnel", {}) if isinstance(config.get("tunnel", {}), dict) else {}
self.previous_master_host = tunnel_cfg.get("previous_master_host", self.previous_master_host)
active_url = tunnel_cfg.get("public_url")
current_master_host = (config.get("master") or {}).get("host")
restore_host = None
if active_url:
active_host = _normalize_host(active_url)
current_host = _normalize_host(current_master_host)
if current_host == active_host:
restore_host = self.previous_master_host or ""
self.status = "stopped"
self.public_url = None
self.pid = None
self.process = None
self.last_error = None
self._url_event = None
if self._reader_thread and self._reader_thread.is_alive():
self._reader_thread.join(timeout=1)
self._reader_thread = None
self._persist_state(
status="stopped",
public_url="",
pid=None,
log_file=self.log_file,
previous_host=self.previous_master_host,
master_host=restore_host
)
return {"status": "stopped"}
def get_status(self):
alive = False
pid = self.process.pid if self.process else self.pid
if pid:
alive = is_process_alive(pid)
if not alive and self.status == "running":
self.status = "stopped"
return {
"status": self.status,
"public_url": self.public_url,
"pid": pid,
"log_file": self.log_file,
"last_error": self.last_error,
"binary_path": self.binary_path or shutil.which("cloudflared"),
"recent_logs": self._recent_logs[-20:],
"previous_master_host": self.previous_master_host,
}
# Singleton tunnel manager
cloudflare_tunnel_manager = CloudflareTunnelManager()
+25 -2
View File
@@ -18,7 +18,16 @@ def get_default_config():
"settings": {
"debug": False,
"auto_launch_workers": False,
"stop_workers_on_master_exit": True
"stop_workers_on_master_exit": True,
"master_delegate_only": False,
"websocket_orchestration": True
},
"tunnel": {
"status": "stopped",
"public_url": "",
"pid": None,
"log_file": "",
"previous_master_host": ""
}
}
@@ -27,7 +36,12 @@ def load_config():
if os.path.exists(CONFIG_FILE):
try:
with open(CONFIG_FILE, 'r') as f:
return json.load(f)
data = json.load(f)
defaults = get_default_config()
for key, value in defaults.items():
if key not in data:
data[key] = value
return data
except Exception as e:
log(f"Error loading config, using defaults: {e}")
return get_default_config()
@@ -69,3 +83,12 @@ def get_worker_timeout_seconds(default: int = HEARTBEAT_TIMEOUT) -> int:
return max(1, val)
except Exception:
return max(1, int(default))
def is_master_delegate_only() -> bool:
"""Returns True when master should skip local workload and act as orchestrator only."""
try:
cfg = load_config()
return bool(cfg.get('settings', {}).get('master_delegate_only', False))
except Exception:
return False
+20 -1
View File
@@ -141,6 +141,25 @@ export function createApiClient(baseUrl) {
return Promise.allSettled(
urls.map(url => this.checkStatus(url))
);
},
// Cloudflare tunnel management
async startTunnel() {
return request('/distributed/tunnel/start', {
method: 'POST',
body: JSON.stringify({})
});
},
async stopTunnel() {
return request('/distributed/tunnel/stop', {
method: 'POST',
body: JSON.stringify({})
});
},
async getTunnelStatus() {
return request('/distributed/tunnel/status');
}
};
}
}
+18 -1
View File
@@ -91,6 +91,23 @@ export const PULSE_ANIMATION_CSS = `
opacity: 1;
padding: 12px 0;
}
/* Cloudflare tunnel spinner */
@keyframes tunnel-spin {
from { transform: rotate(0deg); }
to { transform: rotate(360deg); }
}
.tunnel-spinner {
width: 14px;
height: 14px;
border: 2px solid rgba(255, 255, 255, 0.35);
border-top-color: #fff;
border-radius: 50%;
display: inline-block;
animation: tunnel-spin 0.9s linear infinite;
margin-right: 8px;
vertical-align: middle;
}
`;
export const UI_STYLES = {
@@ -150,4 +167,4 @@ export const TIMEOUTS = {
// Background tasks
LOG_REFRESH: 2000, // log auto-refresh interval
IMAGE_CACHE_CLEAR: 30000 // delay before clearing image cache
};
};
+332 -19
View File
@@ -1,5 +1,5 @@
import { api } from "../../scripts/api.js";
import { findNodesByClass, findImageReferences, hasUpstreamNode, pruneWorkflowForWorker, getCachedWorkerSystemInfo } from './workerUtils.js';
import { findNodesByClass, findImageReferences, hasUpstreamNode, pruneWorkflowForWorker, getCachedWorkerSystemInfo, findCollectorDownstreamNodes } from './workerUtils.js';
import { TIMEOUTS } from './constants.js';
/**
@@ -60,7 +60,8 @@ export function setupInterceptor(extension) {
if (extension.isEnabled) {
const hasCollector = findNodesByClass(prompt.output, "DistributedCollector").length > 0;
const hasDistUpscale = findNodesByClass(prompt.output, "UltimateSDUpscaleDistributed").length > 0;
const hasDistQueue = findNodesByClass(prompt.output, "DistributedQueue").length > 0;
if (hasCollector || hasDistUpscale) {
const result = await executeParallelDistributed(extension, prompt);
// Immediate status check for instant feedback
@@ -69,6 +70,14 @@ export function setupInterceptor(extension) {
setTimeout(() => extension.checkAllWorkerStatuses(), TIMEOUTS.POST_ACTION_DELAY);
return result;
}
// DistributedQueue: route entire workflow to least-busy worker
if (hasDistQueue) {
const result = await executeQueueDistributed(extension, prompt);
extension.checkAllWorkerStatuses();
setTimeout(() => extension.checkAllWorkerStatuses(), TIMEOUTS.POST_ACTION_DELAY);
return result;
}
}
return extension.originalQueuePrompt(number, prompt);
};
@@ -94,11 +103,11 @@ export async function executeParallelDistributed(extension, promptWrapper) {
}
extension.log(`Pre-flight check: ${activeWorkers.length} of ${enabledWorkers.length} workers are active`, "debug");
// Check if master host might be unreachable by workers (cloudflare tunnel down)
const masterHost = extension.config?.master?.host || '';
const isCloudflareHost = /\.(trycloudflare\.com|cloudflare\.dev)$/i.test(masterHost);
if (isCloudflareHost && activeWorkers.length > 0) {
// Try to verify if the cloudflare tunnel is actually up
try {
@@ -109,28 +118,43 @@ export async function executeParallelDistributed(extension, promptWrapper) {
cache: 'no-cache',
signal: AbortSignal.timeout(3000) // 3 second timeout
});
if (!response.ok) {
throw new Error('Master not reachable');
}
} catch (error) {
// Cloudflare tunnel appears to be down
extension.log(`Master host ${masterHost} is not reachable - cloudflare tunnel may be down`, "error");
if (extension.ui?.showCloudflareWarning) {
extension.ui.showCloudflareWarning(extension, masterHost);
}
// Stop execution - workers won't be able to send results back
extension.log("Blocking execution - workers cannot reach master at cloudflare domain", "error");
return null; // This will prevent the workflow from running
}
}
// Find all distributed nodes in the workflow
const collectorNodes = findNodesByClass(promptWrapper.output, "DistributedCollector");
const upscaleNodes = findNodesByClass(promptWrapper.output, "UltimateSDUpscaleDistributed");
const allDistributedNodes = [...collectorNodes, ...upscaleNodes];
const wantsDelegate = Boolean(extension.config?.settings?.master_delegate_only);
let masterDelegateActive = false;
if (wantsDelegate) {
if (activeWorkers.length === 0) {
extension.log("Master delegate-only mode disabled: no active workers detected. Falling back to master execution.", "debug");
} else if (upscaleNodes.length > 0) {
extension.log("Master delegate-only mode is not yet supported for UltimateSDUpscaleDistributed nodes. Falling back to normal execution.", "warn");
} else if (collectorNodes.length === 0) {
extension.log("Master delegate-only mode requested but no DistributedCollector nodes found. Falling back to normal execution.", "debug");
} else {
masterDelegateActive = true;
extension.log("Master delegate-only mode enabled: master will orchestrate without running upstream nodes.", "debug");
}
}
// Map original node IDs to truly unique job IDs for this specific run
const job_id_map = new Map(allDistributedNodes.map(node => [node.id, `${executionPrefix}_${node.id}`]));
@@ -147,7 +171,8 @@ export async function executeParallelDistributed(extension, promptWrapper) {
const options = {
enabled_worker_ids: activeWorkers.map(w => w.id),
workflow: promptWrapper.workflow,
job_id_map: job_id_map // Pass the map of unique IDs
job_id_map: job_id_map, // Pass the map of unique IDs
delegate_master: masterDelegateActive && participantId === 'master'
};
const jobApiPrompt = await prepareApiPromptForParticipant(
@@ -184,14 +209,147 @@ export async function executeParallelDistributed(extension, promptWrapper) {
}
}
/**
* Execute a workflow via DistributedQueue - routes to the least-busy worker.
* The workflow runs ONLY on the selected worker, not on the master.
*/
export async function executeQueueDistributed(extension, promptWrapper) {
try {
const enabledWorkers = extension.enabledWorkers;
if (enabledWorkers.length === 0) {
extension.log("DistributedQueue: No enabled workers. Falling back to local execution.", "warn");
return extension.originalQueuePrompt(0, promptWrapper);
}
// Fetch queue status from all enabled workers
const workerStatuses = await fetchWorkerQueueStatuses(extension, enabledWorkers);
if (workerStatuses.length === 0) {
extension.log("DistributedQueue: No reachable workers. Falling back to local execution.", "warn");
return extension.originalQueuePrompt(0, promptWrapper);
}
// Select the least-busy worker
const selectedWorker = selectLeastBusyWorker(extension, workerStatuses);
const worker = selectedWorker.worker;
const queueRemaining = selectedWorker.queueRemaining;
extension.log(`DistributedQueue: Routing to ${worker.name} (queue_remaining=${queueRemaining})`, "info");
// Prepare the prompt with skip_dispatch=True for DistributedQueue nodes
const promptToSend = JSON.parse(JSON.stringify(promptWrapper.output));
markSkipDispatch(promptToSend);
// Dispatch to the selected worker
const workerUrl = extension.getWorkerUrl(worker);
try {
const response = await fetch(`${workerUrl}/prompt`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
mode: 'cors',
body: JSON.stringify({
prompt: promptToSend,
extra_data: { extra_pnginfo: { workflow: promptWrapper.workflow } },
client_id: api.clientId
})
});
if (!response.ok) {
throw new Error(`Worker returned ${response.status}`);
}
const result = await response.json();
extension.log(`DistributedQueue: Dispatched to ${worker.name}, prompt_id=${result.prompt_id}`, "info");
return { prompt_id: result.prompt_id, worker_id: worker.id };
} catch (error) {
extension.log(`DistributedQueue: Failed to dispatch to ${worker.name}: ${error.message}`, "error");
// Fall back to local execution
return extension.originalQueuePrompt(0, promptWrapper);
}
} catch (error) {
extension.log("DistributedQueue execution failed: " + error.message, "error");
throw error;
}
}
async function fetchWorkerQueueStatuses(extension, workers) {
const statuses = [];
const checkPromises = workers.map(async (worker) => {
const url = extension.getWorkerUrl(worker, '/prompt');
try {
const response = await fetch(url, {
method: 'GET',
mode: 'cors',
cache: 'no-store',
signal: AbortSignal.timeout(TIMEOUTS.STATUS_CHECK)
});
if (response.ok) {
const data = await response.json();
const queueRemaining = data.exec_info?.queue_remaining || 0;
return { worker, queueRemaining, online: true };
}
} catch (error) {
extension.log(`DistributedQueue: Worker ${worker.name} unreachable: ${error.message}`, "debug");
}
return null;
});
const results = await Promise.all(checkPromises);
return results.filter(r => r !== null);
}
function selectLeastBusyWorker(extension, statuses) {
// Find idle workers (queue_remaining == 0)
const idleWorkers = statuses.filter(s => s.queueRemaining === 0);
if (idleWorkers.length > 0) {
// Round-robin among idle workers
if (!extension._distributedQueueRRIndex) {
extension._distributedQueueRRIndex = 0;
}
const index = extension._distributedQueueRRIndex % idleWorkers.length;
extension._distributedQueueRRIndex++;
return idleWorkers[index];
}
// No idle workers - pick the one with the shortest queue
return statuses.reduce((min, s) => s.queueRemaining < min.queueRemaining ? s : min);
}
function markSkipDispatch(promptObj) {
for (const node of Object.values(promptObj)) {
if (node && typeof node === 'object' && node.class_type === 'DistributedQueue') {
node.inputs = node.inputs || {};
node.inputs.skip_dispatch = true;
}
}
}
export async function prepareApiPromptForParticipant(extension, baseApiPrompt, participantId, options = {}) {
let jobApiPrompt = JSON.parse(JSON.stringify(baseApiPrompt));
const isMaster = participantId === 'master';
const delegateMaster = Boolean(options.delegate_master);
// Find all distributed nodes once (before pruning)
const collectorNodes = findNodesByClass(jobApiPrompt, "DistributedCollector");
let collectorNodes = findNodesByClass(jobApiPrompt, "DistributedCollector");
const upscaleNodes = findNodesByClass(jobApiPrompt, "UltimateSDUpscaleDistributed");
const allDistributedNodes = [...collectorNodes, ...upscaleNodes];
let allDistributedNodes = [...collectorNodes, ...upscaleNodes];
if (isMaster && delegateMaster) {
if (upscaleNodes.length > 0) {
extension.log("Delegate-only master mode does not support UltimateSDUpscaleDistributed nodes yet. Using full prompt.", "warn");
} else if (collectorNodes.length === 0) {
extension.log("Delegate-only master mode requested but no collectors found in master prompt. Using full prompt.", "debug");
} else {
jobApiPrompt = prepareMasterDelegatePrompt(extension, jobApiPrompt, collectorNodes);
collectorNodes = findNodesByClass(jobApiPrompt, "DistributedCollector");
allDistributedNodes = [...collectorNodes, ...upscaleNodes];
}
}
// For workers, handle platform-specific path conversion
if (!isMaster) {
@@ -282,6 +440,9 @@ export async function prepareApiPromptForParticipant(extension, baseApiPrompt, p
inputs.is_worker = !isMaster;
if (isMaster) {
inputs.enabled_worker_ids = JSON.stringify(options.enabled_worker_ids || []);
if (delegateMaster) {
inputs.delegate_only = true;
}
} else {
inputs.master_url = extension.getMasterUrl();
// Also make the worker_job_id unique to prevent potential caching issues
@@ -323,6 +484,75 @@ export async function prepareDistributedJob(extension, multi_job_id) {
}
}
function createNumericIdGenerator(promptObj) {
let maxId = 0;
for (const key of Object.keys(promptObj)) {
const numeric = parseInt(key, 10);
if (!Number.isNaN(numeric)) {
maxId = Math.max(maxId, numeric);
}
}
return () => {
maxId += 1;
return String(maxId);
};
}
function prepareMasterDelegatePrompt(extension, apiPrompt, collectorNodes) {
const collectorIds = collectorNodes.map((node) => node.id);
const nodesToKeep = findCollectorDownstreamNodes(apiPrompt, collectorIds);
// Ensure collectors themselves are present
collectorIds.forEach((id) => nodesToKeep.add(id));
const prunedPrompt = {};
nodesToKeep.forEach((nodeId) => {
const nodeData = apiPrompt[nodeId];
if (nodeData) {
prunedPrompt[nodeId] = JSON.parse(JSON.stringify(nodeData));
}
});
const prunedIds = new Set(Object.keys(prunedPrompt));
// Remove dangling references to trimmed nodes
for (const [nodeId, node] of Object.entries(prunedPrompt)) {
if (!node.inputs) continue;
for (const [inputName, inputValue] of Object.entries(node.inputs)) {
if (Array.isArray(inputValue) && inputValue.length === 2) {
const sourceId = String(inputValue[0]);
if (!prunedIds.has(sourceId)) {
delete node.inputs[inputName];
extension.log(`Removed upstream reference '${inputName}' from node ${nodeId} for delegate-only master prompt.`, "debug");
}
}
}
}
const nextId = createNumericIdGenerator(prunedPrompt);
collectorIds.forEach((collectorId) => {
const collectorEntry = prunedPrompt[collectorId];
if (!collectorEntry) return;
const placeholderId = nextId();
prunedPrompt[placeholderId] = {
class_type: "DistributedEmptyImage",
inputs: {
height: 64,
width: 64,
channels: 3
}
};
prunedIds.add(placeholderId);
collectorEntry.inputs = collectorEntry.inputs || {};
collectorEntry.inputs.images = [placeholderId, 0];
});
extension.log(`Prepared delegate-only master prompt with ${Object.keys(prunedPrompt).length} nodes (kept ${nodesToKeep.size}).`, "debug");
return prunedPrompt;
}
export async function executeJobs(extension, jobs) {
let masterPromptId = null;
@@ -395,24 +625,102 @@ async function dispatchToWorker(extension, worker, prompt, workflow, imageRefere
const promptToSend = {
prompt,
extra_data: { extra_pnginfo: { workflow } },
workflow,
client_id: api.clientId
};
extension.log('[Distributed] Prompt data: ' + JSON.stringify(promptToSend), "debug");
try {
await fetch(`${workerUrl}/prompt`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
if (extension.config?.settings?.websocket_orchestration) {
const wsUrl = buildWorkerWebSocketUrl(workerUrl);
await dispatchPromptViaWebSocket(extension, wsUrl, promptToSend);
return;
}
await fetch(`${workerUrl}/prompt`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
mode: 'cors',
body: JSON.stringify(promptToSend)
body: JSON.stringify({
prompt,
extra_data: { extra_pnginfo: { workflow } },
client_id: api.clientId
})
});
} catch (e) {
extension.log(`Failed to connect to worker ${worker.name} at ${workerUrl}: ${e.message}`, "error");
}
}
function buildWorkerWebSocketUrl(workerUrl) {
const url = new URL(workerUrl);
url.protocol = url.protocol === 'https:' ? 'wss:' : 'ws:';
url.pathname = '/distributed/worker_ws';
url.search = '';
return url.toString();
}
function generateRequestId() {
if (crypto.randomUUID) {
return crypto.randomUUID();
}
return 'xxxxxxxx-xxxx-4xxx-yxxx-xxxxxxxxxxxx'.replace(/[xy]/g, (c) => {
const r = Math.random() * 16 | 0;
const v = c === 'x' ? r : (r & 0x3 | 0x8);
return v.toString(16);
});
}
async function dispatchPromptViaWebSocket(extension, wsUrl, payload) {
const requestId = generateRequestId();
const message = {
type: 'dispatch_prompt',
request_id: requestId,
prompt: payload.prompt,
workflow: payload.workflow,
client_id: payload.client_id
};
await new Promise((resolve, reject) => {
const ws = new WebSocket(wsUrl);
const timeoutId = setTimeout(() => {
ws.close();
reject(new Error('Websocket dispatch timed out.'));
}, 60000);
ws.addEventListener('open', () => {
ws.send(JSON.stringify(message));
});
ws.addEventListener('message', (event) => {
try {
const data = JSON.parse(event.data || '{}');
if (data.type !== 'dispatch_ack' || data.request_id !== requestId) {
return;
}
clearTimeout(timeoutId);
ws.close();
if (data.ok) {
resolve();
} else {
reject(new Error(data.error || 'Worker rejected websocket dispatch.'));
}
} catch (error) {
clearTimeout(timeoutId);
ws.close();
reject(error);
}
});
ws.addEventListener('error', () => {
clearTimeout(timeoutId);
ws.close();
reject(new Error('Websocket connection failed.'));
});
});
extension.log(`[Distributed] Websocket dispatch succeeded for ${wsUrl}`, "debug");
}
export async function loadImagesForWorker(extension, imageReferences) {
const images = [];
@@ -565,9 +873,10 @@ export async function performPreflightCheck(extension, workers) {
const response = await fetch(url, {
method: 'GET',
mode: 'cors',
cache: 'no-store',
signal: AbortSignal.timeout(TIMEOUTS.STATUS_CHECK)
});
if (response.ok) {
extension.log(`Worker ${worker.name} is active`, "debug");
return { worker, active: true };
@@ -576,6 +885,10 @@ export async function performPreflightCheck(extension, workers) {
return { worker, active: false };
}
} catch (error) {
if (error?.name === 'AbortError') {
extension.log(`Worker ${worker.name} pre-flight check timed out; assuming active`, "debug");
return { worker, active: true, uncertain: true };
}
extension.log(`Worker ${worker.name} is offline or unreachable: ${error.message}`, "debug");
return { worker, active: false };
}
+74 -72
View File
@@ -1,86 +1,88 @@
import { app } from "/scripts/app.js";
// Configuration for each batch divider node type
const BATCH_DIVIDER_NODES = {
"ImageBatchDivider": { outputPrefix: "batch_", outputType: "IMAGE" },
"AudioBatchDivider": { outputPrefix: "audio_", outputType: "AUDIO" }
};
app.registerExtension({
name: "Distributed.ImageBatchDivider",
name: "Distributed.BatchDividers",
async nodeCreated(node) {
if (node.comfyClass === "ImageBatchDivider") {
try {
const updateOutputs = () => {
if (!node.widgets) return;
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (!divideByWidget) return;
const divideBy = parseInt(divideByWidget.value, 10) || 1;
const totalOutputs = divideBy; // Direct divide by value
// Ensure outputs array exists
if (!node.outputs) node.outputs = [];
// Remove excess outputs
while (node.outputs.length > totalOutputs) {
node.removeOutput(node.outputs.length - 1);
}
// Add missing outputs
while (node.outputs.length < totalOutputs) {
const outputIndex = node.outputs.length + 1;
node.addOutput(`batch_${outputIndex}`, "IMAGE");
}
if (node.setDirty) node.setDirty(true); // Refresh canvas
};
// Initial update with delay to allow workflow loading
setTimeout(updateOutputs, 200);
// Find the widget and set up responsive handlers
const config = BATCH_DIVIDER_NODES[node.comfyClass];
if (!config) return;
try {
const updateOutputs = () => {
if (!node.widgets) return;
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (divideByWidget) {
// Override callback for immediate trigger on value set
const originalCallback = divideByWidget.callback;
divideByWidget.callback = (value) => {
updateOutputs();
if (originalCallback) originalCallback.call(divideByWidget, value); // Preserve 'this' context
};
// Add event listener for real-time input changes (e.g., typing/dragging)
if (divideByWidget.inputEl) {
divideByWidget.inputEl.addEventListener('input', updateOutputs);
}
// Lightweight MutationObserver as fallback (observe attributes on widget element if available)
const observer = new MutationObserver(updateOutputs);
if (divideByWidget.element) {
observer.observe(divideByWidget.element, { attributes: true, childList: true, subtree: true });
}
// Store cleanup function
node._batchDividerCleanup = () => {
observer.disconnect();
if (divideByWidget.inputEl) {
divideByWidget.inputEl.removeEventListener('input', updateOutputs);
}
divideByWidget.callback = originalCallback; // Restore original
};
if (!divideByWidget) return;
const divideBy = parseInt(divideByWidget.value, 10) || 1;
const totalOutputs = divideBy;
// Ensure outputs array exists
if (!node.outputs) node.outputs = [];
// Remove excess outputs
while (node.outputs.length > totalOutputs) {
node.removeOutput(node.outputs.length - 1);
}
// Add post-configure hook for reliable workflow loading
const originalConfigure = node.configure;
node.configure = function(data) {
const result = originalConfigure ? originalConfigure.call(this, data) : undefined;
updateOutputs(); // Re-run after config load
return result;
// Add missing outputs
while (node.outputs.length < totalOutputs) {
const outputIndex = node.outputs.length + 1;
node.addOutput(`${config.outputPrefix}${outputIndex}`, config.outputType);
}
if (node.setDirty) node.setDirty(true);
};
// Initial update with delay to allow workflow loading
setTimeout(updateOutputs, 200);
// Find the widget and set up responsive handlers
const divideByWidget = node.widgets.find(w => w.name === "divide_by");
if (divideByWidget) {
const originalCallback = divideByWidget.callback;
divideByWidget.callback = (value) => {
updateOutputs();
if (originalCallback) originalCallback.call(divideByWidget, value);
};
if (divideByWidget.inputEl) {
divideByWidget.inputEl.addEventListener('input', updateOutputs);
}
const observer = new MutationObserver(updateOutputs);
if (divideByWidget.element) {
observer.observe(divideByWidget.element, { attributes: true, childList: true, subtree: true });
}
node._batchDividerCleanup = () => {
observer.disconnect();
if (divideByWidget.inputEl) {
divideByWidget.inputEl.removeEventListener('input', updateOutputs);
}
divideByWidget.callback = originalCallback;
};
} catch (error) {
console.error("Error in ImageBatchDivider extension:", error);
}
const originalConfigure = node.configure;
node.configure = function(data) {
const result = originalConfigure ? originalConfigure.call(this, data) : undefined;
updateOutputs();
return result;
};
} catch (error) {
console.error(`Error in ${node.comfyClass} extension:`, error);
}
},
nodeBeforeRemove(node) {
if (node.comfyClass === "ImageBatchDivider" && node._batchDividerCleanup) {
if (BATCH_DIVIDER_NODES[node.comfyClass] && node._batchDividerCleanup) {
node._batchDividerCleanup();
}
}
});
});
+223 -16
View File
@@ -17,6 +17,8 @@ class DistributedExtension {
this.logAutoRefreshInterval = null;
this.masterSettingsExpanded = false;
this.app = app; // Store app reference for toast notifications
this.tunnelStatus = { status: "unknown" };
this.tunnelElements = {};
// Initialize centralized state
this.state = createStateManager();
@@ -90,6 +92,34 @@ class DistributedExtension {
return this.enabledWorkers.length > 0;
}
isMasterParticipationEnabled() {
return !Boolean(this.config?.settings?.master_delegate_only);
}
isMasterFallbackActive() {
return Boolean(this.config?.settings?.master_delegate_only) && this.enabledWorkers.length === 0;
}
isMasterParticipating() {
return this.isMasterParticipationEnabled() || this.isMasterFallbackActive();
}
async updateMasterParticipation(enabled) {
if (!this.config?.settings) {
this.config.settings = {};
}
const delegateOnly = !enabled;
if (this.config.settings.master_delegate_only === delegateOnly) {
return;
}
await this._updateSetting('master_delegate_only', delegateOnly);
if (this.panelElement) {
renderSidebarContent(this, this.panelElement);
}
}
async loadConfig() {
try {
this.config = await this.api.getConfig();
@@ -118,6 +148,162 @@ class DistributedExtension {
}
}
_applyMasterHost(host) {
if (!host || !this.config) return;
if (!this.config.master) this.config.master = {};
this.config.master.host = host;
const hostInput = document.getElementById('master-host');
if (hostInput) {
hostInput.value = host;
}
}
_parseHostInput(value) {
if (!value) {
return { host: "", port: null };
}
let cleaned = value.trim().replace(/^https?:\/\//i, "");
cleaned = cleaned.split("/")[0];
try {
const url = new URL(`http://${cleaned}`);
const port = url.port ? parseInt(url.port, 10) : null;
return {
host: url.hostname || cleaned,
port: Number.isFinite(port) ? port : null,
};
} catch (error) {
return { host: cleaned, port: null };
}
}
updateTunnelUIElements() {
const elements = this.tunnelElements || {};
const status = (this.tunnelStatus?.status || "stopped").toLowerCase();
const enableColor = "#665533"; // requested yellow/brown tone
const disableColor = "#7c4a4a"; // match worker delete button
const colors = {
running: disableColor,
starting: enableColor,
stopped: enableColor,
error: disableColor,
unknown: enableColor,
stopping: disableColor
};
if (elements.button) {
elements.button.disabled = status === "starting" || status === "stopping";
if (status === "starting") {
elements.button.innerHTML = `<span class="tunnel-spinner"></span> Starting...`;
elements.button.style.backgroundColor = enableColor;
} else if (status === "running") {
elements.button.textContent = "Disable Cloudflare Tunnel";
elements.button.style.backgroundColor = disableColor;
} else if (status === "error") {
elements.button.textContent = "Retry Cloudflare Tunnel";
elements.button.style.backgroundColor = disableColor;
} else {
elements.button.textContent = "Enable Cloudflare Tunnel";
elements.button.style.backgroundColor = enableColor;
}
}
if (elements.status) {
elements.status.textContent = status.toUpperCase();
elements.status.style.backgroundColor = colors[status] || colors.stopped;
}
if (elements.url) {
const url = this.tunnelStatus?.public_url;
if (url) {
elements.url.innerHTML = `<a href="${url}" target="_blank" style="color: #eee; text-decoration: none;">${url}</a>`;
} else {
elements.url.textContent = status === "starting" ? "Requesting public URL..." : "No tunnel active";
}
}
if (elements.copyBtn) {
const hasUrl = Boolean(this.tunnelStatus?.public_url);
elements.copyBtn.disabled = !hasUrl;
elements.copyBtn.style.opacity = hasUrl ? "1" : "0.5";
}
}
async refreshTunnelStatus() {
try {
const data = await this.api.getTunnelStatus();
this.tunnelStatus = data.tunnel || { status: "stopped" };
if (data.master_host !== undefined) {
this._applyMasterHost(data.master_host);
}
return this.tunnelStatus;
} catch (error) {
this.tunnelStatus = { status: "error", last_error: error.message };
this.log("Failed to fetch tunnel status: " + error.message, "error");
return this.tunnelStatus;
} finally {
this.updateTunnelUIElements();
}
}
async handleTunnelToggle(button) {
const currentStatus = (this.tunnelStatus?.status || "stopped").toLowerCase();
if (currentStatus === "starting" || currentStatus === "stopping") {
return;
}
const setStatus = (status) => {
this.tunnelStatus = { ...(this.tunnelStatus || {}), status };
this.updateTunnelUIElements();
};
if (currentStatus === "running") {
setStatus("stopping");
try {
if (button) {
button.innerHTML = `<span class="tunnel-spinner"></span> Stopping...`;
button.disabled = true;
}
const data = await this.api.stopTunnel();
this.tunnelStatus = data.tunnel || { status: "stopped" };
if (data.master_host !== undefined) {
this._applyMasterHost(data.master_host);
}
this.updateTunnelUIElements();
this.ui.showToast(this.app, "info", "Cloudflare Tunnel Disabled", "Master address restored", 4000);
} catch (error) {
this.tunnelStatus = { status: "error", last_error: error.message };
this.updateTunnelUIElements();
this.ui.showToast(this.app, "error", "Failed to stop tunnel", error.message, 5000);
} finally {
if (button) button.disabled = false;
}
return;
}
// Start tunnel
setStatus("starting");
if (button) {
button.innerHTML = `<span class="tunnel-spinner"></span> Starting...`;
button.disabled = true;
}
try {
const data = await this.api.startTunnel();
this.tunnelStatus = data.tunnel || { status: "running" };
if (data.master_host !== undefined) {
this._applyMasterHost(data.master_host);
}
this.updateTunnelUIElements();
const url = data.tunnel?.public_url || data.master_host;
this.ui.showToast(this.app, "success", "Cloudflare Tunnel Ready", url || "Public URL created", 4500);
} catch (error) {
this.tunnelStatus = { status: "error", last_error: error.message };
this.updateTunnelUIElements();
this.ui.showToast(this.app, "error", "Failed to start tunnel", error.message, 5000);
} finally {
if (button) button.disabled = false;
}
}
async updateWorkerEnabled(workerId, enabled) {
const worker = this.config.workers.find(w => w.id === workerId);
if (worker) {
@@ -143,6 +329,10 @@ class DistributedExtension {
} catch (error) {
this.log("Error updating worker: " + error.message, "error");
}
if (this.panelElement) {
await renderSidebarContent(this, this.panelElement);
}
}
async _updateSetting(key, value) {
@@ -300,11 +490,19 @@ class DistributedExtension {
// Update master status dot
const statusDot = document.getElementById('master-status');
if (statusDot) {
if (isProcessing) {
statusDot.style.backgroundColor = "#f0ad4e";
if (!this.isMasterParticipating()) {
if (isProcessing) {
statusDot.style.backgroundColor = STATUS_COLORS.PROCESSING_YELLOW;
statusDot.title = `Orchestrating (${queueRemaining} in queue)`;
} else {
statusDot.style.backgroundColor = STATUS_COLORS.DISABLED_GRAY;
statusDot.title = "Master orchestrator only";
}
} else if (isProcessing) {
statusDot.style.backgroundColor = STATUS_COLORS.PROCESSING_YELLOW;
statusDot.title = `Processing (${queueRemaining} in queue)`;
} else {
statusDot.style.backgroundColor = "#4CAF50";
statusDot.style.backgroundColor = STATUS_COLORS.ONLINE_GREEN;
statusDot.title = "Online";
}
}
@@ -313,15 +511,17 @@ class DistributedExtension {
// Master is always online (we're running on it), so keep it green
const statusDot = document.getElementById('master-status');
if (statusDot) {
statusDot.style.backgroundColor = "#4CAF50";
statusDot.title = "Online";
statusDot.style.backgroundColor = this.isMasterParticipating() ? STATUS_COLORS.ONLINE_GREEN : STATUS_COLORS.DISABLED_GRAY;
statusDot.title = this.isMasterParticipating() ? "Online" : "Master orchestrator only";
}
}
}
// Helper to build worker URL
getWorkerUrl(worker, endpoint = '') {
const host = worker.host || window.location.hostname;
const parsed = this._parseHostInput(worker.host || window.location.hostname);
const host = parsed.host || window.location.hostname;
const resolvedPort = parsed.port || worker.port;
// Cloud workers always use HTTPS
const isCloud = worker.type === 'cloud';
@@ -344,13 +544,13 @@ class DistributedExtension {
}
// Determine protocol: HTTPS for cloud, Runpod proxies, or port 443
const useHttps = isCloud || isRunpodProxy || worker.port === 443;
const useHttps = isCloud || isRunpodProxy || resolvedPort === 443;
const protocol = useHttps ? 'https' : 'http';
// Only add port if non-standard
const defaultPort = useHttps ? 443 : 80;
const needsPort = !isRunpodProxy && worker.port !== defaultPort;
const portStr = needsPort ? `:${worker.port}` : '';
const needsPort = !isRunpodProxy && resolvedPort !== defaultPort;
const portStr = needsPort ? `:${resolvedPort}` : '';
return `${protocol}://${finalHost}${portStr}${endpoint}`;
}
@@ -436,21 +636,21 @@ class DistributedExtension {
async launchWorker(workerId) {
const worker = this.config.workers.find(w => w.id === workerId);
const launchBtn = document.querySelector(`#controls-${workerId} button`);
// If worker is disabled, enable it first
if (!worker.enabled) {
await this.updateWorkerEnabled(workerId, true);
// Update the checkbox UI
const checkbox = document.getElementById(`gpu-${workerId}`);
if (checkbox) {
checkbox.checked = true;
}
this.updateSummary();
}
// Re-query button AFTER updateWorkerEnabled (which may re-render sidebar)
const launchBtn = document.querySelector(`#controls-${workerId} button`);
this.ui.updateStatusDot(workerId, "#f0ad4e", "Launching...", true);
this.state.setWorkerLaunching(workerId, true);
@@ -794,7 +994,8 @@ class DistributedExtension {
return true;
}
// Otherwise check by host (backward compatibility)
const host = worker.host || window.location.hostname;
const parsed = this._parseHostInput(worker.host || window.location.hostname);
const host = parsed.host || window.location.hostname;
return host !== "localhost" && host !== "127.0.0.1" && host !== window.location.hostname;
}
@@ -1014,10 +1215,16 @@ class DistributedExtension {
const workerType = document.getElementById(`worker-type-${workerId}`).value;
const isRemote = workerType === 'remote' || workerType === 'cloud';
const isCloud = workerType === 'cloud';
const host = isRemote ? document.getElementById(`host-${workerId}`).value : window.location.hostname;
const port = parseInt(document.getElementById(`port-${workerId}`).value);
const rawHost = isRemote ? document.getElementById(`host-${workerId}`).value : window.location.hostname;
const parsedHost = isRemote ? this._parseHostInput(rawHost) : { host: window.location.hostname, port: null };
const host = isRemote ? parsedHost.host : window.location.hostname;
let port = parseInt(document.getElementById(`port-${workerId}`).value);
const cudaDevice = isRemote ? undefined : parseInt(document.getElementById(`cuda-${workerId}`).value);
const extraArgs = isRemote ? undefined : document.getElementById(`args-${workerId}`).value;
if (isRemote && Number.isFinite(parsedHost.port)) {
port = parsedHost.port;
}
// Validate
if (!name.trim()) {
+2
View File
@@ -43,6 +43,8 @@ export async function renderSidebarContent(extension, el) {
// Preload data outside render
await extension.loadConfig();
await extension.loadManagedWorkers();
await extension.refreshTunnelStatus();
extension.tunnelElements = {};
el.innerHTML = '';
+78 -15
View File
@@ -3,28 +3,32 @@ import { BUTTON_STYLES, UI_STYLES, STATUS_COLORS, UI_COLORS, TIMEOUTS } from './
const cardConfigs = {
master: {
checkbox: {
enabled: false,
checked: true,
disabled: true,
opacity: 0.6,
title: "Master node is always enabled"
enabled: true,
masterToggle: true,
title: "Toggle master participation in workloads"
},
statusDot: {
color: STATUS_COLORS.ONLINE_GREEN,
title: 'Online',
id: 'master-status',
initialColor: (_, extension) => extension.isMasterParticipating() ? STATUS_COLORS.ONLINE_GREEN : STATUS_COLORS.DISABLED_GRAY,
initialTitle: (_, extension) => extension.isMasterParticipating() ? 'Master participating' : 'Master orchestrator only',
dynamic: true
},
infoText: (data, extension) => {
const cudaDevice = extension.config?.master?.cuda_device ?? extension.masterCudaDevice;
const cudaInfo = cudaDevice !== undefined ? `CUDA ${cudaDevice} • ` : '';
const port = window.location.port || (window.location.protocol === 'https:' ? '443' : '80');
return `<strong id="master-name-display">${data?.name || extension.config?.master?.name || "Master"}</strong><br><small style="color: ${UI_COLORS.MUTED_TEXT};"><span id="master-cuda-info">${cudaInfo}Port ${port}</span></small>`;
const participationEnabled = extension.isMasterParticipationEnabled();
const fallbackActive = extension.isMasterFallbackActive();
let delegateBadge = '';
if (!participationEnabled) {
delegateBadge = fallbackActive
? `<br><small style="color: #6bd06b;">Fallback active • Master executing</small>`
: `<br><small style="color: ${UI_COLORS.SECONDARY_TEXT};">Orchestrator-only mode</small>`;
}
return `<strong id="master-name-display">${data?.name || extension.config?.master?.name || "Master"}</strong><br><small style="color: ${UI_COLORS.MUTED_TEXT};"><span id="master-cuda-info">${cudaInfo}Port ${port}</span></small>${delegateBadge}`;
},
controls: {
type: 'info',
text: 'Master',
style: "background-color: #333; color: #999;"
type: 'master'
},
settings: {
formType: 'master',
@@ -868,7 +872,36 @@ export class DistributedUI {
if (config?.opacity) checkbox.style.opacity = config.opacity;
if (config?.title) column.title = config.title;
if (config?.enabled && !config?.disabled && data?.id) {
const isMasterToggle = config?.masterToggle && typeof extension.isMasterParticipating === 'function';
if (isMasterToggle) {
const participationEnabled = extension.isMasterParticipationEnabled();
const fallbackActive = extension.isMasterFallbackActive();
const buildTitle = (enabled, fallback) => {
if (enabled) {
return "Master participating • Click to switch to orchestrator-only";
}
if (fallback) {
return "No workers selected • Master fallback execution active";
}
return "Master orchestrator-only • Click to re-enable participation";
};
checkbox.checked = participationEnabled;
checkbox.style.pointerEvents = "none";
column.style.cursor = "pointer";
column.title = buildTitle(participationEnabled, fallbackActive);
column.onclick = async (event) => {
if (event) {
event.stopPropagation();
event.preventDefault();
}
const nextState = !extension.isMasterParticipationEnabled();
const nextFallback = !nextState && extension.enabledWorkers.length === 0;
checkbox.checked = nextState;
column.title = buildTitle(nextState, nextFallback);
await extension.updateMasterParticipation(nextState);
};
} else if (config?.enabled && !config?.disabled && data?.id) {
checkbox.style.pointerEvents = "none";
column.style.cursor = "pointer";
column.onclick = async () => {
@@ -890,10 +923,10 @@ export class DistributedUI {
let id = config.id;
if (typeof config.initialColor === 'function') {
color = config.initialColor(data);
color = config.initialColor(data, extension);
}
if (typeof config.initialTitle === 'function') {
title = config.initialTitle(data);
title = config.initialTitle(data, extension);
}
if (typeof config.id === 'function') {
id = config.id(data);
@@ -940,7 +973,29 @@ export class DistributedUI {
const controlsWrapper = document.createElement("div");
controlsWrapper.style.cssText = this.styles.controlsWrapper;
if (config.dynamic && data) {
if (config.type === 'master') {
const participationEnabled = extension.isMasterParticipationEnabled();
const fallbackActive = extension.isMasterFallbackActive();
let message;
const badge = document.createElement("div");
badge.style.cssText = this.styles.infoBox;
if (fallbackActive) {
message = "No workers selected. Master fallback execution active.";
badge.textContent = message;
badge.style.backgroundColor = "#243024";
badge.style.color = "#6bd06b";
badge.style.border = "1px solid #335533";
} else if (!participationEnabled) {
message = "Master disabled: running as orchestrator only.";
badge.textContent = message;
badge.style.backgroundColor = "#3a3a3a";
badge.style.color = "#ffcc66";
} else {
message = "Master participating in workflows.";
badge.textContent = message;
}
controlsWrapper.appendChild(badge);
} else if (config.dynamic && data) {
if (isRemote) {
const isCloud = data.type === 'cloud';
const workerTypeText = isCloud ? "Cloud worker" : "Remote worker";
@@ -1029,6 +1084,14 @@ export class DistributedUI {
const hostResult = this.createFormGroup("Host:", extension.config?.master?.host || "", "master-host", "text", "Auto-detect if empty");
settingsForm.appendChild(hostResult.group);
// Cloudflare tunnel toggle (simple button inside master settings)
const tunnelBtn = this.createButton("Enable Cloudflare Tunnel", (e) => extension.handleTunnelToggle(e.target), "background-color: #665533;");
tunnelBtn.id = "cloudflare-tunnel-button";
tunnelBtn.style.cssText = BUTTON_STYLES.base + " background-color: #665533; margin: 4px 0 -5px 0;";
settingsForm.appendChild(tunnelBtn);
extension.tunnelElements = { button: tunnelBtn };
extension.updateTunnelUIElements();
const saveBtn = this.createButton("Save", async () => {
const nameInput = document.getElementById('master-name');
+41
View File
@@ -155,6 +155,47 @@ export function findCollectorUpstreamNodes(apiPrompt, collectorIds) {
return connected;
}
/**
* Find all downstream nodes from the provided collector IDs.
* Used to keep only post-collection nodes when master runs in orchestrator mode.
* @param {Object} apiPrompt
* @param {Array<string>} collectorIds
* @returns {Set<string>}
*/
export function findCollectorDownstreamNodes(apiPrompt, collectorIds) {
const adjacency = new Map();
for (const [nodeId, node] of Object.entries(apiPrompt)) {
if (!node.inputs) continue;
for (const inputValue of Object.values(node.inputs)) {
if (Array.isArray(inputValue) && inputValue.length === 2) {
const sourceId = String(inputValue[0]);
if (!adjacency.has(sourceId)) {
adjacency.set(sourceId, new Set());
}
adjacency.get(sourceId).add(String(nodeId));
}
}
}
const connected = new Set(collectorIds);
const queue = [...collectorIds];
while (queue.length > 0) {
const current = queue.shift();
const dependents = adjacency.get(current);
if (!dependents) continue;
dependents.forEach((dependentId) => {
if (!connected.has(dependentId)) {
connected.add(dependentId);
queue.push(dependentId);
}
});
}
return connected;
}
/**
* Prune workflow to only include nodes connected to distributed nodes
* @param {Object} apiPrompt - The full workflow API prompt
+1
View File
@@ -39,6 +39,7 @@ except ImportError:
def monitor_and_run(master_pid, command):
"""Run command and monitor master process."""
# Start the actual worker process
print(f"[Distributed] Launching worker command: {' '.join(command)}")
worker_process = subprocess.Popen(command)
print(f"[Distributed] Started worker PID: {worker_process.pid}")