Compare commits
22
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1d4dcb985e | ||
|
|
b6d4756392 | ||
|
|
044c54b780 | ||
|
|
b768e7fc8a | ||
|
|
9805049268 | ||
|
|
399ed7c7d1 | ||
|
|
64ba4b75db | ||
|
|
d353d118d7 | ||
|
|
5dd10cef39 | ||
|
|
8eba225f5e | ||
|
|
9bc2cfc43e | ||
|
|
370b48438b | ||
|
|
71c7a99e1e | ||
|
|
def7c0df1d | ||
|
|
5f5320b34e | ||
|
|
5a33da828c | ||
|
|
ef8e0542c8 | ||
|
|
89b7812fb4 | ||
|
|
f13b953a82 | ||
|
|
1e3e094cdb | ||
|
|
7bcf9290ba | ||
|
|
6aec567071 |
@@ -0,0 +1,4 @@
|
||||
bin/
|
||||
logs/
|
||||
gpu_config.json
|
||||
__pycache__/
|
||||
@@ -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
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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 = []
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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()) {
|
||||
|
||||
@@ -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 = '';
|
||||
|
||||
|
||||
@@ -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');
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user