Compare commits

...
23 Commits
Author SHA1 Message Date
Robert Wojciechowski 9739f432a5 Refactor: reduce duplication and simplify arg passing across USDU modules
- Pass UpscaleCoreArgs through static mode instead of 11 individual args
- Extract _post_with_retry helper from duplicate retry loops in worker_comms
- Remove 4 passthrough wrapper pairs in job_store
- Extract _process_image_tiles from duplicate master/worker tile loops
- Merge 4 override functions into 2 in prompt_transform (seed+value, collector+upscale)
- Split _async_collect_results into focused static/dynamic methods

Net reduction of ~256 lines across 10 files.
2026-03-08 22:05:29 +00:00
Robert Wojciechowski 5e724665a1 Add USDU delegate-only master mode support
When master is in delegate-only mode, USDU nodes are swapped to a
lightweight USDUDelegateCollector that skips model loading and only
collects results from workers. Workers initialize the dynamic job
queue on the master via a new /distributed/init_dynamic_job endpoint
since the delegate master doesn't know the batch size.
2026-03-08 03:34:17 +00:00
Robert Wojciechowski d387769453 Skip heartbeat timeout check for workers with no incomplete tasks
Prevents false "heartbeat timed out" warnings when workers finish
their work before the master completes local tile processing.
2026-03-08 03:11:50 +00:00
Robert Wojciechowski ab4e1536a5 worker launch bug fix 2026-03-07 05:45:11 +00:00
Robert Wojciechowski ce3b2c63f3 Fix DistributedSeed and DistributedValue not receiving enabled_worker_ids
The orchestration layer never injected enabled_worker_ids into seed and
value nodes, so workers could not resolve their index and fell back to
returning the original seed unchanged.
2026-03-07 05:33:49 +00:00
Robert Wojciechowski dc0945e0f2 Add .claude/ to gitignore and remove from tracking 2026-03-07 05:23:41 +00:00
Robert Wojciechowski 3058b11ff2 Remove list splitter/collector feature and refactor distributed modules
- Remove DistributedListSplitter and DistributedListCollector nodes
- Remove /distributed/init_list_queue and /distributed/request_list_item endpoints
- Clean up all references from prompt_transform, queue_orchestration, constants, and tests
- Refactor and reorganize distributed modules for quality improvements
2026-03-07 05:22:09 +00:00
Robert Wojciechowski 2a25999f4f ignore desloppify artifacts 2026-03-03 20:52:15 +00:00
Robert Wojciechowski 684b276619 refactor distributed modules and reorganize tests for quality improvements 2026-03-03 20:50:40 +00:00
Robert Wojciechowski 785fa05856 harden branch collector local-slot mapping 2026-03-02 08:16:53 +00:00
Robert Wojciechowski 5e8c35b913 guard worker prompts from branch-collector downstream outputs 2026-03-02 08:01:16 +00:00
Robert Wojciechowski 8241bb9585 prune worker nodes downstream of branch collector 2026-03-02 07:54:44 +00:00
Robert Wojciechowski 0e7690b3a8 replace DistributedJoin with branch collector node 2026-03-02 07:35:11 +00:00
Robert Wojciechowski 8d67a5f5e7 Handle missing join worker results with safe fallback outputs 2026-03-02 07:17:31 +00:00
Robert Wojciechowski 2f8b43b130 Add create_time metadata for distributed queued prompts 2026-03-02 07:10:32 +00:00
Robert Wojciechowski 2d5a56a95e Prevent prompt_no_outputs on pruned worker branch prompts 2026-03-02 07:06:17 +00:00
Robert Wojciechowski dd3fbfadbe Fix branch join convergence and remote log visibility gating 2026-03-02 06:57:41 +00:00
Robert Wojciechowski cfed768528 Apply dynamic num_branches sockets to DistributedJoin UI 2026-03-01 23:31:56 +00:00
Robert Wojciechowski 34123a0d30 Add DistributedJoin and dynamic list-queue splitter mode 2026-03-01 23:29:50 +00:00
Robert Wojciechowski e37a001d59 Rank branch participants by worker queue depth 2026-03-01 23:22:55 +00:00
Robert Wojciechowski eb2dc77636 Make DistributedBranch outputs follow num_branches in UI 2026-03-01 22:59:54 +00:00
Robert Wojciechowski c0c24431e5 Fix DistributedBranch to expose all branch outputs 2026-03-01 22:56:05 +00:00
Robert Wojciechowski 910c260857 Add distributed list splitting/collection and branch routing 2026-03-01 22:48:08 +00:00
131 changed files with 10021 additions and 2440 deletions
+5
View File
@@ -7,3 +7,8 @@ __pycache__/
node_modules/
npm-debug.log*
AGENTS.md
.claude/
# Desloppify artifacts
.desloppify/
scorecard.png
+3 -3
View File
@@ -157,9 +157,9 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
Control your distributed cluster programmatically without opening the browser.
* **Endpoint:** `POST /distributed/queue`
* **Functionality:** Accepts a standard ComfyUI workflow JSON, automatically distributes it to available workers, and returns the execution ID.
* **Documentation:** [See API Examples & Scripts](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/comfyui-distributed-api.md)
* **Endpoint:** `POST /distributed/queue`
* **Functionality:** Accepts a distributed queue payload (`prompt`, `client_id`, `enabled_worker_ids`, optional `workflow`/`delegate_master`/`trace_execution_id`), dispatches to healthy workers, and returns `{prompt_id, worker_count}`.
* **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).
+6 -24
View File
@@ -1,29 +1,11 @@
# Import everything needed from the main module
from .distributed import (
NODE_CLASS_MAPPINGS as DISTRIBUTED_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS as DISTRIBUTED_DISPLAY_NAME_MAPPINGS
)
"""ComfyUI-Distributed package entrypoint."""
from __future__ import annotations
# Import utilities
from .utils.config import ensure_config_exists, CONFIG_FILE
from .utils.logging import debug_log
# Import distributed upscale nodes
from .nodes.distributed_upscale import (
NODE_CLASS_MAPPINGS as UPSCALE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS as UPSCALE_DISPLAY_NAME_MAPPINGS
)
from .bootstrap.entrypoint import build_node_mappings, initialize_runtime
WEB_DIRECTORY = "./web"
ensure_config_exists()
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
initialize_runtime()
# Merge node mappings
NODE_CLASS_MAPPINGS = {**DISTRIBUTED_CLASS_MAPPINGS, **UPSCALE_CLASS_MAPPINGS}
NODE_DISPLAY_NAME_MAPPINGS = {**DISTRIBUTED_DISPLAY_NAME_MAPPINGS, **UPSCALE_DISPLAY_NAME_MAPPINGS}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
debug_log("Loaded Distributed nodes.")
debug_log(f"Config file: {CONFIG_FILE}")
debug_log(f"Available nodes: {list(NODE_CLASS_MAPPINGS.keys())}")
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
+21 -5
View File
@@ -1,5 +1,21 @@
from . import config_routes # noqa: F401
from . import tunnel_routes # noqa: F401
from . import worker_routes # noqa: F401
from . import job_routes # noqa: F401
from . import usdu_routes # noqa: F401
"""Distributed API package with explicit route bootstrap."""
from __future__ import annotations
from importlib import import_module
_ROUTE_MODULES = (
"config_routes",
"tunnel_routes",
"worker_routes",
"job_routes",
"usdu_routes",
)
def bootstrap_routes() -> None:
"""Import route modules once so aiohttp decorators register endpoints."""
for module_name in _ROUTE_MODULES:
import_module(f"{__name__}.{module_name}")
__all__ = ["bootstrap_routes"]
+97 -79
View File
@@ -1,32 +1,15 @@
import json
from contextlib import asynccontextmanager
from __future__ import annotations
from aiohttp import web
import server
try:
from ..utils.config import config_transaction, load_config, save_config
except ImportError:
from ..utils.config import load_config
try:
from ..utils.config import save_config
except ImportError:
def save_config(_config):
return True
@asynccontextmanager
async def config_transaction():
config = load_config()
original_snapshot = json.dumps(config, sort_keys=True)
yield config
if json.dumps(config, sort_keys=True) != original_snapshot:
save_config(config)
from ..utils.logging import debug_log, log
from ..utils.config import config_transaction, load_config
from ..utils.logging import debug_log
from ..utils.network import handle_api_error, normalize_host
from .endpoint_policy import run_authorized_endpoint
def _positive_int(value):
def _positive_int(value: int) -> bool:
return value > 0
@@ -86,68 +69,103 @@ def _apply_field_patch(target: dict, data: dict, field_rules: list) -> None:
target[key] = normalizer(value) if (normalizer and value is not None) else value
def _validate_workers_payload(workers_value: object) -> list[str]:
"""Validate top-level workers payload shape before persisting config."""
if not isinstance(workers_value, list):
return ["workers: expected list"]
errors: list[str] = []
seen_ids: set[str] = set()
for index, worker in enumerate(workers_value):
if not isinstance(worker, dict):
errors.append(f"workers[{index}]: expected object")
continue
worker_id = str(worker.get("id", "")).strip()
if not worker_id:
errors.append(f"workers[{index}].id: expected non-empty string")
elif worker_id in seen_ids:
errors.append(f"workers[{index}].id: duplicate id '{worker_id}'")
else:
seen_ids.add(worker_id)
if "port" in worker and worker.get("port") is not None:
try:
port = int(worker.get("port"))
except (TypeError, ValueError):
errors.append(f"workers[{index}].port: expected integer")
else:
if port <= 0:
errors.append(f"workers[{index}].port: expected positive integer")
return errors
@server.PromptServer.instance.routes.get("/distributed/config")
async def get_config_endpoint(request):
config = load_config()
return web.json_response(config)
async def get_config_endpoint(request: web.Request) -> web.StreamResponse:
async def _operation() -> web.StreamResponse:
return web.json_response(load_config())
return await run_authorized_endpoint(request, _operation)
@server.PromptServer.instance.routes.post("/distributed/config")
async def update_config_endpoint(request):
async def update_config_endpoint(request: web.Request) -> web.StreamResponse:
"""Bulk config update with schema validation."""
try:
data = await request.json()
except Exception as e:
return await handle_api_error(request, f"Invalid JSON payload: {e}", 400)
async def _operation() -> web.StreamResponse:
try:
data = await request.json()
except Exception as exc:
return await handle_api_error(request, f"Invalid JSON payload: {exc}", 400)
if not isinstance(data, dict):
return await handle_api_error(request, "Config payload must be an object", 400)
if not isinstance(data, dict):
return await handle_api_error(request, "Config payload must be an object", 400)
validated_settings = {}
validated_root = {}
errors = []
validated_settings = {}
validated_root = {}
errors = []
for key, value in data.items():
if key not in CONFIG_SCHEMA:
errors.append(f"Unknown field: {key}")
continue
for key, value in data.items():
if key not in CONFIG_SCHEMA:
errors.append(f"Unknown field: {key}")
continue
expected_type, validator = CONFIG_SCHEMA[key]
if not isinstance(value, expected_type):
errors.append(f"{key}: expected {expected_type.__name__}")
continue
expected_type, validator = CONFIG_SCHEMA[key]
if not isinstance(value, expected_type):
errors.append(f"{key}: expected {expected_type.__name__}")
continue
if validator and not validator(value):
errors.append(f"{key}: value {value!r} failed validation")
continue
if validator and not validator(value):
errors.append(f"{key}: value {value!r} failed validation")
continue
if key in _SETTINGS_FIELDS:
validated_settings[key] = value
else:
validated_root[key] = value
if key == "workers":
errors.extend(_validate_workers_payload(value))
if errors:
continue
if errors:
return web.json_response({
"status": "error",
"error": errors,
"message": "; ".join(errors),
}, status=400)
if key in _SETTINGS_FIELDS:
validated_settings[key] = value
else:
validated_root[key] = value
if errors:
return await handle_api_error(request, errors, 400)
try:
async with config_transaction() as config:
settings = config.setdefault("settings", {})
settings.update(validated_settings)
for key, value in validated_root.items():
config[key] = value
return web.json_response({"status": "success", "config": config})
except Exception as e:
return await handle_api_error(request, e)
return await run_authorized_endpoint(request, _operation)
@server.PromptServer.instance.routes.get("/distributed/queue_status/{job_id}")
async def queue_status_endpoint(request):
async def queue_status_endpoint(request: web.Request) -> web.StreamResponse:
"""Check if a job queue is initialized."""
try:
async def _operation() -> web.StreamResponse:
job_id = request.match_info['job_id']
# Import to ensure initialization
@@ -159,12 +177,12 @@ async def queue_status_endpoint(request):
debug_log(f"Queue status check for job {job_id}: {'exists' if exists else 'not found'}")
return web.json_response({"exists": exists, "job_id": job_id})
except Exception as e:
return await handle_api_error(request, e, 500)
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
@server.PromptServer.instance.routes.post("/distributed/config/update_worker")
async def update_worker_endpoint(request):
try:
async def update_worker_endpoint(request: web.Request) -> web.StreamResponse:
async def _operation() -> web.StreamResponse:
data = await request.json()
worker_id = data.get("worker_id")
@@ -203,12 +221,12 @@ async def update_worker_endpoint(request):
)
return web.json_response({"status": "success"})
except Exception as e:
return await handle_api_error(request, e, 400)
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
@server.PromptServer.instance.routes.post("/distributed/config/delete_worker")
async def delete_worker_endpoint(request):
try:
async def delete_worker_endpoint(request: web.Request) -> web.StreamResponse:
async def _operation() -> web.StreamResponse:
data = await request.json()
worker_id = data.get("worker_id")
@@ -235,13 +253,13 @@ async def delete_worker_endpoint(request):
"status": "success",
"message": f"Worker {removed_worker.get('name', worker_id)} deleted"
})
except Exception as e:
return await handle_api_error(request, e, 400)
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
@server.PromptServer.instance.routes.post("/distributed/config/update_setting")
async def update_setting_endpoint(request):
async def update_setting_endpoint(request: web.Request) -> web.StreamResponse:
"""Updates a specific key in the settings object."""
try:
async def _operation() -> web.StreamResponse:
data = await request.json()
key = data.get("key")
value = data.get("value")
@@ -258,13 +276,13 @@ async def update_setting_endpoint(request):
config['settings'][key] = value
return web.json_response({"status": "success", "message": f"Setting '{key}' updated."})
except Exception as e:
return await handle_api_error(request, e, 400)
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
@server.PromptServer.instance.routes.post("/distributed/config/update_master")
async def update_master_endpoint(request):
async def update_master_endpoint(request: web.Request) -> web.StreamResponse:
"""Updates master configuration."""
try:
async def _operation() -> web.StreamResponse:
data = await request.json()
async with config_transaction() as config:
@@ -273,5 +291,5 @@ async def update_master_endpoint(request):
_apply_field_patch(config['master'], data, _MASTER_FIELDS)
return web.json_response({"status": "success", "message": "Master configuration updated."})
except Exception as e:
return await handle_api_error(request, e, 400)
return await run_authorized_endpoint(request, _operation, unexpected_status=500)
+25
View File
@@ -0,0 +1,25 @@
"""Shared endpoint policy helpers for distributed API routes."""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from aiohttp import web
from ..utils.network import handle_api_error
from .request_guards import authorization_error_or_none
async def run_authorized_endpoint(
request: web.Request,
operation: Callable[[], Awaitable[web.StreamResponse]],
*,
unexpected_status: int = 500,
) -> web.StreamResponse:
"""Run endpoint operation with shared auth guard and fallback error mapping."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
return await operation()
except Exception as exc:
return await handle_api_error(request, exc, unexpected_status)
+89 -40
View File
@@ -1,3 +1,5 @@
from __future__ import annotations
import json
import asyncio
import io
@@ -5,6 +7,7 @@ import os
import base64
import binascii
import time
from typing import Any
from aiohttp import web
import server
@@ -15,22 +18,21 @@ from ..utils.logging import debug_log
from ..utils.image import pil_to_tensor, ensure_contiguous
from ..utils.network import handle_api_error
from ..utils.constants import JOB_INIT_GRACE_PERIOD, MEMORY_CLEAR_DELAY
try:
from .queue_orchestration import ensure_distributed_state, orchestrate_distributed_execution
except ImportError:
from .queue_orchestration import orchestrate_distributed_execution
def ensure_distributed_state():
return None
from ..utils.runtime_state import ensure_distributed_runtime_state
from .request_guards import authorization_error_or_none
from .schemas import require_bool_literal
from .queue_orchestration import orchestrate_distributed_execution
from .queue_request import parse_queue_request_payload
prompt_server = server.PromptServer.instance
# Canonical worker result envelope accepted by POST /distributed/job_complete:
# { "job_id": str, "worker_id": str, "batch_idx": int, "image": <base64 PNG>, "is_last": bool }
def _decode_image_sync(image_path):
def _runtime_state():
return ensure_distributed_runtime_state()
def _decode_image_sync(image_path: str) -> dict[str, Any]:
"""Decode image/video file and compute hash in a threadpool worker."""
import base64
import hashlib
@@ -76,7 +78,7 @@ def _decode_image_sync(image_path):
}
def _check_file_sync(filename, expected_hash):
def _check_file_sync(filename: str, expected_hash: str) -> dict[str, Any]:
"""Check file presence and hash in a threadpool worker."""
import hashlib
import folder_paths
@@ -101,7 +103,7 @@ def _check_file_sync(filename, expected_hash):
}
def _decode_canonical_png_tensor(image_payload):
def _decode_canonical_png_tensor(image_payload: str) -> torch.Tensor:
"""Decode canonical base64 PNG payload into a contiguous IMAGE tensor."""
if not isinstance(image_payload, str) or not image_payload.strip():
raise ValueError("Field 'image' must be a non-empty base64 PNG string.")
@@ -132,7 +134,7 @@ def _decode_canonical_png_tensor(image_payload):
raise ValueError(f"Failed to decode PNG image payload: {exc}") from exc
def _decode_audio_payload(audio_payload):
def _decode_audio_payload(audio_payload: dict[str, Any]) -> dict[str, Any]:
"""Decode canonical audio payload into an AUDIO dict."""
from ..utils.audio_payload import decode_audio_payload
@@ -140,17 +142,20 @@ def _decode_audio_payload(audio_payload):
@server.PromptServer.instance.routes.post("/distributed/prepare_job")
async def prepare_job_endpoint(request):
async def prepare_job_endpoint(request: web.Request) -> web.StreamResponse:
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
multi_job_id = data.get('multi_job_id')
if not multi_job_id:
return await handle_api_error(request, "Missing multi_job_id", 400)
ensure_distributed_state()
async with prompt_server.distributed_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_jobs:
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
runtime_state = _runtime_state()
async with runtime_state.distributed_jobs_lock:
if multi_job_id not in runtime_state.distributed_pending_jobs:
runtime_state.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
debug_log(f"Prepared queue for job {multi_job_id}")
return web.json_response({"status": "success"})
@@ -158,9 +163,13 @@ async def prepare_job_endpoint(request):
return await handle_api_error(request, e)
@server.PromptServer.instance.routes.post("/distributed/clear_memory")
async def clear_memory_endpoint(request):
async def clear_memory_endpoint(request: web.Request) -> web.StreamResponse:
debug_log("Received request to clear VRAM.")
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
warnings: list[str] = []
# Use ComfyUI's prompt server queue system like the /free endpoint does
if hasattr(server.PromptServer.instance, 'prompt_queue'):
server.PromptServer.instance.prompt_queue.set_flag("unload_models", True)
@@ -177,12 +186,16 @@ async def clear_memory_endpoint(request):
try:
mm.unload_all_models()
except AttributeError as e:
debug_log(f"Warning during model unload: {e}")
warning = f"Model unload warning: {e}"
warnings.append(warning)
debug_log(warning)
try:
mm.soft_empty_cache()
except Exception as e:
debug_log(f"Warning during cache clear: {e}")
warning = f"Cache clear warning: {e}"
warnings.append(warning)
debug_log(warning)
for _ in range(3):
gc.collect()
@@ -191,6 +204,16 @@ async def clear_memory_endpoint(request):
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if warnings:
debug_log("VRAM cleared with warnings.")
return web.json_response(
{
"status": "partial",
"message": "GPU memory cleared with warnings.",
"warnings": warnings,
}
)
debug_log("VRAM cleared successfully.")
return web.json_response({"status": "success", "message": "GPU memory cleared."})
except Exception as e:
@@ -199,13 +222,16 @@ async def clear_memory_endpoint(request):
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
debug_log(f"Partial VRAM clear completed with warning: {e}")
return web.json_response({"status": "success", "message": "GPU memory cleared (with warnings)"})
debug_log(f"VRAM clear failed: {e}")
return await handle_api_error(request, f"GPU memory clear failed: {e}", 500)
@server.PromptServer.instance.routes.post("/distributed/queue")
async def distributed_queue_endpoint(request):
async def distributed_queue_endpoint(request: web.Request) -> web.StreamResponse:
"""Queue a distributed workflow, mirroring the UI orchestration pipeline."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
raw_payload = await request.json()
except Exception as exc:
@@ -228,14 +254,16 @@ async def distributed_queue_endpoint(request):
return web.json_response({
"prompt_id": prompt_id,
"worker_count": worker_count,
"auto_prepare_supported": True,
})
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):
async def load_image_endpoint(request: web.Request) -> web.StreamResponse:
"""Load an image or video file and return it as base64 data with hash."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
image_path = data.get("image_path")
@@ -251,8 +279,11 @@ async def load_image_endpoint(request):
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.post("/distributed/check_file")
async def check_file_endpoint(request):
async def check_file_endpoint(request: web.Request) -> web.StreamResponse:
"""Check if a file exists and matches the given hash."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
filename = data.get("filename")
@@ -269,7 +300,10 @@ async def check_file_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/job_complete")
async def job_complete_endpoint(request):
async def job_complete_endpoint(request: web.Request) -> web.StreamResponse:
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
except Exception as exc:
@@ -284,7 +318,7 @@ async def job_complete_endpoint(request):
batch_idx = data.get("batch_idx")
image_payload = data.get("image")
audio_payload = data.get("audio")
is_last = data.get("is_last")
is_last_raw = data.get("is_last")
errors = []
if not isinstance(job_id, str) or not job_id.strip():
@@ -297,8 +331,12 @@ async def job_complete_endpoint(request):
errors.append("image: expected non-empty base64 PNG string")
if audio_payload is not None and not isinstance(audio_payload, dict):
errors.append("audio: expected object when provided")
if not isinstance(is_last, bool):
errors.append("is_last: expected boolean")
try:
is_last = require_bool_literal(is_last_raw, field_name="is_last")
except ValueError as exc:
errors.append(str(exc))
is_last = False
if errors:
return await handle_api_error(request, errors, 400)
@@ -307,21 +345,32 @@ async def job_complete_endpoint(request):
multi_job_id = job_id.strip()
worker_id = worker_id.strip()
runtime_state = _runtime_state()
allowed_workers = runtime_state.distributed_job_allowed_workers.get(multi_job_id)
if allowed_workers is not None and worker_id not in allowed_workers:
return await handle_api_error(
request,
f"Unauthorized worker_id for job {multi_job_id}",
403,
)
pending = None
queue_size = 0
deadline = time.monotonic() + float(JOB_INIT_GRACE_PERIOD)
while pending is None:
async with prompt_server.distributed_jobs_lock:
pending = prompt_server.distributed_pending_jobs.get(multi_job_id)
async with runtime_state.distributed_jobs_lock:
pending = runtime_state.distributed_pending_jobs.get(multi_job_id)
if pending is not None:
queue_item = {
"tensor": tensor,
"worker_id": worker_id,
"image_index": int(batch_idx),
"is_last": is_last,
}
if decoded_audio is not None:
queue_item["audio"] = decoded_audio
await pending.put(
{
"tensor": tensor,
"worker_id": worker_id,
"image_index": int(batch_idx),
"is_last": is_last,
"audio": decoded_audio,
}
queue_item
)
queue_size = pending.qsize()
break
+7 -4
View File
@@ -1,4 +1,7 @@
# Orchestration helpers split out from queue_orchestration.py:
# - prompt_transform.py: graph pruning + hidden input overrides
# - media_sync.py: remote media/path normalization
# - dispatch.py: worker probe + prompt dispatch
"""Orchestration helpers used by distributed queue execution."""
__all__ = [
"dispatch",
"media_sync",
"prompt_transform",
]
+91 -43
View File
@@ -1,45 +1,33 @@
import asyncio
import json
import uuid
from typing import Any
import aiohttp
from ...utils.auth import distributed_auth_headers
from ...utils.config import load_config
from ...utils.logging import debug_log, log
from ...utils.network import build_worker_url, get_client_session, probe_worker
try:
from ...utils.trace_logger import trace_debug, trace_info
except ImportError:
def trace_debug(*_args, **_kwargs):
return None
def trace_info(*_args, **_kwargs):
return None
try:
from ..schemas import parse_positive_int
except ImportError:
def parse_positive_int(value, default):
try:
parsed = int(value)
return parsed if parsed > 0 else default
except (TypeError, ValueError):
return default
from ...utils.trace_logger import trace_debug, trace_info
from ..schemas import coerce_positive_int
_least_busy_rr_index = 0
async def worker_is_active(worker):
async def worker_is_active(worker: dict[str, Any]) -> bool:
"""Ping worker's /prompt endpoint to confirm it's reachable."""
url = build_worker_url(worker)
return await probe_worker(url, timeout=3.0) is not None
async def worker_ws_is_active(worker):
async def worker_ws_is_active(worker: dict[str, Any]) -> bool:
"""Ping worker's websocket endpoint to confirm it's reachable."""
session = await get_client_session()
url = build_worker_url(worker, "/distributed/worker_ws")
headers = distributed_auth_headers(load_config())
try:
ws = await session.ws_connect(url, heartbeat=20, timeout=3)
ws = await session.ws_connect(url, heartbeat=20, timeout=3, headers=headers)
await ws.close()
return True
except asyncio.TimeoutError:
@@ -59,7 +47,13 @@ async def _probe_worker_active(worker, use_websocket, semaphore):
return worker, is_active
async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
async def _dispatch_via_websocket(
worker_url,
payload,
client_id,
headers: dict[str, str],
timeout=60.0,
):
"""Open a fresh worker websocket, dispatch one prompt, wait for ack, then close."""
request_id = uuid.uuid4().hex
ws_payload = {
@@ -73,7 +67,12 @@ async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
ws_url = f"{ws_url}/distributed/worker_ws"
session = await get_client_session()
async with session.ws_connect(ws_url, heartbeat=20, timeout=timeout) as ws:
async with session.ws_connect(
ws_url,
heartbeat=20,
timeout=timeout,
headers=headers,
) as ws:
await ws.send_json(ws_payload)
async for msg in ws:
if msg.type == aiohttp.WSMsgType.TEXT:
@@ -96,13 +95,13 @@ async def _dispatch_via_websocket(worker_url, payload, client_id, timeout=60.0):
async def dispatch_worker_prompt(
worker,
prompt_obj,
workflow_meta,
client_id=None,
use_websocket=False,
trace_execution_id=None,
):
worker: dict[str, Any],
prompt_obj: dict[str, Any],
workflow_meta: dict[str, Any] | None,
client_id: str | None = None,
use_websocket: bool = False,
trace_execution_id: str | None = None,
) -> None:
"""Send the prepared prompt to a worker ComfyUI instance."""
worker_url = build_worker_url(worker)
url = build_worker_url(worker, "/prompt")
@@ -114,6 +113,7 @@ async def dispatch_worker_prompt(
payload["extra_data"] = extra_data
if use_websocket:
ws_headers = distributed_auth_headers(load_config())
try:
await _dispatch_via_websocket(
worker_url,
@@ -122,6 +122,7 @@ async def dispatch_worker_prompt(
"workflow": workflow_meta,
},
client_id,
ws_headers,
)
return
except Exception as exc:
@@ -142,14 +143,14 @@ async def dispatch_worker_prompt(
async def select_active_workers(
workers,
use_websocket,
delegate_master,
trace_execution_id=None,
probe_concurrency=8,
):
workers: list[dict[str, Any]],
use_websocket: bool,
delegate_master: bool,
trace_execution_id: str | None = None,
probe_concurrency: int = 8,
) -> tuple[list[dict[str, Any]], bool]:
"""Probe workers and return (active_workers, updated_delegate_master)."""
probe_limit = parse_positive_int(probe_concurrency, 8)
probe_limit = coerce_positive_int(probe_concurrency, 8)
probe_semaphore = asyncio.Semaphore(probe_limit)
if trace_execution_id and workers:
@@ -223,16 +224,16 @@ def _select_idle_round_robin(statuses):
async def select_least_busy_worker(
workers,
trace_execution_id=None,
probe_concurrency=8,
probe_timeout=3.0,
):
workers: list[dict[str, Any]],
trace_execution_id: str | None = None,
probe_concurrency: int = 8,
probe_timeout: float = 3.0,
) -> dict[str, Any] | None:
"""Select one worker by queue depth, round-robin among idle workers."""
if not workers:
return None
probe_limit = parse_positive_int(probe_concurrency, 8)
probe_limit = coerce_positive_int(probe_concurrency, 8)
probe_semaphore = asyncio.Semaphore(probe_limit)
statuses = await asyncio.gather(
*[
@@ -266,3 +267,50 @@ async def select_least_busy_worker(
f"Least-busy worker selected: {worker.get('name')} ({worker.get('id')}), queue_remaining={queue_remaining}"
)
return worker
async def rank_workers_by_load(
workers: list[dict[str, Any]],
trace_execution_id: str | None = None,
probe_concurrency: int = 8,
probe_timeout: float = 3.0,
) -> list[dict[str, Any]]:
"""Return all workers sorted by queue depth (ascending), preserving unreachable workers at the end."""
if not workers:
return []
probe_limit = coerce_positive_int(probe_concurrency, 8)
probe_semaphore = asyncio.Semaphore(probe_limit)
statuses = await asyncio.gather(
*[
_probe_worker_queue(worker, probe_semaphore, probe_timeout)
for worker in workers
]
)
ranked_statuses = []
unreachable_workers = []
for worker, status in zip(workers, statuses):
if status is None:
unreachable_workers.append(worker)
continue
ranked_statuses.append(status)
ranked_statuses.sort(key=lambda status: status["queue_remaining"])
ranked_workers = [status["worker"] for status in ranked_statuses] + unreachable_workers
if trace_execution_id:
trace_debug(
trace_execution_id,
"Ranked workers by load: "
+ ", ".join(
f"{status['worker'].get('id')}={status['queue_remaining']}"
for status in ranked_statuses
)
+ (
f"; unreachable={','.join(worker.get('id', '') for worker in unreachable_workers)}"
if unreachable_workers
else ""
),
)
return ranked_workers
+17 -4
View File
@@ -3,9 +3,12 @@ import hashlib
import mimetypes
import os
import re
from typing import Any
import aiohttp
from ...utils.auth import distributed_auth_headers
from ...utils.config import load_config
from ...utils.logging import debug_log, log
from ...utils.network import build_worker_url, get_client_session
from ...utils.trace_logger import trace_debug, trace_info
@@ -33,7 +36,7 @@ def _normalize_media_reference(value):
return None
def convert_paths_for_platform(obj, target_separator):
def convert_paths_for_platform(obj: Any, target_separator: str) -> Any:
"""Recursively normalize likely file paths for the worker platform separator."""
if target_separator not in ("/", "\\"):
return obj
@@ -124,12 +127,16 @@ def _load_media_file_sync(filename):
return file_bytes, file_hash, mime_type
async def fetch_worker_path_separator(worker, trace_execution_id=None):
async def fetch_worker_path_separator(
worker: dict[str, Any],
trace_execution_id: str | None = None,
) -> str | None:
"""Best-effort fetch of a worker's path separator from /distributed/system_info."""
url = build_worker_url(worker, "/distributed/system_info")
session = await get_client_session()
headers = distributed_auth_headers(load_config())
try:
async with session.get(url, timeout=aiohttp.ClientTimeout(total=5)) as resp:
async with session.get(url, headers=headers, timeout=aiohttp.ClientTimeout(total=5)) as resp:
if resp.status != 200:
return None
payload = await resp.json()
@@ -146,6 +153,7 @@ async def fetch_worker_path_separator(worker, trace_execution_id=None):
async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_type, trace_execution_id=None):
"""Upload one media file to worker iff missing or hash-mismatched."""
session = await get_client_session()
headers = distributed_auth_headers(load_config())
normalized = filename.replace("\\", "/")
check_url = build_worker_url(worker, "/distributed/check_file")
@@ -153,6 +161,7 @@ async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_
async with session.post(
check_url,
json={"filename": normalized, "hash": file_hash},
headers=headers,
timeout=aiohttp.ClientTimeout(total=6),
) as resp:
if resp.status == 200:
@@ -193,7 +202,11 @@ async def _upload_media_to_worker(worker, filename, file_bytes, file_hash, mime_
return True, worker_path
async def sync_worker_media(worker, prompt_obj, trace_execution_id=None):
async def sync_worker_media(
worker: dict[str, Any],
prompt_obj: dict[str, Any],
trace_execution_id: str | None = None,
) -> None:
"""Sync referenced media files from master to a remote worker before dispatch."""
media_refs = _find_media_references(prompt_obj)
if not media_refs:
+408 -84
View File
@@ -1,5 +1,6 @@
import json
from collections import deque
from typing import Any
from ...utils.logging import debug_log
@@ -27,7 +28,7 @@ class PromptIndex:
def nodes_for_class(self, class_name):
return self.nodes_by_class.get(class_name, [])
def has_upstream(self, start_node_id, target_class):
def has_upstream(self, start_node_id: str, target_class: str) -> bool:
cache_key = (str(start_node_id), target_class)
if cache_key in self._upstream_cache:
return self._upstream_cache[cache_key]
@@ -88,6 +89,38 @@ def _find_downstream_nodes(prompt_obj, start_ids):
return connected
def _find_downstream_of_output_slot(prompt_obj, node_id, slot_index):
"""Return all nodes reachable from a specific output slot on one node."""
source_id = str(node_id)
try:
slot = int(slot_index)
except (TypeError, ValueError):
return set()
direct_consumers = set()
for candidate_id, candidate_node in _iter_prompt_nodes(prompt_obj):
inputs = candidate_node.get("inputs", {})
for value in inputs.values():
if not (isinstance(value, list) and len(value) == 2):
continue
if str(value[0]) != source_id:
continue
try:
input_slot = int(value[1])
except (TypeError, ValueError):
debug_log(
f"Prompt transform: skipping malformed slot reference for node {candidate_id}: {value}"
)
continue
if input_slot == slot:
direct_consumers.add(str(candidate_id))
break
if not direct_consumers:
return set()
return _find_downstream_nodes(prompt_obj, list(direct_consumers))
def _create_numeric_id_generator(prompt_obj):
"""Return a closure that yields new numeric string IDs."""
max_id = 0
@@ -95,6 +128,7 @@ def _create_numeric_id_generator(prompt_obj):
try:
numeric = int(node_id)
except (TypeError, ValueError):
debug_log(f"Prompt transform: ignoring non-numeric node id while generating ids: {node_id}")
continue
max_id = max(max_id, numeric)
@@ -125,15 +159,193 @@ def _find_upstream_nodes(prompt_obj, start_ids):
return connected
def prune_prompt_for_worker(prompt_obj):
def _resolve_participants(enabled_worker_ids, delegate_master):
worker_ids = [str(worker_id) for worker_id in (enabled_worker_ids or [])]
if delegate_master:
return worker_ids
return ["master"] + worker_ids
def _remove_dangling_input_refs(prompt_obj):
"""Drop input links that point to nodes no longer present in prompt_obj."""
existing_ids = set(prompt_obj.keys())
for _node_id, node in _iter_prompt_nodes(prompt_obj):
inputs = node.get("inputs", {})
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 existing_ids:
inputs.pop(input_name, None)
def _has_terminal_output_nodes(prompt_obj):
"""Return True when the prompt already has at least one terminal output node."""
for _node_id, node in _iter_prompt_nodes(prompt_obj):
if node.get("class_type") in {"PreviewImage", "SaveImage"}:
return True
return False
def _add_preview_node(prompt_obj, next_id_fn, source_node_id, slot_index, title_suffix):
"""Attach a PreviewImage output node to the given source output slot."""
preview_id = next_id_fn()
prompt_obj[preview_id] = {
"inputs": {
"images": [str(source_node_id), int(slot_index)],
},
"class_type": "PreviewImage",
"_meta": {
"title": f"Preview Image ({title_suffix})",
},
}
return preview_id
def _ensure_worker_output_node(prompt_obj):
"""Guarantee worker prompts retain at least one output node after pruning.
Branch pruning can remove all explicit output nodes for worker prompts, which
causes ComfyUI validation to reject the prompt (`prompt_no_outputs`).
"""
if _has_terminal_output_nodes(prompt_obj):
return prompt_obj
next_id = _create_numeric_id_generator(prompt_obj)
# Prefer a DistributedBranchCollector output tied to the assigned branch for this worker.
for node_id, node in _iter_prompt_nodes(prompt_obj):
if node.get("class_type") != "DistributedBranchCollector":
continue
inputs = node.get("inputs", {})
try:
assigned_branch = int(inputs.get("assigned_branch", -1))
except (TypeError, ValueError):
assigned_branch = -1
if assigned_branch >= 0:
_add_preview_node(prompt_obj, next_id, node_id, assigned_branch, "auto-added worker output")
return prompt_obj
# Otherwise, anchor to the assigned DistributedBranch slot if available.
for node_id, node in _iter_prompt_nodes(prompt_obj):
if node.get("class_type") != "DistributedBranch":
continue
inputs = node.get("inputs", {})
try:
assigned_branch = int(inputs.get("assigned_branch", -1))
except (TypeError, ValueError):
assigned_branch = -1
if assigned_branch >= 0:
_add_preview_node(prompt_obj, next_id, node_id, assigned_branch, "auto-added worker output")
return prompt_obj
# Idle participants (assigned_branch=-1) still need a terminal node to pass
# validation; keep this cheap by emitting a tiny synthetic image preview.
empty_id = next_id()
prompt_obj[empty_id] = {
"class_type": "DistributedEmptyImage",
"inputs": {
"height": 64,
"width": 64,
"channels": 3,
},
"_meta": {
"title": "Distributed Empty Image (auto-added worker output)",
},
}
_add_preview_node(prompt_obj, next_id, empty_id, 0, "auto-added worker output")
return prompt_obj
def _prune_worker_downstream_of_branch_collectors(prompt_obj):
"""Drop worker-side nodes downstream of branch collectors."""
collector_ids = find_nodes_by_class(prompt_obj, "DistributedBranchCollector")
if not collector_ids:
return prompt_obj
downstream = _find_downstream_nodes(prompt_obj, collector_ids)
for node_id in downstream:
if node_id in collector_ids:
continue
prompt_obj.pop(node_id, None)
_remove_dangling_input_refs(prompt_obj)
return prompt_obj
def prune_prompt_for_branch_worker(
prompt_obj: dict[str, Any],
branch_node_id: str,
assigned_branch: int | list[int] | tuple[int, ...] | set[int],
num_branches: int,
) -> dict[str, Any]:
"""Prune non-assigned branch paths while keeping shared downstream nodes."""
branch_id = str(branch_node_id)
assigned_slots = set()
if isinstance(assigned_branch, (list, tuple, set)):
for value in assigned_branch:
try:
idx = int(value)
except (TypeError, ValueError):
debug_log(f"Prompt transform: invalid assigned branch slot ignored: {value}")
continue
if idx >= 0:
assigned_slots.add(idx)
else:
try:
idx = int(assigned_branch)
except (TypeError, ValueError):
idx = -1
if idx >= 0:
assigned_slots.add(idx)
try:
total_slots = int(num_branches)
except (TypeError, ValueError):
total_slots = 2
total_slots = max(2, min(total_slots, 10))
downstream_by_slot = {}
for slot_idx in range(total_slots):
downstream = _find_downstream_of_output_slot(prompt_obj, branch_id, slot_idx)
downstream.discard(branch_id)
downstream_by_slot[slot_idx] = downstream
keep = {branch_id}
for slot_idx in assigned_slots:
keep.update(downstream_by_slot.get(slot_idx, set()))
remove = set()
for slot_idx in range(total_slots):
if slot_idx in assigned_slots:
continue
remove.update(downstream_by_slot.get(slot_idx, set()))
remove -= keep
for node_id in remove:
prompt_obj.pop(node_id, None)
_remove_dangling_input_refs(prompt_obj)
return prompt_obj
def prune_prompt_for_worker(prompt_obj: dict[str, Any]) -> dict[str, Any]:
"""Prune worker prompt to distributed nodes and their upstream dependencies."""
collector_ids = find_nodes_by_class(prompt_obj, "DistributedCollector")
branch_collector_ids = find_nodes_by_class(prompt_obj, "DistributedBranchCollector")
upscale_ids = find_nodes_by_class(prompt_obj, "UltimateSDUpscaleDistributed")
distributed_ids = collector_ids + upscale_ids
branch_ids = find_nodes_by_class(prompt_obj, "DistributedBranch")
distributed_ids = collector_ids + branch_collector_ids + upscale_ids + branch_ids
if not distributed_ids:
return prompt_obj
connected = _find_upstream_nodes(prompt_obj, distributed_ids)
if branch_ids:
connected.update(_find_downstream_nodes(prompt_obj, branch_ids))
if branch_collector_ids:
downstream_of_collectors = _find_downstream_nodes(prompt_obj, branch_collector_ids)
connected -= (downstream_of_collectors - set(branch_collector_ids))
pruned_prompt = {}
for node_id in connected:
node = prompt_obj.get(node_id)
@@ -146,7 +358,7 @@ def prune_prompt_for_worker(prompt_obj):
if dist_id not in pruned_prompt:
continue
downstream = _find_downstream_nodes(prompt_obj, [dist_id])
has_removed_downstream = any(node_id != dist_id for node_id in downstream)
has_removed_downstream = any(node_id != dist_id and node_id not in connected for node_id in downstream)
if has_removed_downstream:
preview_id = next_id()
pruned_prompt[preview_id] = {
@@ -162,7 +374,10 @@ def prune_prompt_for_worker(prompt_obj):
return pruned_prompt
def prepare_delegate_master_prompt(prompt_obj, collector_ids):
def prepare_delegate_master_prompt(
prompt_obj: dict[str, Any],
collector_ids: list[str],
) -> dict[str, Any]:
"""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)
@@ -194,6 +409,17 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
collector_entry = pruned_prompt.get(collector_id)
if not collector_entry:
continue
collector_class = collector_entry.get("class_type")
if collector_class == "DistributedBranchCollector":
# Branch collector does not require an image input in delegate-only mode.
continue
if collector_class == "UltimateSDUpscaleDistributed":
# Swap to lightweight delegate collector — no model inputs needed.
collector_entry["class_type"] = "USDUDelegateCollector"
debug_log(
f"Swapped USDU node {collector_id} to USDUDelegateCollector for delegate-only master prompt."
)
continue
placeholder_id = next_id()
pruned_prompt[placeholder_id] = {
"class_type": "DistributedEmptyImage",
@@ -214,32 +440,145 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
return pruned_prompt
def generate_job_id_map(prompt_index, prefix):
def generate_job_id_map(prompt_index: PromptIndex, prefix: str) -> dict[str, str]:
"""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"
distributed_nodes = (
prompt_index.nodes_for_class("DistributedCollector")
+ prompt_index.nodes_for_class("DistributedBranchCollector")
+ prompt_index.nodes_for_class("UltimateSDUpscaleDistributed")
+ prompt_index.nodes_for_class("DistributedBranch")
)
for node_id in distributed_nodes:
job_map[node_id] = f"{prefix}_{node_id}"
return job_map
def _override_seed_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map):
"""Configure DistributedSeed nodes for master or worker role."""
for node_id in prompt_index.nodes_for_class("DistributedSeed"):
def _override_simple_nodes(prompt_copy, prompt_index, is_master, participant_id, enabled_json, class_names):
"""Configure simple distributed nodes (DistributedSeed, DistributedValue)."""
for class_name in class_names:
for node_id in prompt_index.nodes_for_class(class_name):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
inputs = node.setdefault("inputs", {})
inputs["is_worker"] = not is_master
inputs["enabled_worker_ids"] = enabled_json
inputs["worker_id"] = "" if is_master else str(participant_id)
def _override_job_nodes(
prompt_copy, prompt_index, is_master, participant_id,
job_id_map, master_url, enabled_json, delegate_master,
class_names, skip_if_upstream=None,
):
"""Configure job-aware distributed nodes (collector, upscale)."""
for class_name in class_names:
for node_id in prompt_index.nodes_for_class(class_name):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
if skip_if_upstream and prompt_index.has_upstream(node_id, skip_if_upstream):
node.setdefault("inputs", {})["pass_through"] = True
continue
inputs = node.setdefault("inputs", {})
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
inputs["is_worker"] = not is_master
inputs["enabled_worker_ids"] = enabled_json
if is_master:
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["delegate_only"] = False
def _override_branch_nodes(
prompt_copy,
prompt_index,
is_master,
participant_id,
enabled_worker_ids,
job_id_map,
delegate_master,
):
"""Assign branch slots per participant and prune non-assigned branch paths."""
participants = _resolve_participants(enabled_worker_ids, delegate_master)
participant_id = str(participant_id)
participant_count = len(participants)
participant_pos = participants.index(participant_id) if participant_id in participants else None
for node_id in prompt_index.nodes_for_class("DistributedBranch"):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
inputs = node.setdefault("inputs", {})
inputs["is_worker"] = not is_master
if is_master:
inputs["worker_id"] = ""
try:
num_branches = int(inputs.get("num_branches", 2))
except (TypeError, ValueError):
num_branches = 2
num_branches = max(2, min(num_branches, 10))
assigned_slots = []
if participant_pos is not None and participant_count > 0:
assigned_slots = [slot for slot in range(num_branches) if slot % participant_count == participant_pos]
if len(assigned_slots) == 1:
assigned_branch = assigned_slots[0]
else:
inputs["worker_id"] = f"worker_{worker_index_map.get(participant_id, 0)}"
assigned_branch = -1
inputs["is_worker"] = not is_master
inputs["worker_id"] = "" if is_master else participant_id
inputs["assigned_branch"] = assigned_branch
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
if is_master and delegate_master:
continue
prompt_copy = prune_prompt_for_branch_worker(
prompt_copy,
node_id,
assigned_slots,
num_branches,
)
return prompt_copy
def _override_collector_nodes(
def _find_upstream_branch_node(prompt_copy, start_node_id):
"""Locate the nearest upstream DistributedBranch node for the provided node."""
visited = set()
stack = [str(start_node_id)]
while stack:
node_id = stack.pop()
if node_id in visited:
continue
visited.add(node_id)
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
inputs = node.get("inputs", {})
for input_value in inputs.values():
if not (isinstance(input_value, list) and len(input_value) == 2):
continue
source_id = str(input_value[0])
source_node = prompt_copy.get(source_id)
if not isinstance(source_node, dict):
continue
if source_node.get("class_type") == "DistributedBranch":
return source_id
stack.append(source_id)
return None
def _override_branch_collector_nodes(
prompt_copy,
prompt_index,
is_master,
@@ -249,87 +588,77 @@ def _override_collector_nodes(
enabled_json,
delegate_master,
):
"""Configure DistributedCollector nodes for master or worker role."""
for node_id in prompt_index.nodes_for_class("DistributedCollector"):
"""Configure DistributedBranchCollector nodes for branch convergence."""
for node_id in prompt_index.nodes_for_class("DistributedBranchCollector"):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
if prompt_index.has_upstream(node_id, "UltimateSDUpscaleDistributed"):
node.setdefault("inputs", {})["pass_through"] = True
continue
upstream_branch_id = _find_upstream_branch_node(prompt_copy, node_id)
assigned_branch = -1
if upstream_branch_id is not None:
branch_node = prompt_copy.get(upstream_branch_id, {})
branch_inputs = branch_node.get("inputs", {}) if isinstance(branch_node, dict) else {}
try:
assigned_branch = int(branch_inputs.get("assigned_branch", -1))
except (TypeError, ValueError):
assigned_branch = -1
inputs = node.setdefault("inputs", {})
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
# Group all branch collectors under the same DistributedBranch into one queue.
job_source_id = upstream_branch_id or node_id
inputs["multi_job_id"] = job_id_map.get(job_source_id, job_source_id)
inputs["is_worker"] = not is_master
inputs["enabled_worker_ids"] = enabled_json
inputs["assigned_branch"] = assigned_branch
if is_master:
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["worker_id"] = str(participant_id)
inputs["delegate_only"] = False
def _override_upscale_nodes(
prompt_copy,
prompt_index,
is_master,
participant_id,
job_id_map,
master_url,
enabled_json,
):
"""Configure UltimateSDUpscaleDistributed nodes for master or worker role."""
for node_id in prompt_index.nodes_for_class("UltimateSDUpscaleDistributed"):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
inputs = node.setdefault("inputs", {})
inputs["multi_job_id"] = job_id_map.get(node_id, node_id)
inputs["is_worker"] = not is_master
inputs["enabled_worker_ids"] = enabled_json
if is_master:
inputs.pop("master_url", None)
inputs.pop("worker_id", None)
else:
inputs["master_url"] = master_url
inputs["worker_id"] = participant_id
def _override_value_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map):
"""Configure DistributedValue nodes for master or worker role."""
for node_id in prompt_index.nodes_for_class("DistributedValue"):
node = prompt_copy.get(node_id)
if not isinstance(node, dict):
continue
inputs = node.setdefault("inputs", {})
inputs["is_worker"] = not is_master
if is_master:
inputs["worker_id"] = ""
else:
inputs["worker_id"] = f"worker_{worker_index_map.get(participant_id, 0)}"
def apply_participant_overrides(
prompt_copy,
participant_id,
enabled_worker_ids,
job_id_map,
master_url,
delegate_master,
prompt_index,
):
prompt_copy: dict[str, Any],
participant_id: str,
enabled_worker_ids: list[str],
job_id_map: dict[str, str],
master_url: str,
delegate_master: bool,
prompt_index: PromptIndex,
) -> dict[str, Any]:
"""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)
_override_seed_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map)
_override_value_nodes(prompt_copy, prompt_index, is_master, participant_id, worker_index_map)
_override_collector_nodes(
_override_simple_nodes(
prompt_copy, prompt_index, is_master, participant_id, enabled_json,
["DistributedSeed", "DistributedValue"],
)
_override_job_nodes(
prompt_copy, prompt_index, is_master, participant_id,
job_id_map, master_url, enabled_json, delegate_master,
["DistributedCollector"],
skip_if_upstream="UltimateSDUpscaleDistributed",
)
_override_job_nodes(
prompt_copy, prompt_index, is_master, participant_id,
job_id_map, master_url, enabled_json, delegate_master,
["UltimateSDUpscaleDistributed"],
)
prompt_copy = _override_branch_nodes(
prompt_copy,
prompt_index,
is_master,
participant_id,
enabled_worker_ids,
job_id_map,
delegate_master,
)
_override_branch_collector_nodes(
prompt_copy,
prompt_index,
is_master,
@@ -339,14 +668,9 @@ def apply_participant_overrides(
enabled_json,
delegate_master,
)
_override_upscale_nodes(
prompt_copy,
prompt_index,
is_master,
participant_id,
job_id_map,
master_url,
enabled_json,
)
if not is_master:
prompt_copy = _prune_worker_downstream_of_branch_collectors(prompt_copy)
prompt_copy = _ensure_worker_output_node(prompt_copy)
return prompt_copy
+240 -135
View File
@@ -1,8 +1,7 @@
import asyncio
import time
import uuid
import server
from typing import Any
from ..utils.async_helpers import queue_prompt_payload
from ..utils.config import load_config
@@ -14,10 +13,12 @@ from ..utils.constants import (
)
from ..utils.logging import debug_log, log
from ..utils.network import build_master_url
from ..utils.runtime_state import ensure_distributed_runtime_state, get_prompt_server_instance
from ..utils.trace_logger import trace_debug
from .schemas import parse_positive_float, parse_positive_int
from .schemas import coerce_positive_float, coerce_positive_int
from .orchestration.dispatch import (
dispatch_worker_prompt,
rank_workers_by_load,
select_active_workers,
select_least_busy_worker,
)
@@ -31,33 +32,21 @@ from .orchestration.prompt_transform import (
prune_prompt_for_worker,
)
prompt_server = server.PromptServer.instance
def _generate_execution_trace_id():
return f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
def ensure_distributed_state(server_instance=None):
"""Ensure prompt_server has the state used by distributed queue orchestration."""
ps = server_instance or prompt_server
if not hasattr(ps, "distributed_pending_jobs"):
ps.distributed_pending_jobs = {}
if not hasattr(ps, "distributed_jobs_lock"):
ps.distributed_jobs_lock = asyncio.Lock()
# Initialize top-level distributed queue state at module import time.
ensure_distributed_state()
ensure_distributed_runtime_state(server_instance)
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()
runtime_state = ensure_distributed_runtime_state()
async with runtime_state.distributed_jobs_lock:
if job_id not in runtime_state.distributed_pending_jobs:
runtime_state.distributed_pending_jobs[job_id] = asyncio.Queue()
def _resolve_enabled_workers(config, requested_ids=None):
@@ -96,19 +85,19 @@ def _resolve_enabled_workers(config, requested_ids=None):
def _resolve_orchestration_limits(config):
"""Resolve bounded concurrency/timeouts for worker preparation pipeline."""
settings = (config or {}).get("settings", {}) or {}
worker_probe_concurrency = parse_positive_int(
worker_probe_concurrency = coerce_positive_int(
settings.get("worker_probe_concurrency"),
ORCHESTRATION_WORKER_PROBE_CONCURRENCY,
)
worker_prep_concurrency = parse_positive_int(
worker_prep_concurrency = coerce_positive_int(
settings.get("worker_prep_concurrency"),
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
)
media_sync_concurrency = parse_positive_int(
media_sync_concurrency = coerce_positive_int(
settings.get("media_sync_concurrency"),
ORCHESTRATION_MEDIA_SYNC_CONCURRENCY,
)
media_sync_timeout_seconds = parse_positive_float(
media_sync_timeout_seconds = coerce_positive_float(
settings.get("media_sync_timeout_seconds"),
ORCHESTRATION_MEDIA_SYNC_TIMEOUT,
)
@@ -191,60 +180,16 @@ async def _prepare_worker_payload(
return worker, worker_prompt
async def orchestrate_distributed_execution(
prompt_obj,
workflow_meta,
client_id,
enabled_worker_ids=None,
delegate_master=None,
trace_execution_id=None,
async def _select_execution_workers(
workers,
use_websocket,
delegate_master,
load_balance_requested,
has_branch_nodes,
master_url,
execution_trace_id,
worker_probe_concurrency,
):
"""Core orchestration logic for the /distributed/queue endpoint.
Returns:
tuple[str, int]: (prompt_id, worker_count)
"""
ensure_distributed_state()
execution_trace_id = trace_execution_id or _generate_execution_trace_id()
config = load_config()
use_websocket = bool(config.get("settings", {}).get("websocket_orchestration", False))
master_url = build_master_url(config=config, prompt_server_instance=prompt_server)
(
worker_probe_concurrency,
worker_prep_concurrency,
media_sync_concurrency,
media_sync_timeout_seconds,
) = _resolve_orchestration_limits(config)
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)
load_balance_requested = _prompt_requests_load_balance(prompt_index)
trace_debug(
execution_trace_id,
(
f"Orchestration start: requested_workers={len(workers)}, "
f"requested_ids={requested_ids if requested_ids is not None else 'enabled_only'}, "
f"websocket={use_websocket}, "
f"probe_concurrency={worker_probe_concurrency}, "
f"prep_concurrency={worker_prep_concurrency}, "
f"media_sync_concurrency={media_sync_concurrency}, "
f"media_sync_timeout={media_sync_timeout_seconds:.1f}s, "
f"load_balance={load_balance_requested}"
),
)
# 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:
trace_debug(
execution_trace_id,
"Delegate-only requested but no workers are enabled. Falling back to master execution.",
)
delegate_master = False
active_workers, delegate_master = await select_active_workers(
workers,
use_websocket,
@@ -256,7 +201,6 @@ async def orchestrate_distributed_execution(
if load_balance_requested:
candidate_workers = list(active_workers)
if not delegate_master:
# Include master in load balancing only when master participation is enabled.
candidate_workers.append(
{
"id": "master",
@@ -281,7 +225,6 @@ async def orchestrate_distributed_execution(
selected_worker = candidate_workers[0]
if selected_worker is not None and str(selected_worker.get("id")) == "master":
# Master selected as least busy; run master workload only.
active_workers = []
delegate_master = False
trace_debug(
@@ -290,7 +233,6 @@ async def orchestrate_distributed_execution(
)
elif selected_worker is not None:
active_workers = [selected_worker]
# Worker selected as least busy; keep master orchestrator-only for this run.
delegate_master = True
trace_debug(
execution_trace_id,
@@ -304,19 +246,47 @@ async def orchestrate_distributed_execution(
active_workers = []
delegate_master = False
enabled_ids = [worker["id"] for worker in active_workers]
if has_branch_nodes and len(active_workers) > 1:
active_workers = await rank_workers_by_load(
active_workers,
trace_execution_id=execution_trace_id,
probe_concurrency=worker_probe_concurrency,
)
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
return active_workers, delegate_master
if not job_id_map:
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
prompt_id = await queue_prompt_payload(prompt_obj, workflow_meta, client_id)
return prompt_id, 0
def _distributed_queue_nodes(prompt_index):
return (
prompt_index.nodes_for_class("DistributedCollector")
+ prompt_index.nodes_for_class("DistributedBranchCollector")
+ prompt_index.nodes_for_class("DistributedBranch")
+ prompt_index.nodes_for_class("UltimateSDUpscaleDistributed")
)
def _register_job_allowed_workers(job_id_map, enabled_ids):
runtime_state = ensure_distributed_runtime_state()
allowed = {str(worker_id) for worker_id in (enabled_ids or []) if str(worker_id).strip()}
for job_id in job_id_map.values():
await _ensure_distributed_queue(job_id)
if job_id:
runtime_state.distributed_job_allowed_workers[str(job_id)] = set(allowed)
async def _ensure_job_queues(prompt_index, job_id_map):
for node_id in _distributed_queue_nodes(prompt_index):
job_id = job_id_map.get(node_id)
if job_id:
await _ensure_distributed_queue(job_id)
def _build_master_prompt(
prompt_index,
enabled_ids,
job_id_map,
master_url,
delegate_master,
):
master_prompt = prompt_index.copy_prompt()
master_prompt = apply_participant_overrides(
master_prompt,
@@ -328,19 +298,171 @@ async def orchestrate_distributed_execution(
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."
if not delegate_master:
return master_prompt
collector_ids = (
find_nodes_by_class(master_prompt, "DistributedCollector")
+ find_nodes_by_class(master_prompt, "DistributedBranchCollector")
)
upscale_nodes = find_nodes_by_class(master_prompt, "UltimateSDUpscaleDistributed")
# Include USDU nodes as collector-like for delegate pruning
collector_ids.extend(upscale_nodes)
if not collector_ids:
debug_log(
"Delegate-only master mode requested but no collector/branch-collector nodes found in master prompt. Running full prompt on master."
)
return master_prompt
return prepare_delegate_master_prompt(master_prompt, collector_ids)
async def _prepare_worker_payloads_for_dispatch(
active_workers,
prompt_index,
enabled_ids,
job_id_map,
master_url,
delegate_master,
execution_trace_id,
worker_prep_concurrency,
media_sync_concurrency,
media_sync_timeout_seconds,
):
if not active_workers:
return []
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
return await asyncio.gather(
*[
_prepare_worker_payload(
worker,
prompt_index,
enabled_ids,
job_id_map,
master_url,
delegate_master,
execution_trace_id,
worker_prep_semaphore,
media_sync_semaphore,
media_sync_timeout_seconds,
)
elif not collector_ids:
debug_log(
"Delegate-only master mode requested but no collectors found in master prompt. Running full prompt on master."
for worker in active_workers
]
)
async def _dispatch_worker_payloads(
worker_payloads,
workflow_meta,
client_id,
use_websocket,
execution_trace_id,
):
if not worker_payloads:
return
await asyncio.gather(
*[
dispatch_worker_prompt(
worker,
wprompt,
workflow_meta,
client_id,
use_websocket=use_websocket,
trace_execution_id=execution_trace_id,
)
else:
master_prompt = prepare_delegate_master_prompt(master_prompt, collector_ids)
for worker, wprompt in worker_payloads
]
)
async def orchestrate_distributed_execution(
prompt_obj: dict[str, Any],
workflow_meta: dict[str, Any] | None,
client_id: str | None,
enabled_worker_ids: list[str] | set[str] | None = None,
delegate_master: bool | None = None,
trace_execution_id: str | None = None,
) -> tuple[str, int]:
"""Core orchestration logic for the /distributed/queue endpoint.
Returns:
tuple[str, int]: (prompt_id, worker_count)
"""
ensure_distributed_state()
execution_trace_id = trace_execution_id or _generate_execution_trace_id()
config = load_config()
use_websocket = bool(config.get("settings", {}).get("websocket_orchestration", False))
master_url = build_master_url(config=config, prompt_server_instance=get_prompt_server_instance())
(
worker_probe_concurrency,
worker_prep_concurrency,
media_sync_concurrency,
media_sync_timeout_seconds,
) = _resolve_orchestration_limits(config)
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)
load_balance_requested = _prompt_requests_load_balance(prompt_index)
has_branch_nodes = bool(prompt_index.nodes_for_class("DistributedBranch"))
trace_debug(
execution_trace_id,
(
f"Orchestration start: requested_workers={len(workers)}, "
f"requested_ids={requested_ids if requested_ids is not None else 'enabled_only'}, "
f"websocket={use_websocket}, "
f"probe_concurrency={worker_probe_concurrency}, "
f"prep_concurrency={worker_prep_concurrency}, "
f"media_sync_concurrency={media_sync_concurrency}, "
f"media_sync_timeout={media_sync_timeout_seconds:.1f}s, "
f"load_balance={load_balance_requested}, has_branch_nodes={has_branch_nodes}"
),
)
# 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:
trace_debug(
execution_trace_id,
"Delegate-only requested but no workers are enabled. Falling back to master execution.",
)
delegate_master = False
active_workers, delegate_master = await _select_execution_workers(
workers=workers,
use_websocket=use_websocket,
delegate_master=delegate_master,
load_balance_requested=load_balance_requested,
has_branch_nodes=has_branch_nodes,
master_url=master_url,
execution_trace_id=execution_trace_id,
worker_probe_concurrency=worker_probe_concurrency,
)
enabled_ids = [worker["id"] for worker in active_workers]
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
if not job_id_map:
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
prompt_id = await queue_prompt_payload(prompt_obj, workflow_meta, client_id)
return prompt_id, 0
_register_job_allowed_workers(job_id_map, enabled_ids)
await _ensure_job_queues(prompt_index, job_id_map)
master_prompt = _build_master_prompt(
prompt_index=prompt_index,
enabled_ids=enabled_ids,
job_id_map=job_id_map,
master_url=master_url,
delegate_master=delegate_master,
)
if active_workers:
trace_debug(
@@ -348,42 +470,25 @@ async def orchestrate_distributed_execution(
"Active distributed workers: "
+ ", ".join(f"{worker['name']} ({worker['id']})" for worker in active_workers),
)
worker_payloads = []
if active_workers:
worker_prep_semaphore = asyncio.Semaphore(worker_prep_concurrency)
media_sync_semaphore = asyncio.Semaphore(media_sync_concurrency)
worker_payloads = await asyncio.gather(
*[
_prepare_worker_payload(
worker,
prompt_index,
enabled_ids,
job_id_map,
master_url,
delegate_master,
execution_trace_id,
worker_prep_semaphore,
media_sync_semaphore,
media_sync_timeout_seconds,
)
for worker in active_workers
]
)
if worker_payloads:
await asyncio.gather(
*[
dispatch_worker_prompt(
worker,
wprompt,
workflow_meta,
client_id,
use_websocket=use_websocket,
trace_execution_id=execution_trace_id,
)
for worker, wprompt in worker_payloads
]
)
worker_payloads = await _prepare_worker_payloads_for_dispatch(
active_workers=active_workers,
prompt_index=prompt_index,
enabled_ids=enabled_ids,
job_id_map=job_id_map,
master_url=master_url,
delegate_master=delegate_master,
execution_trace_id=execution_trace_id,
worker_prep_concurrency=worker_prep_concurrency,
media_sync_concurrency=media_sync_concurrency,
media_sync_timeout_seconds=media_sync_timeout_seconds,
)
await _dispatch_worker_payloads(
worker_payloads=worker_payloads,
workflow_meta=workflow_meta,
client_id=client_id,
use_websocket=use_websocket,
execution_trace_id=execution_trace_id,
)
prompt_id = await queue_prompt_payload(master_prompt, workflow_meta, client_id)
trace_debug(
+20 -15
View File
@@ -1,16 +1,20 @@
from dataclasses import dataclass
from typing import Any, Dict, List, Optional
from typing import Any
try:
from ..utils.worker_ids import require_enabled_worker_ids
except ImportError: # pragma: no cover - supports direct module loading in isolated tests.
from utils.worker_ids import require_enabled_worker_ids
@dataclass(frozen=True)
class QueueRequestPayload:
prompt: Dict[str, Any]
workflow_meta: Any
prompt: dict[str, Any]
workflow_meta: dict[str, Any] | None
client_id: str
delegate_master: Optional[bool]
enabled_worker_ids: List[str]
auto_prepare: bool
trace_execution_id: Optional[str]
delegate_master: bool | None
enabled_worker_ids: list[str]
trace_execution_id: str | None
def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
@@ -18,11 +22,6 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
if not isinstance(data, dict):
raise ValueError("Expected a JSON object body")
auto_prepare_raw = data.get("auto_prepare", True)
if not isinstance(auto_prepare_raw, bool):
raise ValueError("auto_prepare must be a boolean when provided")
auto_prepare = auto_prepare_raw
prompt = data.get("prompt")
# Auto-prepare is always on server-side; keep the field for wire compatibility.
if prompt is None:
@@ -35,6 +34,10 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
if not isinstance(prompt, dict):
raise ValueError("Field 'prompt' must be an object")
workflow_meta = data.get("workflow")
if workflow_meta is not None and not isinstance(workflow_meta, dict):
raise ValueError("Field 'workflow' must be an object when provided")
enabled_ids_raw = data.get("enabled_worker_ids")
workers_field = data.get("workers")
if enabled_ids_raw is None and workers_field is not None:
@@ -51,7 +54,10 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
else:
if not isinstance(enabled_ids_raw, list):
raise ValueError("enabled_worker_ids must be a list of worker IDs")
enabled_ids = [str(worker_id).strip() for worker_id in enabled_ids_raw if str(worker_id).strip()]
try:
enabled_ids = require_enabled_worker_ids(enabled_ids_raw)
except ValueError as exc:
raise ValueError(f"enabled_worker_ids invalid: {exc}") from exc
delegate_master = data.get("delegate_master")
if delegate_master is not None and not isinstance(delegate_master, bool):
@@ -70,10 +76,9 @@ def parse_queue_request_payload(data: Any) -> QueueRequestPayload:
return QueueRequestPayload(
prompt=prompt,
workflow_meta=data.get("workflow"),
workflow_meta=workflow_meta,
client_id=client_id,
delegate_master=delegate_master,
enabled_worker_ids=enabled_ids,
auto_prepare=auto_prepare,
trace_execution_id=trace_execution_id,
)
+15
View File
@@ -0,0 +1,15 @@
"""Shared request guards for distributed API endpoints."""
from __future__ import annotations
from aiohttp import web
from ..utils.config import load_config
from ..utils.network import handle_api_error
from .schemas import is_authorized_request
async def authorization_error_or_none(request: web.Request) -> web.StreamResponse | None:
"""Return a 403 response when distributed API auth fails, else None."""
if is_authorized_request(request, load_config()):
return None
return await handle_api_error(request, "Unauthorized", 403)
+61 -15
View File
@@ -1,3 +1,29 @@
from __future__ import annotations
from typing import Any
from ..utils.auth import (
AUTH_HEADER_NAME,
distributed_auth_headers,
is_authorized_request,
)
from ..utils.parsing import coerce_bool, require_bool_literal
__all__ = [
"AUTH_HEADER_NAME",
"coerce_bool",
"coerce_positive_float",
"coerce_positive_int",
"distributed_auth_headers",
"is_authorized_request",
"require_bool_literal",
"require_fields",
"require_worker_id",
"require_positive_int",
"validate_worker_id",
]
def require_fields(data: dict, *fields) -> list[str]:
"""Return field names that are missing or empty in a JSON object."""
if not isinstance(data, dict):
@@ -18,26 +44,43 @@ def require_fields(data: dict, *fields) -> list[str]:
return missing
def validate_worker_id(worker_id: str, config: dict) -> bool:
def validate_worker_id(worker_id: str, config: dict[str, Any]) -> bool:
"""Return True when worker_id exists in config['workers']."""
worker_id_str = str(worker_id)
try:
require_worker_id(worker_id, config)
return True
except ValueError:
return False
def require_worker_id(worker_id: str, config: dict[str, Any], field_name: str = "worker_id") -> str:
"""Strict worker-id validator for request boundaries."""
worker_id_str = str(worker_id).strip()
if not worker_id_str:
raise ValueError(f"Field '{field_name}' must be a non-empty worker id.")
workers = (config or {}).get("workers", [])
return any(str(worker.get("id")) == worker_id_str for worker in workers)
for worker in workers:
if not isinstance(worker, dict):
continue
if str(worker.get("id")).strip() == worker_id_str:
return worker_id_str
raise ValueError(f"Worker {worker_id_str} not found.")
def validate_positive_int(value, field_name: str) -> str | None:
"""Validate positive integers and return an error string when invalid."""
def require_positive_int(value, field_name: str = "value") -> int:
"""Strict positive-int parser for request boundary validation."""
try:
parsed = int(value)
except (TypeError, ValueError):
return f"Field '{field_name}' must be a positive integer."
except (TypeError, ValueError) as exc:
raise ValueError(f"Field '{field_name}' must be a positive integer.") from exc
if parsed <= 0:
return f"Field '{field_name}' must be a positive integer."
return None
raise ValueError(f"Field '{field_name}' must be a positive integer.")
return parsed
def parse_positive_int(value, default: int) -> int:
"""Parse value as positive int, returning default on failure."""
def coerce_positive_int(value, default: int) -> int:
"""Permissive positive-int coercion used for non-boundary defaults."""
try:
parsed = int(value)
except (TypeError, ValueError):
@@ -45,10 +88,13 @@ def parse_positive_int(value, default: int) -> int:
return max(1, parsed)
def parse_positive_float(value, default: float) -> float:
"""Parse value as positive float, returning default on failure."""
def coerce_positive_float(value, default: float) -> float:
"""Permissive positive-float coercion used for non-boundary defaults."""
fallback = float(default)
if fallback <= 0.0:
fallback = 0.1
try:
parsed = float(value)
except (TypeError, ValueError):
return max(0.0, float(default))
return max(0.0, parsed)
return fallback
return parsed if parsed > 0.0 else fallback
+41 -29
View File
@@ -1,51 +1,63 @@
from __future__ import annotations
from aiohttp import web
import server
from .endpoint_policy import run_authorized_endpoint
from ..utils.cloudflare import cloudflare_tunnel_manager
from ..utils.config import load_config
from ..utils.logging import debug_log, log
from ..utils.network import handle_api_error
@server.PromptServer.instance.routes.get("/distributed/tunnel/status")
async def tunnel_status_endpoint(request):
async def tunnel_status_endpoint(request: web.Request) -> web.StreamResponse:
"""Return Cloudflare tunnel status and last known details."""
try:
async def _operation() -> web.StreamResponse:
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)
return web.json_response(
{
"status": "success",
"tunnel": status,
"master_host": master_host,
}
)
return await run_authorized_endpoint(request, _operation)
@server.PromptServer.instance.routes.post("/distributed/tunnel/start")
async def tunnel_start_endpoint(request):
async def tunnel_start_endpoint(request: web.Request) -> web.StreamResponse:
"""Start a Cloudflare tunnel pointing at the current ComfyUI server."""
try:
async def _operation() -> web.StreamResponse:
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)
return web.json_response(
{
"status": "success",
"tunnel": result,
"master_host": (config.get("master") or {}).get("host"),
}
)
return await run_authorized_endpoint(request, _operation)
@server.PromptServer.instance.routes.post("/distributed/tunnel/stop")
async def tunnel_stop_endpoint(request):
async def tunnel_stop_endpoint(request: web.Request) -> web.StreamResponse:
"""Stop the managed Cloudflare tunnel if running."""
try:
async def _operation() -> web.StreamResponse:
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)
return web.json_response(
{
"status": "success",
"tunnel": result,
"master_host": (config.get("master") or {}).get("host"),
}
)
return await run_authorized_endpoint(request, _operation)
+137 -21
View File
@@ -1,3 +1,5 @@
from __future__ import annotations
import asyncio
import io
import time
@@ -6,15 +8,78 @@ from aiohttp import web
from PIL import Image
import server
from .request_guards import authorization_error_or_none
from .schemas import require_bool_literal
from ..upscale.job_models import BaseJobState, ImageJobState, TileJobState
from ..upscale.job_store import MAX_PAYLOAD_SIZE, ensure_tile_jobs_initialized
from ..upscale.payload_parsers import _parse_tiles_from_form
from ..upscale.job_store import ensure_tile_jobs_initialized, init_dynamic_job
from ..upscale.payload_parsers import parse_tiles_from_form
from ..utils.logging import debug_log
from ..utils.network import handle_api_error
from ..utils.usdu_management import MAX_PAYLOAD_SIZE
def _parse_int_field(value, field_name: str, *, minimum: int | None = None) -> int:
try:
parsed = int(value)
except (TypeError, ValueError) as exc:
raise ValueError(f"{field_name} must be an integer") from exc
if minimum is not None and parsed < minimum:
raise ValueError(f"{field_name} must be >= {minimum}")
return parsed
def _parse_bool_field(
value,
field_name: str,
*,
default: bool = False,
) -> bool:
if value is None:
return default
return require_bool_literal(value, field_name=field_name)
def _is_worker_allowed(job_data: BaseJobState, worker_id: str) -> bool:
known_workers = {str(wid) for wid in getattr(job_data, "worker_status", {}).keys()}
known_workers.update(str(wid) for wid in getattr(job_data, "assigned_to_workers", {}).keys())
if not known_workers:
return True
return str(worker_id) in known_workers
@server.PromptServer.instance.routes.post("/distributed/init_dynamic_job")
async def init_dynamic_job_endpoint(request: web.Request) -> web.StreamResponse:
"""Allow workers to initialize a dynamic job queue on the master.
Called by the first worker to reach the USDU node in delegate-only mode.
Idempotent — subsequent calls are no-ops if the queue already exists.
"""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
multi_job_id = data.get("multi_job_id")
if not multi_job_id:
return await handle_api_error(request, "Missing multi_job_id", 400)
batch_size = _parse_int_field(data.get("batch_size"), "batch_size", minimum=1)
enabled_workers = data.get("enabled_workers") or []
if isinstance(enabled_workers, str):
import json as _json
enabled_workers = _json.loads(enabled_workers)
await init_dynamic_job(multi_job_id, batch_size, enabled_workers)
return web.json_response({"status": "success"})
except Exception as e:
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.post("/distributed/heartbeat")
async def heartbeat_endpoint(request):
async def heartbeat_endpoint(request: web.Request) -> web.StreamResponse:
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
worker_id = data.get('worker_id')
@@ -28,6 +93,8 @@ async def heartbeat_endpoint(request):
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if isinstance(job_data, BaseJobState):
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
job_data.worker_status[worker_id] = time.time()
debug_log(f"Heartbeat from worker {worker_id}")
return web.json_response({"status": "success"})
@@ -38,24 +105,38 @@ async def heartbeat_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/submit_tiles")
async def submit_tiles_endpoint(request):
async def submit_tiles_endpoint(request: web.Request) -> web.StreamResponse:
"""Endpoint for workers to submit processed tiles in static mode."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
content_length = request.headers.get('content-length')
if content_length and int(content_length) > MAX_PAYLOAD_SIZE:
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
if content_length:
try:
payload_size = int(content_length)
except ValueError:
return await handle_api_error(request, "Invalid content-length header", 400)
if payload_size > MAX_PAYLOAD_SIZE:
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
data = await request.post()
multi_job_id = data.get('multi_job_id')
worker_id = data.get('worker_id')
is_last = data.get('is_last', 'False').lower() == 'true'
try:
is_last = _parse_bool_field(data.get('is_last', False), "is_last", default=False)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
if multi_job_id is None or worker_id is None:
return await handle_api_error(request, "Missing multi_job_id or worker_id", 400)
prompt_server = ensure_tile_jobs_initialized()
batch_size = int(data.get('batch_size', 0))
try:
batch_size = _parse_int_field(data.get('batch_size', 0), "batch_size", minimum=0)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
# Handle completion signal
if batch_size == 0 and is_last:
@@ -64,6 +145,8 @@ async def submit_tiles_endpoint(request):
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if not isinstance(job_data, TileJobState):
return await handle_api_error(request, "Job not configured for tile submissions", 400)
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
await job_data.queue.put({
'worker_id': worker_id,
'is_last': True,
@@ -73,7 +156,7 @@ async def submit_tiles_endpoint(request):
return web.json_response({"status": "success"})
try:
tiles = _parse_tiles_from_form(data)
tiles = parse_tiles_from_form(data)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
@@ -83,6 +166,8 @@ async def submit_tiles_endpoint(request):
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if not isinstance(job_data, TileJobState):
return await handle_api_error(request, "Job not configured for tile submissions", 400)
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
q = job_data.queue
if batch_size > 0 or len(tiles) > 0:
@@ -106,17 +191,28 @@ async def submit_tiles_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/submit_image")
async def submit_image_endpoint(request):
async def submit_image_endpoint(request: web.Request) -> web.StreamResponse:
"""Endpoint for workers to submit processed images in dynamic mode."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
content_length = request.headers.get('content-length')
if content_length and int(content_length) > MAX_PAYLOAD_SIZE:
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
if content_length:
try:
payload_size = int(content_length)
except ValueError:
return await handle_api_error(request, "Invalid content-length header", 400)
if payload_size > MAX_PAYLOAD_SIZE:
return await handle_api_error(request, f"Payload too large: {content_length} bytes", 413)
data = await request.post()
multi_job_id = data.get('multi_job_id')
worker_id = data.get('worker_id')
is_last = data.get('is_last', 'False').lower() == 'true'
try:
is_last = _parse_bool_field(data.get('is_last', False), "is_last", default=False)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
if multi_job_id is None or worker_id is None:
return await handle_api_error(request, "Missing multi_job_id or worker_id", 400)
@@ -125,7 +221,10 @@ async def submit_image_endpoint(request):
# Handle image submission
if 'full_image' in data and 'image_idx' in data:
image_idx = int(data.get('image_idx'))
try:
image_idx = _parse_int_field(data.get('image_idx'), "image_idx", minimum=0)
except ValueError as e:
return await handle_api_error(request, str(e), 400)
img_data = data['full_image'].file.read()
img = Image.open(io.BytesIO(img_data)).convert("RGB")
@@ -136,6 +235,8 @@ async def submit_image_endpoint(request):
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if not isinstance(job_data, ImageJobState):
return await handle_api_error(request, "Job not configured for image submissions", 400)
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
await job_data.queue.put({
'worker_id': worker_id,
'image_idx': image_idx,
@@ -151,6 +252,8 @@ async def submit_image_endpoint(request):
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if not isinstance(job_data, ImageJobState):
return await handle_api_error(request, "Job not configured for image submissions", 400)
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
await job_data.queue.put({
'worker_id': worker_id,
'is_last': True,
@@ -166,8 +269,11 @@ async def submit_image_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/request_image")
async def request_image_endpoint(request):
async def request_image_endpoint(request: web.Request) -> web.StreamResponse:
"""Endpoint for workers to request tasks (images in dynamic mode, tiles in static mode)."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
data = await request.json()
worker_id = data.get('worker_id')
@@ -190,6 +296,8 @@ async def request_image_endpoint(request):
pending_queue = job_data.pending_tasks
else:
return await handle_api_error(request, "Invalid job configuration", 400)
if not _is_worker_allowed(job_data, worker_id):
return await handle_api_error(request, "Unauthorized worker_id", 403)
try:
task_idx = await asyncio.wait_for(pending_queue.get(), timeout=0.1)
@@ -199,25 +307,33 @@ async def request_image_endpoint(request):
if mode == 'dynamic':
debug_log(f"UltimateSDUpscale API - Assigned image {task_idx} to worker {worker_id}")
return web.json_response({"image_idx": task_idx, "estimated_remaining": remaining})
return web.json_response(
{
"kind": "image",
"task_idx": task_idx,
"estimated_remaining": remaining,
}
)
debug_log(f"UltimateSDUpscale API - Assigned tile {task_idx} to worker {worker_id}")
return web.json_response({
"tile_idx": task_idx,
"kind": "tile",
"task_idx": task_idx,
"estimated_remaining": remaining,
"batched_static": job_data.batched_static,
})
except asyncio.TimeoutError:
if mode == 'dynamic':
return web.json_response({"image_idx": None})
return web.json_response({"tile_idx": None})
return web.json_response({"kind": "none", "task_idx": None, "estimated_remaining": 0})
return await handle_api_error(request, "Job not found", 404)
except Exception as e:
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.get("/distributed/job_status")
async def job_status_endpoint(request):
async def job_status_endpoint(request: web.Request) -> web.StreamResponse:
"""Endpoint to check if a job is ready."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
multi_job_id = request.query.get('multi_job_id')
if not multi_job_id:
return web.json_response({"ready": False})
+108 -47
View File
@@ -1,10 +1,13 @@
from __future__ import annotations
import json
import asyncio
import os
import time
import platform
import subprocess
import subprocess # nosec B404 - commands are fixed and never shell-expanded
import socket
from typing import Any
import torch
import aiohttp
@@ -22,27 +25,27 @@ from ..utils.network import (
)
from ..utils.constants import CHUNK_SIZE
from ..workers import get_worker_manager
from .schemas import require_fields, validate_worker_id
from .request_guards import authorization_error_or_none
from .schemas import (
distributed_auth_headers,
require_fields,
require_worker_id,
)
from ..workers.detection import (
get_machine_id,
is_docker_environment,
is_runpod_environment,
)
try:
from ..utils.async_helpers import PromptValidationError, queue_prompt_payload
except ImportError:
from ..utils.async_helpers import queue_prompt_payload
class PromptValidationError(RuntimeError):
def __init__(self, message, validation_error=None, node_errors=None):
super().__init__(str(message))
self.validation_error = validation_error if isinstance(validation_error, dict) else {}
self.node_errors = node_errors if isinstance(node_errors, dict) else {}
from ..utils.async_helpers import PromptValidationError, queue_prompt_payload
@server.PromptServer.instance.routes.get("/distributed/worker_ws")
async def worker_ws_endpoint(request):
async def worker_ws_endpoint(request: web.Request) -> web.StreamResponse:
"""WebSocket endpoint for worker prompt dispatch."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
ws = web.WebSocketResponse(heartbeat=30)
await ws.prepare(request)
@@ -113,8 +116,11 @@ async def worker_ws_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/worker/clear_launching")
async def clear_launching_state(request):
async def clear_launching_state(request: web.Request) -> web.StreamResponse:
"""Clear the launching flag when worker is confirmed running."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
wm = get_worker_manager()
data = await request.json()
@@ -124,8 +130,10 @@ async def clear_launching_state(request):
worker_id = str(data.get("worker_id")).strip()
config = load_config()
if not validate_worker_id(worker_id, config):
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
try:
worker_id = require_worker_id(worker_id, config)
except ValueError as exc:
return await handle_api_error(request, exc, 404)
# Clear launching flag in managed processes
if worker_id in wm.processes:
@@ -139,8 +147,9 @@ async def clear_launching_state(request):
return await handle_api_error(request, e, 500)
def get_network_ips():
def get_network_ips() -> list[str]:
"""Get all network IPs, trying multiple methods."""
command_timeout = 5.0
ips = []
hostname = socket.gethostname()
@@ -151,8 +160,8 @@ def get_network_ips():
ip = info[4][0]
if ip and ip not in ips and not ip.startswith('::'): # Skip IPv6 for now
ips.append(ip)
except (socket.gaierror, OSError):
pass
except (socket.gaierror, OSError) as exc:
debug_log(f"get_network_ips: getaddrinfo failed for hostname {hostname}: {exc}")
# Method 2: Try to connect to external server and get local IP
try:
@@ -162,14 +171,19 @@ def get_network_ips():
s.close()
if local_ip not in ips:
ips.append(local_ip)
except (OSError, socket.error):
pass
except (OSError, socket.error) as exc:
debug_log(f"get_network_ips: UDP local IP probe failed: {exc}")
# Method 3: Platform-specific commands
try:
if platform.system() == "Windows":
# Windows ipconfig
result = subprocess.run(["ipconfig"], capture_output=True, text=True)
result = subprocess.run( # nosec B603 - static command, no user input
["ipconfig"],
capture_output=True,
text=True,
timeout=command_timeout,
)
lines = result.stdout.split('\n')
for i, line in enumerate(lines):
if 'IPv4' in line and i + 1 < len(lines):
@@ -179,10 +193,20 @@ def get_network_ips():
else:
# Unix/Linux/Mac ifconfig or ip addr
try:
result = subprocess.run(["ip", "addr"], capture_output=True, text=True)
result = subprocess.run( # nosec B603 - static command, no user input
["ip", "addr"],
capture_output=True,
text=True,
timeout=command_timeout,
)
except (FileNotFoundError, OSError):
try:
result = subprocess.run(["ifconfig"], capture_output=True, text=True)
result = subprocess.run( # nosec B603 - static command, no user input
["ifconfig"],
capture_output=True,
text=True,
timeout=command_timeout,
)
except (FileNotFoundError, OSError):
result = None
@@ -193,13 +217,13 @@ def get_network_ips():
ip = match.group(1)
if ip and ip not in ips:
ips.append(ip)
except (OSError, subprocess.SubprocessError):
pass
except (OSError, subprocess.SubprocessError) as exc:
debug_log(f"get_network_ips: platform command probe failed: {exc}")
return ips
def get_recommended_ip(ips):
def get_recommended_ip(ips: list[str]) -> str | None:
"""Choose the best IP for master-worker communication."""
# Priority order:
# 1. Private network ranges (192.168.x.x, 10.x.x.x, 172.16-31.x.x)
@@ -234,7 +258,7 @@ def get_recommended_ip(ips):
return None
def _get_cuda_info():
def _get_cuda_info() -> tuple[int | None, int, int]:
"""Detect CUDA device index and total physical GPU count.
Returns (cuda_device, cuda_device_count, physical_device_count).
@@ -250,7 +274,7 @@ def _get_cuda_info():
if visible_devices:
cuda_device = visible_devices[0]
try:
result = subprocess.run(
result = subprocess.run( # nosec B603 - static command, no user input
['nvidia-smi', '--query-gpu=name', '--format=csv,noheader'],
capture_output=True,
text=True,
@@ -274,7 +298,7 @@ def _get_cuda_info():
return None, 0, 0
def _collect_network_info_sync():
def _collect_network_info_sync() -> dict[str, Any]:
"""Collect network/cuda info in a worker thread to avoid blocking route handlers."""
cuda_device, cuda_device_count, physical_device_count = _get_cuda_info()
hostname = socket.gethostname()
@@ -289,7 +313,7 @@ def _collect_network_info_sync():
}
def _read_worker_log_sync(log_file, lines_to_read):
def _read_worker_log_sync(log_file: str, lines_to_read: int) -> dict[str, Any]:
"""Read worker log content from disk in a threadpool worker."""
file_size = os.path.getsize(log_file)
@@ -325,7 +349,12 @@ def _read_worker_log_sync(log_file, lines_to_read):
}
def _parse_positive_int_query(value, default, minimum=1, maximum=10000):
def _parse_positive_int_query(
value: Any,
default: int,
minimum: int = 1,
maximum: int | None = 10000,
) -> int:
"""Parse bounded positive integer query params with sane fallback."""
try:
parsed = int(value)
@@ -337,7 +366,7 @@ def _parse_positive_int_query(value, default, minimum=1, maximum=10000):
return parsed
def _find_worker_by_id(config, worker_id):
def _find_worker_by_id(config: dict[str, Any], worker_id: str) -> dict[str, Any] | None:
worker_id_str = str(worker_id).strip()
for worker in config.get("workers", []):
if str(worker.get("id")).strip() == worker_id_str:
@@ -346,8 +375,11 @@ def _find_worker_by_id(config, worker_id):
@server.PromptServer.instance.routes.get("/distributed/local_log")
async def get_local_log_endpoint(request):
async def get_local_log_endpoint(request: web.Request) -> web.StreamResponse:
"""Return this instance's in-memory ComfyUI log buffer."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
from app.logger import get_logs
except Exception as e:
@@ -391,8 +423,11 @@ async def get_local_log_endpoint(request):
@server.PromptServer.instance.routes.get("/distributed/network_info")
async def get_network_info_endpoint(request):
async def get_network_info_endpoint(request: web.Request) -> web.StreamResponse:
"""Get network interfaces and recommend best IP for master."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
loop = asyncio.get_running_loop()
info = await loop.run_in_executor(None, _collect_network_info_sync)
@@ -406,8 +441,11 @@ async def get_network_info_endpoint(request):
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.get("/distributed/system_info")
async def get_system_info_endpoint(request):
async def get_system_info_endpoint(request: web.Request) -> web.StreamResponse:
"""Get system information including machine ID for local worker detection."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
import socket
@@ -430,8 +468,11 @@ async def get_system_info_endpoint(request):
return await handle_api_error(request, e, 500)
@server.PromptServer.instance.routes.post("/distributed/launch_worker")
async def launch_worker_endpoint(request):
async def launch_worker_endpoint(request: web.Request) -> web.StreamResponse:
"""Launch a worker process from the UI."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
wm = get_worker_manager()
data = await request.json()
@@ -443,8 +484,10 @@ async def launch_worker_endpoint(request):
# Find worker config
config = load_config()
if not validate_worker_id(worker_id, config):
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
try:
worker_id = require_worker_id(worker_id, config)
except ValueError as exc:
return await handle_api_error(request, exc, 404)
worker = next((w for w in config.get("workers", []) if str(w.get("id")) == worker_id), None)
if not worker:
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
@@ -463,7 +506,7 @@ async def launch_worker_endpoint(request):
is_running = process.poll() is None
else:
# Restored process without subprocess object
is_running = wm._is_process_running(proc_info['pid'])
is_running = wm.is_process_running(proc_info['pid'])
if is_running:
return await handle_api_error(request, "Worker already running (managed by UI)", 409)
@@ -491,8 +534,11 @@ async def launch_worker_endpoint(request):
@server.PromptServer.instance.routes.post("/distributed/stop_worker")
async def stop_worker_endpoint(request):
async def stop_worker_endpoint(request: web.Request) -> web.StreamResponse:
"""Stop a worker process that was launched from the UI."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
wm = get_worker_manager()
data = await request.json()
@@ -502,8 +548,10 @@ async def stop_worker_endpoint(request):
worker_id = str(data.get("worker_id")).strip()
config = load_config()
if not validate_worker_id(worker_id, config):
return await handle_api_error(request, f"Worker {worker_id} not found", 404)
try:
worker_id = require_worker_id(worker_id, config)
except ValueError as exc:
return await handle_api_error(request, exc, 404)
success, message = wm.stop_worker(worker_id)
@@ -521,8 +569,11 @@ async def stop_worker_endpoint(request):
@server.PromptServer.instance.routes.get("/distributed/managed_workers")
async def get_managed_workers_endpoint(request):
async def get_managed_workers_endpoint(request: web.Request) -> web.StreamResponse:
"""Get list of workers managed by this UI instance."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
managed = get_worker_manager().get_managed_workers()
return web.json_response({
@@ -534,8 +585,11 @@ async def get_managed_workers_endpoint(request):
@server.PromptServer.instance.routes.get("/distributed/local-worker-status")
async def get_local_worker_status_endpoint(request):
async def get_local_worker_status_endpoint(request: web.Request) -> web.StreamResponse:
"""Check status of all local workers (localhost/no host specified)."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
config = load_config()
worker_statuses = {}
@@ -604,8 +658,11 @@ async def get_local_worker_status_endpoint(request):
@server.PromptServer.instance.routes.get("/distributed/worker_log/{worker_id}")
async def get_worker_log_endpoint(request):
async def get_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
"""Get log content for a specific worker."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
wm = get_worker_manager()
worker_id = request.match_info['worker_id']
@@ -647,8 +704,11 @@ async def get_worker_log_endpoint(request):
@server.PromptServer.instance.routes.get("/distributed/remote_worker_log/{worker_id}")
async def get_remote_worker_log_endpoint(request):
async def get_remote_worker_log_endpoint(request: web.Request) -> web.StreamResponse:
"""Proxy a remote worker log request to the worker's local in-memory log endpoint."""
auth_error = await authorization_error_or_none(request)
if auth_error is not None:
return auth_error
try:
worker_id = str(request.match_info["worker_id"]).strip()
config = load_config()
@@ -671,6 +731,7 @@ async def get_remote_worker_log_endpoint(request):
async with session.get(
worker_url,
params={"lines": str(lines_to_read)},
headers=distributed_auth_headers(config),
timeout=aiohttp.ClientTimeout(total=5),
) as resp:
if resp.status >= 400:
+5
View File
@@ -0,0 +1,5 @@
"""Bootstrap entrypoints for ComfyUI-Distributed."""
from .entrypoint import build_node_mappings, initialize_runtime
__all__ = ["build_node_mappings", "initialize_runtime"]
+66
View File
@@ -0,0 +1,66 @@
"""Single startup owner for node surface and runtime initialization."""
from __future__ import annotations
import atexit
import os
from functools import lru_cache
from typing import Any
@lru_cache(maxsize=1)
def build_node_mappings() -> tuple[dict[str, Any], dict[str, str]]:
"""Build and cache the full exported node mapping surface."""
from ..nodes import (
NODE_CLASS_MAPPINGS as distributed_class_mappings,
NODE_DISPLAY_NAME_MAPPINGS as distributed_display_name_mappings,
)
from ..nodes.distributed_upscale import (
NODE_CLASS_MAPPINGS as upscale_class_mappings,
NODE_DISPLAY_NAME_MAPPINGS as upscale_display_name_mappings,
)
return (
{**distributed_class_mappings, **upscale_class_mappings},
{**distributed_display_name_mappings, **upscale_display_name_mappings},
)
@lru_cache(maxsize=1)
def _runtime_initialized() -> dict[str, bool]:
return {"done": False}
def initialize_runtime(prompt_server: Any | None = None) -> None:
"""Initialize distributed runtime side effects once."""
state = _runtime_initialized()
if state["done"]:
return
import server
from ..api import bootstrap_routes
from ..api.queue_orchestration import ensure_distributed_state
from ..upscale.job_store import ensure_tile_jobs_initialized
from ..utils.config import CONFIG_FILE, ensure_config_exists
from ..utils.logging import debug_log
from ..workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
ensure_config_exists()
bootstrap_routes()
server_instance = prompt_server if prompt_server is not None else server.PromptServer.instance
ensure_distributed_state(server_instance)
ensure_tile_jobs_initialized()
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
if not os.environ.get("COMFYUI_IS_WORKER"):
atexit.register(sync_cleanup)
delayed_auto_launch()
register_async_signals()
state["done"] = True
node_mappings, _ = build_node_mappings()
debug_log("Loaded Distributed nodes.")
debug_log(f"Config file: {CONFIG_FILE}")
debug_log(f"Available nodes: {list(node_mappings.keys())}")
+1 -25
View File
@@ -1,28 +1,4 @@
# conftest.py — project-level pytest configuration.
#
# Problem: custom_nodes/ComfyUI-Distributed/__init__.py uses relative imports
# (from .distributed import ...) that fail when pytest tries to import it as a
# standalone module during Package.setup() for the root package node.
#
# Fix: patch Package.setup() to skip the root-package's __init__.py import.
# All actual package context is provided by each test module via
# importlib.util.spec_from_file_location with synthetic stub packages.
from _pytest.python import Package
_orig_pkg_setup = Package.setup
def _patched_pkg_setup(self) -> None:
# Skip the root package setup — its __init__.py uses relative imports
# that require a parent package (ComfyUI's plugin loader) which is not
# available in the test environment.
if self.path == self.config.rootpath:
return
_orig_pkg_setup(self)
Package.setup = _patched_pkg_setup
"""Project-level pytest collection rules for plugin-style package layout."""
collect_ignore = [
"__init__.py",
+11 -47
View File
@@ -1,51 +1,15 @@
"""
ComfyUI-Distributed: thin entry point.
All implementation lives in workers/, nodes/, api/.
"""
import atexit
import os
"""Compatibility facade for legacy imports from `distributed`."""
from __future__ import annotations
import server
from .utils.config import ensure_config_exists
from .utils.logging import debug_log
from .utils.network import cleanup_client_session
from .workers import get_worker_manager
from .workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
from .upscale.job_store import ensure_tile_jobs_initialized
from .nodes import (
NODE_CLASS_MAPPINGS,
NODE_DISPLAY_NAME_MAPPINGS,
ImageBatchDivider,
DistributedCollectorNode,
DistributedSeed,
DistributedModelName,
DistributedValue,
AudioBatchDivider,
DistributedEmptyImage,
AnyType,
ByPassTypeTuple,
any_type,
from .bootstrap.entrypoint import (
build_node_mappings,
initialize_runtime,
)
from . import api # noqa: F401 - triggers all @routes.* registrations
from .api.queue_orchestration import ensure_distributed_state
ensure_config_exists()
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
# Aiohttp session cleanup
async def _cleanup_session():
await cleanup_client_session()
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
# Initialize distributed job state on prompt_server
prompt_server = server.PromptServer.instance
ensure_distributed_state(prompt_server)
ensure_tile_jobs_initialized()
# Worker startup
if not os.environ.get('COMFYUI_IS_WORKER'):
atexit.register(sync_cleanup)
delayed_auto_launch()
register_async_signals()
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"initialize_runtime",
]
+1 -7
View File
@@ -53,7 +53,6 @@ Queue a workflow for distributed execution.
"delegate_master": false,
"enabled_worker_ids": ["1", "2"],
"workers": ["1", "2"],
"auto_prepare": true,
"trace_execution_id": "exec_1700000000_ab12cd"
}
```
@@ -74,10 +73,6 @@ Queue a workflow for distributed execution.
- The explicit worker IDs to consider for this run.
- `workers` (optional, array of strings or objects with `id`)
- Transitional alias for `enabled_worker_ids` used by older clients.
- `auto_prepare` (optional, boolean)
- Kept for wire compatibility.
- Backend orchestration always runs with auto-prepare semantics.
- If top-level `prompt` is omitted, backend will attempt `workflow.prompt`.
- `trace_execution_id` (optional, string)
- Passed through to orchestration logs.
- Server log lines include the marker as `[exec:<trace_execution_id>]`.
@@ -105,8 +100,7 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
```json
{
"prompt_id": "<uuid>",
"worker_count": 2,
"auto_prepare_supported": true
"worker_count": 2
}
```
+6
View File
@@ -10,9 +10,13 @@ from .utilities import (
any_type,
)
from .collector import DistributedCollectorNode
from .branch import DistributedBranch
from .branch_collector import DistributedBranchCollector
NODE_CLASS_MAPPINGS = {
"DistributedCollector": DistributedCollectorNode,
"DistributedBranch": DistributedBranch,
"DistributedBranchCollector": DistributedBranchCollector,
"DistributedSeed": DistributedSeed,
"DistributedModelName": DistributedModelName,
"DistributedValue": DistributedValue,
@@ -22,6 +26,8 @@ NODE_CLASS_MAPPINGS = {
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DistributedCollector": "Distributed Collector",
"DistributedBranch": "Distributed Branch",
"DistributedBranchCollector": "Distributed Branch Collector",
"DistributedSeed": "Distributed Seed",
"DistributedModelName": "Distributed Model Name",
"DistributedValue": "Distributed Value",
+33
View File
@@ -0,0 +1,33 @@
from typing import Any
from .utilities import any_type
class DistributedBranch:
@classmethod
def INPUT_TYPES(cls: type["DistributedBranch"]) -> dict[str, Any]:
return {
"required": {
"input": (any_type,),
"num_branches": ("INT", {"default": 2, "min": 2, "max": 10, "step": 1}),
},
"hidden": {
"is_worker": ("BOOLEAN", {"default": False}),
"worker_id": ("STRING", {"default": ""}),
"assigned_branch": ("INT", {"default": -1, "min": -1, "max": 9}),
"multi_job_id": ("STRING", {"default": ""}),
},
}
RETURN_TYPES = tuple([any_type] * 10)
RETURN_NAMES = tuple([f"branch_{idx + 1}" for idx in range(10)])
FUNCTION = "branch"
CATEGORY = "utils"
def branch(
self,
input: Any,
num_branches: int = 2,
**_kwargs: Any,
) -> tuple[Any, ...]:
return tuple([input] * 10)
+298
View File
@@ -0,0 +1,298 @@
import asyncio
from dataclasses import dataclass, field
from typing import Any
import aiohttp
import torch
from ..utils.auth import distributed_auth_headers
from ..utils.async_helpers import run_async_in_server_loop
from ..utils.config import is_master_delegate_only, load_config
from ..utils.image import encode_tensor_png_data_url, ensure_contiguous
from ..utils.logging import log
from ..utils.network import get_client_session
from ..utils.worker_ids import coerce_enabled_worker_ids
from .context_kwargs import parse_distributed_hidden_context
from .hidden_inputs import build_distributed_hidden_inputs
from .queue_wait import collect_worker_queue_results
from .runtime_helpers import (
get_prompt_server_instance as _get_prompt_server_instance,
throw_if_processing_interrupted as _throw_if_processing_interrupted,
)
from .utilities import any_type
MAX_BRANCH_OUTPUTS = 10
@dataclass(frozen=True)
class BranchRunContext:
multi_job_id: str = ""
is_worker: bool = False
master_url: str = ""
enabled_worker_ids: list[str] = field(default_factory=list)
worker_id: str = ""
assigned_branch: int = -1
delegate_only: bool = False
class DistributedBranchCollector:
@classmethod
def INPUT_TYPES(cls: type["DistributedBranchCollector"]) -> dict[str, Any]:
optional_inputs = {
f"branch_{idx + 1}": (any_type,)
for idx in range(MAX_BRANCH_OUTPUTS)
}
return {
"required": {
"num_branches": ("INT", {"default": 2, "min": 2, "max": MAX_BRANCH_OUTPUTS, "step": 1}),
},
"optional": optional_inputs,
"hidden": build_distributed_hidden_inputs(
include_assigned_branch=True,
assigned_branch_max=MAX_BRANCH_OUTPUTS - 1,
),
}
RETURN_TYPES = tuple([any_type] * MAX_BRANCH_OUTPUTS)
RETURN_NAMES = tuple([f"branch_{idx + 1}" for idx in range(MAX_BRANCH_OUTPUTS)])
FUNCTION = "run"
CATEGORY = "utils"
def _branch_values(self, **kwargs):
values = []
for idx in range(MAX_BRANCH_OUTPUTS):
values.append(kwargs.get(f"branch_{idx + 1}", None))
return values
def _build_run_context(self, **kwargs: Any) -> BranchRunContext:
try:
assigned_branch = int(kwargs.get("assigned_branch", -1))
except (TypeError, ValueError):
assigned_branch = -1
common_context = parse_distributed_hidden_context(kwargs)
return BranchRunContext(
assigned_branch=assigned_branch,
**common_context,
)
def _resolve_local_branch_value(self, branch_values, assigned_branch):
try:
assigned_idx = int(assigned_branch)
except (TypeError, ValueError):
assigned_idx = -1
if 0 <= assigned_idx < MAX_BRANCH_OUTPUTS:
assigned_value = branch_values[assigned_idx]
if assigned_value is not None:
return assigned_value, assigned_idx
# In participant-pruned prompts, each participant should typically have
# exactly one local branch input. If socket names drifted (e.g. branch_2
# connected for a participant assigned to slot 0), remap that sole value
# to the assigned slot so the participant still contributes correctly.
non_none_indices = [idx for idx, value in enumerate(branch_values) if value is not None]
if len(non_none_indices) == 1:
return branch_values[non_none_indices[0]], assigned_idx
for idx, value in enumerate(branch_values):
if value is not None:
return value, idx
return None, assigned_idx
def _build_outputs(self, branch_values, num_branches):
outputs = [None] * MAX_BRANCH_OUTPUTS
try:
branch_count = int(num_branches)
except (TypeError, ValueError):
branch_count = 2
branch_count = max(2, min(branch_count, MAX_BRANCH_OUTPUTS))
for idx in range(branch_count):
value = branch_values[idx]
if value is not None:
outputs[idx] = value
return tuple(outputs)
def run(
self,
num_branches: int = 2,
**kwargs: Any,
) -> tuple[Any, ...]:
branch_values = self._branch_values(**kwargs)
context = self._build_run_context(**kwargs)
if not context.multi_job_id:
return self._build_outputs(branch_values, num_branches)
return run_async_in_server_loop(
self.execute(
branch_values,
num_branches=num_branches,
context=context,
)
)
def _participant_branch_slots(self, enabled_workers, delegate_mode, num_branches):
participants = list(enabled_workers) if delegate_mode else ["master"] + list(enabled_workers)
participant_count = len(participants)
if participant_count <= 0:
return {}
slots_by_participant = {}
for branch_slot in range(int(num_branches)):
participant_id = str(participants[branch_slot % participant_count])
slots_by_participant.setdefault(participant_id, []).append(branch_slot)
return slots_by_participant
def _expected_worker_ids_for_branches(self, enabled_workers, delegate_mode, num_branches):
slots_by_participant = self._participant_branch_slots(enabled_workers, delegate_mode, num_branches)
expected = []
for participant_id, slots in slots_by_participant.items():
if participant_id == "master" or not slots:
continue
expected.append(str(participant_id))
return expected
def _fallback_for_missing_branch(self, source_value):
if isinstance(source_value, torch.Tensor):
tensor = source_value
if tensor.ndim == 3:
tensor = tensor.unsqueeze(0)
if tensor.ndim == 4:
fallback = torch.zeros_like(tensor)
if fallback.is_cuda:
fallback = fallback.cpu()
return ensure_contiguous(fallback)
return torch.zeros((1, 64, 64, 3), dtype=torch.float32)
async def _send_branch_result_to_master(self, value, branch_idx, multi_job_id, master_url, worker_id):
if not isinstance(value, torch.Tensor):
raise ValueError("DistributedBranchCollector currently supports torch.Tensor results for worker transfer.")
image_batch = value
if image_batch.ndim == 3:
image_batch = image_batch.unsqueeze(0)
if image_batch.ndim != 4 or image_batch.shape[0] <= 0:
raise ValueError(
"DistributedBranchCollector tensor result must be IMAGE-like with shape [B,H,W,C] or [H,W,C]."
)
payload = {
"job_id": str(multi_job_id),
"worker_id": str(worker_id),
"batch_idx": int(branch_idx),
"image": encode_tensor_png_data_url(image_batch, 0),
"is_last": True,
}
session = await get_client_session()
url = f"{master_url}/distributed/job_complete"
async with session.post(
url,
json=payload,
headers=distributed_auth_headers(load_config()),
timeout=aiohttp.ClientTimeout(total=60),
) as response:
response.raise_for_status()
async def execute(self, *args: Any, **kwargs: Any) -> tuple[Any, ...]:
"""Compatibility wrapper with a uniform execute() signature across collectors."""
return await self._execute_branch(*args, **kwargs)
async def _execute_branch(
self,
branch_values: list[Any],
num_branches: int = 2,
context: BranchRunContext | None = None,
) -> tuple[Any, ...]:
run_context = context or BranchRunContext()
try:
branch_count = int(num_branches)
except (TypeError, ValueError):
branch_count = 2
branch_count = max(2, min(branch_count, MAX_BRANCH_OUTPUTS))
outputs = [None] * MAX_BRANCH_OUTPUTS
for idx in range(branch_count):
if branch_values[idx] is not None:
outputs[idx] = branch_values[idx]
local_value, local_branch_idx = self._resolve_local_branch_value(branch_values, run_context.assigned_branch)
if 0 <= local_branch_idx < MAX_BRANCH_OUTPUTS and local_value is not None:
outputs[local_branch_idx] = local_value
if run_context.is_worker:
if 0 <= local_branch_idx < MAX_BRANCH_OUTPUTS and local_value is not None:
try:
await self._send_branch_result_to_master(
local_value,
local_branch_idx,
run_context.multi_job_id,
run_context.master_url,
run_context.worker_id,
)
outputs[local_branch_idx] = local_value
except Exception as exc:
log(f"Worker - DistributedBranchCollector failed to send branch result: {exc}")
return tuple(outputs)
delegate_mode = bool(run_context.delegate_only or is_master_delegate_only())
enabled_workers = coerce_enabled_worker_ids(run_context.enabled_worker_ids)
slots_by_participant = self._participant_branch_slots(enabled_workers, delegate_mode, branch_count)
expected_worker_ids = self._expected_worker_ids_for_branches(enabled_workers, delegate_mode, branch_count)
expected_workers = set(expected_worker_ids)
if not expected_workers:
return tuple(outputs)
prompt_server = _get_prompt_server_instance()
async with prompt_server.distributed_jobs_lock:
if run_context.multi_job_id not in prompt_server.distributed_pending_jobs:
prompt_server.distributed_pending_jobs[run_context.multi_job_id] = asyncio.Queue()
workers_done: set[str] = set()
def _handle_queue_result(result: dict[str, Any]) -> None:
image_index = result.get("image_index")
tensor = result.get("tensor")
try:
slot_idx = int(image_index)
except (TypeError, ValueError):
slot_idx = -1
if 0 <= slot_idx < MAX_BRANCH_OUTPUTS and tensor is not None:
if isinstance(tensor, torch.Tensor):
if tensor.is_cuda:
tensor = tensor.cpu()
tensor = ensure_contiguous(tensor)
outputs[slot_idx] = tensor
try:
workers_done = await collect_worker_queue_results(
prompt_server=prompt_server,
multi_job_id=run_context.multi_job_id,
expected_workers=expected_workers,
on_result=_handle_queue_result,
timeout_log_prefix=(
"Master - DistributedBranchCollector heartbeat timeout. "
"Still waiting for workers: "
),
throw_if_interrupted=_throw_if_processing_interrupted,
)
finally:
async with prompt_server.distributed_jobs_lock:
prompt_server.distributed_pending_jobs.pop(run_context.multi_job_id, None)
if hasattr(prompt_server, "distributed_job_allowed_workers"):
prompt_server.distributed_job_allowed_workers.pop(run_context.multi_job_id, None)
missing_workers = sorted(expected_workers - workers_done)
if missing_workers:
fallback = self._fallback_for_missing_branch(local_value)
for missing_worker_id in missing_workers:
for slot_idx in slots_by_participant.get(str(missing_worker_id), []):
if 0 <= slot_idx < MAX_BRANCH_OUTPUTS and outputs[slot_idx] is None:
outputs[slot_idx] = fallback
return tuple(outputs)
+312 -248
View File
@@ -1,31 +1,49 @@
import torch
import io
import json
import asyncio
import time
import base64
import io
import time
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
import aiohttp
import server as _server
import comfy.model_management
import torch
from comfy.utils import ProgressBar
from ..utils.auth import distributed_auth_headers
from ..utils.logging import debug_log, log
from ..utils.config import get_worker_timeout_seconds, load_config, is_master_delegate_only
from ..utils.constants import HEARTBEAT_INTERVAL
from ..utils.image import tensor_to_pil, pil_to_tensor, ensure_contiguous
from ..utils.network import build_worker_url, get_client_session, probe_worker
from ..utils.worker_ids import coerce_enabled_worker_ids
from ..utils.audio_payload import encode_audio_payload
from ..utils.async_helpers import run_async_in_server_loop
from .context_kwargs import parse_distributed_hidden_context
from .hidden_inputs import build_distributed_hidden_inputs
prompt_server = _server.PromptServer.instance
@dataclass(frozen=True)
class CollectorRunContext:
multi_job_id: str = ""
is_worker: bool = False
master_url: str = ""
enabled_worker_ids: list[str] = field(default_factory=list)
worker_batch_size: int = 1
worker_id: str = ""
pass_through: bool = False
delegate_only: bool = False
class DistributedCollectorNode:
EMPTY_AUDIO = {"waveform": torch.zeros(1, 2, 1), "sample_rate": 44100}
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls: type["DistributedCollectorNode"]) -> dict[str, Any]:
return {
"required": {
"images": ("IMAGE",),
@@ -38,50 +56,65 @@ class DistributedCollectorNode:
),
},
"optional": { "audio": ("AUDIO",) },
"hidden": {
"multi_job_id": ("STRING", {"default": ""}),
"is_worker": ("BOOLEAN", {"default": False}),
"master_url": ("STRING", {"default": ""}),
"enabled_worker_ids": ("STRING", {"default": "[]"}),
"worker_batch_size": ("INT", {"default": 1, "min": 1, "max": 1024}),
"worker_id": ("STRING", {"default": ""}),
"pass_through": ("BOOLEAN", {"default": False}),
"delegate_only": ("BOOLEAN", {"default": False}),
},
"hidden": build_distributed_hidden_inputs(
include_worker_batch_size=True,
include_pass_through=True,
),
}
RETURN_TYPES = ("IMAGE", "AUDIO")
RETURN_NAMES = ("images", "audio")
FUNCTION = "run"
CATEGORY = "image"
def run(self, images, load_balance=False, audio=None, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False, delegate_only=False):
def _build_run_context(self, **kwargs: Any) -> CollectorRunContext:
try:
worker_batch_size = int(kwargs.get("worker_batch_size", 1))
except (TypeError, ValueError):
worker_batch_size = 1
common_context = parse_distributed_hidden_context(kwargs)
return CollectorRunContext(
worker_batch_size=max(worker_batch_size, 1),
pass_through=bool(kwargs.get("pass_through", False)),
**common_context,
)
def run(
self,
images: torch.Tensor,
load_balance: bool = False,
audio: dict[str, Any] | None = None,
**kwargs: Any,
) -> tuple[torch.Tensor, dict[str, Any]]:
context = self._build_run_context(**kwargs)
# 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:
if not context.multi_job_id or context.pass_through:
if context.pass_through:
debug_log("Collector: pass-through mode enabled, returning images unchanged")
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,
audio,
load_balance,
multi_job_id,
is_worker,
master_url,
enabled_worker_ids,
worker_batch_size,
worker_id,
delegate_only,
images=images,
audio=audio,
load_balance=load_balance,
context=context,
)
)
return result
async def send_batch_to_master(self, image_batch, audio, multi_job_id, master_url, worker_id):
async def send_batch_to_master(
self,
image_batch: torch.Tensor,
audio: dict[str, Any] | None,
multi_job_id: str,
master_url: str,
worker_id: str,
) -> None:
"""Send image batch to master via canonical JSON envelopes."""
batch_size = image_batch.shape[0]
if batch_size == 0:
@@ -110,6 +143,7 @@ class DistributedCollectorNode:
async with session.post(
url,
json=payload,
headers=distributed_auth_headers(load_config()),
timeout=aiohttp.ClientTimeout(total=60),
) as response:
response.raise_for_status()
@@ -235,235 +269,265 @@ class DistributedCollectorNode:
else:
raise ValueError("No image data collected from master or workers")
async def execute(self, images, audio, load_balance=False, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", delegate_only=False):
if is_worker:
# Worker mode: send images and audio to master in a single batch
debug_log(f"Worker - Job {multi_job_id} complete. Sending {images.shape[0]} image(s) to master")
await self.send_batch_to_master(images, audio, multi_job_id, master_url, worker_id)
return (images, audio if audio is not None else self.EMPTY_AUDIO)
else:
delegate_mode = delegate_only or is_master_delegate_only()
# Master mode: collect images and audio from workers
enabled_workers_raw = json.loads(enabled_worker_ids)
enabled_workers = []
seen_enabled = set()
for worker_id in enabled_workers_raw:
worker_id_str = str(worker_id)
if worker_id_str in seen_enabled:
continue
seen_enabled.add(worker_id_str)
enabled_workers.append(worker_id_str)
expected_workers = set(enabled_workers)
num_workers = len(expected_workers)
if num_workers == 0:
return (images, audio if audio is not None else self.EMPTY_AUDIO)
# Create the queue before any expensive local work to avoid job_complete race.
async with prompt_server.distributed_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_jobs:
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
debug_log(f"Master - Initialized queue early for job {multi_job_id}")
else:
existing_size = prompt_server.distributed_pending_jobs[multi_job_id].qsize()
debug_log(f"Master - Using existing queue for job {multi_job_id} (current size: {existing_size})")
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")
async def _ensure_pending_queue(self, multi_job_id):
async with prompt_server.distributed_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_jobs:
prompt_server.distributed_pending_jobs[multi_job_id] = asyncio.Queue()
debug_log(f"Master - Initialized queue early for job {multi_job_id}")
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...")
existing_size = prompt_server.distributed_pending_jobs[multi_job_id].qsize()
debug_log(f"Master - Using existing queue for job {multi_job_id} (current size: {existing_size})")
# Ensure master images are contiguous
images_on_cpu = ensure_contiguous(images_on_cpu)
async def _cleanup_pending_queue(self, multi_job_id):
async with prompt_server.distributed_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_jobs:
del prompt_server.distributed_pending_jobs[multi_job_id]
if hasattr(prompt_server, "distributed_job_allowed_workers"):
prompt_server.distributed_job_allowed_workers.pop(multi_job_id, None)
def _build_master_inputs(self, images, audio, delegate_mode, multi_job_id, num_workers):
if delegate_mode:
debug_log(
f"Master - Job {multi_job_id}: Delegate-only mode enabled, collecting exclusively from {num_workers} workers"
)
return 0, None, None
# 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
# Collect images until all workers report they're done
collected_count = 0
workers_done = set()
# Use unified worker timeout from config/UI with simple sliced waits
base_timeout = float(get_worker_timeout_seconds())
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
last_activity = time.time()
# Get queue size before starting
async with prompt_server.distributed_jobs_lock:
q = prompt_server.distributed_pending_jobs[multi_job_id]
initial_size = q.qsize()
images_on_cpu = ensure_contiguous(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..."
)
return master_batch_size, images_on_cpu, audio
# NEW: Initialize progress bar for workers (total = num_workers)
p = ProgressBar(num_workers)
def mark_worker_done(done_worker_id):
done_worker_id = str(done_worker_id)
if done_worker_id not in expected_workers:
async def _probe_missing_workers_busy(self, missing_workers):
any_busy = False
try:
cfg = load_config()
cfg_workers = cfg.get('workers', [])
for wid in list(missing_workers):
wrec = next((w for w in cfg_workers if str(w.get('id')) == str(wid)), None)
if not wrec:
debug_log(f"Collector probe: worker {wid} not found in config")
continue
worker_url = build_worker_url(wrec)
try:
payload = await probe_worker(worker_url, timeout=2.0)
queue_remaining = None
if payload is not None:
queue_remaining = int(payload.get('exec_info', {}).get('queue_remaining', 0))
debug_log(
f"Master - Ignoring completion from unexpected worker {done_worker_id} for job {multi_job_id}"
"Collector probe: worker "
f"{wid} online={payload is not None} queue_remaining={queue_remaining}"
)
return
if done_worker_id in workers_done:
debug_log(
f"Master - Ignoring duplicate completion from worker {done_worker_id} for job {multi_job_id}"
)
return
workers_done.add(done_worker_id)
p.update(1) # +1 per completed expected worker
try:
while len(workers_done) < num_workers:
# Check for user interruption to abort collection promptly
comfy.model_management.throw_exception_if_processing_interrupted()
try:
# Get the queue again each time to ensure we have the right reference
async with prompt_server.distributed_jobs_lock:
q = prompt_server.distributed_pending_jobs[multi_job_id]
current_size = q.qsize()
result = await asyncio.wait_for(q.get(), timeout=slice_timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
count = self._store_worker_result(worker_images, result)
collected_count += count
debug_log(
f"Master - Got canonical result from worker {worker_id}, "
f"image {result.get('image_index', 0)}, is_last={is_last}"
)
# 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:
mark_worker_done(worker_id)
except asyncio.TimeoutError:
# If we still have time, continue polling; otherwise handle timeout
if (time.time() - last_activity) < base_timeout:
comfy.model_management.throw_exception_if_processing_interrupted()
continue
# Re-check for user interruption after timeout expiry
comfy.model_management.throw_exception_if_processing_interrupted()
missing_workers = set(str(w) for w in enabled_workers) - workers_done
elapsed = time.time() - last_activity
for missing_worker_id in sorted(missing_workers):
log(
"Master - Heartbeat timeout: "
f"worker={missing_worker_id}, elapsed={elapsed:.1f}s"
)
if payload is not None and queue_remaining and queue_remaining > 0:
any_busy = True
log(
f"Master - Heartbeat timeout. Still waiting for workers: {list(missing_workers)} "
f"(elapsed={elapsed:.1f}s)"
f"Master - Probe grace: worker {wid} appears busy "
f"(queue_remaining={queue_remaining}). Continuing to wait."
)
# Probe missing workers' /prompt endpoints to check if they are actively processing
any_busy = False
try:
cfg = load_config()
cfg_workers = cfg.get('workers', [])
for wid in list(missing_workers):
wrec = next((w for w in cfg_workers if str(w.get('id')) == str(wid)), None)
if not wrec:
debug_log(f"Collector probe: worker {wid} not found in config")
continue
worker_url = build_worker_url(wrec)
try:
payload = await probe_worker(worker_url, timeout=2.0)
queue_remaining = None
if payload is not None:
queue_remaining = int(payload.get('exec_info', {}).get('queue_remaining', 0))
debug_log(
"Collector probe: worker "
f"{wid} online={payload is not None} queue_remaining={queue_remaining}"
)
if payload is not None and queue_remaining and queue_remaining > 0:
any_busy = True
log(
f"Master - Probe grace: worker {wid} appears busy "
f"(queue_remaining={queue_remaining}). Continuing to wait."
)
break
except Exception as e:
debug_log(f"Collector probe failed for worker {wid}: {e}")
except Exception as e:
debug_log(f"Collector probe setup error: {e}")
if any_busy:
# Refresh last_activity and continue waiting
last_activity = time.time()
# Refresh base timeout in case the user changed it in UI
base_timeout = float(get_worker_timeout_seconds())
continue
# Check queue size again with lock
async with prompt_server.distributed_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_jobs:
final_q = prompt_server.distributed_pending_jobs[multi_job_id]
final_size = final_q.qsize()
# Try to drain any remaining items
remaining_items = []
while not final_q.empty():
try:
item = final_q.get_nowait()
remaining_items.append(item)
except asyncio.QueueEmpty:
break
if remaining_items:
# Process them
for item in remaining_items:
worker_id = item['worker_id']
is_last = item.get('is_last', False)
collected_count += self._store_worker_result(worker_images, item)
if is_last:
mark_worker_done(worker_id)
else:
log(f"Master - Queue {multi_job_id} no longer exists!")
break
except comfy.model_management.InterruptProcessingException:
# Cleanup queue on interruption and re-raise to abort prompt cleanly
async with prompt_server.distributed_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_jobs:
del prompt_server.distributed_pending_jobs[multi_job_id]
raise
total_collected = sum(len(imgs) for imgs in worker_images.values())
# Clean up job queue
async with prompt_server.distributed_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_jobs:
del prompt_server.distributed_pending_jobs[multi_job_id]
except Exception as e:
debug_log(f"Collector probe failed for worker {wid}: {e}")
except Exception as e:
debug_log(f"Collector probe setup error: {e}")
return any_busy
try:
combined = self._reorder_and_combine_tensors(
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
async def _drain_remaining_queue_items(
self,
multi_job_id: str,
worker_images: dict[str, dict[int, torch.Tensor]],
mark_worker_done: Callable[[str], None],
) -> int:
collected_count = 0
async with prompt_server.distributed_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_jobs:
log(f"Master - Queue {multi_job_id} no longer exists!")
return collected_count
final_q = prompt_server.distributed_pending_jobs[multi_job_id]
remaining_items = []
while not final_q.empty():
try:
remaining_items.append(final_q.get_nowait())
except asyncio.QueueEmpty:
break
for item in remaining_items:
worker_id = item['worker_id']
is_last = item.get('is_last', False)
collected_count += self._store_worker_result(worker_images, item)
if is_last:
mark_worker_done(worker_id)
return collected_count
async def _collect_worker_results(
self,
multi_job_id: str,
enabled_workers: list[str],
expected_workers: set[str],
worker_images: dict[str, dict[int, torch.Tensor]],
worker_audio: dict[str, dict[str, Any]],
) -> int:
num_workers = len(expected_workers)
workers_done = set()
collected_count = 0
base_timeout = float(get_worker_timeout_seconds())
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
last_activity = time.time()
progress = ProgressBar(num_workers)
def mark_worker_done(done_worker_id: str) -> None:
done_worker_id = str(done_worker_id)
if done_worker_id not in expected_workers:
debug_log(
f"Master - Ignoring completion from unexpected worker {done_worker_id} for job {multi_job_id}"
)
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
return
if done_worker_id in workers_done:
debug_log(
f"Master - Ignoring duplicate completion from worker {done_worker_id} for job {multi_job_id}"
)
return
workers_done.add(done_worker_id)
progress.update(1)
# Combine audio from master and workers
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
while len(workers_done) < num_workers:
comfy.model_management.throw_exception_if_processing_interrupted()
try:
async with prompt_server.distributed_jobs_lock:
q = prompt_server.distributed_pending_jobs[multi_job_id]
return (combined, combined_audio)
except Exception as e:
log(f"Master - Error combining images: {e}")
# Return just the master images as fallback
return (images, audio if audio is not None else self.EMPTY_AUDIO)
result = await asyncio.wait_for(q.get(), timeout=slice_timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
collected_count += self._store_worker_result(worker_images, result)
debug_log(
f"Master - Got canonical result from worker {worker_id}, "
f"image {result.get('image_index', 0)}, is_last={is_last}"
)
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}")
last_activity = time.time()
base_timeout = float(get_worker_timeout_seconds())
if is_last:
mark_worker_done(worker_id)
except asyncio.TimeoutError:
if (time.time() - last_activity) < base_timeout:
comfy.model_management.throw_exception_if_processing_interrupted()
continue
comfy.model_management.throw_exception_if_processing_interrupted()
missing_workers = set(str(w) for w in enabled_workers) - workers_done
elapsed = time.time() - last_activity
for missing_worker_id in sorted(missing_workers):
log(
"Master - Heartbeat timeout: "
f"worker={missing_worker_id}, elapsed={elapsed:.1f}s"
)
log(
f"Master - Heartbeat timeout. Still waiting for workers: {list(missing_workers)} "
f"(elapsed={elapsed:.1f}s)"
)
if await self._probe_missing_workers_busy(missing_workers):
last_activity = time.time()
base_timeout = float(get_worker_timeout_seconds())
continue
collected_count += await self._drain_remaining_queue_items(
multi_job_id,
worker_images,
mark_worker_done,
)
break
return collected_count
async def _execute_master(
self,
images,
audio,
multi_job_id,
enabled_worker_ids,
delegate_only,
):
delegate_mode = delegate_only or is_master_delegate_only()
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
expected_workers = set(enabled_workers)
num_workers = len(expected_workers)
if num_workers == 0:
return (images, audio if audio is not None else self.EMPTY_AUDIO)
await self._ensure_pending_queue(multi_job_id)
master_batch_size, images_on_cpu, master_audio = self._build_master_inputs(
images, audio, delegate_mode, multi_job_id, num_workers
)
worker_images = {}
worker_audio = {}
try:
await self._collect_worker_results(
multi_job_id=multi_job_id,
enabled_workers=enabled_workers,
expected_workers=expected_workers,
worker_images=worker_images,
worker_audio=worker_audio,
)
except comfy.model_management.InterruptProcessingException:
await self._cleanup_pending_queue(multi_job_id)
raise
await self._cleanup_pending_queue(multi_job_id)
try:
combined = self._reorder_and_combine_tensors(
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
)
debug_log(
f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})"
)
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
return (combined, combined_audio)
except Exception as e:
log(f"Master - Error combining images: {e}")
return (images, audio if audio is not None else self.EMPTY_AUDIO)
async def execute(self, *args: Any, **kwargs: Any) -> tuple[torch.Tensor, dict[str, Any]]:
"""Compatibility wrapper with a uniform execute() signature across collectors."""
return await self._execute_collector(*args, **kwargs)
async def _execute_collector(
self,
images: torch.Tensor,
audio: dict[str, Any] | None,
load_balance: bool = False,
context: CollectorRunContext | None = None,
) -> tuple[torch.Tensor, dict[str, Any]]:
run_context = context or CollectorRunContext()
_ = load_balance
_ = run_context.worker_batch_size
if run_context.is_worker:
debug_log(
"Worker - Job "
f"{run_context.multi_job_id} complete. Sending {images.shape[0]} image(s) to master"
)
await self.send_batch_to_master(
images,
audio,
run_context.multi_job_id,
run_context.master_url,
run_context.worker_id,
)
return (images, audio if audio is not None else self.EMPTY_AUDIO)
return await self._execute_master(
images=images,
audio=audio,
multi_job_id=run_context.multi_job_id,
enabled_worker_ids=run_context.enabled_worker_ids,
delegate_only=run_context.delegate_only,
)
+26
View File
@@ -0,0 +1,26 @@
from typing import Any
from ..utils.parsing import coerce_bool
from ..utils.worker_ids import coerce_enabled_worker_ids
def _parse_bool(value: Any, default: bool = False) -> bool:
return coerce_bool(value, default=default)
def _parse_enabled_workers(value: Any) -> list[str]:
if value is None:
return []
return coerce_enabled_worker_ids(value)
def parse_distributed_hidden_context(kwargs: dict[str, Any]) -> dict[str, Any]:
"""Parse common distributed hidden inputs from node kwargs."""
return {
"multi_job_id": str(kwargs.get("multi_job_id", "") or ""),
"is_worker": _parse_bool(kwargs.get("is_worker", False)),
"master_url": str(kwargs.get("master_url", "") or ""),
"enabled_worker_ids": _parse_enabled_workers(kwargs.get("enabled_worker_ids", "[]")),
"worker_id": str(kwargs.get("worker_id", "") or ""),
"delegate_only": _parse_bool(kwargs.get("delegate_only", False)),
}
+320 -82
View File
@@ -1,25 +1,28 @@
import json
import math
from functools import wraps
from typing import Any, Callable
import comfy.samplers
from ..utils.logging import debug_log, log
from ..utils.async_helpers import run_async_in_server_loop
from ..upscale.job_store import ensure_tile_jobs_initialized
from .hidden_inputs import build_distributed_hidden_inputs
from ..upscale.tile_ops import TileOpsMixin
from ..upscale.result_collector import ResultCollectorMixin
from ..upscale.worker_comms import WorkerCommsMixin
from ..upscale.job_state import JobStateMixin
from ..upscale.processing_args import UpscaleCoreArgs
from ..upscale.modes.single_gpu import SingleGpuModeMixin
from ..upscale.modes.static import StaticModeMixin
from ..upscale.modes.dynamic import DynamicModeMixin
from ..utils.worker_ids import coerce_enabled_worker_ids
def sync_wrapper(async_func):
def sync_wrapper(async_func: Callable[..., Any]) -> Callable[..., Any]:
"""Decorator to wrap async methods for synchronous execution."""
@wraps(async_func)
def sync_func(self, *args, **kwargs):
def sync_func(self, *args: Any, **kwargs: Any) -> Any:
# Use run_async_in_server_loop for ComfyUI compatibility
return run_async_in_server_loop(
async_func(self, *args, **kwargs),
@@ -27,21 +30,10 @@ def sync_wrapper(async_func):
)
return sync_func
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 []
def _parse_enabled_worker_ids(enabled_worker_ids: str | list[str] | None) -> list[str]:
"""Backward-compatible alias for enabled-worker normalization."""
return coerce_enabled_worker_ids(enabled_worker_ids)
class UltimateSDUpscaleDistributed(
DynamicModeMixin,
@@ -81,7 +73,14 @@ class UltimateSDUpscaleDistributed(
debug_log("UltimateSDUpscaleDistributed - Node initialized")
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls: type["UltimateSDUpscaleDistributed"]) -> dict[str, Any]:
hidden_inputs = build_distributed_hidden_inputs()
hidden_inputs.update(
{
"tile_indices": ("STRING", {"default": ""}), # Unused - kept for compatibility
"dynamic_threshold": ("INT", {"default": 8, "min": 1, "max": 64}),
}
)
return {
"required": {
"upscaled_image": ("IMAGE",),
@@ -102,15 +101,7 @@ class UltimateSDUpscaleDistributed(
"force_uniform_tiles": ("BOOLEAN", {"default": True}),
"tiled_decode": ("BOOLEAN", {"default": False}),
},
"hidden": {
"multi_job_id": ("STRING", {"default": ""}),
"is_worker": ("BOOLEAN", {"default": False}),
"master_url": ("STRING", {"default": ""}),
"enabled_worker_ids": ("STRING", {"default": "[]"}),
"worker_id": ("STRING", {"default": ""}),
"tile_indices": ("STRING", {"default": ""}), # Unused - kept for compatibility
"dynamic_threshold": ("INT", {"default": 8, "min": 1, "max": 64}),
},
"hidden": hidden_inputs,
}
RETURN_TYPES = ("IMAGE",)
@@ -122,12 +113,47 @@ class UltimateSDUpscaleDistributed(
"""Force re-execution."""
return float("nan") # Always re-execute
def run(self, upscaled_image, model, positive, negative, vae, seed, steps, cfg,
sampler_name, scheduler, denoise, tile_width, tile_height, padding,
mask_blur, force_uniform_tiles, tiled_decode,
multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]",
worker_id="", tile_indices="", dynamic_threshold=8):
def run(
self,
upscaled_image: Any,
model: Any,
positive: Any,
negative: Any,
vae: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
denoise: float,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
tiled_decode: bool,
multi_job_id: str = "",
is_worker: bool = False,
master_url: str = "",
enabled_worker_ids: str = "[]",
worker_id: str = "",
tile_indices: str = "",
dynamic_threshold: int = 8,
) -> tuple[Any, ...]:
"""Entry point - runs SYNCHRONOUSLY like Ultimate SD Upscaler."""
core_args = UpscaleCoreArgs(
model=model,
positive=positive,
negative=negative,
vae=vae,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
denoise=denoise,
tiled_decode=tiled_decode,
)
# Strict WAN/FLOW batching: error if batch is not 4n+1 (except allow 1)
try:
batch_size = int(getattr(upscaled_image, 'shape', [1])[0])
@@ -142,9 +168,15 @@ class UltimateSDUpscaleDistributed(
)
if not multi_job_id:
# No distributed processing, run single GPU version
return self.process_single_gpu(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode)
return self.process_single_gpu(
upscaled_image=upscaled_image,
core_args=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mask_blur=mask_blur,
force_uniform_tiles=force_uniform_tiles,
)
if is_worker:
# Worker mode: process tiles synchronously
@@ -155,24 +187,64 @@ class UltimateSDUpscaleDistributed(
worker_id, enabled_worker_ids, dynamic_threshold)
else:
# Master mode: distribute and collect synchronously
return self.process_master(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_worker_ids,
dynamic_threshold)
return self.process_master(
upscaled_image=upscaled_image,
core_args=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mask_blur=mask_blur,
force_uniform_tiles=force_uniform_tiles,
multi_job_id=multi_job_id,
enabled_worker_ids=enabled_worker_ids,
dynamic_threshold=dynamic_threshold,
)
def process_worker(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
worker_id, enabled_worker_ids, dynamic_threshold):
def process_worker(
self,
upscaled_image: Any,
model: Any,
positive: Any,
negative: Any,
vae: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
denoise: float,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
tiled_decode: bool,
multi_job_id: str,
master_url: str,
worker_id: str,
enabled_worker_ids: str,
dynamic_threshold: int,
) -> tuple[Any, ...]:
"""Unified worker processing - handles both static and dynamic modes."""
core_args = UpscaleCoreArgs(
model=model,
positive=positive,
negative=negative,
vae=vae,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name=sampler_name,
scheduler=scheduler,
denoise=denoise,
tiled_decode=tiled_decode,
)
# Get batch size to determine mode
batch_size = upscaled_image.shape[0]
# Ensure mode consistency across master/workers via shared threshold
# Determine mode (must match master's logic)
enabled_workers = json.loads(enabled_worker_ids)
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
num_workers = len(enabled_workers)
# Compute number of tiles for this image to decide if tile distribution makes sense
_, height, width, _ = upscaled_image.shape
@@ -181,31 +253,47 @@ class UltimateSDUpscaleDistributed(
mode = self._determine_processing_mode(batch_size, num_workers, dynamic_threshold)
# For USDU-style processing, we want tile distribution whenever workers are available
# and there is more than one tile to process, even if batch == 1.
if num_workers > 0 and num_tiles_per_image > 1:
# and there is more than one tile to process for single-image runs.
if num_workers > 0 and batch_size <= 1 and num_tiles_per_image > 1:
mode = "static"
debug_log(f"USDU Dist Worker - Batch size {batch_size}")
if mode == "dynamic":
return self.process_worker_dynamic(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
worker_id, enabled_worker_ids, dynamic_threshold)
return self.process_worker_dynamic(
upscaled_image=upscaled_image,
core_args=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mask_blur=mask_blur,
force_uniform_tiles=force_uniform_tiles,
multi_job_id=multi_job_id,
master_url=master_url,
worker_id=worker_id,
enabled_worker_ids=enabled_worker_ids,
dynamic_threshold=dynamic_threshold,
)
# Static mode - enhanced with health monitoring and retry logic
return self._process_worker_static_sync(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
return self._process_worker_static_sync(upscaled_image, core_args,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
force_uniform_tiles, multi_job_id, master_url,
worker_id, enabled_workers)
def process_master(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_worker_ids,
dynamic_threshold):
def process_master(
self,
upscaled_image: Any,
core_args: UpscaleCoreArgs,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
multi_job_id: str,
enabled_worker_ids: str,
dynamic_threshold: int,
) -> tuple[Any, ...]:
"""Unified master processing with enhanced monitoring and failure handling."""
# Round tile dimensions
tile_width = self.round_to_multiple(tile_width)
@@ -224,56 +312,206 @@ class UltimateSDUpscaleDistributed(
)
# Parse enabled workers
enabled_workers = json.loads(enabled_worker_ids)
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
num_workers = len(enabled_workers)
# Determine processing mode
mode = self._determine_processing_mode(batch_size, num_workers, dynamic_threshold)
# Prefer tile-based static distribution when workers are available and there are multiple tiles,
# even for batch == 1, to spread tiles across GPUs like the legacy dynamic tile queue.
if num_workers > 0 and num_tiles_per_image > 1:
# for single-image jobs to spread tiles across GPUs like the legacy tile queue.
if num_workers > 0 and batch_size <= 1 and num_tiles_per_image > 1:
mode = "static"
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,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode)
return self.process_single_gpu(
upscaled_image=upscaled_image,
core_args=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mask_blur=mask_blur,
force_uniform_tiles=force_uniform_tiles,
)
elif mode == "dynamic":
# Dynamic mode for large batches
return self.process_master_dynamic(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers)
return self.process_master_dynamic(
upscaled_image=upscaled_image,
core_args=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mask_blur=mask_blur,
force_uniform_tiles=force_uniform_tiles,
multi_job_id=multi_job_id,
enabled_workers=enabled_workers,
)
# Static mode - enhanced with unified job management
return self._process_master_static_sync(upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
return self._process_master_static_sync(upscaled_image, core_args,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers,
force_uniform_tiles, multi_job_id, enabled_workers,
all_tiles, num_tiles_per_image)
def _determine_processing_mode(self, batch_size: int, num_workers: int, dynamic_threshold: int) -> str:
"""Determines processing mode per requested policy:
- any workers => prefer static (tile-based) for USDU
- no workers => single_gpu
"""
"""Determine mode from worker availability and configured threshold."""
if num_workers == 0:
return "single_gpu"
# Default to static when distributed; master/worker may still override if special cases arise
threshold = max(1, int(dynamic_threshold))
if int(batch_size) >= threshold:
return "dynamic"
return "static"
# Ensure initialization before registering routes
ensure_tile_jobs_initialized()
class USDUDelegateCollector:
"""Lightweight stand-in used automatically in delegate-only master prompts.
When the master is in delegate-only mode, the orchestration code swaps the
full ``UltimateSDUpscaleDistributed`` class for this one so that no upstream
model/image nodes execute on the master. Workers initialise the dynamic job
queue via ``/distributed/init_dynamic_job`` and this node simply waits for
all images to arrive, then assembles the output tensor.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {},
"hidden": {
"multi_job_id": ("STRING", {"default": ""}),
"enabled_worker_ids": ("STRING", {"default": "[]"}),
"delegate_only": ("BOOLEAN", {"default": True}),
"is_worker": ("BOOLEAN", {"default": False}),
"worker_id": ("STRING", {"default": ""}),
"master_url": ("STRING", {"default": ""}),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "run"
CATEGORY = "image/upscaling"
@classmethod
def IS_CHANGED(cls, **kwargs):
return float("nan")
def run(self, multi_job_id="", enabled_worker_ids="[]", **kwargs):
import time
import torch
from ..utils.image import pil_to_tensor
from ..upscale.job_models import ImageJobState
if not multi_job_id:
log("USDUDelegateCollector: no multi_job_id, returning empty image")
return (torch.zeros(1, 64, 64, 3),)
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
num_workers = len(enabled_workers)
log(f"USDU delegate-only: waiting for workers to init job {multi_job_id}")
prompt_server = ensure_tile_jobs_initialized()
# Wait for workers to create the job via /distributed/init_dynamic_job
max_wait = 120.0
poll_interval = 1.0
start = time.time()
job_data = None
while time.time() - start < max_wait:
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
if isinstance(job_data, ImageJobState) and job_data.batch_size > 0:
break
job_data = None
time.sleep(poll_interval)
if job_data is None:
log(f"USDUDelegateCollector: job {multi_job_id} was not initialized by workers within {max_wait}s")
return (torch.zeros(1, 64, 64, 3),)
batch_size = job_data.batch_size
log(f"USDU delegate-only: job ready, collecting {batch_size} images from {num_workers} workers")
# Collect all images using the existing result collector
collected = run_async_in_server_loop(
self._collect_all_images(multi_job_id, batch_size, num_workers, prompt_server),
timeout=600.0,
)
# Assemble tensor from collected PIL images
result_images = []
for idx in range(batch_size):
img = collected.get(idx)
if img is not None:
result_images.append(pil_to_tensor(img))
else:
log(f"USDUDelegateCollector: missing image {idx}, using blank")
result_images.append(torch.zeros(1, 64, 64, 3))
result_tensor = torch.cat(result_images, dim=0)
log(f"USDU delegate-only: collected all {batch_size} images")
return (result_tensor,)
@staticmethod
async def _collect_all_images(multi_job_id, batch_size, num_workers, prompt_server):
"""Wait for all images to arrive from workers."""
import asyncio
from ..upscale.job_models import ImageJobState
from ..upscale.job_timeout import check_and_requeue_timed_out_workers
timeout_seconds = 300.0
poll_interval = 2.0
start = asyncio.get_event_loop().time()
while True:
# Drain results from the queue
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
if isinstance(job_data, ImageJobState):
# Drain queue items into completed_images
while True:
try:
result = job_data.queue.get_nowait()
worker_id = result.get("worker_id")
if "image_idx" in result and "image" in result:
idx = result["image_idx"]
if idx not in job_data.completed_images:
job_data.completed_images[idx] = result["image"]
debug_log(f"Delegate collected image {idx} from worker {worker_id}")
except asyncio.QueueEmpty:
break
completed_count = len(job_data.completed_images)
if completed_count >= batch_size:
debug_log(f"Delegate: all {batch_size} images collected")
completed = dict(job_data.completed_images)
# Cleanup
async with prompt_server.distributed_tile_jobs_lock:
prompt_server.distributed_pending_tile_jobs.pop(multi_job_id, None)
return completed
# Check for timeouts periodically
await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
elapsed = asyncio.get_event_loop().time() - start
if elapsed > timeout_seconds:
log(f"USDUDelegateCollector: timed out after {elapsed:.0f}s with {len(getattr(job_data, 'completed_images', {}))} of {batch_size} images")
if isinstance(job_data, ImageJobState):
completed = dict(job_data.completed_images)
async with prompt_server.distributed_tile_jobs_lock:
prompt_server.distributed_pending_tile_jobs.pop(multi_job_id, None)
return completed
return {}
await asyncio.sleep(poll_interval)
# Node registration
NODE_CLASS_MAPPINGS = {
"UltimateSDUpscaleDistributed": UltimateSDUpscaleDistributed,
"USDUDelegateCollector": USDUDelegateCollector,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"UltimateSDUpscaleDistributed": "Ultimate SD Upscale Distributed (No Upscale)",
# No display name for USDUDelegateCollector — it's an internal node
}
+38
View File
@@ -0,0 +1,38 @@
from typing import Any
def build_worker_identity_hidden_inputs() -> dict[str, tuple[str, dict[str, Any]]]:
"""Hidden-input contract for worker identity context."""
return {
"is_worker": ("BOOLEAN", {"default": False}),
"enabled_worker_ids": ("STRING", {"default": "[]"}),
"worker_id": ("STRING", {"default": ""}),
}
def build_distributed_hidden_inputs(
*,
include_assigned_branch: bool = False,
assigned_branch_max: int = 9,
include_worker_batch_size: bool = False,
include_pass_through: bool = False,
) -> dict[str, tuple[str, dict[str, Any]]]:
hidden_inputs: dict[str, tuple[str, dict[str, Any]]] = {
"multi_job_id": ("STRING", {"default": ""}),
"master_url": ("STRING", {"default": ""}),
"delegate_only": ("BOOLEAN", {"default": False}),
}
hidden_inputs.update(build_worker_identity_hidden_inputs())
if include_assigned_branch:
hidden_inputs["assigned_branch"] = (
"INT",
{"default": -1, "min": -1, "max": assigned_branch_max},
)
if include_worker_batch_size:
hidden_inputs["worker_batch_size"] = (
"INT",
{"default": 1, "min": 1, "max": 1024},
)
if include_pass_through:
hidden_inputs["pass_through"] = ("BOOLEAN", {"default": False})
return hidden_inputs
+46
View File
@@ -0,0 +1,46 @@
import asyncio
import time
from collections.abc import Callable
from typing import Any
from ..utils.config import get_worker_timeout_seconds
from ..utils.constants import HEARTBEAT_INTERVAL
from ..utils.logging import log
async def collect_worker_queue_results(
*,
prompt_server: Any,
multi_job_id: str,
expected_workers: set[str],
on_result: Callable[[dict[str, Any]], None],
timeout_log_prefix: str,
throw_if_interrupted: Callable[[], None],
) -> set[str]:
workers_done: set[str] = set()
base_timeout = float(get_worker_timeout_seconds())
slice_timeout = min(max(0.1, HEARTBEAT_INTERVAL / 20.0), base_timeout)
last_activity = time.time()
while len(workers_done) < len(expected_workers):
throw_if_interrupted()
try:
async with prompt_server.distributed_jobs_lock:
queue = prompt_server.distributed_pending_jobs[multi_job_id]
result = await asyncio.wait_for(queue.get(), timeout=slice_timeout)
except asyncio.TimeoutError:
if (time.time() - last_activity) < base_timeout:
continue
missing_workers = sorted(expected_workers - workers_done)
log(f"{timeout_log_prefix}{missing_workers}")
break
on_result(result)
worker_id_value = str(result.get("worker_id", ""))
is_last = bool(result.get("is_last", False))
last_activity = time.time()
base_timeout = float(get_worker_timeout_seconds())
if is_last and worker_id_value in expected_workers:
workers_done.add(worker_id_value)
return workers_done
+16
View File
@@ -0,0 +1,16 @@
from typing import Any
def get_prompt_server_instance() -> Any:
import server as _server
return _server.PromptServer.instance
def throw_if_processing_interrupted() -> None:
try:
import comfy.model_management as model_management
model_management.throw_exception_if_processing_interrupted()
except Exception:
return
+73 -48
View File
@@ -1,7 +1,10 @@
import torch
import json
from typing import Any
from ..utils.logging import debug_log, log
from ..utils.worker_ids import coerce_enabled_worker_ids, parse_worker_index, worker_value_key
from .hidden_inputs import build_worker_identity_hidden_inputs
def _chunk_bounds(total_items: int, n_splits: int) -> list[tuple[int, int]]:
@@ -28,7 +31,7 @@ class DistributedSeed:
"""
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls: type["DistributedSeed"]) -> dict[str, Any]:
return {
"required": {
"seed": ("INT", {
@@ -39,8 +42,7 @@ class DistributedSeed:
}),
},
"hidden": {
"is_worker": ("BOOLEAN", {"default": False}),
"worker_id": ("STRING", {"default": ""}),
**build_worker_identity_hidden_inputs(),
},
}
@@ -49,30 +51,29 @@ class DistributedSeed:
FUNCTION = "distribute"
CATEGORY = "utils"
def distribute(self, seed, is_worker=False, worker_id=""):
def distribute(
self,
seed: int,
is_worker: bool = False,
worker_id: str = "",
enabled_worker_ids: str = "[]",
) -> tuple[int]:
if not is_worker:
# Master node: pass through original values
debug_log(f"Distributor - Master: seed={seed}")
return (seed,)
else:
# Worker node: apply offset based on worker index
# Find worker index from enabled_worker_ids
try:
# Worker IDs are passed as "worker_0", "worker_1", etc.
if worker_id.startswith("worker_"):
worker_index = int(worker_id.split("_")[1])
else:
# Fallback: try to parse as direct index
worker_index = int(worker_id)
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
worker_index = parse_worker_index(worker_id, enabled_workers)
if worker_index is not None:
offset = worker_index + 1
new_seed = seed + offset
debug_log(f"Distributor - Worker {worker_index}: seed={seed} → {new_seed}")
return (new_seed,)
except (ValueError, IndexError) as e:
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': {e}")
# Fallback: return original seed
return (seed,)
debug_log(f"Distributor - Error parsing worker_id '{worker_id}': no worker index resolved")
# Fallback: return original seed
return (seed,)
# Define ByPassTypeTuple for flexible return types
@@ -92,15 +93,14 @@ class DistributedValue:
"""
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls: type["DistributedValue"]) -> dict[str, Any]:
return {
"required": {
"default_value": ("STRING", {"default": ""}),
"worker_values": ("STRING", {"default": "{}"}),
},
"hidden": {
"is_worker": ("BOOLEAN", {"default": False}),
"worker_id": ("STRING", {"default": ""}),
**build_worker_identity_hidden_inputs(),
},
}
@@ -110,7 +110,7 @@ class DistributedValue:
CATEGORY = "utils"
@staticmethod
def _coerce(value, value_type):
def _coerce(value: Any, value_type: str) -> Any:
"""Convert a string value to the requested type."""
if value_type == "INT":
return int(float(value))
@@ -119,51 +119,71 @@ class DistributedValue:
return value # STRING and COMBO stay as strings
@staticmethod
def _coerce_safe(value, value_type):
def _coerce_safe(value: Any, value_type: str) -> Any:
"""Best-effort coercion with graceful fallback to original value."""
try:
return DistributedValue._coerce(value, value_type)
except (TypeError, ValueError):
return value
def distribute(self, default_value, worker_values="{}", is_worker=False, worker_id=""):
@staticmethod
def _infer_value_type(value: Any) -> str:
"""Infer coercion type from the provided default value."""
if isinstance(value, bool):
return "STRING"
if isinstance(value, int):
return "INT"
if isinstance(value, float):
return "FLOAT"
return "STRING"
def distribute(
self,
default_value: Any,
worker_values: str | dict[str, Any] = "{}",
is_worker: bool = False,
worker_id: str = "",
enabled_worker_ids: str = "[]",
) -> tuple[Any]:
values = {}
value_type = "STRING"
try:
values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
raw_values = json.loads(worker_values) if isinstance(worker_values, str) else worker_values
values = dict(raw_values) if isinstance(raw_values, dict) else {}
if not isinstance(values, dict):
values = {}
except json.JSONDecodeError as e:
debug_log(f"DistributedValue - Error parsing worker_values: {e}")
values = {}
value_type = values.get("_type", "STRING")
inferred_type = self._infer_value_type(default_value)
value_type = values.get("_type", inferred_type)
if value_type not in {"STRING", "COMBO", "INT", "FLOAT"}:
value_type = inferred_type
coerced_default = self._coerce_safe(default_value, value_type)
if not is_worker:
debug_log(f"DistributedValue - Master: returning default '{coerced_default}'")
return (coerced_default,)
try:
if worker_id.startswith("worker_"):
idx = int(worker_id.split("_")[1])
else:
idx = int(worker_id)
key = str(idx + 1) # worker_0 → key "1" (1-indexed)
raw = values.get(key, "")
if raw:
enabled_workers = coerce_enabled_worker_ids(enabled_worker_ids)
direct_key = str(worker_id).strip()
lookup_key = direct_key if direct_key in values else worker_value_key(worker_id, enabled_workers)
raw = values.get(lookup_key)
if raw is not None and raw != "":
try:
coerced = self._coerce(raw, value_type)
debug_log(f"DistributedValue - Worker {idx}: returning '{coerced}'")
debug_log(f"DistributedValue - Worker key {lookup_key}: returning '{coerced}'")
return (coerced,)
except (ValueError, IndexError) as e:
debug_log(f"DistributedValue - Error: {e}")
except (TypeError, ValueError) as e:
debug_log(f"DistributedValue - Error coercing worker value for key {lookup_key}: {e}")
debug_log(f"DistributedValue - Worker fallback: returning default '{coerced_default}'")
return (coerced_default,)
class DistributedModelName:
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls: type["DistributedModelName"]) -> dict[str, Any]:
return {
"required": {
"text": ("STRING", {"default": ""}),
@@ -180,7 +200,7 @@ class DistributedModelName:
OUTPUT_NODE = True
CATEGORY = "utils"
def _stringify(self, value):
def _stringify(self, value: Any) -> str:
if isinstance(value, str):
return value
if isinstance(value, (int, float, bool)):
@@ -190,7 +210,7 @@ class DistributedModelName:
except Exception:
return str(value)
def _update_workflow(self, extra_pnginfo, unique_id, values):
def _update_workflow(self, extra_pnginfo: Any, unique_id: Any, values: list[str]) -> None:
if not extra_pnginfo:
return
info = extra_pnginfo[0] if isinstance(extra_pnginfo, list) else extra_pnginfo
@@ -208,7 +228,12 @@ class DistributedModelName:
if node:
node["widgets_values"] = [values]
def log_input(self, text, unique_id=None, extra_pnginfo=None):
def log_input(
self,
text: Any,
unique_id: Any = None,
extra_pnginfo: Any = None,
) -> dict[str, Any]:
values = []
if isinstance(text, list):
for val in text:
@@ -224,7 +249,7 @@ class DistributedModelName:
return {"ui": {"text": values}, "result": (values,)}
class ByPassTypeTuple(tuple):
def __getitem__(self, index):
def __getitem__(self, index: int) -> Any:
if index > 0:
index = 0
item = super().__getitem__(index)
@@ -234,7 +259,7 @@ class ByPassTypeTuple(tuple):
class ImageBatchDivider:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls: type["ImageBatchDivider"]) -> dict[str, Any]:
return {
"required": {
"images": ("IMAGE",),
@@ -255,7 +280,7 @@ class ImageBatchDivider:
OUTPUT_NODE = True
CATEGORY = "image"
def divide_batch(self, images, divide_by):
def divide_batch(self, images: torch.Tensor, divide_by: int) -> tuple[torch.Tensor, ...]:
total_splits = max(1, min(int(divide_by), 10))
total_frames = images.shape[0]
empty_tensor = images[:0]
@@ -272,7 +297,7 @@ class AudioBatchDivider:
"""Divides an audio waveform into multiple parts along the time/samples dimension."""
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls: type["AudioBatchDivider"]) -> dict[str, Any]:
return {
"required": {
"audio": ("AUDIO",),
@@ -293,7 +318,7 @@ class AudioBatchDivider:
OUTPUT_NODE = True
CATEGORY = "audio"
def divide_audio(self, audio, divide_by):
def divide_audio(self, audio: dict[str, Any], divide_by: int) -> tuple[dict[str, Any], ...]:
import torch
waveform = audio.get("waveform")
@@ -333,7 +358,7 @@ class DistributedEmptyImage:
"""Produces an empty IMAGE batch used when the master delegates all work."""
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls: type["DistributedEmptyImage"]) -> dict[str, Any]:
return {
"required": {
"height": ("INT", {"default": 64, "min": 1, "max": 4096, "step": 1}),
@@ -346,7 +371,7 @@ class DistributedEmptyImage:
FUNCTION = "create"
CATEGORY = "image"
def create(self, height, width, channels):
def create(self, height: int, width: int, channels: int) -> tuple[torch.Tensor]:
import torch
shape = (0, height, width, channels)
+4 -1
View File
@@ -3,7 +3,10 @@ name = "ComfyUI-Distributed"
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
version = "1.4.1"
license = {file = "LICENSE"}
dependencies = []
dependencies = [
"aiohttp>=3.9,<4",
"Pillow>=10,<12",
]
[project.urls]
Repository = "https://github.com/robertvoy/ComfyUI-Distributed"
+163
View File
@@ -0,0 +1,163 @@
"""Shared stubs/utilities for API route module unit tests."""
from __future__ import annotations
import sys
import types
from typing import Any, Callable
class RoutesStub:
"""Simple decorator-style route registrar used in unit tests."""
def get(self, _path: str):
return lambda fn: fn
def post(self, _path: str):
return lambda fn: fn
class ClientTimeoutStub:
def __init__(self, total: float | None = None):
self.total = total
class WSMsgTypeStub:
TEXT = "TEXT"
ERROR = "ERROR"
CLOSED = "CLOSED"
class WebSocketResponseStub:
def __init__(self, *args: Any, **kwargs: Any):
self.args = args
self.kwargs = kwargs
async def prepare(self, _request: Any) -> None:
return None
async def send_json(self, _payload: Any) -> None:
return None
def __aiter__(self):
async def _empty():
if False:
yield None
return _empty()
class FormDataStub:
def add_field(self, *_args: Any, **_kwargs: Any) -> None:
return None
def reset_package_namespace(package_name: str) -> None:
"""Clear test package modules from sys.modules."""
prefix = f"{package_name}."
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(prefix):
del sys.modules[mod_name]
def ensure_namespace_package(module_name: str) -> types.ModuleType:
"""Create a namespace-like package module in sys.modules."""
module = types.ModuleType(module_name)
module.__path__ = []
sys.modules[module_name] = module
return module
def bootstrap_test_package(
package_name: str,
*,
with_api: bool = True,
with_utils: bool = True,
with_workers: bool = False,
with_upscale: bool = False,
with_orchestration: bool = False,
) -> None:
"""Create a clean package scaffold for module-level import tests."""
reset_package_namespace(package_name)
ensure_namespace_package(package_name)
if with_api:
ensure_namespace_package(f"{package_name}.api")
if with_utils:
ensure_namespace_package(f"{package_name}.utils")
if with_workers:
ensure_namespace_package(f"{package_name}.workers")
if with_upscale:
ensure_namespace_package(f"{package_name}.upscale")
if with_orchestration:
ensure_namespace_package(f"{package_name}.api.orchestration")
def install_server_stub(prompt_server_instance: Any | None = None) -> Any:
"""Install a minimal `server` module exposing PromptServer.instance.routes."""
if prompt_server_instance is None:
prompt_server_instance = types.SimpleNamespace()
if not hasattr(prompt_server_instance, "routes"):
prompt_server_instance.routes = RoutesStub()
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
sys.modules["server"] = server_module
return prompt_server_instance
def install_aiohttp_stub(
json_response_factory: Callable[..., Any],
) -> bool:
"""Install/augment `aiohttp` with common members used by API tests."""
created = False
if "aiohttp" not in sys.modules:
created = True
aiohttp_module = types.ModuleType("aiohttp")
sys.modules["aiohttp"] = aiohttp_module
aiohttp_module = sys.modules["aiohttp"]
if not hasattr(aiohttp_module, "ClientTimeout"):
aiohttp_module.ClientTimeout = ClientTimeoutStub
if not hasattr(aiohttp_module, "WSMsgType"):
aiohttp_module.WSMsgType = WSMsgTypeStub
if not hasattr(aiohttp_module, "FormData"):
aiohttp_module.FormData = FormDataStub
web_obj = getattr(aiohttp_module, "web", None)
if web_obj is None:
web_obj = types.SimpleNamespace()
aiohttp_module.web = web_obj
web_obj.json_response = json_response_factory
if not hasattr(web_obj, "WebSocketResponse"):
web_obj.WebSocketResponse = WebSocketResponseStub
return created
def cleanup_optional_module(module_name: str, created: bool) -> None:
"""Remove stubbed module only when this test created it."""
if created:
sys.modules.pop(module_name, None)
def install_request_guards_stub(package_name: str) -> None:
"""Install permissive request guard for isolated route tests."""
request_guards_module = types.ModuleType(f"{package_name}.api.request_guards")
async def _authorization_error_or_none(_request):
return None
request_guards_module.authorization_error_or_none = _authorization_error_or_none
sys.modules[f"{package_name}.api.request_guards"] = request_guards_module
def install_endpoint_policy_passthrough_stub(package_name: str) -> None:
"""Install endpoint policy stub that executes operation directly."""
endpoint_policy_module = types.ModuleType(f"{package_name}.api.endpoint_policy")
async def _run_authorized_endpoint(_request, operation, unexpected_status=500):
_ = unexpected_status
return await operation()
endpoint_policy_module.run_authorized_endpoint = _run_authorized_endpoint
sys.modules[f"{package_name}.api.endpoint_policy"] = endpoint_policy_module
+423
View File
@@ -0,0 +1,423 @@
import asyncio
import copy
import importlib.util
import sys
import types
import unittest
from contextlib import asynccontextmanager
from pathlib import Path
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
install_server_stub,
)
AUTH_TOKEN = "integration-secret"
class _FakeResponse:
def __init__(self, payload, status=200):
self.payload = payload
self.status = status
class _FakeRequest:
def __init__(self, *, headers=None, json_payload=None, post_payload=None, query=None):
self.headers = headers or {}
self._json_payload = json_payload or {}
self._post_payload = post_payload or {}
self.query = query or {}
async def json(self):
return self._json_payload
async def post(self):
return self._post_payload
def _load_module(package_name: str, rel_path: str, module_name: str):
module_path = Path(__file__).resolve().parents[2] / rel_path
spec = importlib.util.spec_from_file_location(f"{package_name}.{module_name}", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _install_common_utils(package_name: str, config_data: dict):
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: copy.deepcopy(config_data)
config_module.save_config = lambda _cfg: True
@asynccontextmanager
async def _config_transaction():
yield config_module.load_config()
config_module.config_transaction = _config_transaction
sys.modules[f"{package_name}.utils.config"] = config_module
network_module = types.ModuleType(f"{package_name}.utils.network")
async def _handle_api_error(_request, error, status=500):
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
network_module.handle_api_error = _handle_api_error
network_module.normalize_host = lambda value: value
network_module.build_worker_url = lambda worker, path="": f"http://{worker.get('host', '127.0.0.1')}{path}"
network_module.get_client_session = lambda: None
network_module.probe_worker = lambda *_args, **_kwargs: {"ok": True}
sys.modules[f"{package_name}.utils.network"] = network_module
def _load_real_guard_stack(package_name: str):
_load_module(package_name, "utils/auth.py", "utils.auth")
_load_module(package_name, "utils/parsing.py", "utils.parsing")
_load_module(package_name, "api/schemas.py", "api.schemas")
_load_module(package_name, "api/request_guards.py", "api.request_guards")
_load_module(package_name, "api/endpoint_policy.py", "api.endpoint_policy")
def _load_request_guards_only_module():
package_name = "dist_auth_guard_testpkg"
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
bootstrap_test_package(package_name, with_api=True, with_utils=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
request_guards = sys.modules[f"{package_name}.api.request_guards"]
schemas = sys.modules[f"{package_name}.api.schemas"]
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return request_guards, schemas
def _load_config_routes_real_guard_module():
package_name = "dist_config_auth_integration_testpkg"
config_data = {
"workers": [],
"master": {"host": "127.0.0.1"},
"settings": {"distributed_api_token": AUTH_TOKEN},
"tunnel": {},
}
bootstrap_test_package(package_name, with_api=True, with_utils=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
module = _load_module(package_name, "api/config_routes.py", "api.config_routes")
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return module
def _load_tunnel_routes_real_guard_module():
package_name = "dist_tunnel_auth_integration_testpkg"
config_data = {
"workers": [],
"master": {"host": "127.0.0.1"},
"settings": {"distributed_api_token": AUTH_TOKEN},
"tunnel": {},
}
bootstrap_test_package(package_name, with_api=True, with_utils=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
cloudflare_module = types.ModuleType(f"{package_name}.utils.cloudflare")
cloudflare_module.cloudflare_tunnel_manager = types.SimpleNamespace(
get_status=lambda: {"active": False},
start_tunnel=lambda: {"active": True},
stop_tunnel=lambda: {"active": False},
)
sys.modules[f"{package_name}.utils.cloudflare"] = cloudflare_module
module = _load_module(package_name, "api/tunnel_routes.py", "api.tunnel_routes")
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return module
def _load_usdu_routes_real_guard_module():
package_name = "dist_usdu_auth_integration_testpkg"
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_upscale=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
usdu_management_module = types.ModuleType(f"{package_name}.utils.usdu_management")
usdu_management_module.MAX_PAYLOAD_SIZE = 1024 * 1024
sys.modules[f"{package_name}.utils.usdu_management"] = usdu_management_module
prompt_server_holder = {
"value": types.SimpleNamespace(
distributed_tile_jobs_lock=asyncio.Lock(),
distributed_pending_tile_jobs={},
)
}
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server_holder["value"]
job_store_module.init_dynamic_job = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
class _BaseJobState:
pass
class _ImageJobState(_BaseJobState):
def __init__(self):
self.mode = "dynamic"
self.queue = asyncio.Queue()
self.pending_images = asyncio.Queue()
self.completed_images = {}
self.worker_status = {}
self.assigned_to_workers = {}
class _TileJobState(_BaseJobState):
def __init__(self):
self.mode = "static"
self.queue = asyncio.Queue()
self.pending_tasks = asyncio.Queue()
self.completed_tasks = {}
self.worker_status = {}
self.assigned_to_workers = {}
self.batched_static = False
job_models_module.BaseJobState = _BaseJobState
job_models_module.ImageJobState = _ImageJobState
job_models_module.TileJobState = _TileJobState
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
parsers_module = types.ModuleType(f"{package_name}.upscale.payload_parsers")
parsers_module.parse_tiles_from_form = lambda _data: []
sys.modules[f"{package_name}.upscale.payload_parsers"] = parsers_module
module = _load_module(package_name, "api/usdu_routes.py", "api.usdu_routes")
module._prompt_server_holder = prompt_server_holder
module._ImageJobState = _ImageJobState
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return module
def _install_torch_stub_if_missing():
if "torch" in sys.modules:
return False
torch_module = types.ModuleType("torch")
torch_module.cuda = types.SimpleNamespace(
is_available=lambda: False,
empty_cache=lambda: None,
ipc_collect=lambda: None,
)
sys.modules["torch"] = torch_module
return True
def _load_worker_routes_real_guard_module():
package_name = "dist_worker_auth_integration_testpkg"
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}, "workers": []}
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_workers=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
created_torch_stub = _install_torch_stub_if_missing()
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.CHUNK_SIZE = 65536
sys.modules[f"{package_name}.utils.constants"] = constants_module
workers_module = types.ModuleType(f"{package_name}.workers")
workers_module.get_worker_manager = lambda: types.SimpleNamespace(
processes={},
save_processes=lambda: None,
is_process_running=lambda _pid: False,
)
sys.modules[f"{package_name}.workers"] = workers_module
detection_module = types.ModuleType(f"{package_name}.workers.detection")
detection_module.get_machine_id = lambda: "machine-id"
detection_module.is_docker_environment = lambda: False
detection_module.is_runpod_environment = lambda: False
sys.modules[f"{package_name}.workers.detection"] = detection_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
class _PromptValidationError(Exception):
def __init__(self, message="validation error"):
super().__init__(message)
self.validation_error = {}
self.node_errors = {}
async def _queue_prompt_payload(*_args, **_kwargs):
return "prompt-id"
async_helpers_module.PromptValidationError = _PromptValidationError
async_helpers_module.queue_prompt_payload = _queue_prompt_payload
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
module = _load_module(package_name, "api/worker_routes.py", "api.worker_routes")
cleanup_optional_module("aiohttp", created_aiohttp_stub)
if created_torch_stub:
sys.modules.pop("torch", None)
return module
def _load_job_routes_real_guard_module():
package_name = "dist_job_auth_integration_testpkg"
config_data = {"settings": {"distributed_api_token": AUTH_TOKEN}}
bootstrap_test_package(package_name, with_api=True, with_utils=True)
install_server_stub()
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
created_torch_stub = _install_torch_stub_if_missing()
_install_common_utils(package_name, config_data)
_load_real_guard_stack(package_name)
image_module = types.ModuleType(f"{package_name}.utils.image")
image_module.pil_to_tensor = lambda _img: None
image_module.ensure_contiguous = lambda tensor: tensor
sys.modules[f"{package_name}.utils.image"] = image_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.JOB_INIT_GRACE_PERIOD = 0.1
constants_module.MEMORY_CLEAR_DELAY = 0.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
runtime_state = types.SimpleNamespace(
distributed_jobs_lock=asyncio.Lock(),
distributed_pending_jobs={},
)
runtime_state_module.ensure_distributed_runtime_state = lambda: runtime_state
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
orchestration_module = types.ModuleType(f"{package_name}.api.queue_orchestration")
orchestration_module.orchestrate_distributed_execution = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.api.queue_orchestration"] = orchestration_module
queue_request_module = types.ModuleType(f"{package_name}.api.queue_request")
queue_request_module.parse_queue_request_payload = lambda data: data
sys.modules[f"{package_name}.api.queue_request"] = queue_request_module
module = _load_module(package_name, "api/job_routes.py", "api.job_routes")
cleanup_optional_module("aiohttp", created_aiohttp_stub)
if created_torch_stub:
sys.modules.pop("torch", None)
return module
request_guards_module, schemas_module = _load_request_guards_only_module()
config_routes_module = _load_config_routes_real_guard_module()
tunnel_routes_module = _load_tunnel_routes_real_guard_module()
usdu_routes_module = _load_usdu_routes_real_guard_module()
worker_routes_module = _load_worker_routes_real_guard_module()
job_routes_module = _load_job_routes_real_guard_module()
class AuthGuardContractTests(unittest.IsolatedAsyncioTestCase):
async def test_is_authorized_request_missing_token_fails(self):
request = _FakeRequest(headers={})
config = {"settings": {"distributed_api_token": AUTH_TOKEN}}
self.assertFalse(schemas_module.is_authorized_request(request, config))
async def test_is_authorized_request_matching_token_succeeds(self):
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN})
config = {"settings": {"distributed_api_token": AUTH_TOKEN}}
self.assertTrue(schemas_module.is_authorized_request(request, config))
async def test_authorization_error_or_none_returns_403_for_invalid_token(self):
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: "wrong-token"})
response = await request_guards_module.authorization_error_or_none(request)
self.assertIsNotNone(response)
self.assertEqual(response.status, 403)
async def test_authorization_error_or_none_returns_none_for_valid_token(self):
request = _FakeRequest(headers={schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN})
response = await request_guards_module.authorization_error_or_none(request)
self.assertIsNone(response)
class RouteAuthIntegrationTests(unittest.IsolatedAsyncioTestCase):
def _auth_headers(self):
return {schemas_module.AUTH_HEADER_NAME: AUTH_TOKEN}
async def test_config_endpoint_enforces_real_guard(self):
unauthorized = await config_routes_module.get_config_endpoint(_FakeRequest(headers={}))
self.assertEqual(unauthorized.status, 403)
authorized = await config_routes_module.get_config_endpoint(_FakeRequest(headers=self._auth_headers()))
self.assertEqual(authorized.status, 200)
self.assertIn("settings", authorized.payload)
async def test_tunnel_endpoint_enforces_real_guard(self):
unauthorized = await tunnel_routes_module.tunnel_status_endpoint(_FakeRequest(headers={}))
self.assertEqual(unauthorized.status, 403)
authorized = await tunnel_routes_module.tunnel_status_endpoint(_FakeRequest(headers=self._auth_headers()))
self.assertEqual(authorized.status, 200)
self.assertEqual(authorized.payload.get("status"), "success")
async def test_usdu_endpoint_enforces_real_guard(self):
unauthorized = await usdu_routes_module.job_status_endpoint(_FakeRequest(headers={}))
self.assertEqual(unauthorized.status, 403)
authorized = await usdu_routes_module.job_status_endpoint(
_FakeRequest(headers=self._auth_headers(), query={"multi_job_id": "job-1"})
)
self.assertEqual(authorized.status, 200)
self.assertIn("ready", authorized.payload)
async def test_worker_endpoint_enforces_real_guard(self):
worker_routes_module._collect_network_info_sync = lambda: {
"interfaces": [],
"recommended_host": "127.0.0.1",
}
unauthorized = await worker_routes_module.get_network_info_endpoint(_FakeRequest(headers={}))
self.assertEqual(unauthorized.status, 403)
authorized = await worker_routes_module.get_network_info_endpoint(
_FakeRequest(headers=self._auth_headers())
)
self.assertEqual(authorized.status, 200)
self.assertEqual(authorized.payload.get("status"), "success")
async def test_job_endpoint_enforces_real_guard(self):
unauthorized = await job_routes_module.prepare_job_endpoint(
_FakeRequest(headers={}, json_payload={"multi_job_id": "job-1"})
)
self.assertEqual(unauthorized.status, 403)
authorized = await job_routes_module.prepare_job_endpoint(
_FakeRequest(headers=self._auth_headers(), json_payload={"multi_job_id": "job-1"})
)
self.assertEqual(authorized.status, 200)
self.assertEqual(authorized.payload.get("status"), "success")
if __name__ == "__main__":
unittest.main()
+38 -43
View File
@@ -3,9 +3,19 @@ import importlib.util
import sys
import types
import unittest
from contextlib import asynccontextmanager
from pathlib import Path
from unittest.mock import patch
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
install_endpoint_policy_passthrough_stub,
install_request_guards_stub,
install_server_stub,
)
class _FakeResponse:
def __init__(self, payload, status=200):
@@ -25,45 +35,18 @@ def _load_config_routes_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "config_routes.py"
package_name = "dist_api_config_testpkg"
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
bootstrap_test_package(package_name, with_api=True, with_utils=True)
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
schemas_module.is_authorized_request = lambda _request, _config: True
sys.modules[f"{package_name}.api.schemas"] = schemas_module
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
created_aiohttp_stub = False
if "aiohttp" not in sys.modules:
created_aiohttp_stub = True
aiohttp_module = types.ModuleType("aiohttp")
aiohttp_module.web = types.SimpleNamespace(
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
)
sys.modules["aiohttp"] = aiohttp_module
class _Routes:
def get(self, _path):
def _decorator(fn):
return fn
return _decorator
def post(self, _path):
def _decorator(fn):
return fn
return _decorator
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
sys.modules["server"] = server_module
install_request_guards_stub(package_name)
install_endpoint_policy_passthrough_stub(package_name)
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
install_server_stub()
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
@@ -73,7 +56,16 @@ def _load_config_routes_module():
network_module = types.ModuleType(f"{package_name}.utils.network")
async def _handle_api_error(_request, error, status=500):
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
if isinstance(error, list):
message = "; ".join(str(item) for item in error)
error_payload = [str(item) for item in error]
else:
message = str(error)
error_payload = str(error)
return _FakeResponse(
{"status": "error", "error": error_payload, "message": message},
status=status,
)
network_module.handle_api_error = _handle_api_error
network_module.normalize_host = lambda value: value
@@ -89,6 +81,12 @@ def _load_config_routes_module():
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: copy.deepcopy(default_config)
config_module.save_config = lambda _cfg: True
@asynccontextmanager
async def _config_transaction():
yield config_module.load_config()
config_module.config_transaction = _config_transaction
sys.modules[f"{package_name}.utils.config"] = config_module
spec = importlib.util.spec_from_file_location(f"{package_name}.api.config_routes", module_path)
@@ -96,8 +94,7 @@ def _load_config_routes_module():
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
if created_aiohttp_stub:
sys.modules.pop("aiohttp", None)
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return module
@@ -118,9 +115,7 @@ class ConfigRoutesTests(unittest.IsolatedAsyncioTestCase):
async def test_update_config_valid_field_persists(self):
cfg = {"workers": [], "master": {}, "settings": {"debug": False}, "tunnel": {}}
with patch.object(config_routes, "load_config", return_value=cfg), patch.object(
config_routes, "save_config", return_value=True
):
with patch.object(config_routes, "load_config", return_value=cfg):
response = await config_routes.update_config_endpoint(_FakeRequest({"debug": True}))
self.assertEqual(response.status, 200)
@@ -11,6 +11,14 @@ from unittest.mock import AsyncMock, patch
import numpy as np
import torch
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
install_request_guards_stub,
install_server_stub,
)
class _FakeResponse:
def __init__(self, payload, status=200):
@@ -19,8 +27,9 @@ class _FakeResponse:
class _FakeRequest:
def __init__(self, payload):
def __init__(self, payload, headers=None):
self._payload = payload
self.headers = headers or {}
async def json(self):
return self._payload
@@ -30,53 +39,19 @@ def _load_job_routes_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "job_routes.py"
package_name = "dist_api_queue_testpkg"
# Reset package namespace to avoid stale module state across test runs.
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
bootstrap_test_package(package_name, with_api=True, with_utils=True)
# aiohttp.web stub
created_aiohttp_stub = False
if "aiohttp" not in sys.modules:
created_aiohttp_stub = True
aiohttp_module = types.ModuleType("aiohttp")
aiohttp_module.web = types.SimpleNamespace(
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
)
sys.modules["aiohttp"] = aiohttp_module
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
# server module stub with route decorators
class _Routes:
def get(self, _path):
def _decorator(fn):
return fn
return _decorator
def post(self, _path):
def _decorator(fn):
return fn
return _decorator
prompt_server_instance = types.SimpleNamespace(
routes=_Routes(),
distributed_jobs_lock=None,
distributed_pending_jobs={},
)
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server_instance)
sys.modules["server"] = server_module
install_server_stub(prompt_server_instance)
# torch stub (only needed to satisfy import)
created_torch_stub = False
@@ -157,6 +132,47 @@ def _load_job_routes_module():
constants_module.JOB_INIT_GRACE_PERIOD = 10.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"settings": {}}
sys.modules[f"{package_name}.utils.config"] = config_module
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
def _ensure_distributed_runtime_state(server_instance=None):
ps = server_instance or prompt_server_instance
if not hasattr(ps, "distributed_pending_jobs"):
ps.distributed_pending_jobs = {}
if not hasattr(ps, "distributed_jobs_lock") or ps.distributed_jobs_lock is None:
ps.distributed_jobs_lock = asyncio.Lock()
if not hasattr(ps, "distributed_job_allowed_workers"):
ps.distributed_job_allowed_workers = {}
if not hasattr(ps, "distributed_pending_tile_jobs"):
ps.distributed_pending_tile_jobs = {}
if not hasattr(ps, "distributed_tile_jobs_lock"):
ps.distributed_tile_jobs_lock = asyncio.Lock()
return types.SimpleNamespace(
distributed_pending_jobs=ps.distributed_pending_jobs,
distributed_jobs_lock=ps.distributed_jobs_lock,
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
)
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
runtime_state_module.get_prompt_server_instance = lambda: prompt_server_instance
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
schemas_module.is_authorized_request = lambda _request, _config: True
schemas_module.require_bool_literal = (
lambda value, field_name="value": value
if isinstance(value, bool)
else (_ for _ in ()).throw(ValueError(f"Field '{field_name}' must be a boolean literal."))
)
sys.modules[f"{package_name}.api.schemas"] = schemas_module
install_request_guards_stub(package_name)
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.queue_prompt_payload = AsyncMock(return_value="prompt_local")
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
@@ -172,7 +188,6 @@ def _load_job_routes_module():
client_id: str
delegate_master: object
enabled_worker_ids: list
auto_prepare: bool
trace_execution_id: object
def _parse_queue_request_payload(data):
@@ -193,7 +208,6 @@ def _load_job_routes_module():
client_id=client_id,
delegate_master=data.get("delegate_master"),
enabled_worker_ids=enabled,
auto_prepare=bool(data.get("auto_prepare", True)),
trace_execution_id=data.get("trace_execution_id"),
)
@@ -206,13 +220,11 @@ def _load_job_routes_module():
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
if created_aiohttp_stub:
sys.modules.pop("aiohttp", None)
if created_torch_stub:
sys.modules.pop("torch", None)
cleanup_optional_module("aiohttp", created_aiohttp_stub)
cleanup_optional_module("torch", created_torch_stub)
if created_pil_stub:
sys.modules.pop("PIL.Image", None)
sys.modules.pop("PIL", None)
cleanup_optional_module("PIL.Image", True)
cleanup_optional_module("PIL", True)
return module
@@ -239,7 +251,6 @@ class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(response.status, 200)
self.assertEqual(response.payload.get("prompt_id"), "prompt_123")
self.assertTrue(response.payload.get("auto_prepare_supported"))
async def test_distributed_queue_missing_prompt_returns_400(self):
request = _FakeRequest(
@@ -276,8 +287,10 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
async def test_job_complete_accepts_audio_payload(self):
queue = asyncio.Queue()
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
job_routes.prompt_server.distributed_pending_jobs = {"job-1": queue}
prompt_server = job_routes.server.PromptServer.instance
prompt_server.distributed_jobs_lock = asyncio.Lock()
prompt_server.distributed_pending_jobs = {"job-1": queue}
prompt_server.distributed_job_allowed_workers = {"job-1": {"worker-1"}}
request = _FakeRequest(
{
"job_id": "job-1",
@@ -300,6 +313,28 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(queued["audio"]["sample_rate"], 44100)
self.assertEqual(tuple(queued["audio"]["waveform"].shape), (1, 2, 4))
async def test_job_complete_rejects_worker_not_in_job_allowlist(self):
queue = asyncio.Queue()
prompt_server = job_routes.server.PromptServer.instance
prompt_server.distributed_jobs_lock = asyncio.Lock()
prompt_server.distributed_pending_jobs = {"job-allow": queue}
prompt_server.distributed_job_allowed_workers = {"job-allow": {"worker-expected"}}
request = _FakeRequest(
{
"job_id": "job-allow",
"worker_id": "worker-unexpected",
"batch_idx": 0,
"image": "data:image/png;base64,AAAA",
"is_last": True,
}
)
with patch.object(job_routes, "_decode_canonical_png_tensor", return_value="tensor-data"):
response = await job_routes.job_complete_endpoint(request)
self.assertEqual(response.status, 403)
self.assertIn("unauthorized", response.payload.get("message", "").lower())
def test_decode_audio_payload_rejects_bad_shape(self):
bad = {
"sample_rate": 44100,
+32 -35
View File
@@ -4,30 +4,29 @@ import types
import unittest
from pathlib import Path
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
)
class _AiohttpResponse:
def __init__(self, payload, status=200):
self.payload = payload
self.status = status
def _load_media_sync_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "orchestration" / "media_sync.py"
package_name = "dist_ms_testpkg"
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
orch_pkg = types.ModuleType(f"{package_name}.api.orchestration")
orch_pkg.__path__ = []
sys.modules[f"{package_name}.api.orchestration"] = orch_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
bootstrap_test_package(
package_name,
with_api=True,
with_utils=True,
with_orchestration=True,
)
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
@@ -48,22 +47,17 @@ def _load_media_sync_module():
trace_module.trace_info = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.trace_logger"] = trace_module
created_aiohttp_stub = False
if "aiohttp" not in sys.modules:
created_aiohttp_stub = True
aiohttp_module = types.ModuleType("aiohttp")
auth_module = types.ModuleType(f"{package_name}.utils.auth")
auth_module.distributed_auth_headers = lambda _config=None: {}
sys.modules[f"{package_name}.utils.auth"] = auth_module
class _ClientTimeout:
def __init__(self, total=None):
pass
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"settings": {}}
sys.modules[f"{package_name}.utils.config"] = config_module
class _FormData:
def add_field(self, *args, **kwargs):
pass
aiohttp_module.ClientTimeout = _ClientTimeout
aiohttp_module.FormData = _FormData
sys.modules["aiohttp"] = aiohttp_module
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _AiohttpResponse(payload, status=status)
)
spec = importlib.util.spec_from_file_location(
f"{package_name}.api.orchestration.media_sync",
@@ -73,8 +67,7 @@ def _load_media_sync_module():
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
if created_aiohttp_stub:
sys.modules.pop("aiohttp", None)
cleanup_optional_module("aiohttp", created_aiohttp_stub)
return module
@@ -260,3 +253,7 @@ class RewritePromptMediaInputsTests(unittest.TestCase):
if __name__ == "__main__":
unittest.main()
class _AiohttpResponse:
def __init__(self, payload, status=200):
self.payload = payload
self.status = status
+195
View File
@@ -0,0 +1,195 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
from tests.api.harness import (
bootstrap_test_package,
install_server_stub,
)
def _bootstrap_package(package_name):
bootstrap_test_package(
package_name,
with_api=True,
with_utils=True,
with_orchestration=True,
)
prompt_server_instance = types.SimpleNamespace(address="127.0.0.1", port=8188)
install_server_stub(prompt_server_instance)
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async def _queue_prompt_payload(prompt_obj, workflow_meta, client_id):
_ = (prompt_obj, workflow_meta, client_id)
return "prompt-id"
async_helpers_module.queue_prompt_payload = _queue_prompt_payload
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"settings": {}, "workers": []}
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.ORCHESTRATION_MEDIA_SYNC_CONCURRENCY = 2
constants_module.ORCHESTRATION_MEDIA_SYNC_TIMEOUT = 120.0
constants_module.ORCHESTRATION_WORKER_PROBE_CONCURRENCY = 8
constants_module.ORCHESTRATION_WORKER_PREP_CONCURRENCY = 4
sys.modules[f"{package_name}.utils.constants"] = constants_module
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.build_master_url = lambda **_kwargs: "http://127.0.0.1:8188"
sys.modules[f"{package_name}.utils.network"] = network_module
trace_logger_module = types.ModuleType(f"{package_name}.utils.trace_logger")
trace_logger_module.trace_debug = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.trace_logger"] = trace_logger_module
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
def _ensure_distributed_runtime_state(server_instance=None):
ps = server_instance or prompt_server_instance
if not hasattr(ps, "distributed_pending_jobs"):
ps.distributed_pending_jobs = {}
if not hasattr(ps, "distributed_jobs_lock"):
ps.distributed_jobs_lock = asyncio.Lock()
if not hasattr(ps, "distributed_job_allowed_workers"):
ps.distributed_job_allowed_workers = {}
if not hasattr(ps, "distributed_pending_tile_jobs"):
ps.distributed_pending_tile_jobs = {}
if not hasattr(ps, "distributed_tile_jobs_lock"):
ps.distributed_tile_jobs_lock = asyncio.Lock()
return types.SimpleNamespace(
distributed_pending_jobs=ps.distributed_pending_jobs,
distributed_jobs_lock=ps.distributed_jobs_lock,
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
)
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
runtime_state_module.get_prompt_server_instance = lambda: prompt_server_instance
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
schemas_module.coerce_positive_int = (
lambda value, default: int(value) if str(value).isdigit() and int(value) > 0 else default
)
schemas_module.coerce_positive_float = (
lambda value, default: float(value) if isinstance(value, (int, float, str)) and float(value) > 0 else default
)
sys.modules[f"{package_name}.api.schemas"] = schemas_module
dispatch_module = types.ModuleType(f"{package_name}.api.orchestration.dispatch")
dispatch_module.dispatch_worker_prompt = lambda *args, **kwargs: None
dispatch_module.rank_workers_by_load = lambda workers, **kwargs: workers
dispatch_module.select_active_workers = lambda workers, *_args, **_kwargs: (workers, False)
dispatch_module.select_least_busy_worker = lambda workers, **_kwargs: workers[0] if workers else None
sys.modules[f"{package_name}.api.orchestration.dispatch"] = dispatch_module
media_sync_module = types.ModuleType(f"{package_name}.api.orchestration.media_sync")
media_sync_module.convert_paths_for_platform = lambda prompt, _sep: prompt
media_sync_module.fetch_worker_path_separator = lambda *_args, **_kwargs: "/"
media_sync_module.sync_worker_media = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.api.orchestration.media_sync"] = media_sync_module
prompt_transform_module = types.ModuleType(f"{package_name}.api.orchestration.prompt_transform")
prompt_transform_module.PromptIndex = object
prompt_transform_module.apply_participant_overrides = lambda prompt, *_args, **_kwargs: prompt
prompt_transform_module.find_nodes_by_class = lambda _prompt, _class_name: []
prompt_transform_module.generate_job_id_map = lambda _index, _prefix: {}
prompt_transform_module.prepare_delegate_master_prompt = lambda prompt, _ids: prompt
prompt_transform_module.prune_prompt_for_worker = lambda prompt: prompt
sys.modules[f"{package_name}.api.orchestration.prompt_transform"] = prompt_transform_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[2] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_queue_orchestration_module():
package_name = "dist_queue_orch_testpkg"
_bootstrap_package(package_name)
return _load_module(package_name, "api/queue_orchestration.py", "api.queue_orchestration")
queue_orchestration = _load_queue_orchestration_module()
class QueueOrchestrationHelpersTests(unittest.TestCase):
def test_resolve_enabled_workers_filters_disabled_and_requested(self):
config = {
"workers": [
{"id": "1", "enabled": True, "host": "a", "port": 8188},
{"id": "2", "enabled": False, "host": "b", "port": 8189},
{"id": "3", "enabled": True, "host": "c", "port": "bad-port"},
]
}
enabled_only = queue_orchestration._resolve_enabled_workers(config)
requested = queue_orchestration._resolve_enabled_workers(config, requested_ids={"2", "3"})
self.assertEqual([worker["id"] for worker in enabled_only], ["1", "3"])
self.assertEqual([worker["id"] for worker in requested], ["2", "3"])
self.assertEqual(requested[1]["port"], 8188)
def test_resolve_orchestration_limits_uses_defaults_for_invalid_values(self):
config = {
"settings": {
"worker_probe_concurrency": "x",
"worker_prep_concurrency": "3",
"media_sync_concurrency": "-1",
"media_sync_timeout_seconds": "42.5",
}
}
limits = queue_orchestration._resolve_orchestration_limits(config)
self.assertEqual(limits[0], 8) # default
self.assertEqual(limits[1], 3) # parsed
self.assertEqual(limits[2], 2) # default
self.assertEqual(limits[3], 42.5)
def test_is_load_balance_enabled_accepts_common_forms(self):
self.assertTrue(queue_orchestration._is_load_balance_enabled(True))
self.assertTrue(queue_orchestration._is_load_balance_enabled(1))
self.assertTrue(queue_orchestration._is_load_balance_enabled(" yes "))
self.assertFalse(queue_orchestration._is_load_balance_enabled(0))
self.assertFalse(queue_orchestration._is_load_balance_enabled("off"))
def test_prompt_requests_load_balance_detects_collector_flag(self):
prompt_index = types.SimpleNamespace(
nodes_for_class=lambda class_name: ["10"] if class_name == "DistributedCollector" else [],
inputs_by_node={"10": {"load_balance": "true"}},
)
self.assertTrue(queue_orchestration._prompt_requests_load_balance(prompt_index))
def test_ensure_distributed_state_initializes_attrs(self):
prompt_server = types.SimpleNamespace()
queue_orchestration.ensure_distributed_state(prompt_server)
self.assertTrue(hasattr(prompt_server, "distributed_pending_jobs"))
self.assertTrue(hasattr(prompt_server, "distributed_jobs_lock"))
self.assertIsInstance(prompt_server.distributed_jobs_lock, asyncio.Lock)
if __name__ == "__main__":
unittest.main()
+95
View File
@@ -0,0 +1,95 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
from tests.api.harness import (
bootstrap_test_package,
install_aiohttp_stub,
install_endpoint_policy_passthrough_stub,
install_request_guards_stub,
install_server_stub,
)
class _FakeResponse:
def __init__(self, payload, status=200):
self.payload = payload
self.status = status
def _load_tunnel_routes_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "tunnel_routes.py"
package_name = "dist_tunnel_routes_testpkg"
bootstrap_test_package(package_name, with_api=True, with_utils=True)
install_server_stub()
install_aiohttp_stub(lambda payload, status=200: _FakeResponse(payload, status=status))
cloudflare_module = types.ModuleType(f"{package_name}.utils.cloudflare")
class _TunnelManager:
def get_status(self):
return {"status": "running"}
async def start_tunnel(self):
return {"status": "running", "public_url": "https://example.trycloudflare.com"}
async def stop_tunnel(self):
return {"status": "stopped"}
cloudflare_module.cloudflare_tunnel_manager = _TunnelManager()
sys.modules[f"{package_name}.utils.cloudflare"] = cloudflare_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"master": {"host": "https://master.example.com"}}
sys.modules[f"{package_name}.utils.config"] = config_module
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
network_module = types.ModuleType(f"{package_name}.utils.network")
async def _handle_api_error(_request, error, status=500):
return _FakeResponse({"status": "error", "message": str(error)}, status=status)
network_module.handle_api_error = _handle_api_error
sys.modules[f"{package_name}.utils.network"] = network_module
install_request_guards_stub(package_name)
install_endpoint_policy_passthrough_stub(package_name)
spec = importlib.util.spec_from_file_location(f"{package_name}.api.tunnel_routes", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
tunnel_routes = _load_tunnel_routes_module()
class TunnelRoutesTests(unittest.IsolatedAsyncioTestCase):
async def test_status_endpoint_returns_tunnel_and_master_host(self):
response = await tunnel_routes.tunnel_status_endpoint(request=None)
self.assertEqual(response.status, 200)
self.assertEqual(response.payload["status"], "success")
self.assertEqual(response.payload["tunnel"]["status"], "running")
self.assertEqual(response.payload["master_host"], "https://master.example.com")
async def test_start_and_stop_endpoints_return_success_payload(self):
start_response = await tunnel_routes.tunnel_start_endpoint(request=None)
stop_response = await tunnel_routes.tunnel_stop_endpoint(request=None)
self.assertEqual(start_response.status, 200)
self.assertEqual(start_response.payload["tunnel"]["status"], "running")
self.assertEqual(stop_response.status, 200)
self.assertEqual(stop_response.payload["tunnel"]["status"], "stopped")
if __name__ == "__main__":
unittest.main()
+57 -45
View File
@@ -9,6 +9,14 @@ from pathlib import Path
from PIL import Image
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
install_request_guards_stub,
install_server_stub,
)
class _FakeResponse:
def __init__(self, payload, status=200):
@@ -30,43 +38,30 @@ class _FakeRequest:
return self._post_payload
class _Routes:
def post(self, _path):
def _decorator(fn):
return fn
return _decorator
def get(self, _path):
def _decorator(fn):
return fn
return _decorator
def _load_usdu_routes_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "usdu_routes.py"
package_name = "dist_api_usdu_testpkg"
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_upscale=True)
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
schemas_module = types.ModuleType(f"{package_name}.api.schemas")
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
def _require_bool_literal(value, field_name: str = "value"):
if isinstance(value, bool):
return value
if isinstance(value, str):
normalized = value.strip().lower()
if normalized == "true":
return True
if normalized == "false":
return False
raise ValueError(f"Field '{field_name}' must be a boolean literal.")
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
schemas_module.require_bool_literal = _require_bool_literal
schemas_module.is_authorized_request = lambda _request, _config: True
sys.modules[f"{package_name}.api.schemas"] = schemas_module
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
install_request_guards_stub(package_name)
prompt_server_holder = {
"value": types.SimpleNamespace(
@@ -75,23 +70,23 @@ def _load_usdu_routes_module():
)
}
created_aiohttp_stub = False
if "aiohttp" not in sys.modules:
created_aiohttp_stub = True
aiohttp_module = types.ModuleType("aiohttp")
aiohttp_module.web = types.SimpleNamespace(
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
)
sys.modules["aiohttp"] = aiohttp_module
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
sys.modules["server"] = server_module
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
install_server_stub()
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"settings": {}}
sys.modules[f"{package_name}.utils.config"] = config_module
usdu_management_module = types.ModuleType(f"{package_name}.utils.usdu_management")
usdu_management_module.MAX_PAYLOAD_SIZE = 1024
sys.modules[f"{package_name}.utils.usdu_management"] = usdu_management_module
network_module = types.ModuleType(f"{package_name}.utils.network")
async def _handle_api_error(_request, error, status=500):
@@ -103,6 +98,21 @@ def _load_usdu_routes_module():
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
job_store_module.MAX_PAYLOAD_SIZE = 1024
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server_holder["value"]
async def _init_dynamic_job(multi_job_id, batch_size, enabled_workers, all_indices=None):
prompt_server = prompt_server_holder["value"]
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
return
job_data = ImageJobState(multi_job_id=multi_job_id)
job_data.worker_status = {str(worker_id): 0 for worker_id in (enabled_workers or [])}
job_data.assigned_to_workers = {str(worker_id): [] for worker_id in (enabled_workers or [])}
indices = all_indices if all_indices is not None else list(range(int(batch_size or 0)))
for idx in indices:
await job_data.pending_images.put(int(idx))
prompt_server.distributed_pending_tile_jobs[multi_job_id] = job_data
job_store_module.init_dynamic_job = _init_dynamic_job
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
@@ -150,6 +160,7 @@ def _load_usdu_routes_module():
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
parsers_module = types.ModuleType(f"{package_name}.upscale.payload_parsers")
parsers_module.parse_tiles_from_form = lambda _data: []
parsers_module._parse_tiles_from_form = lambda _data: []
sys.modules[f"{package_name}.upscale.payload_parsers"] = parsers_module
@@ -158,8 +169,7 @@ def _load_usdu_routes_module():
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
if created_aiohttp_stub:
sys.modules.pop("aiohttp", None)
cleanup_optional_module("aiohttp", created_aiohttp_stub)
module.web = types.SimpleNamespace(
json_response=lambda payload, status=200: _FakeResponse(payload, status=status)
@@ -222,7 +232,8 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
response = await usdu_routes.request_image_endpoint(request)
self.assertEqual(response.status, 200)
self.assertEqual(response.payload.get("image_idx"), 7)
self.assertEqual(response.payload.get("kind"), "image")
self.assertEqual(response.payload.get("task_idx"), 7)
self.assertEqual(response.payload.get("estimated_remaining"), 0)
self.assertEqual(job_data.assigned_to_workers["worker-a"], [7])
self.assertIn("worker-a", job_data.worker_status)
@@ -238,7 +249,8 @@ class USDURoutesTests(unittest.IsolatedAsyncioTestCase):
response = await usdu_routes.request_image_endpoint(request)
self.assertEqual(response.status, 200)
self.assertEqual(response.payload.get("tile_idx"), 4)
self.assertEqual(response.payload.get("kind"), "tile")
self.assertEqual(response.payload.get("task_idx"), 4)
self.assertTrue(response.payload.get("batched_static"))
self.assertEqual(job_data.assigned_to_workers["worker-a"], [4])
+33 -77
View File
@@ -8,6 +8,14 @@ from collections import deque
from pathlib import Path
from unittest.mock import patch
from tests.api.harness import (
bootstrap_test_package,
cleanup_optional_module,
install_aiohttp_stub,
install_request_guards_stub,
install_server_stub,
)
class _FakeResponse:
def __init__(self, payload, status=200):
@@ -16,10 +24,11 @@ class _FakeResponse:
class _FakeRequest:
def __init__(self, payload=None, match_info=None, query=None):
def __init__(self, payload=None, match_info=None, query=None, headers=None):
self._payload = payload
self.match_info = match_info or {}
self.query = query or {}
self.headers = headers or {}
async def json(self):
return self._payload
@@ -49,8 +58,8 @@ class _FakeHTTPClientSession:
self._status = status
self.calls = []
def get(self, url, params=None, timeout=None):
self.calls.append({"url": url, "params": params, "timeout": timeout})
def get(self, url, params=None, headers=None, timeout=None):
self.calls.append({"url": url, "params": params, "headers": headers, "timeout": timeout})
return _FakeHTTPClientResponse(self._payload, status=self._status)
@@ -62,12 +71,12 @@ class _DummyWorkerManager:
worker_id = str(worker["id"])
self.processes[worker_id] = {
"pid": 12345,
"log_file": f"/tmp/distributed_worker_{worker_id}.log",
"log_file": f"/tmp/distributed_worker_{worker_id}.log", # nosec B108 - deterministic fake test path
"process": None,
}
return 12345
def _is_process_running(self, _pid):
def is_process_running(self, _pid):
return False
def save_processes(self):
@@ -89,21 +98,7 @@ def _load_worker_routes_module():
module_path = Path(__file__).resolve().parents[2] / "api" / "worker_routes.py"
package_name = "dist_api_worker_testpkg"
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
api_pkg = types.ModuleType(f"{package_name}.api")
api_pkg.__path__ = []
sys.modules[f"{package_name}.api"] = api_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
bootstrap_test_package(package_name, with_api=True, with_utils=True, with_workers=True)
workers_pkg = types.ModuleType(f"{package_name}.workers")
workers_pkg.__path__ = []
@@ -119,59 +114,10 @@ def _load_worker_routes_module():
detection_module.get_comms_channel = lambda *_args, **_kwargs: "lan"
sys.modules[f"{package_name}.workers.detection"] = detection_module
created_aiohttp_stub = False
if "aiohttp" not in sys.modules:
created_aiohttp_stub = True
aiohttp_module = types.ModuleType("aiohttp")
class _ClientTimeout:
def __init__(self, total=None):
self.total = total
class _WSMsgType:
TEXT = "TEXT"
ERROR = "ERROR"
CLOSED = "CLOSED"
class _WebSocketResponse:
def __init__(self, *args, **kwargs):
self.args = args
self.kwargs = kwargs
async def prepare(self, _request):
return None
async def send_json(self, _payload):
return None
def __aiter__(self):
async def _empty():
if False:
yield None
return _empty()
aiohttp_module.ClientTimeout = _ClientTimeout
aiohttp_module.WSMsgType = _WSMsgType
aiohttp_module.web = types.SimpleNamespace(
json_response=lambda payload, status=200: _FakeResponse(payload, status=status),
WebSocketResponse=_WebSocketResponse,
)
sys.modules["aiohttp"] = aiohttp_module
class _Routes:
def get(self, _path):
def _decorator(fn):
return fn
return _decorator
def post(self, _path):
def _decorator(fn):
return fn
return _decorator
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=types.SimpleNamespace(routes=_Routes()))
sys.modules["server"] = server_module
created_aiohttp_stub = install_aiohttp_stub(
lambda payload, status=200: _FakeResponse(payload, status=status)
)
install_server_stub()
created_torch_stub = False
if "torch" not in sys.modules:
@@ -248,19 +194,29 @@ def _load_worker_routes_module():
def _validate_worker_id(worker_id, config):
return any(str(worker.get("id")) == str(worker_id) for worker in config.get("workers", []))
def _require_worker_id(worker_id, config, field_name="worker_id"):
_ = field_name
worker_id_str = str(worker_id).strip()
if any(str(worker.get("id")) == worker_id_str for worker in config.get("workers", [])):
return worker_id_str
raise ValueError(f"Worker {worker_id_str} not found")
schemas_module.require_fields = _require_fields
schemas_module.validate_worker_id = _validate_worker_id
schemas_module.require_worker_id = _require_worker_id
schemas_module.is_authorized_request = lambda _request, _config: True
schemas_module.distributed_auth_headers = lambda _config: {}
sys.modules[f"{package_name}.api.schemas"] = schemas_module
install_request_guards_stub(package_name)
spec = importlib.util.spec_from_file_location(f"{package_name}.api.worker_routes", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
if created_aiohttp_stub:
sys.modules.pop("aiohttp", None)
if created_torch_stub:
sys.modules.pop("torch", None)
cleanup_optional_module("aiohttp", created_aiohttp_stub)
cleanup_optional_module("torch", created_torch_stub)
return module
-2
View File
@@ -1,2 +0,0 @@
# This conftest.py marks the tests/ directory as the pytest collection root,
# preventing pytest from traversing into the parent package's __init__.py.
+1
View File
@@ -0,0 +1 @@
"""Unit tests."""
+1
View File
@@ -0,0 +1 @@
"""Unit tests."""
@@ -8,7 +8,7 @@ import torch
def _load_utilities_module():
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
package_name = "dist_divider_testpkg"
for mod_name in list(sys.modules):
+80
View File
@@ -0,0 +1,80 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
nodes_pkg.__path__ = []
sys.modules[f"{package_name}.nodes"] = nodes_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_branch_module():
package_name = "dist_branch_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "nodes/utilities.py", "nodes.utilities")
return _load_module(package_name, "nodes/branch.py", "nodes.branch")
branch_module = _load_branch_module()
class DistributedBranchTests(unittest.TestCase):
def test_node_declares_ten_output_slots(self):
self.assertEqual(len(branch_module.DistributedBranch.RETURN_TYPES), 10)
self.assertEqual(len(branch_module.DistributedBranch.RETURN_NAMES), 10)
def test_no_distribution_outputs_all_branches(self):
node = branch_module.DistributedBranch()
token = {"k": "v"}
outputs = node.branch(token, num_branches=2, assigned_branch=-1)
self.assertEqual(len(outputs), 10)
self.assertTrue(all(output is token for output in outputs))
def test_assigned_branch_still_passes_input(self):
node = branch_module.DistributedBranch()
token = "payload"
outputs = node.branch(token, num_branches=3, is_worker=True, worker_id="worker-a", assigned_branch=1)
self.assertEqual(outputs[0], "payload")
self.assertEqual(outputs[1], "payload")
self.assertEqual(outputs[9], "payload")
if __name__ == "__main__":
unittest.main()
+282
View File
@@ -0,0 +1,282 @@
import asyncio
import base64
import importlib.util
import io
import json
import sys
import types
import unittest
from pathlib import Path
import torch
from PIL import Image
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
nodes_pkg.__path__ = []
sys.modules[f"{package_name}.nodes"] = nodes_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.get_worker_timeout_seconds = lambda: 0.2
config_module.is_master_delegate_only = lambda: False
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.HEARTBEAT_INTERVAL = 1.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
image_module = types.ModuleType(f"{package_name}.utils.image")
def ensure_contiguous(tensor):
return tensor.contiguous()
def tensor_to_pil(image_batch, index):
tensor = image_batch[index].detach().cpu().clamp(0, 1)
array = (tensor.numpy() * 255.0).astype("uint8")
return Image.fromarray(array)
def encode_tensor_png_data_url(image_batch, batch_index=0):
image = tensor_to_pil(image_batch, batch_index)
byte_io = io.BytesIO()
image.save(byte_io, format="PNG", compress_level=0)
encoded = base64.b64encode(byte_io.getvalue()).decode("utf-8")
return f"data:image/png;base64,{encoded}"
image_module.ensure_contiguous = ensure_contiguous
image_module.tensor_to_pil = tensor_to_pil
image_module.encode_tensor_png_data_url = encode_tensor_png_data_url
sys.modules[f"{package_name}.utils.image"] = image_module
network_module = types.ModuleType(f"{package_name}.utils.network")
async def get_client_session():
raise RuntimeError("test should patch get_client_session")
network_module.get_client_session = get_client_session
sys.modules[f"{package_name}.utils.network"] = network_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_branch_collector_module():
package_name = "dist_join_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "nodes/utilities.py", "nodes.utilities")
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
_load_module(package_name, "nodes/runtime_helpers.py", "nodes.runtime_helpers")
_load_module(package_name, "nodes/queue_wait.py", "nodes.queue_wait")
_load_module(package_name, "utils/worker_ids.py", "utils.worker_ids")
return _load_module(package_name, "nodes/branch_collector.py", "nodes.branch_collector")
branch_collector_module = _load_branch_collector_module()
def _context(**overrides):
defaults = {
"multi_job_id": "",
"is_worker": False,
"master_url": "",
"enabled_worker_ids": "[]",
"worker_id": "",
"assigned_branch": -1,
"delegate_only": False,
}
defaults.update(overrides)
return branch_collector_module.BranchRunContext(**defaults)
class _FakeResponseCtx:
def __init__(self, recorder, payload, url):
self._recorder = recorder
self._payload = payload
self._url = url
async def __aenter__(self):
self._recorder.append({"url": self._url, "payload": self._payload})
return self
async def __aexit__(self, exc_type, exc, tb):
return False
def raise_for_status(self):
return None
class _FakeSession:
def __init__(self):
self.calls = []
def post(self, url, json=None, timeout=None):
return _FakeResponseCtx(self.calls, json, url)
class _FakePromptServer:
def __init__(self):
self.distributed_pending_jobs = {}
self.distributed_jobs_lock = asyncio.Lock()
class DistributedBranchCollectorTests(unittest.IsolatedAsyncioTestCase):
def setUp(self):
self.node = branch_collector_module.DistributedBranchCollector()
async def test_worker_mode_sends_result_with_assigned_branch_index(self):
fake_session = _FakeSession()
async def _get_session():
return fake_session
branch_collector_module.get_client_session = _get_session
tensor = torch.ones((1, 4, 4, 3), dtype=torch.float32)
outputs = await self.node.execute(
[None, tensor, None, None, None, None, None, None, None, None],
num_branches=2,
context=_context(
multi_job_id="join-job",
is_worker=True,
master_url="http://master.local:8188",
worker_id="worker-a",
assigned_branch=1,
),
)
self.assertEqual(len(fake_session.calls), 1)
payload = fake_session.calls[0]["payload"]
self.assertEqual(payload["batch_idx"], 1)
self.assertTrue(payload["is_last"])
self.assertEqual(payload["worker_id"], "worker-a")
self.assertTrue(torch.equal(outputs[1], tensor))
async def test_master_mode_collects_worker_results_into_branch_slots(self):
fake_server = _FakePromptServer()
branch_collector_module._get_prompt_server_instance = lambda: fake_server
queue = asyncio.Queue()
worker_1 = torch.full((1, 4, 4, 3), 0.5, dtype=torch.float32)
worker_2 = torch.full((1, 4, 4, 3), 0.8, dtype=torch.float32)
await queue.put({"tensor": worker_1, "worker_id": "worker-a", "image_index": 1, "is_last": True})
await queue.put({"tensor": worker_2, "worker_id": "worker-b", "image_index": 2, "is_last": True})
fake_server.distributed_pending_jobs["join-job"] = queue
master_tensor = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
outputs = await self.node.execute(
[master_tensor, None, None, None, None, None, None, None, None, None],
num_branches=3,
context=_context(
multi_job_id="join-job",
enabled_worker_ids=json.dumps(["worker-a", "worker-b"]),
assigned_branch=0,
),
)
self.assertTrue(torch.equal(outputs[0], master_tensor))
self.assertTrue(torch.equal(outputs[1], worker_1))
self.assertTrue(torch.equal(outputs[2], worker_2))
async def test_unassigned_slots_remain_none(self):
fake_server = _FakePromptServer()
branch_collector_module._get_prompt_server_instance = lambda: fake_server
queue = asyncio.Queue()
worker_1 = torch.full((1, 4, 4, 3), 0.5, dtype=torch.float32)
await queue.put({"tensor": worker_1, "worker_id": "worker-a", "image_index": 1, "is_last": True})
fake_server.distributed_pending_jobs["join-job"] = queue
master_tensor = torch.zeros((1, 4, 4, 3), dtype=torch.float32)
outputs = await self.node.execute(
[master_tensor, None, None, None, None, None, None, None, None, None],
num_branches=2,
context=_context(
multi_job_id="join-job",
enabled_worker_ids=json.dumps(["worker-a"]),
assigned_branch=0,
),
)
self.assertIsNone(outputs[2])
self.assertIsNone(outputs[9])
async def test_missing_worker_results_use_safe_fallback_for_expected_slots(self):
fake_server = _FakePromptServer()
branch_collector_module._get_prompt_server_instance = lambda: fake_server
queue = asyncio.Queue()
# No worker payloads enqueued; master should timeout and fill expected worker slots.
fake_server.distributed_pending_jobs["join-job"] = queue
master_tensor = torch.full((1, 4, 4, 3), 0.7, dtype=torch.float32)
outputs = await self.node.execute(
[master_tensor, None, None, None, None, None, None, None, None, None],
num_branches=3,
context=_context(
multi_job_id="join-job",
enabled_worker_ids=json.dumps(["worker-a", "worker-b"]),
assigned_branch=0,
),
)
self.assertTrue(torch.equal(outputs[0], master_tensor))
self.assertIsInstance(outputs[1], torch.Tensor)
self.assertIsInstance(outputs[2], torch.Tensor)
self.assertEqual(float(outputs[1].abs().sum().item()), 0.0)
self.assertEqual(float(outputs[2].abs().sum().item()), 0.0)
async def test_assigned_slot_can_be_filled_from_single_mismatched_local_input(self):
fake_server = _FakePromptServer()
branch_collector_module._get_prompt_server_instance = lambda: fake_server
# No worker results expected for this check.
fake_server.distributed_pending_jobs["join-job"] = asyncio.Queue()
local_tensor = torch.full((1, 4, 4, 3), 0.33, dtype=torch.float32)
outputs = await self.node.execute(
# Local value arrives on slot 1 even though participant is assigned slot 0.
[None, local_tensor, None, None, None, None, None, None, None, None],
num_branches=3,
context=_context(
multi_job_id="join-job",
enabled_worker_ids=json.dumps([]),
assigned_branch=0,
),
)
self.assertTrue(torch.equal(outputs[0], local_tensor))
if __name__ == "__main__":
unittest.main()
+228
View File
@@ -0,0 +1,228 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
import numpy as np
import torch
from PIL import Image
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
nodes_pkg.__path__ = []
sys.modules[f"{package_name}.nodes"] = nodes_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.get_worker_timeout_seconds = lambda: 0.2
config_module.load_config = lambda: {"workers": []}
config_module.is_master_delegate_only = lambda: False
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.HEARTBEAT_INTERVAL = 1.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
image_module = types.ModuleType(f"{package_name}.utils.image")
def ensure_contiguous(tensor):
return tensor.contiguous()
def tensor_to_pil(image_batch, index):
tensor = image_batch[index].detach().cpu().clamp(0, 1)
array = (tensor.numpy() * 255.0).astype("uint8")
return Image.fromarray(array)
def pil_to_tensor(image):
array = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(array).unsqueeze(0)
image_module.ensure_contiguous = ensure_contiguous
image_module.tensor_to_pil = tensor_to_pil
image_module.pil_to_tensor = pil_to_tensor
sys.modules[f"{package_name}.utils.image"] = image_module
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.build_worker_url = lambda _worker: "http://worker.local:8188"
network_module.get_client_session = lambda: None
async def probe_worker(_url, timeout=2.0):
_ = timeout
return None
network_module.probe_worker = probe_worker
sys.modules[f"{package_name}.utils.network"] = network_module
audio_payload_module = types.ModuleType(f"{package_name}.utils.audio_payload")
audio_payload_module.encode_audio_payload = lambda payload: payload
sys.modules[f"{package_name}.utils.audio_payload"] = audio_payload_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.run_async_in_server_loop = lambda coro: asyncio.run(coro)
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
if "aiohttp" not in sys.modules:
try:
import aiohttp as _aiohttp # noqa: F401
except Exception:
aiohttp_module = types.ModuleType("aiohttp")
class _ClientTimeout:
def __init__(self, total=None):
self.total = total
aiohttp_module.ClientTimeout = _ClientTimeout
sys.modules["aiohttp"] = aiohttp_module
class _FakePromptServer:
def __init__(self):
self.distributed_pending_jobs = {}
self.distributed_jobs_lock = asyncio.Lock()
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=_FakePromptServer())
sys.modules["server"] = server_module
comfy_module = types.ModuleType("comfy")
model_mgmt = types.ModuleType("comfy.model_management")
class _InterruptProcessingException(Exception):
pass
model_mgmt.throw_exception_if_processing_interrupted = lambda: None
model_mgmt.InterruptProcessingException = _InterruptProcessingException
comfy_module.model_management = model_mgmt
sys.modules["comfy"] = comfy_module
sys.modules["comfy.model_management"] = model_mgmt
comfy_utils_module = types.ModuleType("comfy.utils")
class _ProgressBar:
def __init__(self, _total):
self.total = _total
def update(self, _step):
return None
comfy_utils_module.ProgressBar = _ProgressBar
sys.modules["comfy.utils"] = comfy_utils_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_collector_module():
package_name = "dist_collector_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
_load_module(package_name, "utils/worker_ids.py", "utils.worker_ids")
return _load_module(package_name, "nodes/collector.py", "nodes.collector")
collector_module = _load_collector_module()
class DistributedCollectorTests(unittest.TestCase):
def setUp(self):
self.node = collector_module.DistributedCollectorNode()
def test_store_worker_result_tracks_by_worker_and_index(self):
worker_images = {}
stored = self.node._store_worker_result(
worker_images,
{
"worker_id": "worker-a",
"image_index": 2,
"tensor": torch.full((1, 1, 1, 1), 0.25),
},
)
self.assertEqual(stored, 1)
self.assertIn("worker-a", worker_images)
self.assertIn(2, worker_images["worker-a"])
def test_combine_audio_honors_master_then_worker_order(self):
master_audio = {"waveform": torch.ones((1, 2, 1)), "sample_rate": 48000}
worker_audio = {
"worker-b": {"waveform": torch.full((1, 2, 1), 3.0), "sample_rate": 48000},
"worker-a": {"waveform": torch.full((1, 2, 2), 2.0), "sample_rate": 48000},
}
empty_audio = {"waveform": torch.zeros((1, 2, 1)), "sample_rate": 44100}
combined = self.node._combine_audio(
master_audio=master_audio,
worker_audio=worker_audio,
empty_audio=empty_audio,
worker_order=["worker-a", "worker-b"],
)
expected = torch.cat(
[
master_audio["waveform"],
worker_audio["worker-a"]["waveform"],
worker_audio["worker-b"]["waveform"],
],
dim=-1,
)
self.assertEqual(combined["sample_rate"], 48000)
self.assertTrue(torch.equal(combined["waveform"], expected))
def test_reorder_and_combine_tensors_uses_enabled_worker_priority(self):
worker_images = {
"worker-b": {
1: torch.full((1, 1, 1, 1), 4.0),
0: torch.full((1, 1, 1, 1), 3.0),
},
"worker-a": {
0: torch.full((1, 1, 1, 1), 2.0),
},
}
master_images = torch.full((1, 1, 1, 1), 1.0)
combined = self.node._reorder_and_combine_tensors(
worker_images=worker_images,
worker_order=["worker-a", "worker-b"],
master_batch_size=1,
images_on_cpu=master_images,
delegate_mode=False,
fallback_images=None,
)
self.assertEqual(combined.shape, (4, 1, 1, 1))
self.assertEqual(combined[0, 0, 0, 0].item(), 1.0)
self.assertEqual(combined[1, 0, 0, 0].item(), 2.0)
self.assertEqual(combined[2, 0, 0, 0].item(), 3.0)
self.assertEqual(combined[3, 0, 0, 0].item(), 4.0)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,248 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
import torch
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
nodes_pkg = types.ModuleType(f"{package_name}.nodes")
nodes_pkg.__path__ = []
sys.modules[f"{package_name}.nodes"] = nodes_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
modes_pkg.__path__ = []
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
worker_ids_module = types.ModuleType(f"{package_name}.utils.worker_ids")
worker_ids_module.coerce_enabled_worker_ids = lambda value: (
[str(v) for v in value]
if isinstance(value, list)
else []
)
sys.modules[f"{package_name}.utils.worker_ids"] = worker_ids_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.run_async_in_server_loop = lambda coro, timeout=None: (_ for _ in ()).throw(
RuntimeError(f"run_async_in_server_loop should not be used in these tests (timeout={timeout})")
)
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
job_store_module.ensure_tile_jobs_initialized = lambda: None
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
tile_ops_module = types.ModuleType(f"{package_name}.upscale.tile_ops")
class _TileOpsMixin:
def round_to_multiple(self, value):
return int(value)
def calculate_tiles(self, width, height, tile_width, tile_height, force_uniform_tiles):
_ = force_uniform_tiles
tiles = []
for y in range(0, int(height), max(1, int(tile_height))):
for x in range(0, int(width), max(1, int(tile_width))):
tiles.append((x, y))
return tiles or [(0, 0)]
tile_ops_module.TileOpsMixin = _TileOpsMixin
sys.modules[f"{package_name}.upscale.tile_ops"] = tile_ops_module
result_collector_module = types.ModuleType(f"{package_name}.upscale.result_collector")
result_collector_module.ResultCollectorMixin = type("ResultCollectorMixin", (), {})
sys.modules[f"{package_name}.upscale.result_collector"] = result_collector_module
worker_comms_module = types.ModuleType(f"{package_name}.upscale.worker_comms")
worker_comms_module.WorkerCommsMixin = type("WorkerCommsMixin", (), {})
sys.modules[f"{package_name}.upscale.worker_comms"] = worker_comms_module
job_state_module = types.ModuleType(f"{package_name}.upscale.job_state")
job_state_module.JobStateMixin = type("JobStateMixin", (), {})
sys.modules[f"{package_name}.upscale.job_state"] = job_state_module
single_gpu_module = types.ModuleType(f"{package_name}.upscale.modes.single_gpu")
class _SingleGpuModeMixin:
def process_single_gpu(self, *_args, **_kwargs):
return ("single_gpu_result",)
single_gpu_module.SingleGpuModeMixin = _SingleGpuModeMixin
sys.modules[f"{package_name}.upscale.modes.single_gpu"] = single_gpu_module
static_mode_module = types.ModuleType(f"{package_name}.upscale.modes.static")
class _StaticModeMixin:
def _process_worker_static_sync(self, *_args, **_kwargs):
return ("worker_static_result",)
def _process_master_static_sync(self, *_args, **_kwargs):
return ("master_static_result",)
static_mode_module.StaticModeMixin = _StaticModeMixin
sys.modules[f"{package_name}.upscale.modes.static"] = static_mode_module
dynamic_mode_module = types.ModuleType(f"{package_name}.upscale.modes.dynamic")
class _DynamicModeMixin:
def process_worker_dynamic(self, *_args, **_kwargs):
return ("worker_dynamic_result",)
def process_master_dynamic(self, *_args, **_kwargs):
return ("master_dynamic_result",)
dynamic_mode_module.DynamicModeMixin = _DynamicModeMixin
sys.modules[f"{package_name}.upscale.modes.dynamic"] = dynamic_mode_module
comfy_module = types.ModuleType("comfy")
samplers_module = types.ModuleType("comfy.samplers")
class _KSampler:
SAMPLERS = ("euler",)
SCHEDULERS = ("normal",)
samplers_module.KSampler = _KSampler
comfy_module.samplers = samplers_module
sys.modules["comfy"] = comfy_module
sys.modules["comfy.samplers"] = samplers_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_distributed_upscale_module():
package_name = "dist_distributed_upscale_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "nodes/hidden_inputs.py", "nodes.hidden_inputs")
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
return _load_module(package_name, "nodes/distributed_upscale.py", "nodes.distributed_upscale")
upscale_module = _load_distributed_upscale_module()
class DistributedUpscaleTests(unittest.TestCase):
def setUp(self):
self.node = upscale_module.UltimateSDUpscaleDistributed()
self.base_image = torch.zeros((1, 64, 64, 3), dtype=torch.float32)
self.common_kwargs = {
"model": object(),
"positive": object(),
"negative": object(),
"vae": object(),
"seed": 1,
"steps": 10,
"cfg": 7.5,
"sampler_name": "euler",
"scheduler": "normal",
"denoise": 0.5,
"tile_width": 32,
"tile_height": 32,
"padding": 16,
"mask_blur": 4,
"force_uniform_tiles": True,
"tiled_decode": False,
}
def test_parse_enabled_worker_ids_supports_json_and_lists(self):
self.assertEqual(
upscale_module._parse_enabled_worker_ids(["worker-a", 2]),
["worker-a", "2"],
)
self.assertEqual(
upscale_module._parse_enabled_worker_ids('["worker-a","worker-b"]'),
["worker-a", "worker-b"],
)
self.assertEqual(upscale_module._parse_enabled_worker_ids("invalid-json"), [])
self.assertEqual(upscale_module._parse_enabled_worker_ids(None), [])
def test_determine_processing_mode(self):
self.assertEqual(self.node._determine_processing_mode(batch_size=1, num_workers=0, dynamic_threshold=8), "single_gpu")
self.assertEqual(self.node._determine_processing_mode(batch_size=9, num_workers=2, dynamic_threshold=8), "static")
def test_run_raises_for_non_4n_plus_1_master_batches(self):
bad_batch = torch.zeros((2, 64, 64, 3), dtype=torch.float32)
with self.assertRaises(ValueError):
self.node.run(
upscaled_image=bad_batch,
multi_job_id="job-1",
is_worker=False,
**self.common_kwargs,
)
def test_run_dispatches_single_gpu_when_not_distributed(self):
self.node.process_single_gpu = lambda *_args, **_kwargs: ("single",)
result = self.node.run(
upscaled_image=self.base_image,
multi_job_id="",
is_worker=False,
**self.common_kwargs,
)
self.assertEqual(result, ("single",))
def test_run_dispatches_worker_or_master_paths(self):
self.node.process_worker = lambda *_args, **_kwargs: ("worker",)
self.node.process_master = lambda *_args, **_kwargs: ("master",)
worker_result = self.node.run(
upscaled_image=self.base_image,
multi_job_id="job-2",
is_worker=True,
master_url="http://master.local:8188",
enabled_worker_ids='["worker-a"]',
worker_id="worker-a",
tile_indices="",
dynamic_threshold=8,
**self.common_kwargs,
)
master_result = self.node.run(
upscaled_image=self.base_image,
multi_job_id="job-2",
is_worker=False,
enabled_worker_ids='["worker-a"]',
worker_id="",
tile_indices="",
dynamic_threshold=8,
**self.common_kwargs,
)
self.assertEqual(worker_result, ("worker",))
self.assertEqual(master_result, ("master",))
if __name__ == "__main__":
unittest.main()
@@ -13,7 +13,7 @@ class DistributedValueTests(unittest.TestCase):
from pathlib import Path
from unittest.mock import MagicMock
module_path = Path(__file__).resolve().parents[1] / "nodes" / "utilities.py"
module_path = Path(__file__).resolve().parents[3] / "nodes" / "utilities.py"
pkg_name = "dv_test_pkg"
for mod_name in list(sys.modules):
@@ -28,6 +28,18 @@ class DistributedValueTests(unittest.TestCase):
root_pkg.__path__ = []
sys.modules[pkg_name] = root_pkg
nodes_pkg = types.ModuleType(f"{pkg_name}.nodes")
nodes_pkg.__path__ = []
sys.modules[f"{pkg_name}.nodes"] = nodes_pkg
hidden_inputs_mod = types.ModuleType(f"{pkg_name}.nodes.hidden_inputs")
hidden_inputs_mod.build_worker_identity_hidden_inputs = lambda: {
"is_worker": ("BOOLEAN", {"default": False}),
"enabled_worker_ids": ("STRING", {"default": "[]"}),
"worker_id": ("STRING", {"default": ""}),
}
sys.modules[f"{pkg_name}.nodes.hidden_inputs"] = hidden_inputs_mod
utils_pkg = types.ModuleType(f"{pkg_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{pkg_name}.utils"] = utils_pkg
@@ -37,6 +49,16 @@ class DistributedValueTests(unittest.TestCase):
logging_mod.log = lambda *_a, **_k: None
sys.modules[f"{pkg_name}.utils.logging"] = logging_mod
worker_ids_path = Path(__file__).resolve().parents[3] / "utils" / "worker_ids.py"
worker_ids_spec = importlib.util.spec_from_file_location(
f"{pkg_name}.utils.worker_ids",
worker_ids_path,
)
worker_ids_mod = importlib.util.module_from_spec(worker_ids_spec)
assert worker_ids_spec is not None and worker_ids_spec.loader is not None
worker_ids_spec.loader.exec_module(worker_ids_mod)
sys.modules[f"{pkg_name}.utils.worker_ids"] = worker_ids_mod
spec = importlib.util.spec_from_file_location(
f"{pkg_name}.nodes.utilities", module_path
)
+1
View File
@@ -0,0 +1 @@
"""Unit tests."""
+104
View File
@@ -0,0 +1,104 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
class _FakePromptQueue:
def __init__(self):
self.items = []
def put(self, item):
self.items.append(item)
class _FakePromptServer:
def __init__(self):
self.number = 0
self.prompt_queue = _FakePromptQueue()
def trigger_on_prompt(self, payload):
return payload
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.get_server_loop = asyncio.get_event_loop
sys.modules[f"{package_name}.utils.network"] = network_module
fake_prompt_server = _FakePromptServer()
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=fake_prompt_server)
sys.modules["server"] = server_module
execution_module = types.ModuleType("execution")
async def _validate_prompt(_prompt_id, _prompt, _partial_targets):
return (True, None, ["1"], {})
execution_module.validate_prompt = _validate_prompt
execution_module.SENSITIVE_EXTRA_DATA_KEYS = ()
sys.modules["execution"] = execution_module
return fake_prompt_server
def _load_async_helpers_module():
package_name = "dist_async_helpers_testpkg"
fake_prompt_server = _bootstrap_package(package_name)
module_path = Path(__file__).resolve().parents[3] / "utils/async_helpers.py"
spec = importlib.util.spec_from_file_location(
f"{package_name}.utils.async_helpers",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module, fake_prompt_server
class AsyncHelpersQueuePromptPayloadTests(unittest.TestCase):
def test_queue_prompt_payload_includes_create_time_metadata(self):
async_helpers_module, fake_prompt_server = _load_async_helpers_module()
prompt_id = asyncio.run(
async_helpers_module.queue_prompt_payload(
{"1": {"class_type": "KSampler", "inputs": {}}},
workflow_meta={"id": "workflow-1"},
client_id="client-1",
)
)
self.assertIsInstance(prompt_id, str)
self.assertTrue(fake_prompt_server.prompt_queue.items)
queued_item = fake_prompt_server.prompt_queue.items[-1]
self.assertEqual(len(queued_item), 6)
self.assertEqual(queued_item[1], prompt_id)
extra_data = queued_item[3]
self.assertEqual(extra_data["client_id"], "client-1")
self.assertEqual(extra_data["extra_pnginfo"]["workflow"]["id"], "workflow-1")
self.assertIn("create_time", extra_data)
self.assertIsInstance(extra_data["create_time"], int)
self.assertGreater(extra_data["create_time"], 0)
if __name__ == "__main__":
unittest.main()
@@ -10,7 +10,7 @@ from unittest.mock import patch
def _load_config_module():
module_path = Path(__file__).resolve().parents[1] / "utils" / "config.py"
module_path = Path(__file__).resolve().parents[3] / "utils" / "config.py"
package_name = "dist_cfg_testpkg"
for mod_name in list(sys.modules):
@@ -21,13 +21,13 @@ def _load_config_module():
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
logging_module = types.ModuleType(f"{package_name}.logging")
logging_module.log = lambda *_args, **_kwargs: None
logging_module.debug_log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.logging"] = logging_module
constants_module = types.ModuleType(f"{package_name}.constants")
constants_module.GPU_CONFIG_FILE = "gpu_config.json"
constants_module.HEARTBEAT_TIMEOUT = 30
constants_module.ORCHESTRATION_WORKER_PROBE_CONCURRENCY = 8
constants_module.ORCHESTRATION_WORKER_PREP_CONCURRENCY = 4
constants_module.ORCHESTRATION_MEDIA_SYNC_CONCURRENCY = 2
constants_module.ORCHESTRATION_MEDIA_SYNC_TIMEOUT = 120
sys.modules[f"{package_name}.constants"] = constants_module
spec = importlib.util.spec_from_file_location(f"{package_name}.config", module_path)
@@ -213,7 +213,7 @@ class SaveConfigTests(unittest.TestCase):
data = config.get_default_config()
config.save_config(data)
# Cache is now None; load_config should re-read
self.assertIsNone(config._config_cache)
self.assertIsNone(config._config_state().cache)
def test_written_file_is_valid_json(self):
data = config.get_default_config()
@@ -7,7 +7,7 @@ from unittest.mock import patch
def _load_detection_module():
module_path = Path(__file__).resolve().parents[1] / "workers" / "detection.py"
module_path = Path(__file__).resolve().parents[3] / "workers" / "detection.py"
package_name = "dist_det_testpkg"
for mod_name in list(sys.modules):
@@ -32,6 +32,7 @@ def _load_detection_module():
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.normalize_host = lambda value: value
network_module.build_worker_url = lambda worker, endpoint="": f"http://{worker.get('host', 'localhost')}{endpoint}"
async def _fake_session():
raise RuntimeError("network calls not used in these tests")
@@ -156,7 +157,7 @@ class IsLocalWorkerTests(unittest.IsolatedAsyncioTestCase):
self.assertTrue(result)
async def test_true_for_0_0_0_0(self):
result = await detection.is_local_worker({"host": "0.0.0.0", "port": 8188})
result = await detection.is_local_worker({"host": "0.0.0.0", "port": 8188}) # nosec B104 - explicit wildcard host test case
self.assertTrue(result)
async def test_true_when_type_is_local(self):
@@ -8,7 +8,7 @@ from unittest.mock import patch
def _load_dispatch_module():
module_path = Path(__file__).resolve().parents[1] / "api" / "orchestration" / "dispatch.py"
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "dispatch.py"
package_name = "dist_dispatch_testpkg"
root_pkg = types.ModuleType(package_name)
@@ -249,6 +249,50 @@ class DispatchSelectionTests(unittest.IsolatedAsyncioTestCase):
self.assertIsNone(selected)
async def test_rank_workers_by_load_returns_sorted_by_queue_depth(self):
workers = [
{"id": "w1", "name": "Worker 1"},
{"id": "w2", "name": "Worker 2"},
{"id": "w3", "name": "Worker 3"},
]
queue_map = {"w1": 4, "w2": 0, "w3": 2}
async def fake_probe(worker_url, timeout=3.0):
worker_id = worker_url.rsplit("/", 1)[-1]
return {"exec_info": {"queue_remaining": queue_map[worker_id]}}
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
dispatch,
"probe_worker",
side_effect=fake_probe,
):
ranked = await dispatch.rank_workers_by_load(workers, probe_concurrency=3)
self.assertEqual([worker["id"] for worker in ranked], ["w2", "w3", "w1"])
async def test_rank_workers_by_load_handles_unreachable_workers(self):
workers = [
{"id": "w1", "name": "Worker 1"},
{"id": "w2", "name": "Worker 2"},
{"id": "w3", "name": "Worker 3"},
]
async def fake_probe(worker_url, timeout=3.0):
worker_id = worker_url.rsplit("/", 1)[-1]
if worker_id == "w2":
return None
queue_map = {"w1": 1, "w3": 0}
return {"exec_info": {"queue_remaining": queue_map[worker_id]}}
with patch.object(dispatch, "build_worker_url", side_effect=lambda worker: f"http://host/{worker['id']}"), patch.object(
dispatch,
"probe_worker",
side_effect=fake_probe,
):
ranked = await dispatch.rank_workers_by_load(workers, probe_concurrency=3)
self.assertEqual([worker["id"] for worker in ranked], ["w3", "w1", "w2"])
if __name__ == "__main__":
unittest.main()
+71
View File
@@ -0,0 +1,71 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
def _load_distributed_module():
module_path = Path(__file__).resolve().parents[3] / "distributed.py"
package_name = "dist_distributed_entry_testpkg"
calls = {
"build_node_mappings": 0,
"initialize_runtime": 0,
"initialize_runtime_args": [],
}
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
bootstrap_pkg = types.ModuleType(f"{package_name}.bootstrap")
bootstrap_pkg.__path__ = []
sys.modules[f"{package_name}.bootstrap"] = bootstrap_pkg
entrypoint_module = types.ModuleType(f"{package_name}.bootstrap.entrypoint")
def _build_node_mappings():
calls["build_node_mappings"] += 1
return ({"MockNode": object}, {"MockNode": "Mock Node"})
def _initialize_runtime(prompt_server=None):
calls["initialize_runtime"] += 1
calls["initialize_runtime_args"].append(prompt_server)
entrypoint_module.build_node_mappings = _build_node_mappings
entrypoint_module.initialize_runtime = _initialize_runtime
sys.modules[f"{package_name}.bootstrap.entrypoint"] = entrypoint_module
spec = importlib.util.spec_from_file_location(f"{package_name}.distributed", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module, calls
class DistributedEntryTests(unittest.TestCase):
def test_import_builds_mappings_without_runtime_bootstrap(self):
module, calls = _load_distributed_module()
self.assertIn("MockNode", module.NODE_CLASS_MAPPINGS)
self.assertEqual(module.NODE_DISPLAY_NAME_MAPPINGS["MockNode"], "Mock Node")
self.assertEqual(calls["build_node_mappings"], 1)
self.assertEqual(calls["initialize_runtime"], 0)
def test_initialize_runtime_delegates_to_bootstrap_entrypoint(self):
module, calls = _load_distributed_module()
sentinel = object()
module.initialize_runtime(sentinel)
self.assertEqual(calls["initialize_runtime"], 1)
self.assertEqual(calls["initialize_runtime_args"], [sentinel])
if __name__ == "__main__":
unittest.main()
+40
View File
@@ -0,0 +1,40 @@
import importlib.util
import unittest
from pathlib import Path
def _load_exceptions_module():
module_path = Path(__file__).resolve().parents[3] / "utils" / "exceptions.py"
spec = importlib.util.spec_from_file_location("dist_test_exceptions", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
return module
exceptions = _load_exceptions_module()
class ExceptionsModuleTests(unittest.TestCase):
def test_worker_error_tracks_worker_id(self):
err = exceptions.WorkerError("worker failed", worker_id="w-1")
self.assertIsInstance(err, exceptions.DistributedError)
self.assertEqual(str(err), "worker failed")
self.assertEqual(err.worker_id, "w-1")
def test_process_error_tracks_pid_and_worker_id(self):
err = exceptions.ProcessError("process failed", pid=1001, worker_id="worker-x")
self.assertIsInstance(err, exceptions.DistributedError)
self.assertEqual(err.pid, 1001)
self.assertEqual(err.worker_id, "worker-x")
def test_specialized_errors_inherit_expected_base_types(self):
self.assertTrue(issubclass(exceptions.WorkerTimeoutError, exceptions.WorkerError))
self.assertTrue(issubclass(exceptions.WorkerNotAvailableError, exceptions.WorkerError))
self.assertTrue(issubclass(exceptions.JobQueueError, exceptions.DistributedError))
self.assertTrue(issubclass(exceptions.TileCollectionError, exceptions.DistributedError))
self.assertTrue(issubclass(exceptions.TunnelError, exceptions.DistributedError))
if __name__ == "__main__":
unittest.main()
+228
View File
@@ -0,0 +1,228 @@
import importlib.util
import sys
import tempfile
import types
import unittest
from pathlib import Path
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
workers_pkg = types.ModuleType(f"{package_name}.workers")
workers_pkg.__path__ = []
sys.modules[f"{package_name}.workers"] = workers_pkg
process_pkg = types.ModuleType(f"{package_name}.workers.process")
process_pkg.__path__ = []
sys.modules[f"{package_name}.workers.process"] = process_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_lifecycle_module():
package_name = "dist_lifecycle_testpkg"
_bootstrap_package(package_name)
state = {
"config": {"settings": {"stop_workers_on_master_exit": True}, "managed_processes": {}},
"saved_configs": [],
"launch_calls": [],
"terminated": [],
"alive_pids": set(),
}
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: state["config"]
config_module.save_config = lambda cfg: state["saved_configs"].append(dict(cfg))
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.PROCESS_TERMINATION_TIMEOUT = 5.0
constants_module.PROCESS_WAIT_TIMEOUT = 1.0
constants_module.WORKER_CHECK_INTERVAL = 0.01
sys.modules[f"{package_name}.utils.constants"] = constants_module
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
process_module = types.ModuleType(f"{package_name}.utils.process")
class _FakeProcess:
def __init__(self, pid=1234, poll_result=None):
self.pid = pid
self._poll_result = poll_result
self.returncode = poll_result
def poll(self):
return self._poll_result
def _launch_process_with_timeout(command, timeout_seconds=10.0, **kwargs):
state["launch_calls"].append((list(command), float(timeout_seconds), dict(kwargs)))
return _FakeProcess(pid=4321, poll_result=None)
def _terminate_process(process, timeout=5):
state["terminated"].append((process.pid, timeout))
process._poll_result = 0
process.returncode = 0
process_module.launch_process_with_timeout = _launch_process_with_timeout
process_module.terminate_process = _terminate_process
process_module.is_process_alive = lambda pid: int(pid) in state["alive_pids"]
process_module.get_python_executable = lambda: "/usr/bin/python3"
process_module._FakeProcess = _FakeProcess
sys.modules[f"{package_name}.utils.process"] = process_module
lifecycle_module = _load_module(package_name, "workers/process/lifecycle.py", "workers.process.lifecycle")
return lifecycle_module, state, process_module
lifecycle_module, test_state, process_stub_module = _load_lifecycle_module()
class _FakeManager:
def __init__(self, comfy_root):
self._comfy_root = comfy_root
self.processes = {}
self.save_calls = 0
def find_comfy_root(self):
return self._comfy_root
def build_launch_command(self, worker_config, _comfy_root):
return ["python", "main.py", "--port", str(worker_config["port"])]
def save_processes(self):
self.save_calls += 1
class ProcessLifecycleTests(unittest.TestCase):
def setUp(self):
test_state["config"] = {"settings": {"stop_workers_on_master_exit": True}, "managed_processes": {}}
test_state["saved_configs"].clear()
test_state["launch_calls"].clear()
test_state["terminated"].clear()
test_state["alive_pids"].clear()
def test_launch_worker_uses_monitor_wrapper_when_enabled(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = _FakeManager(tmpdir)
lifecycle = lifecycle_module.ProcessLifecycle(manager)
pid = lifecycle.launch_worker(
{
"id": "1",
"name": "Worker A",
"port": 8189,
"cuda_device": 2,
}
)
self.assertEqual(pid, 4321)
self.assertEqual(manager.save_calls, 1)
self.assertIn("1", manager.processes)
self.assertTrue(manager.processes["1"]["is_monitor"])
self.assertIn("logs/workers", manager.processes["1"]["log_file"])
self.assertEqual(len(test_state["launch_calls"]), 1)
launched_command, timeout_seconds, launch_kwargs = test_state["launch_calls"][0]
self.assertEqual(timeout_seconds, 1.0)
self.assertEqual(launched_command[0], "/usr/bin/python3")
self.assertTrue(launched_command[1].endswith("workers/worker_monitor.py"))
self.assertEqual(launch_kwargs["env"]["CUDA_VISIBLE_DEVICES"], "2")
def test_launch_worker_skips_monitor_when_disabled(self):
test_state["config"] = {"settings": {"stop_workers_on_master_exit": False}, "managed_processes": {}}
with tempfile.TemporaryDirectory() as tmpdir:
manager = _FakeManager(tmpdir)
lifecycle = lifecycle_module.ProcessLifecycle(manager)
lifecycle.launch_worker(
{
"id": "2",
"name": "Worker-B",
"port": 8190,
"cuda_device": 0,
}
)
launched_command, _, _ = test_state["launch_calls"][0]
self.assertEqual(launched_command[0], "python")
self.assertFalse(manager.processes["2"]["is_monitor"])
def test_stop_worker_uses_fallback_terminate_when_tree_kill_fails(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = _FakeManager(tmpdir)
lifecycle = lifecycle_module.ProcessLifecycle(manager)
running_process = process_stub_module._FakeProcess(pid=88, poll_result=None)
manager.processes["7"] = {
"pid": 88,
"process": running_process,
"started_at": 0.0,
"config": {},
}
lifecycle._kill_process_tree = lambda _pid: False
ok, message = lifecycle.stop_worker("7")
self.assertTrue(ok)
self.assertIn("fallback", message.lower())
self.assertEqual(test_state["terminated"], [(88, 5.0)])
self.assertNotIn("7", manager.processes)
def test_stop_worker_without_process_uses_tree_kill(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = _FakeManager(tmpdir)
lifecycle = lifecycle_module.ProcessLifecycle(manager)
manager.processes["9"] = {
"pid": 99,
"process": None,
"started_at": 0.0,
"config": {},
}
lifecycle._kill_process_tree = lambda _pid: True
ok, message = lifecycle.stop_worker("9")
self.assertTrue(ok)
self.assertEqual(message, "Worker stopped")
self.assertNotIn("9", manager.processes)
def test_check_worker_process_falls_back_to_pid_probe(self):
with tempfile.TemporaryDirectory() as tmpdir:
manager = _FakeManager(tmpdir)
lifecycle = lifecycle_module.ProcessLifecycle(manager)
lifecycle._is_process_running = lambda pid: int(pid) == 42
running, from_subprocess = lifecycle._check_worker_process(
"any",
{"pid": 42, "process": None},
)
self.assertTrue(running)
self.assertFalse(from_subprocess)
if __name__ == "__main__":
unittest.main()
+32
View File
@@ -0,0 +1,32 @@
import io
import unittest
from contextlib import redirect_stdout
from unittest.mock import patch
from utils import logging as logging_utils
class LoggingUtilsTests(unittest.TestCase):
def test_is_debug_enabled_delegates_to_config(self):
with patch.object(logging_utils, "get_debug_enabled", return_value=True) as mocked:
self.assertTrue(logging_utils.is_debug_enabled())
mocked.assert_called_once_with(default=False)
def test_is_debug_enabled_handles_config_errors(self):
with patch.object(logging_utils, "get_debug_enabled", side_effect=RuntimeError("boom")):
self.assertFalse(logging_utils.is_debug_enabled())
def test_debug_log_emits_only_when_enabled(self):
output = io.StringIO()
with patch.object(logging_utils, "is_debug_enabled", return_value=True), redirect_stdout(output):
logging_utils.debug_log("hello")
self.assertIn("[Distributed] hello", output.getvalue())
output = io.StringIO()
with patch.object(logging_utils, "is_debug_enabled", return_value=False), redirect_stdout(output):
logging_utils.debug_log("hidden")
self.assertEqual(output.getvalue(), "")
if __name__ == "__main__":
unittest.main()
@@ -6,7 +6,7 @@ from pathlib import Path
def _load_network_module():
module_path = Path(__file__).resolve().parents[1] / "utils" / "network.py"
module_path = Path(__file__).resolve().parents[3] / "utils" / "network.py"
package_name = "dist_utils_testpkg"
package_module = types.ModuleType(package_name)
@@ -17,6 +17,10 @@ def _load_network_module():
logging_module.debug_log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.logging"] = logging_module
exceptions_module = types.ModuleType(f"{package_name}.exceptions")
exceptions_module.DistributedError = type("DistributedError", (Exception,), {})
sys.modules[f"{package_name}.exceptions"] = exceptions_module
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(
instance=types.SimpleNamespace(address="127.0.0.1", port=8188, loop=None)
@@ -92,7 +96,7 @@ class NetworkHelpersTests(unittest.TestCase):
def test_build_master_url_falls_back_to_server_address(self):
cfg = {"master": {"host": ""}}
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190)
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190) # nosec B104 - wildcard bind normalization test
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://127.0.0.1:8190",
@@ -7,7 +7,7 @@ from pathlib import Path
def _load_prompt_transform_module():
module_path = Path(__file__).resolve().parents[1] / "api" / "orchestration" / "prompt_transform.py"
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "prompt_transform.py"
package_name = "dist_pt_testpkg"
for mod_name in list(sys.modules):
@@ -81,13 +81,38 @@ def _delegate_prompt():
}
def _branch_prompt():
"""1 → 2(DistributedBranch) -> slot0:3->4, slot1:5->6."""
return {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
"4": {"class_type": "SaveImage", "inputs": {"images": ["3", 0]}},
"5": {"class_type": "Sharpen", "inputs": {"image": ["2", 1]}},
"6": {"class_type": "SaveImage", "inputs": {"images": ["5", 0]}},
}
def _branch_collector_prompt():
"""1 → 2(DistributedBranch) -> 3(branch0) and 4(branch1) -> 5(DistributedBranchCollector) -> 6(SaveImage)."""
return {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
"4": {"class_type": "Sharpen", "inputs": {"image": ["2", 1]}},
"5": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["3", 0], "num_branches": 2}},
"6": {"class_type": "SaveImage", "inputs": {"images": ["5", 0]}},
}
def _apply(prompt, participant_id, enabled_worker_ids=None, delegate_master=False):
if enabled_worker_ids is None:
enabled_worker_ids = ["worker-a", "worker-b"]
idx = pt.PromptIndex(prompt)
prompt_copy = json.loads(json.dumps(prompt))
idx = pt.PromptIndex(prompt_copy)
job_id_map = pt.generate_job_id_map(idx, "run")
return pt.apply_participant_overrides(
prompt,
prompt_copy,
participant_id=participant_id,
enabled_worker_ids=enabled_worker_ids,
job_id_map=job_id_map,
@@ -250,6 +275,21 @@ class PrunePromptForWorkerTests(unittest.TestCase):
self.assertIn("2", result)
self.assertNotIn("3", result)
def test_branch_anchor_keeps_downstream_for_later_branch_pruning(self):
prompt = _branch_prompt()
result = pt.prune_prompt_for_worker(prompt)
self.assertIn("2", result)
self.assertIn("3", result)
self.assertIn("4", result)
self.assertIn("5", result)
self.assertIn("6", result)
def test_branch_collector_anchor_keeps_collector_but_prunes_downstream(self):
prompt = _branch_collector_prompt()
result = pt.prune_prompt_for_worker(prompt)
self.assertIn("5", result)
self.assertNotIn("6", result)
# ---------------------------------------------------------------------------
# prepare_delegate_master_prompt
@@ -327,6 +367,22 @@ class GenerateJobIdMapTests(unittest.TestCase):
job_map = pt.generate_job_id_map(idx, "run")
self.assertEqual(job_map["5"], "run_5")
def test_maps_branch_collector_nodes(self):
prompt = {
"10": {"class_type": "DistributedBranchCollector", "inputs": {}},
}
idx = pt.PromptIndex(prompt)
job_map = pt.generate_job_id_map(idx, "run")
self.assertEqual(job_map["10"], "run_10")
def test_maps_branch_nodes(self):
prompt = {
"9": {"class_type": "DistributedBranch", "inputs": {"num_branches": 2}},
}
idx = pt.PromptIndex(prompt)
job_map = pt.generate_job_id_map(idx, "run")
self.assertEqual(job_map["9"], "run_9")
def test_empty_prompt_returns_empty_map(self):
idx = pt.PromptIndex({})
self.assertEqual(pt.generate_job_id_map(idx, "prefix"), {})
@@ -418,9 +474,9 @@ class ApplyOverridesSeedTests(unittest.TestCase):
result = _apply(self._seed_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
def test_worker_id_reflects_index_in_enabled_list(self):
def test_worker_id_uses_canonical_worker_id(self):
result = _apply(self._seed_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker_1")
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
def test_master_sets_is_worker_false(self):
result = _apply(self._seed_prompt(), "master")
@@ -430,6 +486,11 @@ class ApplyOverridesSeedTests(unittest.TestCase):
result = _apply(self._seed_prompt(), "master")
self.assertEqual(result["1"]["inputs"]["worker_id"], "")
def test_enabled_worker_ids_is_set(self):
enabled = ["worker-a", "worker-b"]
result = _apply(self._seed_prompt(), "worker-a", enabled_worker_ids=enabled)
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
# ---------------------------------------------------------------------------
# apply_participant_overrides – UltimateSDUpscaleDistributed
@@ -476,9 +537,9 @@ class ApplyOverridesValueTests(unittest.TestCase):
result = _apply(self._value_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
def test_worker_id_reflects_index_in_enabled_list(self):
def test_worker_id_uses_canonical_worker_id(self):
result = _apply(self._value_prompt(), "worker-b", enabled_worker_ids=["worker-a", "worker-b"])
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker_1")
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-b")
def test_master_sets_is_worker_false(self):
result = _apply(self._value_prompt(), "master")
@@ -488,6 +549,157 @@ class ApplyOverridesValueTests(unittest.TestCase):
result = _apply(self._value_prompt(), "master")
self.assertEqual(result["1"]["inputs"]["worker_id"], "")
def test_enabled_worker_ids_is_set(self):
enabled = ["worker-a", "worker-b"]
result = _apply(self._value_prompt(), "worker-a", enabled_worker_ids=enabled)
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
class ApplyOverridesBranchTests(unittest.TestCase):
def test_master_gets_branch_zero_worker_gets_branch_one(self):
prompt = _branch_prompt()
master_result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
self.assertEqual(master_result["2"]["inputs"]["assigned_branch"], 0)
self.assertEqual(worker_result["2"]["inputs"]["assigned_branch"], 1)
self.assertIn("3", master_result)
self.assertIn("4", master_result)
self.assertNotIn("5", master_result)
self.assertNotIn("6", master_result)
self.assertIn("5", worker_result)
self.assertIn("6", worker_result)
self.assertNotIn("3", worker_result)
self.assertNotIn("4", worker_result)
def test_delegate_mode_assigns_worker_a_to_branch_zero(self):
prompt = _branch_prompt()
worker_result = _apply(
prompt,
"worker-a",
enabled_worker_ids=["worker-a", "worker-b"],
delegate_master=True,
)
self.assertEqual(worker_result["2"]["inputs"]["assigned_branch"], 0)
def test_pruning_keeps_shared_nodes_between_branches(self):
prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "NodeA", "inputs": {"image": ["2", 0]}},
"4": {"class_type": "NodeB", "inputs": {"image": ["2", 1]}},
"5": {"class_type": "SaveImage", "inputs": {"images": ["3", 0], "aux": ["4", 0]}},
}
worker_result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
self.assertIn("3", worker_result)
self.assertNotIn("4", worker_result)
self.assertIn("5", worker_result)
self.assertNotIn("aux", worker_result["5"]["inputs"])
def test_more_participants_than_branches_sets_unassigned_to_minus_one(self):
prompt = _branch_prompt()
worker_b_result = _apply(
prompt,
"worker-b",
enabled_worker_ids=["worker-a", "worker-b", "worker-c"],
)
self.assertEqual(worker_b_result["2"]["inputs"]["assigned_branch"], -1)
class_types = {node.get("class_type") for node in worker_b_result.values()}
self.assertNotIn("Blur", class_types)
self.assertNotIn("Sharpen", class_types)
def test_worker_with_pruned_outputs_gets_auto_preview_for_assigned_branch(self):
prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 3}},
"3": {"class_type": "Blur", "inputs": {"image": ["2", 0]}},
"4": {"class_type": "PreviewImage", "inputs": {"images": ["3", 0]}},
}
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a", "worker-b"])
preview_nodes = [node for node in worker_result.values() if node.get("class_type") == "PreviewImage"]
self.assertTrue(preview_nodes)
self.assertIn(["2", 1], [node.get("inputs", {}).get("images") for node in preview_nodes])
def test_unassigned_worker_gets_idle_fallback_output(self):
prompt = _branch_prompt()
worker_result = _apply(
prompt,
"worker-c",
enabled_worker_ids=["worker-a", "worker-b", "worker-c"],
)
preview_nodes = [node for node in worker_result.values() if node.get("class_type") == "PreviewImage"]
empty_nodes = [node for node in worker_result.values() if node.get("class_type") == "DistributedEmptyImage"]
self.assertTrue(preview_nodes)
self.assertTrue(empty_nodes)
class ApplyOverridesBranchCollectorTests(unittest.TestCase):
def test_branch_collector_inherits_assigned_branch_from_upstream_branch_node(self):
master_prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0], "num_branches": 2}},
}
worker_prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 1], "num_branches": 2}},
}
master_result = _apply(master_prompt, "master", enabled_worker_ids=["worker-a"])
worker_result = _apply(worker_prompt, "worker-a", enabled_worker_ids=["worker-a"])
self.assertEqual(master_result["3"]["inputs"]["assigned_branch"], 0)
self.assertEqual(worker_result["3"]["inputs"]["assigned_branch"], 1)
def test_branch_collector_sets_multi_job_id(self):
prompt = {
"1": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0]}},
"2": {"class_type": "KSampler", "inputs": {}},
}
result = _apply(prompt, "master", enabled_worker_ids=["worker-a"])
self.assertEqual(result["1"]["inputs"]["multi_job_id"], "run_1")
def test_branch_collector_uses_upstream_branch_job_id_for_grouped_convergence(self):
master_prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 0], "num_branches": 2}},
}
worker_prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "DistributedBranchCollector", "inputs": {"input": ["2", 1], "num_branches": 2}},
}
master_result = _apply(master_prompt, "master", enabled_worker_ids=["worker-a"])
worker_result = _apply(worker_prompt, "worker-a", enabled_worker_ids=["worker-a"])
self.assertEqual(master_result["3"]["inputs"]["multi_job_id"], "run_2")
self.assertEqual(worker_result["3"]["inputs"]["multi_job_id"], "run_2")
def test_worker_prunes_nodes_downstream_of_branch_collector(self):
prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedBranch", "inputs": {"input": ["1", 0], "num_branches": 2}},
"3": {"class_type": "Blur", "inputs": {"image": ["2", 1]}},
"4": {
"class_type": "DistributedBranchCollector",
"inputs": {"branch_2": ["3", 0], "num_branches": 2},
},
"5": {"class_type": "SaveImage", "inputs": {"images": ["4", 0]}},
}
worker_result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
class_types = {node.get("class_type") for node in worker_result.values()}
self.assertIn("DistributedBranchCollector", class_types)
self.assertNotIn("SaveImage", class_types)
self.assertIn("PreviewImage", class_types)
if __name__ == "__main__":
unittest.main()
@@ -4,7 +4,7 @@ from pathlib import Path
def _load_queue_request_module():
module_path = Path(__file__).resolve().parents[1] / "api" / "queue_request.py"
module_path = Path(__file__).resolve().parents[3] / "api" / "queue_request.py"
spec = importlib.util.spec_from_file_location("queue_request", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
@@ -26,27 +26,26 @@ class QueueRequestPayloadTests(unittest.TestCase):
def test_normalizes_enabled_worker_ids(self):
payload_data = self._base_payload()
payload_data["enabled_worker_ids"] = ["a", 2, 3]
payload_data["enabled_worker_ids"] = [" worker-a ", "worker-b", "worker-a"]
payload_data["delegate_master"] = True
payload = parse_queue_request_payload(
payload_data
)
self.assertEqual(payload.enabled_worker_ids, ["a", "2", "3"])
self.assertEqual(payload.enabled_worker_ids, ["worker-a", "worker-b"])
self.assertTrue(payload.delegate_master)
def test_supports_legacy_workers_field(self):
payload_data = self._base_payload()
payload_data.pop("enabled_worker_ids", None)
payload_data["workers"] = [{"id": "w1"}, "w2", {"id": 3}, {"name": "no-id"}]
payload_data["workers"] = [{"id": "w1"}, "w2", {"id": "w3"}, {"name": "no-id"}]
payload = parse_queue_request_payload(
payload_data
)
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "3"])
self.assertEqual(payload.enabled_worker_ids, ["w1", "w2", "w3"])
def test_supports_auto_prepare_prompt_fallback(self):
def test_supports_workflow_prompt_fallback(self):
payload_data = self._base_payload()
payload_data.pop("prompt", None)
payload_data["auto_prepare"] = True
payload_data["workflow"] = {
"prompt": {
"10": {"class_type": "DistributedCollector"},
@@ -56,7 +55,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
payload_data
)
self.assertIn("10", payload.prompt)
self.assertTrue(payload.auto_prepare)
def test_normalizes_trace_execution_id(self):
payload_data = self._base_payload()
@@ -74,10 +72,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
)
self.assertIsNone(payload.trace_execution_id)
def test_auto_prepare_defaults_true(self):
payload = parse_queue_request_payload(self._base_payload())
self.assertTrue(payload.auto_prepare)
def test_workers_field_must_be_list(self):
payload_data = self._base_payload()
payload_data.pop("enabled_worker_ids", None)
@@ -91,7 +85,7 @@ class QueueRequestPayloadTests(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "trace_execution_id must be a string"):
parse_queue_request_payload(payload_data)
def test_auto_prepare_false_still_falls_back_to_workflow_prompt(self):
def test_auto_prepare_is_ignored_for_backward_compat(self):
payload_data = self._base_payload()
payload_data.pop("prompt", None)
payload_data["auto_prepare"] = False
@@ -100,13 +94,6 @@ class QueueRequestPayloadTests(unittest.TestCase):
}
payload = parse_queue_request_payload(payload_data)
self.assertIn("10", payload.prompt)
self.assertFalse(payload.auto_prepare)
def test_auto_prepare_must_be_boolean(self):
payload_data = self._base_payload()
payload_data["auto_prepare"] = "true"
with self.assertRaisesRegex(ValueError, "auto_prepare must be a boolean"):
parse_queue_request_payload(payload_data)
def test_invalid_delegate_master_type_raises(self):
payload_data = self._base_payload()
@@ -120,6 +107,12 @@ class QueueRequestPayloadTests(unittest.TestCase):
with self.assertRaisesRegex(ValueError, "enabled_worker_ids must be a list"):
parse_queue_request_payload(payload_data)
def test_legacy_index_worker_tokens_are_rejected(self):
payload_data = self._base_payload()
payload_data["enabled_worker_ids"] = ["worker-a", "0", "worker_1"]
with self.assertRaisesRegex(ValueError, "legacy index token"):
parse_queue_request_payload(payload_data)
def test_invalid_top_level_payload_raises(self):
with self.assertRaisesRegex(ValueError, "Expected a JSON object body"):
parse_queue_request_payload(["not", "an", "object"])
+154
View File
@@ -0,0 +1,154 @@
import importlib.util
import os
import tempfile
import types
import unittest
from pathlib import Path
def _load_worker_monitor_module():
module_path = Path(__file__).resolve().parents[3] / "workers" / "worker_monitor.py"
spec = importlib.util.spec_from_file_location("dist_test_worker_monitor", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
return module
worker_monitor = _load_worker_monitor_module()
class _FakeWorkerProcess:
def __init__(self, pid=9001, poll_sequence=None):
self.pid = pid
self._poll_sequence = list(poll_sequence or [None])
self._poll_calls = 0
self.returncode = None
self.terminate_calls = 0
self.kill_calls = 0
self.wait_calls = 0
def poll(self):
if self._poll_calls < len(self._poll_sequence):
value = self._poll_sequence[self._poll_calls]
else:
value = self._poll_sequence[-1]
self._poll_calls += 1
if value is not None:
self.returncode = value
return value
def terminate(self):
self.terminate_calls += 1
self.returncode = 0
def kill(self):
self.kill_calls += 1
self.returncode = -9
def wait(self, timeout=None):
_ = timeout
self.wait_calls += 1
if self.returncode is None:
self.returncode = 0
return self.returncode
class WorkerMonitorTests(unittest.TestCase):
def setUp(self):
self._orig_launch = worker_monitor.launch_process_with_timeout
self._orig_is_alive = worker_monitor.is_process_alive
self._orig_terminate = getattr(worker_monitor, "terminate_process", None)
self._orig_sleep = worker_monitor.time.sleep
self._orig_signal = worker_monitor.signal.signal
self._env_backup = dict(os.environ)
worker_monitor.time.sleep = lambda _seconds: None
worker_monitor.signal.signal = lambda *_args, **_kwargs: None
def tearDown(self):
worker_monitor.launch_process_with_timeout = self._orig_launch
worker_monitor.is_process_alive = self._orig_is_alive
if self._orig_terminate is not None:
worker_monitor.terminate_process = self._orig_terminate
worker_monitor.time.sleep = self._orig_sleep
worker_monitor.signal.signal = self._orig_signal
os.environ.clear()
os.environ.update(self._env_backup)
def test_main_validates_required_env_and_args(self):
os.environ.pop("COMFYUI_MASTER_PID", None)
self.assertEqual(worker_monitor.main(["python", "main.py"]), 1)
os.environ["COMFYUI_MASTER_PID"] = "not-an-int"
self.assertEqual(worker_monitor.main(["python", "main.py"]), 1)
os.environ["COMFYUI_MASTER_PID"] = "101"
self.assertEqual(worker_monitor.main([]), 1)
def test_main_delegates_to_monitor_and_run(self):
calls = []
def _monitor(master_pid, command):
calls.append((master_pid, list(command)))
return 7
os.environ["COMFYUI_MASTER_PID"] = "202"
original_monitor = worker_monitor.monitor_and_run
worker_monitor.monitor_and_run = _monitor
try:
exit_code = worker_monitor.main(["python", "worker.py"])
finally:
worker_monitor.monitor_and_run = original_monitor
self.assertEqual(exit_code, 7)
self.assertEqual(calls, [(202, ["python", "worker.py"])])
def test_monitor_and_run_returns_worker_exit_code(self):
fake_process = _FakeWorkerProcess(poll_sequence=[None, 5])
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
worker_monitor.is_process_alive = lambda _pid: True
worker_monitor.terminate_process = lambda *_args, **_kwargs: None
exit_code = worker_monitor.monitor_and_run(100, ["python", "worker.py"])
self.assertEqual(exit_code, 5)
def test_monitor_and_run_terminates_worker_when_master_dies(self):
fake_process = _FakeWorkerProcess(poll_sequence=[None, None])
term_calls = []
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
worker_monitor.is_process_alive = lambda _pid: False
def _terminate(process, timeout):
term_calls.append((process.pid, timeout))
process.returncode = 0
worker_monitor.terminate_process = _terminate
exit_code = worker_monitor.monitor_and_run(333, ["python", "worker.py"])
self.assertEqual(exit_code, 0)
self.assertEqual(term_calls, [(9001, worker_monitor.PROCESS_TERMINATION_TIMEOUT)])
def test_monitor_and_run_writes_pid_file_when_requested(self):
fake_process = _FakeWorkerProcess(poll_sequence=[0])
worker_monitor.launch_process_with_timeout = lambda *_args, **_kwargs: fake_process
worker_monitor.is_process_alive = lambda _pid: True
worker_monitor.terminate_process = lambda *_args, **_kwargs: None
with tempfile.TemporaryDirectory() as tmpdir:
pid_file = Path(tmpdir) / "worker.pid"
os.environ["WORKER_PID_FILE"] = str(pid_file)
exit_code = worker_monitor.monitor_and_run(444, ["python", "worker.py"])
self.assertEqual(exit_code, 0)
contents = pid_file.read_text(encoding="utf-8")
monitor_pid, worker_pid = contents.split(",")
self.assertTrue(monitor_pid.isdigit())
self.assertEqual(worker_pid, "9001")
if __name__ == "__main__":
unittest.main()
+1
View File
@@ -0,0 +1 @@
"""Unit tests."""
+53
View File
@@ -0,0 +1,53 @@
import unittest
import torch
from upscale.conditioning import clone_conditioning, clone_control_chain
class _FakeControl:
def __init__(self, hint, previous=None):
self.cond_hint_original = hint
self.previous_controlnet = previous
class ConditioningUtilsTests(unittest.TestCase):
def test_clone_control_chain_clones_hints_and_links(self):
tail = _FakeControl(torch.ones((1, 2)))
head = _FakeControl(torch.zeros((1, 2)), previous=tail)
cloned = clone_control_chain(head, clone_hint=True)
self.assertIsNot(cloned, head)
self.assertIsNot(cloned.previous_controlnet, tail)
self.assertTrue(torch.equal(cloned.cond_hint_original, head.cond_hint_original))
self.assertTrue(torch.equal(cloned.previous_controlnet.cond_hint_original, tail.cond_hint_original))
self.assertIsNot(cloned.cond_hint_original, head.cond_hint_original)
def test_clone_conditioning_clones_emb_and_known_fields(self):
control = _FakeControl(torch.ones((1, 3)))
cond = [
[
torch.tensor([[1.0, 2.0, 3.0]]),
{
"control": control,
"mask": torch.ones((1, 1, 1)),
"pooled_output": torch.zeros((1, 4)),
"area": [0, 1, 2, 3],
},
]
]
cloned = clone_conditioning(cond, clone_hints=True)
self.assertTrue(torch.equal(cloned[0][0], cond[0][0]))
self.assertIsNot(cloned[0][0], cond[0][0])
self.assertIsNot(cloned[0][1]["control"], cond[0][1]["control"])
self.assertIsNot(cloned[0][1]["mask"], cond[0][1]["mask"])
self.assertIsNot(cloned[0][1]["pooled_output"], cond[0][1]["pooled_output"])
self.assertEqual(cloned[0][1]["area"], cond[0][1]["area"])
self.assertIsNot(cloned[0][1]["area"], cond[0][1]["area"])
if __name__ == "__main__":
unittest.main()
+247
View File
@@ -0,0 +1,247 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
from types import SimpleNamespace
import numpy as np
from PIL import Image
try:
import torch
except ModuleNotFoundError: # pragma: no cover - optional dependency in CI/runtime
torch = None
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
modes_pkg.__path__ = []
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
image_module = types.ModuleType(f"{package_name}.utils.image")
def tensor_to_pil(image_batch, index):
tensor = image_batch[index].detach().cpu().clamp(0, 1)
array = (tensor.numpy() * 255.0).astype("uint8")
return Image.fromarray(array)
def pil_to_tensor(image):
if torch is None:
raise RuntimeError("torch is required for this test helper")
array = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(array).unsqueeze(0)
image_module.tensor_to_pil = tensor_to_pil
image_module.pil_to_tensor = pil_to_tensor
sys.modules[f"{package_name}.utils.image"] = image_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
async_helpers_module.run_async_in_server_loop = lambda _coro, timeout=None: (_ for _ in ()).throw(
RuntimeError(f"run_async_in_server_loop unexpectedly called with timeout={timeout}")
)
sys.modules[f"{package_name}.utils.async_helpers"] = async_helpers_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.get_worker_timeout_seconds = lambda: 5.0
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.JOB_POLL_INTERVAL = 0.0
constants_module.JOB_POLL_MAX_ATTEMPTS = 3
constants_module.TILE_WAIT_TIMEOUT = 5.0
constants_module.TILE_SEND_TIMEOUT = 5.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
if torch is None:
torch_module = types.ModuleType("torch")
torch_module.cat = lambda *_args, **_kwargs: None
sys.modules["torch"] = torch_module
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
job_store_module.ensure_tile_jobs_initialized = lambda: None
async def init_dynamic_job(*_args, **_kwargs):
return None
job_store_module.init_dynamic_job = init_dynamic_job
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
comfy_module = types.ModuleType("comfy")
model_mgmt = types.ModuleType("comfy.model_management")
model_mgmt.throw_exception_if_processing_interrupted = lambda: None
comfy_module.model_management = model_mgmt
sys.modules["comfy"] = comfy_module
sys.modules["comfy.model_management"] = model_mgmt
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_dynamic_mode_module():
package_name = "dist_dynamic_mode_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
return _load_module(package_name, "upscale/modes/dynamic.py", "upscale.modes.dynamic")
dynamic_module = _load_dynamic_mode_module()
def _run_sync(value, timeout=None):
_ = timeout
if asyncio.iscoroutine(value):
return asyncio.run(value)
return value
dynamic_module.run_async_in_server_loop = _run_sync
class _DummyDynamicNode(dynamic_module.DynamicModeMixin):
def __init__(self, *, poll_ready=True, assignments=None):
self._poll_ready = poll_ready
self._assignments = list(assignments or [])
self.sent_images = []
self.heartbeat_count = 0
self.completion_sent = False
def round_to_multiple(self, value):
return int(value)
def calculate_tiles(self, _width, _height, _tile_width, _tile_height, _force_uniform_tiles):
return [(0, 0)]
def _poll_job_ready(self, *_args, **_kwargs):
return self._poll_ready
async def check_job_status(self, *_args, **_kwargs):
return self._poll_ready
async def request_assignment(self, *_args, **_kwargs):
if self._assignments:
image_idx, estimated_remaining = self._assignments.pop(0)
return SimpleNamespace(
kind="image" if image_idx is not None else "none",
task_idx=image_idx,
estimated_remaining=estimated_remaining,
batched_static=False,
)
return SimpleNamespace(kind="none", task_idx=None, estimated_remaining=0, batched_static=False)
def slice_conditioning(self, positive, negative, _image_idx):
return positive, negative
def process_and_blend_tile(self, _tile_idx, _pos, _source_tensor, local_image, *_args, **_kwargs):
return local_image
async def send_heartbeat(self, *_args, **_kwargs):
self.heartbeat_count += 1
async def send_full_image(self, local_image, image_idx, _multi_job_id, _master_url, _worker_id, is_last):
self.sent_images.append((image_idx, is_last, local_image.size))
async def send_worker_complete_signal(self, *_args, **_kwargs):
self.completion_sent = True
@unittest.skipIf(torch is None, "torch is not installed")
class DynamicModeTests(unittest.TestCase):
def setUp(self):
self.image_batch = torch.zeros((1, 8, 8, 3), dtype=torch.float32)
self.core_args = dynamic_module.UpscaleCoreArgs(
model=None,
positive=None,
negative=None,
vae=None,
seed=1,
steps=10,
cfg=7.0,
sampler_name="euler",
scheduler="normal",
denoise=0.5,
tiled_decode=False,
)
def test_worker_dynamic_returns_early_if_job_not_ready(self):
node = _DummyDynamicNode(poll_ready=False)
result = node.process_worker_dynamic(
upscaled_image=self.image_batch,
core_args=self.core_args,
tile_width=8,
tile_height=8,
padding=4,
mask_blur=2,
force_uniform_tiles=True,
multi_job_id="job-1",
master_url="http://master.local:8188",
worker_id="worker-a",
enabled_worker_ids='["worker-a"]',
dynamic_threshold=8,
)
self.assertEqual(result[0].shape, self.image_batch.shape)
self.assertEqual(node.sent_images, [])
self.assertFalse(node.completion_sent)
def test_worker_dynamic_processes_assigned_image_and_sends_completion(self):
node = _DummyDynamicNode(poll_ready=True, assignments=[(0, 0)])
result = node.process_worker_dynamic(
upscaled_image=self.image_batch,
core_args=self.core_args,
tile_width=8,
tile_height=8,
padding=4,
mask_blur=2,
force_uniform_tiles=True,
multi_job_id="job-2",
master_url="http://master.local:8188",
worker_id="worker-a",
enabled_worker_ids='["worker-a"]',
dynamic_threshold=8,
)
self.assertEqual(result[0].shape, self.image_batch.shape)
self.assertEqual(len(node.sent_images), 1)
self.assertEqual(node.sent_images[0][0], 0)
self.assertTrue(node.sent_images[0][1])
self.assertTrue(node.completion_sent)
self.assertGreaterEqual(node.heartbeat_count, 1)
if __name__ == "__main__":
unittest.main()
+39
View File
@@ -0,0 +1,39 @@
import unittest
import torch
from PIL import Image
from utils.image import ensure_contiguous, pil_to_tensor, tensor_to_pil
class ImageUtilsTests(unittest.TestCase):
def test_tensor_to_pil_converts_batch_item(self):
tensor = torch.tensor(
[[[[0.0, 0.5, 1.0], [1.0, 0.0, 0.0]]]],
dtype=torch.float32,
)
image = tensor_to_pil(tensor, batch_index=0)
self.assertEqual(image.size, (2, 1))
self.assertEqual(image.getpixel((0, 0)), (0, 127, 255))
self.assertEqual(image.getpixel((1, 0)), (255, 0, 0))
def test_pil_to_tensor_adds_channel_for_grayscale(self):
grayscale = Image.new("L", (3, 2), color=128)
tensor = pil_to_tensor(grayscale)
self.assertEqual(tuple(tensor.shape), (1, 2, 3, 1))
self.assertAlmostEqual(float(tensor[0, 0, 0, 0]), 128 / 255.0, places=6)
def test_ensure_contiguous_makes_non_contiguous_tensor_contiguous(self):
base = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
non_contiguous = base.transpose(1, 2)
self.assertFalse(non_contiguous.is_contiguous())
contiguous = ensure_contiguous(non_contiguous)
self.assertTrue(contiguous.is_contiguous())
self.assertTrue(torch.equal(contiguous, non_contiguous))
if __name__ == "__main__":
unittest.main()
+131
View File
@@ -0,0 +1,131 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_job_state_modules():
package_name = "dist_job_state_testpkg"
_bootstrap_package(package_name)
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
prompt_server = types.SimpleNamespace(
distributed_pending_tile_jobs={},
distributed_tile_jobs_lock=asyncio.Lock(),
)
job_store_module = types.ModuleType(f"{package_name}.upscale.job_store")
job_store_module.ensure_tile_jobs_initialized = lambda: prompt_server
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
call_log = {"args": None}
job_timeout_module = types.ModuleType(f"{package_name}.upscale.job_timeout")
async def _requeue(multi_job_id, batch_size):
call_log["args"] = (multi_job_id, batch_size)
return 3
job_timeout_module._check_and_requeue_timed_out_workers = _requeue
sys.modules[f"{package_name}.upscale.job_timeout"] = job_timeout_module
job_state = _load_module(package_name, "upscale/job_state.py", "upscale.job_state")
return job_models, job_state, prompt_server, call_log
job_models_module, job_state_module, prompt_server, timeout_call_log = _load_job_state_modules()
class _Node(job_state_module.JobStateMixin):
pass
class JobStateMixinTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
prompt_server.distributed_pending_tile_jobs = {}
prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
self.node = _Node()
async def test_get_all_completed_tasks_supports_tile_and_image_modes(self):
tile_state = job_models_module.TileJobState("tile-job")
tile_state.completed_tasks[4] = "tile-4"
image_state = job_models_module.ImageJobState("image-job")
image_state.completed_images[2] = "image-2"
prompt_server.distributed_pending_tile_jobs = {
"tile-job": tile_state,
"image-job": image_state,
}
self.assertEqual(await self.node._get_all_completed_tasks("tile-job"), {4: "tile-4"})
self.assertEqual(await self.node._get_all_completed_tasks("image-job"), {2: "image-2"})
async def test_next_index_and_pending_count_from_queues(self):
image_state = job_models_module.ImageJobState("image-job")
await image_state.pending_images.put(7)
tile_state = job_models_module.TileJobState("tile-job")
await tile_state.pending_tasks.put(9)
prompt_server.distributed_pending_tile_jobs = {
"image-job": image_state,
"tile-job": tile_state,
}
self.assertEqual(await self.node._get_next_image_index("image-job"), 7)
self.assertEqual(await self.node._get_next_tile_index("tile-job"), 9)
self.assertEqual(await self.node._get_pending_count("image-job"), 0)
self.assertEqual(await self.node._get_pending_count("tile-job"), 0)
async def test_drain_worker_results_queue_collects_completed_images(self):
image_state = job_models_module.ImageJobState("image-job")
await image_state.queue.put({"worker_id": "w1", "image_idx": 1, "image": "img1"})
await image_state.queue.put({"worker_id": "w2", "image_idx": 2, "image": "img2"})
prompt_server.distributed_pending_tile_jobs = {"image-job": image_state}
drained = await self.node._drain_worker_results_queue("image-job")
self.assertEqual(drained, 2)
self.assertEqual(image_state.completed_images[1], "img1")
self.assertEqual(image_state.completed_images[2], "img2")
async def test_check_and_requeue_delegates_to_timeout_helper(self):
count = await self.node._check_and_requeue_timed_out_workers("job-abc", 5)
self.assertEqual(count, 3)
self.assertEqual(timeout_call_log["args"], ("job-abc", 5))
if __name__ == "__main__":
unittest.main()
+164
View File
@@ -0,0 +1,164 @@
import asyncio
import importlib.util
import sys
import types
import unittest
from pathlib import Path
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
prompt_server = types.SimpleNamespace()
runtime_state_module = types.ModuleType(f"{package_name}.utils.runtime_state")
def _ensure_distributed_runtime_state(server_instance=None):
ps = server_instance or prompt_server
if not hasattr(ps, "distributed_pending_jobs"):
ps.distributed_pending_jobs = {}
if not hasattr(ps, "distributed_jobs_lock"):
ps.distributed_jobs_lock = asyncio.Lock()
if not hasattr(ps, "distributed_job_allowed_workers"):
ps.distributed_job_allowed_workers = {}
if not hasattr(ps, "distributed_pending_tile_jobs"):
ps.distributed_pending_tile_jobs = {}
if not hasattr(ps, "distributed_tile_jobs_lock"):
ps.distributed_tile_jobs_lock = asyncio.Lock()
return types.SimpleNamespace(
distributed_pending_jobs=ps.distributed_pending_jobs,
distributed_jobs_lock=ps.distributed_jobs_lock,
distributed_job_allowed_workers=ps.distributed_job_allowed_workers,
distributed_pending_tile_jobs=ps.distributed_pending_tile_jobs,
distributed_tile_jobs_lock=ps.distributed_tile_jobs_lock,
)
runtime_state_module.ensure_distributed_runtime_state = _ensure_distributed_runtime_state
runtime_state_module.get_prompt_server_instance = lambda: prompt_server
sys.modules[f"{package_name}.utils.runtime_state"] = runtime_state_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_job_store_modules():
package_name = "dist_job_store_testpkg"
_bootstrap_package(package_name)
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
job_store = _load_module(package_name, "upscale/job_store.py", "upscale.job_store")
return job_models, job_store
job_models_module, job_store_module = _load_job_store_modules()
class JobStoreTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.prompt_server = job_store_module.get_prompt_server_instance()
self.prompt_server.distributed_pending_tile_jobs = {}
self.prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
async def test_ensure_tile_jobs_initialized_prunes_invalid_entries(self):
self.prompt_server.distributed_pending_tile_jobs["valid"] = job_models_module.TileJobState("job-valid")
self.prompt_server.distributed_pending_tile_jobs["invalid"] = {"not": "state"}
job_store_module.ensure_tile_jobs_initialized()
self.assertIn("valid", self.prompt_server.distributed_pending_tile_jobs)
self.assertNotIn("invalid", self.prompt_server.distributed_pending_tile_jobs)
async def test_init_dynamic_job_populates_pending_images(self):
await job_store_module.init_dynamic_job(
multi_job_id="job-dynamic",
batch_size=3,
enabled_workers=["worker-a"],
all_indices=[2, 0, 1],
)
job_data = self.prompt_server.distributed_pending_tile_jobs["job-dynamic"]
self.assertIsInstance(job_data, job_models_module.ImageJobState)
self.assertEqual(job_data.batch_size, 3)
self.assertIn("worker-a", job_data.worker_status)
pulled = [job_data.pending_images.get_nowait() for _ in range(3)]
self.assertEqual(pulled, [2, 0, 1])
async def test_init_static_job_batched_queues_tile_ids(self):
await job_store_module.init_static_job_batched(
multi_job_id="job-static",
batch_size=2,
num_tiles_per_image=4,
enabled_workers=["worker-a", "worker-b"],
)
job_data = self.prompt_server.distributed_pending_tile_jobs["job-static"]
self.assertIsInstance(job_data, job_models_module.TileJobState)
self.assertTrue(job_data.batched_static)
pulled = [job_data.pending_tasks.get_nowait() for _ in range(4)]
self.assertEqual(pulled, [0, 1, 2, 3])
async def test_drain_results_queue_collects_images_tiles_and_completion(self):
job_data = job_models_module.TileJobState("job-drain")
job_data.worker_status = {"worker-a": 1.0, "worker-b": 1.0}
self.prompt_server.distributed_pending_tile_jobs["job-drain"] = job_data
await job_data.queue.put(
{
"worker_id": "worker-a",
"is_last": True,
"image_idx": 5,
"image": "img-5",
}
)
await job_data.queue.put(
{
"worker_id": "worker-b",
"is_last": False,
"tiles": [
{"tile_idx": 0, "global_idx": 0, "tensor": "tile-0"},
{"tile_idx": 1, "global_idx": 1, "tensor": "tile-1"},
],
}
)
drained = await job_store_module.drain_results_queue("job-drain")
self.assertEqual(drained, 3)
self.assertEqual(job_data.completed_tasks[5], "img-5")
self.assertEqual(job_data.completed_tasks[0]["tensor"], "tile-0")
self.assertEqual(job_data.completed_tasks[1]["tensor"], "tile-1")
self.assertNotIn("worker-a", job_data.worker_status)
self.assertIn("worker-b", job_data.worker_status)
if __name__ == "__main__":
unittest.main()
@@ -9,7 +9,7 @@ from pathlib import Path
def _load_job_timeout_module():
module_path = Path(__file__).resolve().parents[1] / "upscale" / "job_timeout.py"
module_path = Path(__file__).resolve().parents[3] / "upscale" / "job_timeout.py"
package_name = "dist_job_timeout_testpkg"
for mod_name in list(sys.modules):
@@ -15,7 +15,7 @@ except ImportError:
def _load_payload_parsers_module():
# payload_parsers.py has no relative imports; only stdlib + PIL
module_path = Path(__file__).resolve().parents[1] / "upscale" / "payload_parsers.py"
module_path = Path(__file__).resolve().parents[3] / "upscale" / "payload_parsers.py"
spec = importlib.util.spec_from_file_location("upscale_payload_parsers", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
+248
View File
@@ -0,0 +1,248 @@
import asyncio
import importlib.util
import sys
import time
import types
import unittest
from pathlib import Path
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.get_worker_timeout_seconds = lambda: 0.05
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.DYNAMIC_MODE_MAX_POLL_TIMEOUT = 10.0
constants_module.HEARTBEAT_INTERVAL = 0.0
sys.modules[f"{package_name}.utils.constants"] = constants_module
comfy_module = types.ModuleType("comfy")
model_mgmt = types.ModuleType("comfy.model_management")
class _InterruptProcessingException(Exception):
pass
model_mgmt.processing_interrupted = lambda: False
model_mgmt.InterruptProcessingException = _InterruptProcessingException
comfy_module.model_management = model_mgmt
sys.modules["comfy"] = comfy_module
sys.modules["comfy.model_management"] = model_mgmt
prompt_server = types.SimpleNamespace(
distributed_pending_tile_jobs={},
distributed_tile_jobs_lock=asyncio.Lock(),
)
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
sys.modules["server"] = server_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_result_collector_modules():
package_name = "dist_result_collector_testpkg"
_bootstrap_package(package_name)
job_models = _load_module(package_name, "upscale/job_models.py", "upscale.job_models")
job_store = _load_module(package_name, "upscale/job_store.py", "upscale.job_store")
job_timeout = types.ModuleType(f"{package_name}.upscale.job_timeout")
async def _noop_requeue(_multi_job_id, _batch_size):
return 0
job_timeout._check_and_requeue_timed_out_workers = _noop_requeue
job_timeout.check_and_requeue_timed_out_workers = _noop_requeue
sys.modules[f"{package_name}.upscale.job_timeout"] = job_timeout
result_collector = _load_module(
package_name,
"upscale/result_collector.py",
"upscale.result_collector",
)
return job_models, job_store, result_collector
job_models_module, job_store_module, result_collector_module = _load_result_collector_modules()
class _Node(result_collector_module.ResultCollectorMixin):
def __init__(self):
self.requeue_calls = []
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
self.requeue_calls.append((multi_job_id, batch_size))
return 1
class ResultCollectorMixinTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.prompt_server = job_store_module.server.PromptServer.instance
self.prompt_server.distributed_pending_tile_jobs = {}
self.prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
self.node = _Node()
async def test_static_collection_collects_tiles_and_cleans_up_job(self):
job_data = job_models_module.TileJobState("job-static")
self.prompt_server.distributed_pending_tile_jobs["job-static"] = job_data
await job_data.queue.put(
{
"worker_id": "worker-a",
"is_last": True,
"tiles": [
{"tile_idx": 3},
{
"tile_idx": 2,
"x": 10,
"y": 20,
"extracted_width": 64,
"extracted_height": 64,
"padding": 8,
"batch_idx": 1,
"global_idx": 9,
"tensor": "tile-9",
},
],
}
)
collected = await self.node._async_collect_worker_tiles(
"job-static",
num_workers=1,
)
self.assertEqual(list(collected.keys()), [9])
self.assertEqual(collected[9]["tensor"], "tile-9")
self.assertEqual(collected[9]["worker_id"], "worker-a")
self.assertNotIn("job-static", self.prompt_server.distributed_pending_tile_jobs)
async def test_dynamic_collection_stops_at_remaining_target_and_cleans_up(self):
job_data = job_models_module.ImageJobState("job-dynamic")
self.prompt_server.distributed_pending_tile_jobs["job-dynamic"] = job_data
await job_data.queue.put(
{
"worker_id": "worker-a",
"is_last": False,
"image_idx": 2,
"image": "img-2",
}
)
completed = await self.node.collect_dynamic_images(
"job-dynamic",
remaining_to_collect=1,
num_workers=2,
batch_size=3,
master_processed_count=0,
)
self.assertEqual(completed[2], "img-2")
self.assertNotIn("job-dynamic", self.prompt_server.distributed_pending_tile_jobs)
async def test_dynamic_timeout_calls_requeue_helper(self):
job_data = job_models_module.ImageJobState("job-timeout")
job_data.worker_status = {"worker-a": time.time() - 5}
self.prompt_server.distributed_pending_tile_jobs["job-timeout"] = job_data
old_timeout = result_collector_module.get_worker_timeout_seconds
result_collector_module.get_worker_timeout_seconds = lambda: 0.01
try:
completed = await self.node.collect_dynamic_images(
"job-timeout",
remaining_to_collect=None,
num_workers=1,
batch_size=4,
master_processed_count=0,
)
finally:
result_collector_module.get_worker_timeout_seconds = old_timeout
self.assertEqual(completed, {})
self.assertEqual(self.node.requeue_calls, [("job-timeout", 4)])
self.assertNotIn("job-timeout", self.prompt_server.distributed_pending_tile_jobs)
async def test_log_worker_timeout_status_handles_non_job_state(self):
worker_ids = self.node._log_worker_timeout_status(
{"worker_status": {"worker-a": time.time()}},
current_time=time.time(),
multi_job_id="job-invalid",
)
self.assertEqual(worker_ids, [])
async def test_log_worker_timeout_status_returns_worker_ids_for_job_state(self):
job_data = job_models_module.ImageJobState("job-state")
job_data.worker_status = {"worker-a": time.time() - 1}
worker_ids = self.node._log_worker_timeout_status(
job_data,
current_time=time.time(),
multi_job_id="job-state",
)
self.assertEqual(worker_ids, ["worker-a"])
async def test_static_mode_mismatch_raises(self):
self.prompt_server.distributed_pending_tile_jobs["job-mismatch"] = (
job_models_module.ImageJobState("job-mismatch")
)
with self.assertRaises(RuntimeError):
await self.node._async_collect_worker_tiles(
"job-mismatch",
num_workers=1,
)
async def test_dynamic_mode_mismatch_raises(self):
self.prompt_server.distributed_pending_tile_jobs["job-mismatch"] = (
job_models_module.TileJobState("job-mismatch")
)
with self.assertRaises(RuntimeError):
await self.node.collect_dynamic_images(
"job-mismatch",
remaining_to_collect=1,
num_workers=1,
batch_size=1,
master_processed_count=0,
)
async def test_mark_image_completed_updates_job_state(self):
job_data = job_models_module.ImageJobState("job-mark")
self.prompt_server.distributed_pending_tile_jobs["job-mark"] = job_data
await self.node.mark_image_completed("job-mark", 7, "img-7")
self.assertEqual(job_data.completed_images[7], "img-7")
if __name__ == "__main__":
unittest.main()
+207
View File
@@ -0,0 +1,207 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
import numpy as np
from PIL import Image
try:
import torch
except ModuleNotFoundError: # pragma: no cover - optional dependency in CI/runtime
torch = None
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
modes_pkg = types.ModuleType(f"{package_name}.upscale.modes")
modes_pkg.__path__ = []
sys.modules[f"{package_name}.upscale.modes"] = modes_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
image_module = types.ModuleType(f"{package_name}.utils.image")
def tensor_to_pil(image_batch, index):
tensor = image_batch[index].detach().cpu().clamp(0, 1)
array = (tensor.numpy() * 255.0).astype("uint8")
return Image.fromarray(array)
def pil_to_tensor(image):
if torch is None:
raise RuntimeError("torch is required for this test helper")
array = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(array).unsqueeze(0)
def blend_processed_batch_item(
result_images,
processed_batch,
batch_index,
blend_fn,
x1,
y1,
ew,
eh,
tile_mask,
padding,
):
tile_pil = tensor_to_pil(processed_batch, batch_index)
if tile_pil.size != (ew, eh):
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
result_images[batch_index] = blend_fn(
result_images[batch_index],
tile_pil,
x1,
y1,
(ew, eh),
tile_mask,
padding,
)
image_module.tensor_to_pil = tensor_to_pil
image_module.pil_to_tensor = pil_to_tensor
image_module.blend_processed_batch_item = blend_processed_batch_item
sys.modules[f"{package_name}.utils.image"] = image_module
if torch is None:
torch_module = types.ModuleType("torch")
torch_module.cat = lambda *_args, **_kwargs: None
sys.modules["torch"] = torch_module
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_single_gpu_mode_module():
package_name = "dist_single_gpu_mode_testpkg"
_bootstrap_package(package_name)
_load_module(package_name, "upscale/processing_args.py", "upscale.processing_args")
_load_module(package_name, "upscale/tile_processing.py", "upscale.tile_processing")
_load_module(package_name, "upscale/mode_contexts.py", "upscale.mode_contexts")
return _load_module(package_name, "upscale/modes/single_gpu.py", "upscale.modes.single_gpu")
single_gpu_module = _load_single_gpu_mode_module()
class _DummySingleGpuNode(single_gpu_module.SingleGpuModeMixin):
def __init__(self):
self.round_calls = []
self.process_batch_calls = 0
def round_to_multiple(self, value):
self.round_calls.append(int(value))
return int(value)
def calculate_tiles(self, _width, _height, _tile_width, _tile_height, _force_uniform_tiles):
return [(0, 0)]
def create_tile_mask(self, _width, _height, _tx, _ty, _tile_width, _tile_height, _mask_blur):
return None
def extract_batch_tile_with_padding(self, source_batch, _tx, _ty, _tile_width, _tile_height, _padding, _force_uniform_tiles):
_, h, w, _ = source_batch.shape
return source_batch, 0, 0, w, h
def process_tiles_batch(
self,
tile_batch,
_model,
_positive,
_negative,
_vae,
_seed,
_steps,
_cfg,
_sampler_name,
_scheduler,
_denoise,
_tiled_decode,
_region,
_canvas_shape,
):
self.process_batch_calls += 1
return tile_batch
def blend_tile(self, _base_image, tile_pil, _x1, _y1, _size, _tile_mask, _padding):
return tile_pil
@unittest.skipIf(torch is None, "torch is not installed")
class SingleGpuModeTests(unittest.TestCase):
def setUp(self):
self.node = _DummySingleGpuNode()
self.images = torch.rand((2, 8, 8, 3), dtype=torch.float32)
self.core_args = single_gpu_module.UpscaleCoreArgs(
model=None,
positive=None,
negative=None,
vae=None,
seed=1,
steps=10,
cfg=7.5,
sampler_name="euler",
scheduler="normal",
denoise=0.5,
tiled_decode=False,
)
def test_process_single_gpu_preserves_batch_shape(self):
result = self.node.process_single_gpu(
upscaled_image=self.images,
core_args=self.core_args,
tile_width=8,
tile_height=8,
padding=4,
mask_blur=2,
force_uniform_tiles=True,
)[0]
self.assertEqual(result.shape, self.images.shape)
self.assertTrue(torch.allclose(result, self.images, atol=(1.5 / 255.0)))
def test_process_single_gpu_invokes_rounding_and_batch_processing(self):
self.node.process_single_gpu(
upscaled_image=self.images,
core_args=self.core_args,
tile_width=16,
tile_height=24,
padding=4,
mask_blur=2,
force_uniform_tiles=True,
)
self.assertEqual(self.node.round_calls[:2], [16, 24])
self.assertEqual(self.node.process_batch_calls, 1)
if __name__ == "__main__":
unittest.main()
@@ -4,12 +4,13 @@ import sys
import types
import unittest
from pathlib import Path
from types import SimpleNamespace
import torch
def _load_static_mode_module():
module_path = Path(__file__).resolve().parents[1] / "upscale" / "modes" / "static.py"
module_path = Path(__file__).resolve().parents[3] / "upscale" / "modes" / "static.py"
package_name = "dist_static_mode_testpkg"
for mod_name in list(sys.modules):
@@ -65,8 +66,34 @@ def _load_static_mode_module():
arr = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(arr).unsqueeze(0)
def _blend_processed_batch_item(
result_images,
processed_batch,
batch_index,
blend_fn,
x1,
y1,
ew,
eh,
tile_mask,
padding,
):
tile_pil = _tensor_to_pil(processed_batch, batch_index)
if tile_pil.size != (ew, eh):
tile_pil = tile_pil.resize((ew, eh), PILImage.LANCZOS)
result_images[batch_index] = blend_fn(
result_images[batch_index],
tile_pil,
x1,
y1,
(ew, eh),
tile_mask,
padding,
)
image_module.tensor_to_pil = _tensor_to_pil
image_module.pil_to_tensor = _pil_to_tensor
image_module.blend_processed_batch_item = _blend_processed_batch_item
sys.modules[f"{package_name}.utils.image"] = image_module
async_helpers_module = types.ModuleType(f"{package_name}.utils.async_helpers")
@@ -102,10 +129,10 @@ def _load_static_mode_module():
distributed_pending_tile_jobs={},
)
job_store_module.init_static_job_batched = _noop
job_store_module._mark_task_completed = _noop
job_store_module._cleanup_job = _noop
job_store_module._drain_results_queue = _noop
job_store_module._get_completed_count = _noop
job_store_module.mark_task_completed = _noop
job_store_module.cleanup_job = _noop
job_store_module.drain_results_queue = _noop
job_store_module.get_completed_count = _noop
sys.modules[f"{package_name}.upscale.job_store"] = job_store_module
job_models_module = types.ModuleType(f"{package_name}.upscale.job_models")
@@ -116,6 +143,36 @@ def _load_static_mode_module():
job_models_module.TileJobState = _TileJobState
sys.modules[f"{package_name}.upscale.job_models"] = job_models_module
processing_args_path = Path(__file__).resolve().parents[3] / "upscale" / "processing_args.py"
processing_args_spec = importlib.util.spec_from_file_location(
f"{package_name}.upscale.processing_args",
processing_args_path,
)
processing_args_module = importlib.util.module_from_spec(processing_args_spec)
assert processing_args_spec is not None and processing_args_spec.loader is not None
sys.modules[processing_args_spec.name] = processing_args_module
processing_args_spec.loader.exec_module(processing_args_module)
tile_processing_path = Path(__file__).resolve().parents[3] / "upscale" / "tile_processing.py"
tile_processing_spec = importlib.util.spec_from_file_location(
f"{package_name}.upscale.tile_processing",
tile_processing_path,
)
tile_processing_module = importlib.util.module_from_spec(tile_processing_spec)
assert tile_processing_spec is not None and tile_processing_spec.loader is not None
sys.modules[tile_processing_spec.name] = tile_processing_module
tile_processing_spec.loader.exec_module(tile_processing_module)
mode_contexts_path = Path(__file__).resolve().parents[3] / "upscale" / "mode_contexts.py"
mode_contexts_spec = importlib.util.spec_from_file_location(
f"{package_name}.upscale.mode_contexts",
mode_contexts_path,
)
mode_contexts_module = importlib.util.module_from_spec(mode_contexts_spec)
assert mode_contexts_spec is not None and mode_contexts_spec.loader is not None
sys.modules[mode_contexts_spec.name] = mode_contexts_module
mode_contexts_spec.loader.exec_module(mode_contexts_module)
spec = importlib.util.spec_from_file_location(f"{package_name}.upscale.modes.static", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
@@ -148,14 +205,20 @@ class _FakeStaticWorker(static_mode.StaticModeMixin):
def _poll_job_ready(self, *_args, **_kwargs):
return self.job_ready
async def _request_tile_from_master(self, *_args, **_kwargs):
async def request_assignment(self, *_args, **_kwargs):
self.request_calls += 1
return self.tile_sequence.pop(0)
tile_idx, estimated_remaining, batched_static = self.tile_sequence.pop(0)
return SimpleNamespace(
kind="tile" if tile_idx is not None else "none",
task_idx=tile_idx,
estimated_remaining=estimated_remaining,
batched_static=batched_static,
)
async def _send_heartbeat_to_master(self, *_args, **_kwargs):
async def send_heartbeat(self, *_args, **_kwargs):
self.heartbeat_calls += 1
async def send_tiles_batch_to_master(
async def send_tiles_batch(
self,
processed_tiles,
_multi_job_id,
+179
View File
@@ -0,0 +1,179 @@
import copy
import importlib.util
import sys
import types
import unittest
from contextlib import nullcontext
from pathlib import Path
import numpy as np
import torch
from PIL import Image
def _bootstrap_package(package_name):
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
upscale_pkg = types.ModuleType(f"{package_name}.upscale")
upscale_pkg.__path__ = []
sys.modules[f"{package_name}.upscale"] = upscale_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
image_module = types.ModuleType(f"{package_name}.utils.image")
def tensor_to_pil(image_batch, index):
tensor = image_batch[index].detach().cpu().clamp(0, 1)
array = (tensor.numpy() * 255.0).astype("uint8")
return Image.fromarray(array)
def pil_to_tensor(image):
array = np.array(image).astype(np.float32) / 255.0
return torch.from_numpy(array).unsqueeze(0)
image_module.tensor_to_pil = tensor_to_pil
image_module.pil_to_tensor = pil_to_tensor
sys.modules[f"{package_name}.utils.image"] = image_module
usdu_utils_module = types.ModuleType(f"{package_name}.utils.usdu_utils")
def crop_cond(cond, *_args, **_kwargs):
return cond
def get_crop_region(mask, padding):
bbox = mask.getbbox()
if bbox is None:
return (0, 0, mask.size[0], mask.size[1])
x1, y1, x2, y2 = bbox
return (
max(0, x1 - int(padding)),
max(0, y1 - int(padding)),
min(mask.size[0], x2 + int(padding)),
min(mask.size[1], y2 + int(padding)),
)
def expand_crop(region, width, height, target_w, target_h):
x1, y1, x2, y2 = region
crop_w = x2 - x1
crop_h = y2 - y1
if crop_w >= target_w and crop_h >= target_h:
return (region, (crop_w, crop_h))
new_x2 = min(width, x1 + max(crop_w, target_w))
new_y2 = min(height, y1 + max(crop_h, target_h))
return ((x1, y1, new_x2, new_y2), (new_x2 - x1, new_y2 - y1))
usdu_utils_module.crop_cond = crop_cond
usdu_utils_module.get_crop_region = get_crop_region
usdu_utils_module.expand_crop = expand_crop
sys.modules[f"{package_name}.utils.usdu_utils"] = usdu_utils_module
crop_model_patch_module = types.ModuleType(f"{package_name}.utils.crop_model_patch")
crop_model_patch_module.crop_model_cond = lambda model, *_args, **_kwargs: nullcontext(model)
sys.modules[f"{package_name}.utils.crop_model_patch"] = crop_model_patch_module
conditioning_module = types.ModuleType(f"{package_name}.upscale.conditioning")
conditioning_module.clone_conditioning = lambda cond, clone_hints=False: copy.deepcopy(cond)
sys.modules[f"{package_name}.upscale.conditioning"] = conditioning_module
comfy_module = types.ModuleType("comfy")
comfy_samplers = types.ModuleType("comfy.samplers")
comfy_model_management = types.ModuleType("comfy.model_management")
comfy_module.samplers = comfy_samplers
comfy_module.model_management = comfy_model_management
sys.modules["comfy"] = comfy_module
sys.modules["comfy.samplers"] = comfy_samplers
sys.modules["comfy.model_management"] = comfy_model_management
def _load_module(package_name, module_rel_path, module_name):
module_path = Path(__file__).resolve().parents[3] / module_rel_path
spec = importlib.util.spec_from_file_location(
f"{package_name}.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
def _load_tile_ops_module():
package_name = "dist_tile_ops_testpkg"
_bootstrap_package(package_name)
return _load_module(package_name, "upscale/tile_ops.py", "upscale.tile_ops")
tile_ops_module = _load_tile_ops_module()
class _DummyTileOps(tile_ops_module.TileOpsMixin):
pass
class _FakeControl:
def __init__(self, cond_hint_original, previous_controlnet=None):
self.cond_hint_original = cond_hint_original
self.previous_controlnet = previous_controlnet
class TileOpsTests(unittest.TestCase):
def setUp(self):
self.node = _DummyTileOps()
def test_round_and_calculate_tiles(self):
self.assertEqual(self.node.round_to_multiple(13, 8), 16)
tiles = self.node.calculate_tiles(10, 9, 4, 5, force_uniform_tiles=True)
self.assertEqual(len(tiles), 6)
self.assertEqual(tiles[0], (0, 0))
self.assertEqual(tiles[-1], (8, 5))
def test_create_tile_mask_and_blend_tile(self):
base = Image.new("RGB", (8, 8), (0, 0, 0))
tile = Image.new("RGB", (4, 4), (255, 255, 255))
mask = self.node.create_tile_mask(8, 8, 2, 2, 4, 4, mask_blur=0)
self.assertEqual(mask.getpixel((0, 0)), 0)
self.assertEqual(mask.getpixel((3, 3)), 255)
blended = self.node.blend_tile(
base_image=base,
tile_image=tile,
x=2,
y=2,
extracted_size=(4, 4),
mask=mask,
padding=0,
)
self.assertEqual(blended.size, (8, 8))
self.assertEqual(blended.getpixel((3, 3)), (255, 255, 255))
def test_slice_conditioning_selects_batch_index(self):
control = _FakeControl(cond_hint_original=torch.ones((2, 1, 1)))
positive = [[torch.arange(6, dtype=torch.float32).reshape(2, 3), {"control": control, "mask": torch.ones((2, 1, 1))}]]
negative = [[torch.arange(6, dtype=torch.float32).reshape(2, 3), {"mask": torch.zeros((2, 1, 1))}]]
pos_sliced, neg_sliced = self.node._slice_conditioning(positive, negative, batch_idx=1)
self.assertEqual(tuple(pos_sliced[0][0].shape), (1, 3))
self.assertEqual(tuple(neg_sliced[0][0].shape), (1, 3))
self.assertEqual(tuple(pos_sliced[0][1]["mask"].shape), (1, 1, 1))
self.assertEqual(tuple(neg_sliced[0][1]["mask"].shape), (1, 1, 1))
self.assertEqual(tuple(pos_sliced[0][1]["control"].cond_hint_original.shape), (1, 1, 1))
if __name__ == "__main__":
unittest.main()
+115
View File
@@ -0,0 +1,115 @@
import unittest
from upscale.processing_args import UpscaleCoreArgs
from upscale.tile_processing import TileBatchArgs, extract_and_process_tile_batch
class _DummyTileProcessor:
def __init__(self):
self.extract_calls = []
self.process_calls = []
def extract_batch_tile_with_padding(
self,
upscaled_image,
tx,
ty,
tile_width,
tile_height,
padding,
force_uniform_tiles,
):
self.extract_calls.append(
{
"upscaled_image": upscaled_image,
"tx": tx,
"ty": ty,
"tile_width": tile_width,
"tile_height": tile_height,
"padding": padding,
"force_uniform_tiles": force_uniform_tiles,
}
)
return "tile-batch", 10, 20, 30, 40
def process_tiles_batch(
self,
tile_batch,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
region,
canvas_shape,
):
self.process_calls.append(
{
"tile_batch": tile_batch,
"model": model,
"positive": positive,
"negative": negative,
"vae": vae,
"seed": seed,
"steps": steps,
"cfg": cfg,
"sampler_name": sampler_name,
"scheduler": scheduler,
"denoise": denoise,
"tiled_decode": tiled_decode,
"region": region,
"canvas_shape": canvas_shape,
}
)
return "processed-batch"
class TileProcessingTests(unittest.TestCase):
def test_extract_and_process_tile_batch_forwards_args_and_computes_region(self):
node = _DummyTileProcessor()
tile_args = TileBatchArgs(
core=UpscaleCoreArgs(
model="model",
positive="positive",
negative="negative",
vae="vae",
seed=42,
steps=20,
cfg=7.5,
sampler_name="euler",
scheduler="normal",
denoise=0.4,
tiled_decode=False,
),
tile_width=256,
tile_height=192,
padding=32,
force_uniform_tiles=True,
width=1024,
height=768,
)
result = extract_and_process_tile_batch(
node=node,
upscaled_image="upscaled-image",
tx=3,
ty=4,
args=tile_args,
)
self.assertEqual(result, ("processed-batch", 10, 20, 30, 40))
self.assertEqual(len(node.extract_calls), 1)
self.assertEqual(len(node.process_calls), 1)
self.assertEqual(node.extract_calls[0]["tx"], 3)
self.assertEqual(node.extract_calls[0]["ty"], 4)
self.assertEqual(node.process_calls[0]["region"], (10, 20, 40, 60))
self.assertEqual(node.process_calls[0]["canvas_shape"], (1024, 768))
if __name__ == "__main__":
unittest.main()
+128
View File
@@ -0,0 +1,128 @@
import importlib.util
import sys
import types
import unittest
from pathlib import Path
from PIL import Image
def _ensure_torch_stubs():
if "torch" not in sys.modules:
torch_module = types.ModuleType("torch")
class _Tensor:
pass
torch_module.Tensor = _Tensor
torch_module.float32 = object()
torch_module.zeros = lambda *_args, **_kwargs: None
torch_module.from_numpy = lambda array: array
torch_module.cat = lambda *_args, **_kwargs: None
nn_module = types.ModuleType("torch.nn")
functional_module = types.ModuleType("torch.nn.functional")
functional_module.interpolate = lambda tensor, **_kwargs: tensor
nn_module.functional = functional_module
torch_module.nn = nn_module
sys.modules["torch"] = torch_module
sys.modules["torch.nn"] = nn_module
sys.modules["torch.nn.functional"] = functional_module
if "torchvision" not in sys.modules:
torchvision_module = types.ModuleType("torchvision")
transforms_module = types.ModuleType("torchvision.transforms")
class _GaussianBlur:
def __init__(self, *_args, **_kwargs):
pass
def __call__(self, value):
return value
transforms_module.GaussianBlur = _GaussianBlur
torchvision_module.transforms = transforms_module
sys.modules["torchvision"] = torchvision_module
sys.modules["torchvision.transforms"] = transforms_module
def _load_usdu_utils_module():
_ensure_torch_stubs()
module_path = Path(__file__).resolve().parents[3] / "utils" / "usdu_utils.py"
spec = importlib.util.spec_from_file_location("dist_test_usdu_utils", module_path)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
return module
usdu_utils = _load_usdu_utils_module()
class UsduUtilsGeometryTests(unittest.TestCase):
def test_fix_crop_region_decrements_upper_bounds_when_interior(self):
self.assertEqual(
usdu_utils.fix_crop_region((2, 3, 9, 9), (10, 10)),
(2, 3, 8, 8),
)
self.assertEqual(
usdu_utils.fix_crop_region((1, 1, 10, 10), (10, 10)),
(1, 1, 10, 10),
)
def test_get_crop_region_with_padding_and_empty_mask(self):
mask = Image.new("L", (10, 10), 0)
for x in range(2, 5):
for y in range(3, 7):
mask.putpixel((x, y), 255)
region = usdu_utils.get_crop_region(mask, pad=1)
self.assertEqual(region, (1, 2, 5, 7))
empty_region = usdu_utils.get_crop_region(Image.new("L", (8, 6), 0), pad=2)
self.assertEqual(empty_region, (6, 4, 1, 1))
def test_expand_crop_and_resize_region(self):
expanded_region, expanded_size = usdu_utils.expand_crop(
region=(2, 2, 6, 6),
width=12,
height=12,
target_width=8,
target_height=10,
)
self.assertEqual(expanded_region, (0, 0, 8, 10))
self.assertEqual(expanded_size, (8, 10))
resized = usdu_utils.resize_region(
region=(2, 2, 6, 10),
init_size=(8, 16),
resize_size=(16, 32),
)
self.assertEqual(resized, (4, 4, 12, 20))
def test_region_intersection(self):
self.assertEqual(
usdu_utils.region_intersection((0, 0, 5, 5), (3, 2, 8, 7)),
(3, 2, 5, 5),
)
self.assertIsNone(usdu_utils.region_intersection((0, 0, 1, 1), (2, 2, 4, 4)))
def test_pad_image2_preserves_expected_output_size(self):
image = Image.new("RGB", (4, 3), "black")
padded = usdu_utils.pad_image2(
image,
left_pad=2,
right_pad=1,
top_pad=3,
bottom_pad=2,
fill=False,
blur=False,
)
self.assertEqual(padded.size, (7, 8))
if __name__ == "__main__":
unittest.main()
+6 -2
View File
@@ -1,7 +1,8 @@
import copy
from typing import Any
def clone_control_chain(control, clone_hint=True):
def clone_control_chain(control: Any, clone_hint: bool = True) -> Any:
"""Shallow copy the ControlNet chain, optionally cloning hints but sharing models."""
if control is None:
return None
@@ -14,7 +15,10 @@ def clone_control_chain(control, clone_hint=True):
return new_control
def clone_conditioning(cond_list, clone_hints=True):
def clone_conditioning(
cond_list: list[tuple[Any, dict[str, Any]]] | list[list[Any]],
clone_hints: bool = True,
) -> list[list[Any]]:
"""Clone conditioning without duplicating ControlNet models."""
new_cond = []
for emb, cond_dict in cond_list:
+34 -2
View File
@@ -2,7 +2,7 @@ import asyncio
from ..utils.logging import debug_log
from .job_store import ensure_tile_jobs_initialized
from .job_timeout import _check_and_requeue_timed_out_workers as _requeue_usdu
from .job_timeout import check_and_requeue_timed_out_workers
from .job_models import ImageJobState, TileJobState
@@ -22,6 +22,10 @@ class JobStateMixin:
return dict(job_data.completed_images)
return {}
async def all_completed_tasks(self, multi_job_id):
"""Public completed-task accessor."""
return await self._get_all_completed_tasks(multi_job_id)
async def _get_next_image_index(self, multi_job_id):
"""Get next image index from pending queue for master."""
prompt_server = ensure_tile_jobs_initialized()
@@ -39,6 +43,10 @@ class JobStateMixin:
except asyncio.TimeoutError:
return None
async def next_image_index(self, multi_job_id):
"""Public next-image accessor."""
return await self._get_next_image_index(multi_job_id)
async def _get_next_tile_index(self, multi_job_id):
"""Get next tile index from pending queue for master in static mode."""
prompt_server = ensure_tile_jobs_initialized()
@@ -56,6 +64,10 @@ class JobStateMixin:
except asyncio.TimeoutError:
return None
async def next_tile_index(self, multi_job_id):
"""Public next-tile accessor."""
return await self._get_next_tile_index(multi_job_id)
async def _get_total_completed_count(self, multi_job_id):
"""Get total count of all completed images (master + workers)."""
prompt_server = ensure_tile_jobs_initialized()
@@ -67,6 +79,10 @@ class JobStateMixin:
return len(job_data.completed_tasks)
return 0
async def total_completed_count(self, multi_job_id):
"""Public completed-count accessor."""
return await self._get_total_completed_count(multi_job_id)
async def _get_all_completed_images(self, multi_job_id):
"""Get all completed images."""
prompt_server = ensure_tile_jobs_initialized()
@@ -76,6 +92,10 @@ class JobStateMixin:
return job_data.completed_images.copy()
return {}
async def all_completed_images(self, multi_job_id):
"""Public completed-image accessor."""
return await self._get_all_completed_images(multi_job_id)
async def _get_pending_count(self, multi_job_id):
"""Get count of pending images in the queue."""
prompt_server = ensure_tile_jobs_initialized()
@@ -87,6 +107,10 @@ class JobStateMixin:
return job_data.pending_tasks.qsize()
return 0
async def pending_count(self, multi_job_id):
"""Public pending-count accessor."""
return await self._get_pending_count(multi_job_id)
async def _drain_worker_results_queue(self, multi_job_id):
"""Drain pending worker results from queue and update completed images."""
prompt_server = ensure_tile_jobs_initialized()
@@ -130,6 +154,14 @@ class JobStateMixin:
return collected
async def drain_worker_results_queue(self, multi_job_id):
"""Public worker-result drain API."""
return await self._drain_worker_results_queue(multi_job_id)
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
"""Check for timed out workers and requeue their assigned images."""
return await _requeue_usdu(multi_job_id, batch_size)
return await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
async def check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
"""Public timeout-requeue API."""
return await self._check_and_requeue_timed_out_workers(multi_job_id, batch_size)
+79 -60
View File
@@ -1,45 +1,51 @@
import asyncio
import os
import time
from typing import List, Optional
import server
from dataclasses import dataclass
from typing import Any
from ..utils.logging import debug_log
from ..utils.runtime_state import ensure_distributed_runtime_state, get_prompt_server_instance
from .job_models import BaseJobState, ImageJobState, TileJobState
# Configure maximum payload size (50MB default, configurable via environment variable)
MAX_PAYLOAD_SIZE = int(os.environ.get('COMFYUI_MAX_PAYLOAD_SIZE', str(50 * 1024 * 1024)))
def ensure_tile_jobs_initialized():
@dataclass(frozen=True)
class JobQueueInitConfig:
mode: str
batch_size: int = 0
num_tiles_per_image: int = 0
all_indices: list[int] | None = None
enabled_workers: list[str] | None = None
batched_static: bool = False
def ensure_tile_jobs_initialized() -> Any:
"""Ensure tile job storage is initialized on the server instance."""
prompt_server = server.PromptServer.instance
if not hasattr(prompt_server, 'distributed_pending_tile_jobs'):
debug_log("Initializing persistent tile job queue on server instance.")
prompt_server.distributed_pending_tile_jobs = {}
prompt_server.distributed_tile_jobs_lock = asyncio.Lock()
else:
invalid_job_ids = [
job_id
for job_id, job_data in prompt_server.distributed_pending_tile_jobs.items()
if not isinstance(job_data, BaseJobState)
]
for job_id in invalid_job_ids:
debug_log(f"Removing invalid job state for {job_id}")
del prompt_server.distributed_pending_tile_jobs[job_id]
prompt_server = get_prompt_server_instance()
state = ensure_distributed_runtime_state(prompt_server)
if not isinstance(state.distributed_pending_tile_jobs, dict):
debug_log("Resetting invalid distributed_pending_tile_jobs state to empty dict.")
state.distributed_pending_tile_jobs = {}
ensure_distributed_runtime_state(prompt_server)
invalid_job_ids = [
job_id
for job_id, job_data in state.distributed_pending_tile_jobs.items()
if not isinstance(job_data, BaseJobState)
]
for job_id in invalid_job_ids:
debug_log(f"Removing invalid job state for {job_id}")
del state.distributed_pending_tile_jobs[job_id]
return prompt_server
async def _init_job_queue(
multi_job_id,
mode,
batch_size=None,
num_tiles_per_image=None,
all_indices=None,
enabled_workers=None,
batched_static: bool = False,
):
multi_job_id: str,
config: JobQueueInitConfig,
) -> None:
"""Unified initialization for job queues in static and dynamic modes."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
@@ -47,33 +53,34 @@ async def _init_job_queue(
debug_log(f"Queue already exists for {multi_job_id}")
return
if mode == 'dynamic':
if config.mode == 'dynamic':
job_data = ImageJobState(multi_job_id=multi_job_id)
elif mode == 'static':
elif config.mode == 'static':
job_data = TileJobState(multi_job_id=multi_job_id)
else:
raise ValueError(f"Unknown mode: {mode}")
raise ValueError(f"Unknown mode: {config.mode}")
job_data.worker_status = {w: time.time() for w in enabled_workers or []}
job_data.assigned_to_workers = {w: [] for w in enabled_workers or []}
job_data.worker_status = {w: time.time() for w in config.enabled_workers or []}
job_data.assigned_to_workers = {w: [] for w in config.enabled_workers or []}
if mode == 'dynamic':
job_data.batch_size = int(batch_size or 0)
if config.mode == 'dynamic':
job_data.batch_size = int(config.batch_size or 0)
pending_queue = job_data.pending_images
for i in (all_indices or range(int(batch_size or 0))):
indices = config.all_indices or list(range(int(config.batch_size or 0)))
for i in indices:
await pending_queue.put(i)
debug_log(f"Initialized image queue with {batch_size} pending items")
elif mode == 'static':
job_data.num_tiles_per_image = int(num_tiles_per_image or 0)
job_data.batch_size = int(batch_size or 0)
job_data.batched_static = bool(batched_static)
debug_log(f"Initialized image queue with {config.batch_size} pending items")
elif config.mode == 'static':
job_data.num_tiles_per_image = int(config.num_tiles_per_image or 0)
job_data.batch_size = int(config.batch_size or 0)
job_data.batched_static = bool(config.batched_static)
# For batched static distribution, populate only tile ids [0..num_tiles_per_image-1]
pending_queue = job_data.pending_tasks
if batched_static and num_tiles_per_image is not None:
for i in range(num_tiles_per_image):
if config.batched_static and config.num_tiles_per_image > 0:
for i in range(config.num_tiles_per_image):
await pending_queue.put(i)
else:
total_tiles = int(batch_size or 0) * int(num_tiles_per_image or 0)
total_tiles = int(config.batch_size or 0) * int(config.num_tiles_per_image or 0)
for i in range(total_tiles):
await pending_queue.put(i)
@@ -83,16 +90,18 @@ async def _init_job_queue(
async def init_dynamic_job(
multi_job_id: str,
batch_size: int,
enabled_workers: List[str],
all_indices: Optional[List[int]] = None,
):
enabled_workers: list[str],
all_indices: list[int] | None = None,
) -> None:
"""Initialize queue for dynamic mode (per-image), with collector fields."""
await _init_job_queue(
multi_job_id,
'dynamic',
batch_size=batch_size,
all_indices=all_indices or list(range(batch_size)),
enabled_workers=enabled_workers,
JobQueueInitConfig(
mode='dynamic',
batch_size=batch_size,
all_indices=all_indices or list(range(batch_size)),
enabled_workers=enabled_workers,
),
)
debug_log(f"Job {multi_job_id} initialized with {batch_size} images")
@@ -101,20 +110,22 @@ async def init_static_job_batched(
multi_job_id: str,
batch_size: int,
num_tiles_per_image: int,
enabled_workers: List[str],
):
enabled_workers: list[str],
) -> None:
"""Initialize queue for static mode (batched-per-tile)."""
await _init_job_queue(
multi_job_id,
'static',
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
enabled_workers=enabled_workers,
batched_static=True,
JobQueueInitConfig(
mode='static',
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
enabled_workers=enabled_workers,
batched_static=True,
),
)
async def _drain_results_queue(multi_job_id):
async def drain_results_queue(multi_job_id: str) -> int:
"""Drain pending results from queue and update completed_tasks. Returns count drained."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
@@ -152,7 +163,7 @@ async def _drain_results_queue(multi_job_id):
return collected
async def _get_completed_count(multi_job_id):
async def get_completed_count(multi_job_id: str) -> int:
"""Get count of completed tasks."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
@@ -162,7 +173,7 @@ async def _get_completed_count(multi_job_id):
return 0
async def _mark_task_completed(multi_job_id, task_id, result):
async def mark_task_completed(multi_job_id: str, task_id: int, result: Any) -> None:
"""Mark a task as completed."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
@@ -171,10 +182,18 @@ async def _mark_task_completed(multi_job_id, task_id, result):
job_data.completed_tasks[task_id] = result
async def _cleanup_job(multi_job_id):
async def cleanup_job(multi_job_id: str) -> None:
"""Cleanup the job data."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
del prompt_server.distributed_pending_tile_jobs[multi_job_id]
debug_log(f"Cleaned up job {multi_job_id}")
async def init_job_queue(
multi_job_id: str,
config: JobQueueInitConfig,
) -> None:
"""Public entry point for unified job queue initialization."""
await _init_job_queue(multi_job_id, config)
+9
View File
@@ -99,6 +99,10 @@ async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
f"Probe diagnostics: online={probe_queue is not None} queue_remaining={probe_queue}"
)
if incomplete_assigned == 0:
debug_log(f"Worker {worker} heartbeat stale but all tasks complete; skipping")
continue
if busy:
workers_graced.append(worker)
debug_log(f"Heartbeat grace: worker {worker} busy via probe; skipping requeue")
@@ -148,3 +152,8 @@ async def _check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
job_data.assigned_to_workers[worker] = []
return requeued_count
async def check_and_requeue_timed_out_workers(multi_job_id, total_tasks):
"""Public wrapper for worker-timeout requeue checks."""
return await _check_and_requeue_timed_out_workers(multi_job_id, total_tasks)
+238
View File
@@ -0,0 +1,238 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
class TileOpsCollaborator:
"""Typed facade over tile-related operations used by USDU modes."""
def __init__(self, delegate: Any) -> None:
self._delegate = delegate
def round_to_multiple(self, value: int) -> int:
return self._delegate.round_to_multiple(value)
def calculate_tiles(
self,
width: int,
height: int,
tile_width: int,
tile_height: int,
force_uniform_tiles: bool,
) -> list[tuple[int, int]]:
return self._delegate.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
def create_tile_mask(
self,
width: int,
height: int,
tx: int,
ty: int,
tile_width: int,
tile_height: int,
mask_blur: int,
) -> Any:
return self._delegate.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur)
def blend_tile(
self,
base_image: Any,
tile_pil: Any,
x1: int,
y1: int,
size: tuple[int, int],
tile_mask: Any,
padding: int,
) -> Any:
return self._delegate.blend_tile(base_image, tile_pil, x1, y1, size, tile_mask, padding)
def slice_conditioning(self, positive: Any, negative: Any, image_idx: int) -> tuple[Any, Any]:
return self._delegate.slice_conditioning(positive, negative, image_idx)
def process_and_blend_tile(
self,
tile_idx: int,
tile_pos: tuple[int, int],
upscaled_image: Any,
result_image: Any,
model: Any,
positive: Any,
negative: Any,
vae: Any,
seed: int,
steps: int,
cfg: float,
sampler_name: str,
scheduler: str,
denoise: float,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
image_width: int,
image_height: int,
force_uniform_tiles: bool,
tiled_decode: bool,
batch_idx: int,
) -> Any:
return self._delegate.process_and_blend_tile(
tile_idx,
tile_pos,
upscaled_image,
result_image,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tile_width,
tile_height,
padding,
mask_blur,
image_width,
image_height,
force_uniform_tiles,
tiled_decode,
batch_idx=batch_idx,
)
class JobStateCollaborator:
"""Typed facade over job-state operations used by USDU modes."""
def __init__(self, delegate: Any) -> None:
self._delegate = delegate
async def next_tile_index(self, multi_job_id: str) -> int | None:
return await self._delegate.next_tile_index(multi_job_id)
async def all_completed_tasks(self, multi_job_id: str) -> dict[int, dict[str, Any]]:
return await self._delegate.all_completed_tasks(multi_job_id)
async def check_and_requeue_timed_out_workers(self, multi_job_id: str, expected_total: int) -> int:
return await self._delegate.check_and_requeue_timed_out_workers(multi_job_id, expected_total)
async def next_image_index(self, multi_job_id: str) -> int | None:
return await self._delegate.next_image_index(multi_job_id)
async def drain_worker_results_queue(self, multi_job_id: str) -> int:
return await self._delegate.drain_worker_results_queue(multi_job_id)
async def total_completed_count(self, multi_job_id: str) -> int:
return await self._delegate.total_completed_count(multi_job_id)
async def pending_count(self, multi_job_id: str) -> int:
return await self._delegate.pending_count(multi_job_id)
async def all_completed_images(self, multi_job_id: str) -> dict[int, Any]:
return await self._delegate.all_completed_images(multi_job_id)
class WorkerCommsCollaborator:
"""Typed facade over worker/master communication operations used by USDU modes."""
def __init__(self, delegate: Any) -> None:
self._delegate = delegate
async def request_assignment(self, multi_job_id: str, master_url: str, worker_id: str) -> Any:
return await self._delegate.request_assignment(multi_job_id, master_url, worker_id)
async def send_tiles_batch(
self,
processed_tiles: list[dict[str, Any]],
multi_job_id: str,
master_url: str,
padding: int,
worker_id: str,
is_final_flush: bool = False,
) -> None:
await self._delegate.send_tiles_batch(
processed_tiles,
multi_job_id,
master_url,
padding,
worker_id,
is_final_flush=is_final_flush,
)
async def send_heartbeat(self, multi_job_id: str, master_url: str, worker_id: str) -> None:
await self._delegate.send_heartbeat(multi_job_id, master_url, worker_id)
async def check_job_status(self, multi_job_id: str, master_url: str) -> bool:
return await self._delegate.check_job_status(multi_job_id, master_url)
async def async_yield(self) -> None:
await self._delegate.async_yield()
async def send_full_image(
self,
image_pil: Any,
image_idx: int,
multi_job_id: str,
master_url: str,
worker_id: str,
is_last: bool,
) -> None:
await self._delegate.send_full_image(
image_pil,
image_idx,
multi_job_id,
master_url,
worker_id,
is_last,
)
async def send_worker_complete_signal(self, multi_job_id: str, master_url: str, worker_id: str) -> None:
await self._delegate.send_worker_complete_signal(multi_job_id, master_url, worker_id)
class ResultCollectorCollaborator:
"""Typed facade over result-collection operations used by dynamic mode."""
def __init__(self, delegate: Any) -> None:
self._delegate = delegate
async def mark_image_completed(self, multi_job_id: str, image_idx: int, image_pil: Any) -> None:
await self._delegate.mark_image_completed(multi_job_id, image_idx, image_pil)
async def collect_dynamic_images(
self,
multi_job_id: str,
remaining_to_collect: int,
num_workers: int,
batch_size: int,
master_processed_count: int,
) -> dict[int, Any]:
return await self._delegate.collect_dynamic_images(
multi_job_id,
remaining_to_collect,
num_workers,
batch_size,
master_processed_count,
)
@dataclass(frozen=True)
class SingleGpuModeContext:
tile_ops: TileOpsCollaborator
@dataclass(frozen=True)
class StaticModeContext:
tile_ops: TileOpsCollaborator
job_state: JobStateCollaborator
worker_comms: WorkerCommsCollaborator
@dataclass(frozen=True)
class DynamicModeContext:
tile_ops: TileOpsCollaborator
job_state: JobStateCollaborator
worker_comms: WorkerCommsCollaborator
result_collector: ResultCollectorCollaborator
+309 -150
View File
@@ -1,12 +1,27 @@
import asyncio, torch
from __future__ import annotations
import asyncio, time, torch
from typing import Any
from PIL import Image
import comfy.model_management
from ...utils.logging import debug_log, log
from ...utils.image import tensor_to_pil, pil_to_tensor
from ...utils.async_helpers import run_async_in_server_loop
from ...utils.config import get_worker_timeout_seconds
from ...utils.constants import TILE_WAIT_TIMEOUT, TILE_SEND_TIMEOUT
from ...utils.constants import (
JOB_POLL_INTERVAL,
JOB_POLL_MAX_ATTEMPTS,
TILE_WAIT_TIMEOUT,
TILE_SEND_TIMEOUT,
)
from ..job_store import ensure_tile_jobs_initialized, init_dynamic_job
from ..mode_contexts import (
DynamicModeContext,
JobStateCollaborator,
ResultCollectorCollaborator,
TileOpsCollaborator,
WorkerCommsCollaborator,
)
from ..processing_args import UpscaleCoreArgs
class DynamicModeMixin:
@@ -14,16 +29,205 @@ class DynamicModeMixin:
Dynamic (per-image queue) USDU mode behaviors for master and worker roles.
Expected co-mixins on `self`:
- TileOpsMixin (`calculate_tiles`, `_slice_conditioning`, `_process_and_blend_tile`).
- TileOpsMixin (`calculate_tiles`, `slice_conditioning`, `process_and_blend_tile`).
- JobStateMixin (image queue/task completion helpers).
- WorkerCommsMixin (`_request_image_from_master`, `_send_full_image_to_master`, `_send_heartbeat_to_master`).
- WorkerCommsMixin (`request_assignment`, `send_full_image`, `send_heartbeat`).
"""
def process_master_dynamic(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers):
def _build_dynamic_mode_context(self) -> DynamicModeContext:
"""Build explicit collaborators for dynamic mode execution."""
return DynamicModeContext(
tile_ops=TileOpsCollaborator(self),
job_state=JobStateCollaborator(self),
worker_comms=WorkerCommsCollaborator(self),
result_collector=ResultCollectorCollaborator(self),
)
def _poll_job_ready_dynamic(
self,
multi_job_id: str,
master_url: str,
worker_id: str | None = None,
max_attempts: int = JOB_POLL_MAX_ATTEMPTS,
mode_context: DynamicModeContext | None = None,
) -> bool:
"""Poll master for job readiness using explicit worker-comms collaborator."""
context = mode_context or self._build_dynamic_mode_context()
for attempt in range(max_attempts):
ready = run_async_in_server_loop(
context.worker_comms.check_job_status(multi_job_id, master_url),
timeout=5.0,
)
if ready:
if worker_id:
debug_log(f"Worker[{worker_id[:8]}] job {multi_job_id} ready after {attempt} attempts")
else:
debug_log(f"Job {multi_job_id} ready after {attempt} attempts")
return True
time.sleep(JOB_POLL_INTERVAL)
return False
def _process_image_tiles(
self,
image_idx: int,
local_image: Image.Image,
single_tensor: torch.Tensor,
core_args: UpscaleCoreArgs,
all_tiles: list,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
width: int,
height: int,
force_uniform_tiles: bool,
context: DynamicModeContext,
per_tile_callback: Any | None = None,
) -> Image.Image:
"""Process all tiles for a single image. Returns the processed PIL image."""
positive_sliced, negative_sliced = context.tile_ops.slice_conditioning(
core_args.positive, core_args.negative, image_idx,
)
for tile_idx, pos in enumerate(all_tiles):
source_tensor = pil_to_tensor(local_image)
if single_tensor.is_cuda:
source_tensor = source_tensor.cuda()
local_image = context.tile_ops.process_and_blend_tile(
tile_idx, pos, source_tensor, local_image,
core_args.model,
positive_sliced,
negative_sliced,
core_args.vae,
core_args.seed,
core_args.steps,
core_args.cfg,
core_args.sampler_name,
core_args.scheduler,
core_args.denoise,
tile_width,
tile_height,
padding, mask_blur, width, height, force_uniform_tiles,
core_args.tiled_decode,
batch_idx=image_idx,
)
if per_tile_callback:
per_tile_callback()
return local_image
def _handle_master_dynamic_idle_state(
self,
multi_job_id,
batch_size,
consecutive_retries,
max_consecutive_retries,
mode_context: DynamicModeContext | None = None,
):
"""Handle queue-empty state while master waits for worker progress."""
context = mode_context or self._build_dynamic_mode_context()
drained_count = run_async_in_server_loop(
context.job_state.drain_worker_results_queue(multi_job_id),
timeout=5.0,
)
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
requeued_count = run_async_in_server_loop(
context.job_state.check_and_requeue_timed_out_workers(multi_job_id, batch_size),
timeout=5.0,
)
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
if requeued_count > 0:
log(f"Requeued {requeued_count} images from timed out workers")
return False, True, 0
completed_now = run_async_in_server_loop(
context.job_state.total_completed_count(multi_job_id),
timeout=1.0,
)
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
if completed_now >= batch_size:
return True, False, consecutive_retries
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
pending_count = run_async_in_server_loop(
context.job_state.pending_count(multi_job_id),
timeout=1.0,
)
if pending_count > 0:
return False, True, 0
consecutive_retries += 1
if consecutive_retries >= max_consecutive_retries:
log(f"Max retries ({max_consecutive_retries}) reached. Forcing collection of remaining results.")
return True, False, consecutive_retries
if drained_count > 0:
debug_log(f"Master idle drain picked up {drained_count} images while waiting for workers")
debug_log("Waiting for workers")
run_async_in_server_loop(asyncio.sleep(2), timeout=3.0)
return False, True, consecutive_retries
def _finalize_master_dynamic_results(
self,
multi_job_id,
batch_size,
num_workers,
processed_count,
result_images,
upscaled_image,
mode_context: DynamicModeContext | None = None,
):
"""Collect remaining worker outputs and convert final images back to tensor."""
context = mode_context or self._build_dynamic_mode_context()
all_completed = run_async_in_server_loop(
context.job_state.all_completed_images(multi_job_id),
timeout=5.0,
)
remaining_to_collect = batch_size - len(all_completed)
if remaining_to_collect > 0:
debug_log(f"Waiting for {remaining_to_collect} more images from workers")
collection_timeout = float(get_worker_timeout_seconds())
collected_images = run_async_in_server_loop(
context.result_collector.collect_dynamic_images(
multi_job_id,
remaining_to_collect,
num_workers,
batch_size,
processed_count,
),
timeout=collection_timeout,
)
all_completed.update(collected_images)
for idx, processed_img in all_completed.items():
if idx < batch_size:
result_images[idx] = processed_img
result_tensor = (
torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
if batch_size > 1
else pil_to_tensor(result_images[0])
)
if upscaled_image.is_cuda:
result_tensor = result_tensor.cuda()
return result_tensor
def process_master_dynamic(
self,
upscaled_image: torch.Tensor,
core_args: UpscaleCoreArgs,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
multi_job_id: str,
enabled_workers: list[str],
mode_context: DynamicModeContext | None = None,
) -> tuple[torch.Tensor]:
"""Dynamic mode for large batches - assigns whole images to workers dynamically, including master."""
context = mode_context or self._build_dynamic_mode_context()
# Get batch size and dimensions
batch_size, height, width, _ = upscaled_image.shape
num_workers = len(enabled_workers)
@@ -36,7 +240,7 @@ class DynamicModeMixin:
debug_log(f"Processing {batch_size} images dynamically across master + {num_workers} workers.")
# Calculate tiles for processing
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
# Initialize job queue for communication
try:
@@ -61,7 +265,7 @@ class DynamicModeMixin:
while processed_count < batch_size:
# Try to get an image to process
image_idx = run_async_in_server_loop(
self._get_next_image_index(multi_job_id),
context.job_state.next_image_index(multi_job_id),
timeout=5.0 # Short timeout to allow frequent checks
)
@@ -74,38 +278,25 @@ class DynamicModeMixin:
# Process locally
single_tensor = upscaled_image[image_idx:image_idx+1]
local_image = result_images[image_idx]
image_seed = seed
# Pre-slice conditioning once per image (not per tile)
positive_sliced, negative_sliced = self._slice_conditioning(positive, negative, image_idx)
for tile_idx, pos in enumerate(all_tiles):
source_tensor = pil_to_tensor(local_image)
if single_tensor.is_cuda:
source_tensor = source_tensor.cuda()
local_image = self._process_and_blend_tile(
tile_idx, pos, source_tensor, local_image,
model, positive_sliced, negative_sliced, vae, image_seed, steps, cfg,
sampler_name, scheduler, denoise, tile_width, tile_height,
padding, mask_blur, width, height, force_uniform_tiles,
tiled_decode, batch_idx=image_idx
)
# Yield after each tile to minimize worker downtime
run_async_in_server_loop(self._async_yield(), timeout=0.1)
# Note: No per-tile drain here – that's what makes this "per-image"
local_image = self._process_image_tiles(
image_idx, local_image, single_tensor, core_args, all_tiles,
tile_width, tile_height, padding, mask_blur, width, height,
force_uniform_tiles, context,
per_tile_callback=lambda: run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1),
)
result_images[image_idx] = local_image
# Mark as completed
run_async_in_server_loop(
self._mark_image_completed(multi_job_id, image_idx, local_image),
context.result_collector.mark_image_completed(multi_job_id, image_idx, local_image),
timeout=5.0
)
# NEW: Drain after the full image is marked complete (catches workers who finished during master's processing)
drained_count = run_async_in_server_loop(
self._drain_worker_results_queue(multi_job_id),
context.job_state.drain_worker_results_queue(multi_job_id),
timeout=5.0
)
@@ -114,133 +305,106 @@ class DynamicModeMixin:
# NEW: Log overall progress (includes master's image + any drained workers)
completed_now = run_async_in_server_loop(
self._get_total_completed_count(multi_job_id),
context.job_state.total_completed_count(multi_job_id),
timeout=1.0
)
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
# Yield to allow workers to get new images after completing one
run_async_in_server_loop(self._async_yield(), timeout=0.1)
run_async_in_server_loop(context.worker_comms.async_yield(), timeout=0.1)
else:
# Queue empty: collect any queued worker results to update progress
drained_count = run_async_in_server_loop(
self._drain_worker_results_queue(multi_job_id),
timeout=5.0
should_break, should_continue, consecutive_retries = self._handle_master_dynamic_idle_state(
multi_job_id=multi_job_id,
batch_size=batch_size,
consecutive_retries=consecutive_retries,
max_consecutive_retries=max_consecutive_retries,
mode_context=context,
)
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield after drain
# Check for timed out workers and requeue their images
requeued_count = run_async_in_server_loop(
self._check_and_requeue_timed_out_workers(multi_job_id, batch_size),
timeout=5.0
)
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield after requeue
if requeued_count > 0:
log(f"Requeued {requeued_count} images from timed out workers")
consecutive_retries = 0 # Reset since we have work to do
continue
# Now check total completed (includes newly collected)
completed_now = run_async_in_server_loop(
self._get_total_completed_count(multi_job_id),
timeout=1.0
)
log(f"USDU Dist: Images progress {completed_now}/{batch_size}")
if completed_now >= batch_size:
if should_break:
break
run_async_in_server_loop(self._async_yield(), timeout=0.1) # Yield before pending check
# Check if there are pending images in the queue (could be requeued)
pending_count = run_async_in_server_loop(
self._get_pending_count(multi_job_id),
timeout=1.0
)
if pending_count > 0:
consecutive_retries = 0 # Reset retries since there's work to do
if should_continue:
continue
consecutive_retries += 1
if consecutive_retries >= max_consecutive_retries:
log(f"Max retries ({max_consecutive_retries}) reached. Forcing collection of remaining results.")
break # Force exit to collection phase
debug_log("Waiting for workers")
# Use async sleep to allow event loop to process worker requests
run_async_in_server_loop(asyncio.sleep(2), timeout=3.0)
debug_log(f"Master processed {processed_count} images locally")
# Get all completed images to check what needs to be collected
all_completed = run_async_in_server_loop(
self._get_all_completed_images(multi_job_id),
timeout=5.0
result_tensor = self._finalize_master_dynamic_results(
multi_job_id=multi_job_id,
batch_size=batch_size,
num_workers=num_workers,
processed_count=processed_count,
result_images=result_images,
upscaled_image=upscaled_image,
mode_context=context,
)
# Calculate how many we still need to collect
remaining_to_collect = batch_size - len(all_completed)
if remaining_to_collect > 0:
debug_log(f"Waiting for {remaining_to_collect} more images from workers")
# Use the unified worker timeout for the collection phase
collection_timeout = float(get_worker_timeout_seconds())
collected_images = run_async_in_server_loop(
self._async_collect_dynamic_images(multi_job_id, remaining_to_collect, num_workers, batch_size, processed_count),
timeout=collection_timeout
)
# Merge collected with already completed
all_completed.update(collected_images)
# Update result images with all completed images
for idx, processed_img in all_completed.items():
if idx < batch_size:
result_images[idx] = processed_img
# Convert back to tensor
result_tensor = torch.cat([pil_to_tensor(img) for img in result_images], dim=0) if batch_size > 1 else pil_to_tensor(result_images[0])
if upscaled_image.is_cuda:
result_tensor = result_tensor.cuda()
debug_log(f"UltimateSDUpscale Master - Job {multi_job_id} complete")
log(f"Completed processing all {batch_size} images")
return (result_tensor,)
def process_worker_dynamic(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
worker_id, enabled_worker_ids, dynamic_threshold):
def process_worker_dynamic(
self,
upscaled_image: torch.Tensor,
core_args: UpscaleCoreArgs,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
multi_job_id: str,
master_url: str,
worker_id: str,
enabled_worker_ids: list[str] | str,
dynamic_threshold: int,
mode_context: DynamicModeContext | None = None,
) -> tuple[torch.Tensor]:
"""Worker processing in dynamic mode - processes whole images."""
context = mode_context or self._build_dynamic_mode_context()
# Round tile dimensions
tile_width = self.round_to_multiple(tile_width)
tile_height = self.round_to_multiple(tile_height)
tile_width = context.tile_ops.round_to_multiple(tile_width)
tile_height = context.tile_ops.round_to_multiple(tile_height)
# Get dimensions and tile grid
batch_size, height, width, _ = upscaled_image.shape
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
log(f"USDU Dist Worker[{worker_id[:8]}]: Processing image queue | Batch {batch_size}")
# Keep track of processed images for is_last detection
processed_count = 0
# Ensure the dynamic job queue exists on the master. In delegate-only
# mode the master's USDU node is replaced with a lightweight collector
# that cannot initialise the queue itself, so workers do it instead.
# The call is idempotent — if the master already created the queue this
# is a no-op.
from ...utils.worker_ids import coerce_enabled_worker_ids as _coerce
_enabled = _coerce(enabled_worker_ids)
run_async_in_server_loop(
context.worker_comms.init_dynamic_job_on_master(
multi_job_id, master_url, batch_size, _enabled,
),
timeout=10.0,
)
# Poll for job readiness to avoid races during master init
max_poll_attempts = 20 # ~20s at 1s sleep
if not self._poll_job_ready(multi_job_id, master_url, worker_id=worker_id, max_attempts=max_poll_attempts):
max_poll_attempts = JOB_POLL_MAX_ATTEMPTS
if not self._poll_job_ready_dynamic(
multi_job_id,
master_url,
worker_id=worker_id,
max_attempts=max_poll_attempts,
mode_context=context,
):
log(f"Job {multi_job_id} not ready after {max_poll_attempts} attempts, aborting")
return (upscaled_image,)
# Loop to request and process images
while True:
# Request an image to process
image_idx, estimated_remaining = run_async_in_server_loop(
self._request_image_from_master(multi_job_id, master_url, worker_id),
timeout=TILE_WAIT_TIMEOUT
assignment = run_async_in_server_loop(
context.worker_comms.request_assignment(multi_job_id, master_url, worker_id),
timeout=TILE_WAIT_TIMEOUT,
)
image_idx = assignment.task_idx if assignment.kind == "image" else None
estimated_remaining = assignment.estimated_remaining
if image_idx is None:
debug_log(f"USDU Dist Worker - No more images to process")
@@ -259,39 +423,34 @@ class DynamicModeMixin:
local_image = tensor_to_pil(single_tensor, 0).copy()
# Process all tiles for this image
image_seed = seed
# Pre-slice conditioning once per image (not per tile)
positive_sliced, negative_sliced = self._slice_conditioning(positive, negative, image_idx)
for tile_idx, pos in enumerate(all_tiles):
source_tensor = pil_to_tensor(local_image)
if single_tensor.is_cuda:
source_tensor = source_tensor.cuda()
local_image = self._process_and_blend_tile(
tile_idx, pos, source_tensor, local_image,
model, positive_sliced, negative_sliced, vae, image_seed, steps, cfg,
sampler_name, scheduler, denoise, tile_width, tile_height,
padding, mask_blur, width, height, force_uniform_tiles,
tiled_decode, batch_idx=image_idx
)
run_async_in_server_loop(
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
timeout=5.0
)
local_image = self._process_image_tiles(
image_idx, local_image, single_tensor, core_args, all_tiles,
tile_width, tile_height, padding, mask_blur, width, height,
force_uniform_tiles, context,
per_tile_callback=lambda: run_async_in_server_loop(
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
timeout=5.0,
),
)
# Send processed image back to master
try:
# Use the estimated remaining to determine if this is the last image
is_last = is_last_for_worker
run_async_in_server_loop(
self._send_full_image_to_master(local_image, image_idx, multi_job_id,
master_url, worker_id, is_last),
context.worker_comms.send_full_image(
local_image,
image_idx,
multi_job_id,
master_url,
worker_id,
is_last,
),
timeout=TILE_SEND_TIMEOUT
)
# Send heartbeat after processing
run_async_in_server_loop(
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
timeout=5.0
)
if is_last:
@@ -304,7 +463,7 @@ class DynamicModeMixin:
debug_log(f"Worker[{worker_id[:8]}] processed {processed_count} images, sending completion signal")
try:
run_async_in_server_loop(
self._send_worker_complete_signal(multi_job_id, master_url, worker_id),
context.worker_comms.send_worker_complete_signal(multi_job_id, master_url, worker_id),
timeout=TILE_SEND_TIMEOUT
)
except Exception as e:
+56 -24
View File
@@ -1,29 +1,57 @@
from __future__ import annotations
import math, torch
from PIL import Image
from ...utils.logging import debug_log, log
from ...utils.image import tensor_to_pil, pil_to_tensor
from typing import Any
from ...utils.logging import log
from ...utils.image import blend_processed_batch_item, pil_to_tensor, tensor_to_pil
from ..mode_contexts import SingleGpuModeContext, TileOpsCollaborator
from ..processing_args import UpscaleCoreArgs
from ..tile_processing import TileBatchArgs, extract_and_process_tile_batch
class SingleGpuModeMixin:
def process_single_gpu(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur, force_uniform_tiles, tiled_decode):
def _build_single_gpu_mode_context(self) -> SingleGpuModeContext:
"""Build explicit collaborators for single-GPU mode execution."""
return SingleGpuModeContext(tile_ops=TileOpsCollaborator(self))
def process_single_gpu(
self,
upscaled_image: torch.Tensor,
core_args: UpscaleCoreArgs,
tile_width: int,
tile_height: int,
padding: int,
mask_blur: int,
force_uniform_tiles: bool,
mode_context: SingleGpuModeContext | None = None,
) -> tuple[torch.Tensor]:
"""Process all tiles on a single GPU (no distribution), batching per tile like USDU."""
context = mode_context or self._build_single_gpu_mode_context()
ops = context.tile_ops
# Round tile dimensions
tile_width = self.round_to_multiple(tile_width)
tile_height = self.round_to_multiple(tile_height)
tile_width = ops.round_to_multiple(tile_width)
tile_height = ops.round_to_multiple(tile_height)
# Get image dimensions and batch size
batch_size, height, width, _ = upscaled_image.shape
# Calculate all tiles
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
all_tiles = ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
rows = math.ceil(height / tile_height)
cols = math.ceil(width / tile_width)
log(
f"USDU Dist: Single GPU | Canvas {width}x{height} | Tile {tile_width}x{tile_height} | Grid {rows}x{cols} ({len(all_tiles)} tiles/image) | Batch {batch_size}"
)
tile_batch_args = TileBatchArgs(
core=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
force_uniform_tiles=force_uniform_tiles,
width=width,
height=height,
)
# Prepare result images list
result_images = []
@@ -34,7 +62,7 @@ class SingleGpuModeMixin:
# Precompute tile masks once
tile_masks = []
for tx, ty in all_tiles:
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
tile_masks.append(ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
# Process tiles batched across images
for tile_idx, (tx, ty) in enumerate(all_tiles):
@@ -43,25 +71,29 @@ class SingleGpuModeMixin:
if upscaled_image.is_cuda:
source_batch = source_batch.cuda()
# Extract batched tile
tile_batch, x1, y1, ew, eh = self.extract_batch_tile_with_padding(
source_batch, tx, ty, tile_width, tile_height, padding, force_uniform_tiles
processed_batch, x1, y1, ew, eh = extract_and_process_tile_batch(
node=ops,
upscaled_image=source_batch,
tx=tx,
ty=ty,
args=tile_batch_args,
)
# Process batch
region = (x1, y1, x1 + ew, y1 + eh)
processed_batch = self.process_tiles_batch(tile_batch, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tiled_decode, region, (width, height))
# Blend results back into each image using cached mask
tile_mask = tile_masks[tile_idx]
for b in range(batch_size):
tile_pil = tensor_to_pil(processed_batch, b)
# Resize back to extracted size
if tile_pil.size != (ew, eh):
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
result_images[b] = self.blend_tile(result_images[b], tile_pil, x1, y1, (ew, eh), tile_mask, padding)
blend_processed_batch_item(
result_images,
processed_batch,
b,
ops.blend_tile,
x1,
y1,
ew,
eh,
tile_mask,
padding,
)
# Convert back to tensor
result_tensors = [pil_to_tensor(img) for img in result_images]
+357 -242
View File
@@ -2,7 +2,7 @@ import asyncio, time, torch
from PIL import Image
import comfy.model_management
from ...utils.logging import debug_log, log
from ...utils.image import tensor_to_pil, pil_to_tensor
from ...utils.image import blend_processed_batch_item, pil_to_tensor, tensor_to_pil
from ...utils.async_helpers import run_async_in_server_loop
from ...utils.config import get_worker_timeout_seconds
from ...utils.constants import (
@@ -15,9 +15,16 @@ from ...utils.constants import (
)
from ..job_store import (
ensure_tile_jobs_initialized, init_static_job_batched,
_mark_task_completed, _cleanup_job, _drain_results_queue, _get_completed_count,
cleanup_job, drain_results_queue, get_completed_count, mark_task_completed,
)
from ..job_models import TileJobState
from ..mode_contexts import (
JobStateCollaborator,
StaticModeContext,
TileOpsCollaborator,
WorkerCommsCollaborator,
)
from ..tile_processing import TileBatchArgs, extract_and_process_tile_batch
class StaticModeMixin:
@@ -27,14 +34,30 @@ class StaticModeMixin:
Expected co-mixins on `self`:
- TileOpsMixin (`calculate_tiles`, tile extract/blend helpers).
- JobStateMixin (`_get_next_tile_index`, `_get_all_completed_tasks`, requeue checks).
- WorkerCommsMixin (`send_tiles_batch_to_master`, `_request_tile_from_master`, `_send_heartbeat_to_master`).
- WorkerCommsMixin (`send_tiles_batch`, `request_assignment`, `send_heartbeat`).
"""
def _poll_job_ready(self, multi_job_id, master_url, worker_id=None, max_attempts=JOB_POLL_MAX_ATTEMPTS):
def _build_static_mode_context(self) -> StaticModeContext:
"""Build explicit collaborators for static mode execution."""
return StaticModeContext(
tile_ops=TileOpsCollaborator(self),
job_state=JobStateCollaborator(self),
worker_comms=WorkerCommsCollaborator(self),
)
def _poll_job_ready(
self,
multi_job_id,
master_url,
worker_id=None,
max_attempts=JOB_POLL_MAX_ATTEMPTS,
mode_context: StaticModeContext | None = None,
):
"""Poll master for job readiness to avoid worker/master initialization race."""
context = mode_context or self._build_static_mode_context()
for attempt in range(max_attempts):
ready = run_async_in_server_loop(
self._check_job_status(multi_job_id, master_url),
context.worker_comms.check_job_status(multi_job_id, master_url),
timeout=5.0
)
if ready:
@@ -55,30 +78,27 @@ class StaticModeMixin:
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
core_args,
width,
height,
):
"""Extract one tile position for the whole batch and process it."""
tx, ty = all_tiles[tile_id]
tile_batch, x1, y1, ew, eh = self.extract_batch_tile_with_padding(
upscaled_image, tx, ty, tile_width, tile_height, padding, force_uniform_tiles
tile_batch_args = TileBatchArgs(
core=core_args,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
force_uniform_tiles=force_uniform_tiles,
width=width,
height=height,
)
region = (x1, y1, x1 + ew, y1 + eh)
processed_batch = self.process_tiles_batch(
tile_batch, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise, tiled_decode,
region, (width, height)
processed_batch, x1, y1, ew, eh = extract_and_process_tile_batch(
node=self,
upscaled_image=upscaled_image,
tx=tx,
ty=ty,
args=tile_batch_args,
)
return processed_batch, x1, y1, ew, eh
@@ -90,12 +110,14 @@ class StaticModeMixin:
padding,
worker_id,
is_final_flush=False,
mode_context: StaticModeContext | None = None,
):
"""Send accumulated tile payloads to master and return a fresh accumulator."""
context = mode_context or self._build_static_mode_context()
if not processed_tiles:
if is_final_flush:
run_async_in_server_loop(
self.send_tiles_batch_to_master(
context.worker_comms.send_tiles_batch(
[],
multi_job_id,
master_url,
@@ -107,7 +129,7 @@ class StaticModeMixin:
)
return processed_tiles
run_async_in_server_loop(
self.send_tiles_batch_to_master(
context.worker_comms.send_tiles_batch(
processed_tiles,
multi_job_id,
master_url,
@@ -133,21 +155,13 @@ class StaticModeMixin:
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
core_args,
width,
height,
mode_context: StaticModeContext | None = None,
):
"""Process one tile_id across the batch and blend into result_images."""
context = mode_context or self._build_static_mode_context()
source_batch = torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
if upscaled_image.is_cuda:
source_batch = source_batch.cuda()
@@ -159,17 +173,7 @@ class StaticModeMixin:
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
core_args,
width,
height,
)
@@ -177,30 +181,39 @@ class StaticModeMixin:
out_bs = processed_batch.shape[0] if hasattr(processed_batch, "shape") else batch_size
processed_items = min(batch_size, out_bs)
for b in range(processed_items):
tile_pil = tensor_to_pil(processed_batch, b)
if tile_pil.size != (ew, eh):
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
result_images[b] = self.blend_tile(result_images[b], tile_pil, x1, y1, (ew, eh), tile_mask, padding)
blend_processed_batch_item(
result_images,
processed_batch,
b,
context.tile_ops.blend_tile,
x1,
y1,
ew,
eh,
tile_mask,
padding,
)
global_idx = b * num_tiles_per_image + tile_id
run_async_in_server_loop(
_mark_task_completed(multi_job_id, global_idx, {'batch_idx': b, 'tile_idx': tile_id}),
mark_task_completed(multi_job_id, global_idx, {'batch_idx': b, 'tile_idx': tile_id}),
timeout=5.0
)
return processed_items
def _process_worker_static_sync(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
def _process_worker_static_sync(self, upscaled_image, core_args,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, master_url,
worker_id, enabled_workers):
force_uniform_tiles, multi_job_id, master_url,
worker_id, enabled_workers,
mode_context: StaticModeContext | None = None):
"""Worker static mode processing with optional dynamic queue pulling."""
context = mode_context or self._build_static_mode_context()
# Round tile dimensions
tile_width = self.round_to_multiple(tile_width)
tile_height = self.round_to_multiple(tile_height)
tile_width = context.tile_ops.round_to_multiple(tile_width)
tile_height = context.tile_ops.round_to_multiple(tile_height)
# Get dimensions and calculate tiles
_, height, width, _ = upscaled_image.shape
all_tiles = self.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
all_tiles = context.tile_ops.calculate_tiles(width, height, tile_width, tile_height, force_uniform_tiles)
num_tiles_per_image = len(all_tiles)
batch_size = upscaled_image.shape[0]
total_tiles = batch_size * num_tiles_per_image
@@ -212,24 +225,31 @@ class StaticModeMixin:
working_images.append(image_pil.copy())
tile_masks = []
for tx, ty in all_tiles:
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
tile_masks.append(context.tile_ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
# Dynamic queue mode (static processing): process batched-per-tile
log(f"USDU Dist Worker[{worker_id[:8]}]: Canvas {width}x{height} | Tile {tile_width}x{tile_height} | Tiles/image {num_tiles_per_image} | Batch {batch_size}")
processed_count = 0
max_poll_attempts = JOB_POLL_MAX_ATTEMPTS
if not self._poll_job_ready(multi_job_id, master_url, worker_id=worker_id, max_attempts=max_poll_attempts):
if not self._poll_job_ready(
multi_job_id,
master_url,
worker_id=worker_id,
max_attempts=max_poll_attempts,
mode_context=context,
):
log(f"Job {multi_job_id} not ready after {max_poll_attempts} attempts, aborting")
return (upscaled_image,)
# Main processing loop - pull tile ids from queue
while True:
# Request a tile to process
tile_idx, estimated_remaining, batched_static = run_async_in_server_loop(
self._request_tile_from_master(multi_job_id, master_url, worker_id),
timeout=TILE_WAIT_TIMEOUT
assignment = run_async_in_server_loop(
context.worker_comms.request_assignment(multi_job_id, master_url, worker_id),
timeout=TILE_WAIT_TIMEOUT,
)
tile_idx = assignment.task_idx if assignment.kind == "tile" else None
if tile_idx is None:
debug_log(f"Worker[{worker_id[:8]}] - No more tiles to process")
@@ -250,17 +270,7 @@ class StaticModeMixin:
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
core_args,
width,
height,
)
@@ -269,7 +279,7 @@ class StaticModeMixin:
tile_pil = tensor_to_pil(processed_batch, b)
if tile_pil.size != (ew, eh):
tile_pil = tile_pil.resize((ew, eh), Image.LANCZOS)
working_images[b] = self.blend_tile(
working_images[b] = context.tile_ops.blend_tile(
working_images[b],
tile_pil,
x1,
@@ -293,7 +303,7 @@ class StaticModeMixin:
# Send heartbeat
try:
run_async_in_server_loop(
self._send_heartbeat_to_master(multi_job_id, master_url, worker_id),
context.worker_comms.send_heartbeat(multi_job_id, master_url, worker_id),
timeout=5.0
)
except Exception as e:
@@ -302,20 +312,39 @@ class StaticModeMixin:
# Send tiles in batches within loop
if len(processed_tiles) >= MAX_BATCH:
processed_tiles = self._flush_tiles_to_master(
processed_tiles, multi_job_id, master_url, padding, worker_id, is_final_flush=False
processed_tiles,
multi_job_id,
master_url,
padding,
worker_id,
is_final_flush=False,
mode_context=context,
)
# Send any remaining tiles
processed_tiles = self._flush_tiles_to_master(
processed_tiles, multi_job_id, master_url, padding, worker_id, is_final_flush=True
processed_tiles,
multi_job_id,
master_url,
padding,
worker_id,
is_final_flush=True,
mode_context=context,
)
debug_log(f"Worker {worker_id} completed all assigned and requeued tiles")
return (upscaled_image,)
async def _async_collect_and_monitor_static(self, multi_job_id, total_tiles, expected_total):
async def _async_collect_and_monitor_static(
self,
multi_job_id,
total_tiles,
expected_total,
mode_context: StaticModeContext | None = None,
):
"""Async helper for collection and monitoring in static mode.
Returns collected tasks dict. Caller should check if all tasks are complete."""
context = mode_context or self._build_static_mode_context()
last_progress_log = time.time()
progress_interval = 5.0
last_heartbeat_check = time.time()
@@ -328,18 +357,18 @@ class StaticModeMixin:
raise comfy.model_management.InterruptProcessingException()
# Drain any pending results
collected_count = await _drain_results_queue(multi_job_id)
collected_count = await drain_results_queue(multi_job_id)
# Check and requeue timed-out workers periodically
current_time = time.time()
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
requeued_count = await self._check_and_requeue_timed_out_workers(multi_job_id, expected_total)
requeued_count = await context.job_state.check_and_requeue_timed_out_workers(multi_job_id, expected_total)
if requeued_count > 0:
log(f"Requeued {requeued_count} tasks from timed-out workers")
last_heartbeat_check = current_time
# Get current completion count
completed_count = await _get_completed_count(multi_job_id)
completed_count = await get_completed_count(multi_job_id)
# Progress logging
if current_time - last_progress_log >= progress_interval:
@@ -366,159 +395,175 @@ class StaticModeMixin:
await asyncio.sleep(0.1)
# Get all completed tasks for return
return await self._get_all_completed_tasks(multi_job_id)
return await context.job_state.all_completed_tasks(multi_job_id)
def _process_master_static_sync(self, upscaled_image, model, positive, negative, vae,
seed, steps, cfg, sampler_name, scheduler, denoise,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, tiled_decode, multi_job_id, enabled_workers,
all_tiles, num_tiles_per_image):
"""Static mode master processing with optional dynamic queue pulling."""
batch_size = upscaled_image.shape[0]
def _init_master_static_state(
self,
upscaled_image,
multi_job_id,
batch_size,
num_tiles_per_image,
enabled_workers,
all_tiles,
tile_width,
tile_height,
mask_blur,
mode_context: StaticModeContext | None = None,
):
context = mode_context or self._build_static_mode_context()
_, height, width, _ = upscaled_image.shape
total_tiles = batch_size * num_tiles_per_image
# Convert batch to PIL list for processing
result_images = []
for b in range(batch_size):
image_pil = tensor_to_pil(upscaled_image[b:b+1], 0)
result_images.append(image_pil.copy())
# Initialize queue: pending queue holds tile ids (batched per tile)
result_images = [tensor_to_pil(upscaled_image[b:b + 1], 0).copy() for b in range(batch_size)]
log("USDU Dist: Using tile queue distribution")
run_async_in_server_loop(
init_static_job_batched(multi_job_id, batch_size, num_tiles_per_image, enabled_workers),
timeout=10.0
)
debug_log(
f"Initialized tile-id queue with {num_tiles_per_image} ids for batch {batch_size}"
timeout=10.0,
)
debug_log(f"Initialized tile-id queue with {num_tiles_per_image} ids for batch {batch_size}")
# Precompute masks for all tile positions to avoid repeated Gaussian blur work during blending
tile_masks = []
for idx, (tx, ty) in enumerate(all_tiles):
tile_masks.append(self.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur))
tile_masks = [
context.tile_ops.create_tile_mask(width, height, tx, ty, tile_width, tile_height, mask_blur)
for tx, ty in all_tiles
]
return result_images, tile_masks, width, height
def _process_master_static_initial_tiles(
self,
multi_job_id,
total_tiles,
all_tiles,
upscaled_image,
result_images,
tile_masks,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
core_args,
width,
height,
mode_context: StaticModeContext | None = None,
):
context = mode_context or self._build_static_mode_context()
processed_count = 0
consecutive_no_tile = 0
max_consecutive_no_tile = 2
while processed_count < total_tiles:
comfy.model_management.throw_exception_if_processing_interrupted()
tile_idx = run_async_in_server_loop(
self._get_next_tile_index(multi_job_id),
timeout=5.0
)
if tile_idx is not None:
consecutive_no_tile = 0
tile_id = tile_idx
processed_count += self._master_process_one_tile(
tile_id,
all_tiles,
upscaled_image,
result_images,
tile_masks,
multi_job_id,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
width,
height,
)
log(f"USDU Dist: Tiles progress {processed_count}/{total_tiles} (tile {tile_id})")
else:
tile_id = run_async_in_server_loop(context.job_state.next_tile_index(multi_job_id), timeout=5.0)
if tile_id is None:
consecutive_no_tile += 1
if consecutive_no_tile >= max_consecutive_no_tile:
debug_log(f"Master processed {processed_count} tiles, moving to collection phase")
break
time.sleep(0.1)
master_processed_count = processed_count
# Continue processing any remaining tiles while collecting worker results
remaining_tiles = total_tiles - master_processed_count
if remaining_tiles > 0:
debug_log(f"Master waiting for {remaining_tiles} tiles from workers")
# Collect worker results using async operations
try:
# Wait until either all tasks are collected or there are no active workers left
collected_tasks = run_async_in_server_loop(
self._async_collect_and_monitor_static(multi_job_id, total_tiles, expected_total=total_tiles),
timeout=None
)
except comfy.model_management.InterruptProcessingException:
# Clean up job on interruption
run_async_in_server_loop(_cleanup_job(multi_job_id), timeout=5.0)
raise
# Check if we need to process any remaining tasks locally after collection
completed_count = len(collected_tasks)
if completed_count < total_tiles:
log(f"Processing remaining {total_tiles - completed_count} tasks locally after worker failures")
# Process any remaining pending tasks (batched-per-tile)
while True:
# Check for user interruption
comfy.model_management.throw_exception_if_processing_interrupted()
continue
# Get next tile_id from pending queue
tile_id = run_async_in_server_loop(
self._get_next_tile_index(multi_job_id),
timeout=5.0
)
if tile_id is None:
break
self._master_process_one_tile(
tile_id,
all_tiles,
upscaled_image,
result_images,
tile_masks,
multi_job_id,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tiled_decode,
width,
height,
)
else:
# Master processed all tiles
collected_tasks = run_async_in_server_loop(
self._get_all_completed_tasks(multi_job_id),
timeout=5.0
consecutive_no_tile = 0
processed_count += self._master_process_one_tile(
tile_id,
all_tiles,
upscaled_image,
result_images,
tile_masks,
multi_job_id,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
core_args,
width,
height,
mode_context=context,
)
# Blend worker tiles synchronously in deterministic tile order.
log(f"USDU Dist: Tiles progress {processed_count}/{total_tiles} (tile {tile_id})")
return processed_count
def _collect_remaining_static_tiles(
self,
multi_job_id,
total_tiles,
master_processed_count,
all_tiles,
upscaled_image,
result_images,
tile_masks,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
core_args,
width,
height,
mode_context: StaticModeContext | None = None,
):
context = mode_context or self._build_static_mode_context()
remaining_tiles = total_tiles - master_processed_count
if remaining_tiles <= 0:
return run_async_in_server_loop(context.job_state.all_completed_tasks(multi_job_id), timeout=5.0)
debug_log(f"Master waiting for {remaining_tiles} tiles from workers")
collected_tasks = run_async_in_server_loop(
self._async_collect_and_monitor_static(
multi_job_id,
total_tiles,
expected_total=total_tiles,
mode_context=context,
),
timeout=None,
)
completed_count = len(collected_tasks)
if completed_count >= total_tiles:
return collected_tasks
log(f"Processing remaining {total_tiles - completed_count} tasks locally after worker failures")
while True:
comfy.model_management.throw_exception_if_processing_interrupted()
tile_id = run_async_in_server_loop(context.job_state.next_tile_index(multi_job_id), timeout=5.0)
if tile_id is None:
break
self._master_process_one_tile(
tile_id,
all_tiles,
upscaled_image,
result_images,
tile_masks,
multi_job_id,
batch_size,
num_tiles_per_image,
tile_width,
tile_height,
padding,
force_uniform_tiles,
core_args,
width,
height,
mode_context=context,
)
return collected_tasks
def _blend_static_collected_tiles(
self,
collected_tasks,
result_images,
all_tiles,
tile_masks,
num_tiles_per_image,
batch_size,
tile_width,
tile_height,
padding,
mode_context: StaticModeContext | None = None,
):
context = mode_context or self._build_static_mode_context()
def _sort_key(item):
global_idx, tile_data = item
batch_idx = tile_data.get('batch_idx', global_idx // num_tiles_per_image)
@@ -526,45 +571,115 @@ class StaticModeMixin:
return (tile_idx, batch_idx, global_idx)
for global_idx, tile_data in sorted(collected_tasks.items(), key=_sort_key):
# Skip tiles that don't have tensor data (already processed)
if 'tensor' not in tile_data and 'image' not in tile_data:
continue
batch_idx = tile_data.get('batch_idx', global_idx // num_tiles_per_image)
tile_idx = tile_data.get('tile_idx', global_idx % num_tiles_per_image)
if batch_idx >= batch_size:
continue
# Blend tile synchronously
x = tile_data.get('x', 0)
y = tile_data.get('y', 0)
# Prefer PIL image if present to avoid reconversion
if 'image' in tile_data:
tile_pil = tile_data['image']
else:
tile_tensor = tile_data['tensor']
tile_pil = tensor_to_pil(tile_tensor, 0)
orig_x, orig_y = all_tiles[tile_idx]
tile_pil = tile_data['image'] if 'image' in tile_data else tensor_to_pil(tile_data['tensor'], 0)
tile_mask = tile_masks[tile_idx]
extracted_width = tile_data.get('extracted_width', tile_width + 2 * padding)
extracted_height = tile_data.get('extracted_height', tile_height + 2 * padding)
result_images[batch_idx] = self.blend_tile(result_images[batch_idx], tile_pil,
x, y, (extracted_width, extracted_height), tile_mask, padding)
result_images[batch_idx] = context.tile_ops.blend_tile(
result_images[batch_idx],
tile_pil,
x,
y,
(extracted_width, extracted_height),
tile_mask,
padding,
)
def _result_images_to_tensor(self, result_images, batch_size, upscaled_image):
if batch_size == 1:
result_tensor = pil_to_tensor(result_images[0])
else:
result_tensor = torch.cat([pil_to_tensor(img) for img in result_images], dim=0)
if upscaled_image.is_cuda:
result_tensor = result_tensor.cuda()
return result_tensor
def _process_master_static_sync(self, upscaled_image, core_args,
tile_width, tile_height, padding, mask_blur,
force_uniform_tiles, multi_job_id, enabled_workers,
all_tiles, num_tiles_per_image,
mode_context: StaticModeContext | None = None):
"""Static mode master processing with optional dynamic queue pulling."""
context = mode_context or self._build_static_mode_context()
batch_size = upscaled_image.shape[0]
total_tiles = batch_size * num_tiles_per_image
result_images, tile_masks, width, height = self._init_master_static_state(
upscaled_image=upscaled_image,
multi_job_id=multi_job_id,
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
enabled_workers=enabled_workers,
all_tiles=all_tiles,
tile_width=tile_width,
tile_height=tile_height,
mask_blur=mask_blur,
mode_context=context,
)
try:
# Convert back to tensor
if batch_size == 1:
result_tensor = pil_to_tensor(result_images[0])
else:
result_tensors = [pil_to_tensor(img) for img in result_images]
result_tensor = torch.cat(result_tensors, dim=0)
if upscaled_image.is_cuda:
result_tensor = result_tensor.cuda()
master_processed_count = self._process_master_static_initial_tiles(
multi_job_id=multi_job_id,
total_tiles=total_tiles,
all_tiles=all_tiles,
upscaled_image=upscaled_image,
result_images=result_images,
tile_masks=tile_masks,
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
force_uniform_tiles=force_uniform_tiles,
core_args=core_args,
width=width,
height=height,
mode_context=context,
)
collected_tasks = self._collect_remaining_static_tiles(
multi_job_id=multi_job_id,
total_tiles=total_tiles,
master_processed_count=master_processed_count,
all_tiles=all_tiles,
upscaled_image=upscaled_image,
result_images=result_images,
tile_masks=tile_masks,
batch_size=batch_size,
num_tiles_per_image=num_tiles_per_image,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
force_uniform_tiles=force_uniform_tiles,
core_args=core_args,
width=width,
height=height,
mode_context=context,
)
self._blend_static_collected_tiles(
collected_tasks=collected_tasks,
result_images=result_images,
all_tiles=all_tiles,
tile_masks=tile_masks,
num_tiles_per_image=num_tiles_per_image,
batch_size=batch_size,
tile_width=tile_width,
tile_height=tile_height,
padding=padding,
mode_context=context,
)
result_tensor = self._result_images_to_tensor(result_images, batch_size, upscaled_image)
log(f"UltimateSDUpscale Master - Job {multi_job_id} complete")
return (result_tensor,)
finally:
# Cleanup (async operation) - always execute
run_async_in_server_loop(_cleanup_job(multi_job_id), timeout=5.0)
run_async_in_server_loop(cleanup_job(multi_job_id), timeout=5.0)
+11 -5
View File
@@ -1,10 +1,12 @@
import io
import json
from collections.abc import Mapping
from typing import Any
from PIL import Image
def _parse_tiles_from_form(data):
def parse_tiles_from_form(data: Mapping[str, Any]) -> list[dict[str, Any]]:
"""Parse tiles submitted via multipart/form-data into a list of tile dicts."""
try:
padding = int(data.get('padding', 0)) if data.get('padding') is not None else 0
@@ -51,14 +53,18 @@ def _parse_tiles_from_form(data):
if 'batch_idx' in meta:
try:
tile_info['batch_idx'] = int(meta['batch_idx'])
except Exception:
pass
except Exception as exc:
raise ValueError(f"Invalid batch_idx for tile {i}: {meta.get('batch_idx')} ({exc})")
if 'global_idx' in meta:
try:
tile_info['global_idx'] = int(meta['global_idx'])
except Exception:
pass
except Exception as exc:
raise ValueError(f"Invalid global_idx for tile {i}: {meta.get('global_idx')} ({exc})")
tiles.append(tile_info)
return tiles
# Backward compatibility for existing imports.
_parse_tiles_from_form = parse_tiles_from_form
+19
View File
@@ -0,0 +1,19 @@
from dataclasses import dataclass
from typing import Any
@dataclass(frozen=True)
class UpscaleCoreArgs:
"""Shared denoise/sampling arguments for USDU tile processing."""
model: Any
positive: Any
negative: Any
vae: Any
seed: int
steps: int
cfg: float
sampler_name: str
scheduler: str
denoise: float
tiled_decode: bool
+144 -119
View File
@@ -4,8 +4,8 @@ import server
from ..utils.constants import DYNAMIC_MODE_MAX_POLL_TIMEOUT, HEARTBEAT_INTERVAL
from ..utils.logging import debug_log, log
from ..utils.config import get_worker_timeout_seconds
from .job_store import ensure_tile_jobs_initialized, _mark_task_completed
from .job_timeout import _check_and_requeue_timed_out_workers
from .job_store import ensure_tile_jobs_initialized, mark_task_completed
from .job_timeout import check_and_requeue_timed_out_workers
from .job_models import BaseJobState, ImageJobState, TileJobState
@@ -16,7 +16,7 @@ class ResultCollectorMixin:
Expected co-mixins/attributes:
- JobStateMixin methods for queue/task access.
- `self._check_and_requeue_timed_out_workers(...)` coroutine.
- `self._async_yield(...)` optional helper from WorkerCommsMixin.
- `self.async_yield(...)` optional helper from WorkerCommsMixin.
"""
def _log_worker_timeout_status(self, job_data, current_time: float, multi_job_id: str) -> list[str]:
@@ -33,166 +33,191 @@ class ResultCollectorMixin:
)
return list(worker_status.keys())
async def _async_collect_results(self, multi_job_id, num_workers, mode='static',
remaining_to_collect=None, batch_size=None):
"""Unified async helper to collect results from workers (tiles or images)."""
# Get the already initialized queue
async def _check_and_requeue_timed_out_workers(self, multi_job_id, batch_size):
"""Default timeout requeue hook; override in host mixins when needed."""
return await check_and_requeue_timed_out_workers(multi_job_id, batch_size)
def _get_job_data_snapshot(self, prompt_server, multi_job_id):
"""Get a snapshot of job data for timeout logging (non-async helper)."""
current_job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
if isinstance(current_job_data, BaseJobState):
return current_job_data
return None
async def _async_collect_worker_tiles(self, multi_job_id, num_workers):
"""Collect tiles from workers in static mode."""
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_tile_jobs:
raise RuntimeError(f"Job queue not initialized for {multi_job_id}")
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if mode == 'dynamic':
if not isinstance(job_data, ImageJobState):
raise RuntimeError(
f"Mode mismatch: expected dynamic, got {getattr(job_data, 'mode', 'unknown')}"
)
q = job_data.queue
completed_images = job_data.completed_images
expected_count = remaining_to_collect or batch_size
elif mode == 'static':
if not isinstance(job_data, TileJobState):
raise RuntimeError(
f"Mode mismatch: expected static, got {getattr(job_data, 'mode', 'unknown')}"
)
q = job_data.queue
expected_count = len(job_data.completed_tasks) + job_data.pending_tasks.qsize()
else:
raise RuntimeError(f"Unsupported mode: {mode}")
item_type = "images" if mode == 'dynamic' else "tiles"
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} {item_type} from {num_workers} workers")
if not isinstance(job_data, TileJobState):
raise RuntimeError(
f"Mode mismatch: expected static, got {getattr(job_data, 'mode', 'unknown')}"
)
q = job_data.queue
expected_count = len(job_data.completed_tasks) + job_data.pending_tasks.qsize()
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} tiles from {num_workers} workers")
collected_results = {}
workers_done = set()
# Unify collector/upscaler wait behavior with the UI worker timeout
timeout = float(get_worker_timeout_seconds())
wait_started_at = time.time()
while len(workers_done) < num_workers:
if comfy.model_management.processing_interrupted():
log("Processing interrupted by user")
raise comfy.model_management.InterruptProcessingException()
job_data_snapshot = None
async with prompt_server.distributed_tile_jobs_lock:
job_data_snapshot = self._get_job_data_snapshot(prompt_server, multi_job_id)
try:
result = await asyncio.wait_for(q.get(), timeout=timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
tiles = result.get('tiles', [])
debug_log(
f"UltimateSDUpscale Master - Received batch of {len(tiles)} tiles from worker "
f"'{worker_id}' (is_last={is_last})"
)
for tile_data in tiles:
if 'batch_idx' not in tile_data:
log("UltimateSDUpscale Master - Missing batch_idx in tile data, skipping")
continue
tile_idx = tile_data['tile_idx']
key = tile_data.get('global_idx', tile_idx)
entry = {
**tile_data,
'tile_idx': tile_idx,
'worker_id': worker_id,
'global_idx': key,
}
collected_results[entry['global_idx']] = entry
if is_last:
workers_done.add(worker_id)
debug_log(f"UltimateSDUpscale Master - Worker {worker_id} completed")
except asyncio.TimeoutError:
current_time = time.time()
waiting_workers = type(self)._log_worker_timeout_status(
self, job_data_snapshot, current_time, multi_job_id,
)
elapsed = current_time - wait_started_at
log(
f"UltimateSDUpscale Master - Heartbeat timeout waiting for tiles; "
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
)
break
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {len(collected_results)} tiles from {len(workers_done)} workers")
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
del prompt_server.distributed_pending_tile_jobs[multi_job_id]
return collected_results
async def collect_dynamic_images(
self,
multi_job_id,
remaining_to_collect,
num_workers,
batch_size,
master_processed_count,
):
"""Collect remaining processed images from workers in dynamic mode."""
_ = master_processed_count
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id not in prompt_server.distributed_pending_tile_jobs:
raise RuntimeError(f"Job queue not initialized for {multi_job_id}")
job_data = prompt_server.distributed_pending_tile_jobs[multi_job_id]
if not isinstance(job_data, ImageJobState):
raise RuntimeError(
f"Mode mismatch: expected dynamic, got {getattr(job_data, 'mode', 'unknown')}"
)
q = job_data.queue
completed_images = job_data.completed_images
expected_count = remaining_to_collect or batch_size
debug_log(f"UltimateSDUpscale Master - Starting collection, expecting {expected_count} images from {num_workers} workers")
workers_done = set()
timeout = float(get_worker_timeout_seconds())
last_heartbeat_check = time.time()
wait_started_at = time.time()
collected_count = 0
while len(workers_done) < num_workers:
# Check for user interruption
if comfy.model_management.processing_interrupted():
log("Processing interrupted by user")
raise comfy.model_management.InterruptProcessingException()
# For dynamic mode with remaining_to_collect, check if we've collected enough
if mode == 'dynamic' and remaining_to_collect and collected_count >= remaining_to_collect:
if remaining_to_collect and collected_count >= remaining_to_collect:
break
job_data_snapshot = None
async with prompt_server.distributed_tile_jobs_lock:
current_job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
if isinstance(current_job_data, BaseJobState):
job_data_snapshot = current_job_data
job_data_snapshot = self._get_job_data_snapshot(prompt_server, multi_job_id)
try:
# Shorter poll for dynamic mode, but never exceed the configured timeout
wait_timeout = (min(DYNAMIC_MODE_MAX_POLL_TIMEOUT, timeout) if mode == 'dynamic' else timeout)
wait_timeout = min(DYNAMIC_MODE_MAX_POLL_TIMEOUT, timeout)
result = await asyncio.wait_for(q.get(), timeout=wait_timeout)
worker_id = result['worker_id']
is_last = result.get('is_last', False)
if mode == 'static':
# Handle tiles
tiles = result.get('tiles', [])
debug_log(
f"UltimateSDUpscale Master - Received batch of {len(tiles)} tiles from worker "
f"'{worker_id}' (is_last={is_last})"
)
for tile_data in tiles:
if 'batch_idx' not in tile_data:
log("UltimateSDUpscale Master - Missing batch_idx in tile data, skipping")
continue
if 'image_idx' in result and 'image' in result:
image_idx = result['image_idx']
image_pil = result['image']
completed_images[image_idx] = image_pil
collected_count += 1
debug_log(f"UltimateSDUpscale Master - Received image {image_idx} from worker {worker_id}")
tile_idx = tile_data['tile_idx']
key = tile_data.get('global_idx', tile_idx)
entry = {
'tile_idx': tile_idx,
'x': tile_data['x'],
'y': tile_data['y'],
'extracted_width': tile_data['extracted_width'],
'extracted_height': tile_data['extracted_height'],
'padding': tile_data['padding'],
'worker_id': worker_id,
'batch_idx': tile_data.get('batch_idx', 0),
'global_idx': tile_data.get('global_idx', tile_idx),
}
if 'image' in tile_data:
entry['image'] = tile_data['image']
elif 'tensor' in tile_data:
entry['tensor'] = tile_data['tensor']
collected_results[key] = entry
elif mode == 'dynamic':
# Handle full images
if 'image_idx' in result and 'image' in result:
image_idx = result['image_idx']
image_pil = result['image']
completed_images[image_idx] = image_pil
collected_results[image_idx] = image_pil
collected_count += 1
debug_log(f"UltimateSDUpscale Master - Received image {image_idx} from worker {worker_id}")
if is_last:
workers_done.add(worker_id)
debug_log(f"UltimateSDUpscale Master - Worker {worker_id} completed")
except asyncio.TimeoutError:
current_time = time.time()
waiting_workers = self._log_worker_timeout_status(job_data_snapshot, current_time, multi_job_id)
if mode == 'dynamic':
# Check for worker timeouts periodically
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
# Use the class method to check and requeue
requeued = await self._check_and_requeue_timed_out_workers(multi_job_id, batch_size)
if requeued > 0:
log(f"UltimateSDUpscale Master - Requeued {requeued} images from timed out workers")
last_heartbeat_check = current_time
# Check if we've been waiting too long overall
if current_time - wait_started_at > timeout:
elapsed = current_time - wait_started_at
log(
"UltimateSDUpscale Master - Heartbeat timeout while waiting for images; "
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
)
break
else:
waiting_workers = type(self)._log_worker_timeout_status(
self, job_data_snapshot, current_time, multi_job_id,
)
if current_time - last_heartbeat_check >= HEARTBEAT_INTERVAL:
requeued = await type(self)._check_and_requeue_timed_out_workers(
self, multi_job_id, batch_size,
)
if requeued > 0:
log(f"UltimateSDUpscale Master - Requeued {requeued} images from timed out workers")
last_heartbeat_check = current_time
if current_time - wait_started_at > timeout:
elapsed = current_time - wait_started_at
log(
f"UltimateSDUpscale Master - Heartbeat timeout waiting for {item_type}; "
"UltimateSDUpscale Master - Heartbeat timeout while waiting for images; "
f"workers={waiting_workers}, elapsed={elapsed:.1f}s"
)
break
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {len(collected_results)} {item_type} from {len(workers_done)} workers")
# Clean up job queue
debug_log(f"UltimateSDUpscale Master - Collection complete. Got {collected_count} images from {len(workers_done)} workers")
async with prompt_server.distributed_tile_jobs_lock:
if multi_job_id in prompt_server.distributed_pending_tile_jobs:
del prompt_server.distributed_pending_tile_jobs[multi_job_id]
return collected_results if mode == 'static' else completed_images
async def _async_collect_worker_tiles(self, multi_job_id, num_workers):
"""Async helper to collect tiles from workers."""
return await self._async_collect_results(multi_job_id, num_workers, mode='static')
return completed_images
async def _mark_image_completed(self, multi_job_id, image_idx, image_pil):
async def mark_image_completed(self, multi_job_id, image_idx, image_pil):
"""Mark an image as completed in the job data."""
# Mark the image as completed with the image data
await _mark_task_completed(multi_job_id, image_idx, {'image': image_pil})
await mark_task_completed(multi_job_id, image_idx, {'image': image_pil})
prompt_server = ensure_tile_jobs_initialized()
async with prompt_server.distributed_tile_jobs_lock:
job_data = prompt_server.distributed_pending_tile_jobs.get(multi_job_id)
if isinstance(job_data, ImageJobState):
job_data.completed_images[image_idx] = image_pil
async def _async_collect_dynamic_images(self, multi_job_id, remaining_to_collect, num_workers, batch_size, master_processed_count):
"""Collect remaining processed images from workers."""
return await self._async_collect_results(multi_job_id, num_workers, mode='dynamic',
remaining_to_collect=remaining_to_collect,
batch_size=batch_size)
+57
View File
@@ -371,6 +371,10 @@ class TileOpsMixin:
return positive_sliced, negative_sliced
def slice_conditioning(self, positive, negative, batch_idx):
"""Public conditioning-slice API."""
return self._slice_conditioning(positive, negative, batch_idx)
def _process_and_blend_tile(self, tile_idx, tile_pos, upscaled_image, result_image,
model, positive, negative, vae, seed, steps, cfg,
sampler_name, scheduler, denoise, tile_width, tile_height,
@@ -399,6 +403,59 @@ class TileOpsMixin:
return result_image
def process_and_blend_tile(
self,
tile_idx,
tile_pos,
upscaled_image,
result_image,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tile_width,
tile_height,
padding,
mask_blur,
image_width,
image_height,
force_uniform_tiles,
tiled_decode,
batch_idx: int = 0,
):
"""Public tile-processing API."""
return self._process_and_blend_tile(
tile_idx,
tile_pos,
upscaled_image,
result_image,
model,
positive,
negative,
vae,
seed,
steps,
cfg,
sampler_name,
scheduler,
denoise,
tile_width,
tile_height,
padding,
mask_blur,
image_width,
image_height,
force_uniform_tiles,
tiled_decode,
batch_idx=batch_idx,
)
def _process_single_tile(self, global_idx, num_tiles_per_image, upscaled_image, all_tiles,
model, positive, negative, vae, seed, steps, cfg, sampler_name,
scheduler, denoise, tiled_decode, tile_width, tile_height, padding,
+56
View File
@@ -0,0 +1,56 @@
from dataclasses import dataclass
from typing import Any
from .processing_args import UpscaleCoreArgs
@dataclass(frozen=True)
class TileBatchArgs:
"""Tile/canvas parameters layered on top of shared core processing args."""
core: UpscaleCoreArgs
tile_width: int
tile_height: int
padding: int
force_uniform_tiles: bool
width: int
height: int
def extract_and_process_tile_batch(
*,
node: Any,
upscaled_image: Any,
tx: int,
ty: int,
args: TileBatchArgs,
) -> tuple[Any, int, int, int, int]:
"""Extract one tile position for the whole batch and process it."""
tile_batch, x1, y1, ew, eh = node.extract_batch_tile_with_padding(
upscaled_image,
tx,
ty,
args.tile_width,
args.tile_height,
args.padding,
args.force_uniform_tiles,
)
region = (x1, y1, x1 + ew, y1 + eh)
core = args.core
processed_batch = node.process_tiles_batch(
tile_batch,
core.model,
core.positive,
core.negative,
core.vae,
core.seed,
core.steps,
core.cfg,
core.sampler_name,
core.scheduler,
core.denoise,
core.tiled_decode,
region,
(args.width, args.height),
)
return processed_batch, x1, y1, ew, eh
+209 -91
View File
@@ -1,20 +1,84 @@
import asyncio, io, json, time
from dataclasses import dataclass
from typing import Any, Literal
import aiohttp
from PIL import Image
from ..utils.logging import debug_log, log
from ..utils.auth import distributed_auth_headers
from ..utils.config import load_config
from ..utils.network import get_client_session
from ..utils.constants import TILE_SEND_TIMEOUT
from ..utils.usdu_managment import MAX_PAYLOAD_SIZE, _send_heartbeat_to_master
from ..utils.usdu_management import MAX_PAYLOAD_SIZE, send_heartbeat_to_master
from ..utils.image import tensor_to_pil
class WorkerCommsMixin:
async def _send_heartbeat_to_master(self, multi_job_id, master_url, worker_id):
"""Proxy heartbeat helper used by worker processing mixins."""
await _send_heartbeat_to_master(multi_job_id, master_url, worker_id)
WorkAssignmentKind = Literal["image", "tile", "none"]
async def send_tiles_batch_to_master(self, processed_tiles, multi_job_id, master_url,
padding, worker_id, is_final_flush=False):
@dataclass(frozen=True, slots=True)
class WorkAssignment:
"""Canonical typed representation of a single work-item assignment."""
kind: WorkAssignmentKind
task_idx: int | None
estimated_remaining: int = 0
batched_static: bool = False
class WorkerCommsMixin:
@staticmethod
def _master_auth_headers() -> dict[str, str]:
return distributed_auth_headers(load_config())
async def _post_with_retry(
self,
url: str,
*,
build_form: callable,
max_retries: int = 5,
initial_delay: float = 0.5,
max_delay: float = 5.0,
error_context: str = "",
) -> None:
"""POST form data with exponential backoff. build_form is called fresh each attempt."""
retry_delay = initial_delay
for attempt in range(max_retries):
try:
session = await get_client_session()
async with session.post(
url,
data=build_form(),
headers=self._master_auth_headers(),
) as response:
response.raise_for_status()
return
except Exception as e:
if attempt < max_retries - 1:
debug_log(f"Retry {attempt + 1}/{max_retries} after error: {e}")
await asyncio.sleep(retry_delay)
retry_delay = min(retry_delay * 2, max_delay)
else:
log(f"{error_context}: {e}")
raise
async def send_heartbeat(
self,
multi_job_id: str,
master_url: str,
worker_id: str,
) -> None:
"""Send worker heartbeat to the master."""
await send_heartbeat_to_master(multi_job_id, master_url, worker_id)
async def send_tiles_batch(
self,
processed_tiles: list[dict[str, Any]],
multi_job_id: str,
master_url: str,
padding: int,
worker_id: str,
is_final_flush: bool = False,
) -> None:
"""Send all processed tiles to master, chunked if large."""
if not processed_tiles:
if is_final_flush:
@@ -50,12 +114,8 @@ class WorkerCommsMixin:
i = 0
chunk_index = 0
while i < total_tiles:
data = aiohttp.FormData()
data.add_field('multi_job_id', multi_job_id)
data.add_field('worker_id', str(worker_id))
data.add_field('padding', str(padding))
metadata = []
chunk_images: list[tuple[int, int, bytes]] = []
used = 0
j = i
while j < total_tiles:
@@ -67,7 +127,7 @@ class WorkerCommsMixin:
break
# Accept this tile in this chunk
metadata.append(meta)
data.add_field(f'tile_{j - i}', io.BytesIO(img_bytes), filename=f'tile_{j}.png', content_type='image/png')
chunk_images.append((j - i, j, img_bytes))
used += len(img_bytes) + overhead
j += 1
@@ -76,32 +136,34 @@ class WorkerCommsMixin:
# Single oversized tile, send anyway
meta = encoded[j]['meta']
metadata.append(meta)
data.add_field('tile_0', io.BytesIO(encoded[j]['bytes']), filename=f'tile_{j}.png', content_type='image/png')
chunk_images.append((0, j, encoded[j]['bytes']))
j += 1
chunk_size = j - i
is_chunk_last = (j >= total_tiles)
data.add_field('is_last', str(bool(is_final_flush and is_chunk_last)))
data.add_field('batch_size', str(chunk_size))
data.add_field('tiles_metadata', json.dumps(metadata), content_type='application/json')
# Retry logic with exponential backoff
max_retries = 5
retry_delay = 0.5
for attempt in range(max_retries):
try:
session = await get_client_session()
url = f"{master_url}/distributed/submit_tiles"
async with session.post(url, data=data) as response:
response.raise_for_status()
break
except Exception as e:
if attempt < max_retries - 1:
await asyncio.sleep(retry_delay)
retry_delay = min(retry_delay * 2, 5.0)
else:
log(f"UltimateSDUpscale Worker - Failed to send chunk {chunk_index} after {max_retries} attempts: {e}")
raise
def _build_chunk_form() -> aiohttp.FormData:
data = aiohttp.FormData()
data.add_field('multi_job_id', multi_job_id)
data.add_field('worker_id', str(worker_id))
data.add_field('padding', str(padding))
data.add_field('is_last', str(bool(is_final_flush and is_chunk_last)))
data.add_field('batch_size', str(chunk_size))
data.add_field('tiles_metadata', json.dumps(metadata), content_type='application/json')
for relative_idx, source_idx, img_bytes in chunk_images:
data.add_field(
f'tile_{relative_idx}',
io.BytesIO(img_bytes),
filename=f'tile_{source_idx}.png',
content_type='image/png',
)
return data
await self._post_with_retry(
f"{master_url}/distributed/submit_tiles",
build_form=_build_chunk_form,
error_context=f"UltimateSDUpscale Worker - Failed to send chunk {chunk_index}",
)
debug_log(f"Worker[{worker_id[:8]}] - Sent chunk {chunk_index} ({chunk_size} tiles, ~{used/1e6:.2f} MB)")
chunk_index += 1
@@ -117,7 +179,7 @@ class WorkerCommsMixin:
session = await get_client_session()
url = f"{master_url}/distributed/submit_tiles"
async with session.post(url, data=data) as response:
async with session.post(url, data=data, headers=self._master_auth_headers()) as response:
response.raise_for_status()
debug_log(f"Worker {worker_id} sent static completion signal")
@@ -144,7 +206,7 @@ class WorkerCommsMixin:
async with session.post(url, json={
'worker_id': str(worker_id),
'multi_job_id': multi_job_id
}) as response:
}, headers=self._master_auth_headers()) as response:
if response.status == 200:
return await response.json()
if response.status == 404:
@@ -168,66 +230,92 @@ class WorkerCommsMixin:
return None
async def _request_image_from_master(self, multi_job_id, master_url, worker_id):
"""Request an image index to process from master in dynamic mode."""
data = await self._request_work_item_from_master(multi_job_id, master_url, worker_id)
@staticmethod
def _parse_work_assignment(data: dict[str, Any] | None) -> WorkAssignment:
"""Normalize assignment payloads into a single discriminated contract."""
if not data:
return None, 0
image_idx = data.get('image_idx')
estimated_remaining = data.get('estimated_remaining', 0)
return image_idx, estimated_remaining
return WorkAssignment(kind="none", task_idx=None)
async def _request_tile_from_master(self, multi_job_id, master_url, worker_id):
"""Request a tile index to process from master in static mode (reusing dynamic infrastructure)."""
kind_value = data.get("kind")
if kind_value is None:
# Backward-compatible parsing for older masters.
if "image_idx" in data:
kind_value = "image"
data = {
"kind": "image",
"task_idx": data.get("image_idx"),
"estimated_remaining": data.get("estimated_remaining", 0),
}
elif "tile_idx" in data:
kind_value = "tile"
data = {
"kind": "tile",
"task_idx": data.get("tile_idx"),
"estimated_remaining": data.get("estimated_remaining", 0),
"batched_static": data.get("batched_static", False),
}
else:
kind_value = "none"
if kind_value not in {"image", "tile", "none"}:
return WorkAssignment(kind="none", task_idx=None)
task_idx_raw = data.get("task_idx")
if task_idx_raw is None:
task_idx = None
else:
try:
task_idx = int(task_idx_raw)
except (TypeError, ValueError):
task_idx = None
try:
estimated_remaining = int(data.get("estimated_remaining", 0) or 0)
except (TypeError, ValueError):
estimated_remaining = 0
return WorkAssignment(
kind=kind_value,
task_idx=task_idx,
estimated_remaining=estimated_remaining,
batched_static=bool(data.get("batched_static", False)),
)
async def request_assignment(self, multi_job_id, master_url, worker_id) -> WorkAssignment:
"""Request one assignment and parse into the canonical discriminated contract."""
data = await self._request_work_item_from_master(multi_job_id, master_url, worker_id)
if not data:
return None, 0, False
tile_idx = data.get('tile_idx')
estimated_remaining = data.get('estimated_remaining', 0)
batched_static = data.get('batched_static', False)
return tile_idx, estimated_remaining, batched_static
return self._parse_work_assignment(data)
async def _send_full_image_to_master(self, image_pil, image_idx, multi_job_id,
master_url, worker_id, is_last):
async def send_full_image(self, image_pil, image_idx, multi_job_id,
master_url, worker_id, is_last):
"""Send a processed full image back to master in dynamic mode."""
# Serialize image to PNG
byte_io = io.BytesIO()
image_pil.save(byte_io, format='PNG', compress_level=0)
byte_io.seek(0)
# Prepare form data
data = aiohttp.FormData()
data.add_field('multi_job_id', multi_job_id)
data.add_field('worker_id', str(worker_id))
data.add_field('image_idx', str(image_idx))
data.add_field('is_last', str(is_last))
data.add_field('full_image', byte_io, filename=f'image_{image_idx}.png',
content_type='image/png')
# Retry logic
max_retries = 5
retry_delay = 0.5
for attempt in range(max_retries):
try:
session = await get_client_session()
url = f"{master_url}/distributed/submit_image"
async with session.post(url, data=data) as response:
response.raise_for_status()
debug_log(f"Successfully sent image {image_idx} to master")
return
except Exception as e:
if attempt < max_retries - 1:
debug_log(f"Retry {attempt + 1}/{max_retries} after error: {e}")
await asyncio.sleep(retry_delay)
retry_delay *= 2
else:
log(f"Failed to send image {image_idx} after {max_retries} attempts: {e}")
raise
image_bytes = byte_io.getvalue()
async def _send_worker_complete_signal(self, multi_job_id, master_url, worker_id):
def _build_image_form() -> aiohttp.FormData:
data = aiohttp.FormData()
data.add_field('multi_job_id', multi_job_id)
data.add_field('worker_id', str(worker_id))
data.add_field('image_idx', str(image_idx))
data.add_field('is_last', str(is_last))
data.add_field(
'full_image',
io.BytesIO(image_bytes),
filename=f'image_{image_idx}.png',
content_type='image/png',
)
return data
await self._post_with_retry(
f"{master_url}/distributed/submit_image",
build_form=_build_image_form,
error_context=f"Failed to send image {image_idx}",
)
debug_log(f"Successfully sent image {image_idx} to master")
async def send_worker_complete_signal(self, multi_job_id, master_url, worker_id):
"""Send completion signal to master in dynamic mode."""
# Send a dummy request with is_last=True
data = aiohttp.FormData()
@@ -239,16 +327,16 @@ class WorkerCommsMixin:
session = await get_client_session()
url = f"{master_url}/distributed/submit_image"
async with session.post(url, data=data) as response:
async with session.post(url, data=data, headers=self._master_auth_headers()) as response:
response.raise_for_status()
debug_log(f"Worker {worker_id} sent completion signal")
async def _check_job_status(self, multi_job_id, master_url):
async def check_job_status(self, multi_job_id, master_url):
"""Check if job is ready on the master."""
try:
session = await get_client_session()
url = f"{master_url}/distributed/job_status?multi_job_id={multi_job_id}"
async with session.get(url) as response:
async with session.get(url, headers=self._master_auth_headers()) as response:
if response.status == 200:
data = await response.json()
return data.get('ready', False)
@@ -257,6 +345,36 @@ class WorkerCommsMixin:
debug_log(f"Job status check failed: {e}")
return False
async def _async_yield(self):
async def init_dynamic_job_on_master(
self,
multi_job_id: str,
master_url: str,
batch_size: int,
enabled_workers: list[str],
) -> bool:
"""Tell the master to create the dynamic job queue (idempotent)."""
try:
session = await get_client_session()
url = f"{master_url}/distributed/init_dynamic_job"
async with session.post(
url,
json={
"multi_job_id": multi_job_id,
"batch_size": batch_size,
"enabled_workers": enabled_workers,
},
headers=self._master_auth_headers(),
) as response:
if response.status == 200:
debug_log(f"Worker initialized dynamic job {multi_job_id} on master (batch_size={batch_size})")
return True
text = await response.text()
debug_log(f"init_dynamic_job failed ({response.status}): {text}")
return False
except Exception as e:
debug_log(f"init_dynamic_job request failed: {e}")
return False
async def async_yield(self):
"""Simple async yield to allow event loop processing."""
await asyncio.sleep(0)
-2
View File
@@ -1,5 +1,3 @@
"""
Utility modules for ComfyUI-Distributed extension.
"""
# Make utils importable as a package
+15 -4
View File
@@ -3,9 +3,9 @@ Async helper utilities for ComfyUI-Distributed.
"""
import asyncio
import threading
import time
import uuid
import execution
import server
from typing import Optional, Any, Coroutine
from .network import get_server_loop
@@ -42,10 +42,11 @@ def run_async_in_server_loop(coro: Coroutine, timeout: Optional[float] = None) -
# Schedule on server's event loop
loop = get_server_loop()
asyncio.run_coroutine_threadsafe(wrapper(), loop)
task_future = asyncio.run_coroutine_threadsafe(wrapper(), loop)
# Wait for completion
if not event.wait(timeout):
task_future.cancel()
raise TimeoutError(f"Async operation timed out after {timeout} seconds")
if error:
@@ -53,7 +54,10 @@ def run_async_in_server_loop(coro: Coroutine, timeout: Optional[float] = None) -
return result
prompt_server = server.PromptServer.instance
def _prompt_server_instance():
import server
return server.PromptServer.instance
def _summarize_node_errors(node_errors: dict) -> str:
@@ -104,8 +108,13 @@ class PromptValidationError(RuntimeError):
super().__init__(f"Invalid prompt: {merged}")
async def queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
async def queue_prompt_payload(
prompt_obj: dict[str, Any],
workflow_meta: dict[str, Any] | None = None,
client_id: str | None = None,
) -> str:
"""Validate and queue a prompt via ComfyUI's prompt queue."""
prompt_server = _prompt_server_instance()
payload = {"prompt": prompt_obj}
payload = prompt_server.trigger_on_prompt(payload)
prompt = payload["prompt"]
@@ -122,6 +131,8 @@ async def queue_prompt_payload(prompt_obj, workflow_meta=None, client_id=None):
extra_data.setdefault("extra_pnginfo", {})["workflow"] = workflow_meta
if client_id:
extra_data["client_id"] = client_id
# Keep parity with ComfyUI /prompt endpoint so Jobs API metadata stays valid.
extra_data.setdefault("create_time", int(time.time() * 1000))
sensitive = {}
for key in getattr(execution, "SENSITIVE_EXTRA_DATA_KEYS", []):
+3 -2
View File
@@ -1,6 +1,7 @@
import base64
import binascii
import os
from typing import Any
import numpy as np
import torch
@@ -13,7 +14,7 @@ MAX_AUDIO_PAYLOAD_BYTES = int(
)
def encode_audio_payload(audio_payload):
def encode_audio_payload(audio_payload: dict[str, Any] | None) -> dict[str, Any] | None:
"""Serialize an AUDIO dict into JSON-safe canonical envelope payload."""
if not isinstance(audio_payload, dict):
return None
@@ -43,7 +44,7 @@ def encode_audio_payload(audio_payload):
}
def decode_audio_payload(audio_payload):
def decode_audio_payload(audio_payload: dict[str, Any] | None) -> dict[str, Any] | None:
"""Decode canonical envelope audio payload into an AUDIO dict."""
if audio_payload is None:
return None
+36
View File
@@ -0,0 +1,36 @@
from __future__ import annotations
import hmac
from typing import Any
AUTH_HEADER_NAME = "X-Distributed-Token"
def get_distributed_api_token(config: dict[str, Any] | None) -> str:
"""Return configured distributed API token or empty string when unset."""
settings = (config or {}).get("settings", {}) if isinstance(config, dict) else {}
token = settings.get("distributed_api_token") or settings.get("api_token")
return str(token).strip() if token is not None else ""
def distributed_auth_headers(config: dict[str, Any] | None) -> dict[str, str]:
"""Build outbound auth headers for distributed internal API calls."""
token = get_distributed_api_token(config)
if not token:
return {}
return {AUTH_HEADER_NAME: token}
def is_authorized_request(request: Any, config: dict[str, Any] | None) -> bool:
"""Validate a request against the configured shared distributed token."""
expected_token = get_distributed_api_token(config)
if not expected_token:
return True
provided_token = ""
headers = getattr(request, "headers", None)
if headers is not None:
provided_token = str(headers.get(AUTH_HEADER_NAME, "")).strip()
return hmac.compare_digest(provided_token, expected_token)
+18 -2
View File
@@ -6,6 +6,7 @@ import shutil
import stat
from urllib import error as urlerror
from urllib import request
from urllib.parse import urlsplit
from ..logging import debug_log
@@ -44,9 +45,24 @@ def _get_binary_path(bin_dir=None):
return os.path.join(bin_dir, binary_name)
def _get_release_download_base() -> str:
"""Resolve cloudflared release download base URL from environment/config."""
configured = (os.environ.get("CLOUDFLARED_RELEASE_BASE") or "").strip()
if configured:
return configured.rstrip("/")
return "github.com/cloudflare/cloudflared/releases/latest/download"
def _download_cloudflared():
asset = _get_platform_binary_name()
url = f"https://github.com/cloudflare/cloudflared/releases/latest/download/{asset}"
base = _get_release_download_base()
if "://" not in base:
scheme = (os.environ.get("CLOUDFLARED_RELEASE_SCHEME") or "https").strip() or "https"
base = f"{scheme}://{base}"
url = f"{base}/{asset}"
parsed = urlsplit(url)
if parsed.scheme not in {"http", "https"}:
raise RuntimeError(f"Unsupported cloudflared download URL scheme: {parsed.scheme}")
bin_dir = _get_cloudflared_dir()
os.makedirs(bin_dir, exist_ok=True)
@@ -54,7 +70,7 @@ def _download_cloudflared():
debug_log(f"Downloading cloudflared from {url}")
try:
with request.urlopen(url, timeout=30) as resp:
with request.urlopen(url, timeout=30) as resp: # nosec B310 - scheme is allowlisted above
with open(target_path, "wb") as f:
shutil.copyfileobj(resp, f)
except urlerror.URLError as exc:
+70 -15
View File
@@ -3,6 +3,7 @@
import asyncio
import re
import threading
from typing import Any
from ..constants import CLOUDFLARE_LOG_BUFFER_SIZE
from ..logging import debug_log
@@ -15,6 +16,7 @@ class ProcessReader:
def __init__(self, log_file=None):
self._process = None
self._thread = None
self._task = None
self._loop = None
self._url_event = None
self._public_url = None
@@ -37,41 +39,90 @@ class ProcessReader:
if len(self._recent_logs) > CLOUDFLARE_LOG_BUFFER_SIZE:
self._recent_logs = self._recent_logs[-CLOUDFLARE_LOG_BUFFER_SIZE:]
@staticmethod
def _normalize_line(raw_line):
if isinstance(raw_line, bytes):
return raw_line.decode("utf-8", errors="replace").strip()
return str(raw_line).strip()
def _process_log_line(self, line):
self._append_log(line)
match = PUBLIC_URL_PATTERN.search(line)
if match and not self._public_url:
self._public_url = match.group(1).rstrip("/")
return True
if "error" in line.lower() and not self._last_error:
self._last_error = line
return False
def _reader(self):
process = self._process
if process is None:
return
loop = self._loop
for raw_line in iter(process.stdout.readline, ""):
line = raw_line.strip()
while True:
raw_line = process.stdout.readline()
if not raw_line:
break
line = self._normalize_line(raw_line)
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("/")
if self._url_event and loop:
loop.call_soon_threadsafe(self._url_event.set)
if "error" in line.lower() and not self._last_error:
self._last_error = line
found_url = self._process_log_line(line)
if found_url and self._url_event and loop:
loop.call_soon_threadsafe(self._url_event.set)
if self._url_event and loop:
if not self._last_error and not self._public_url:
self._last_error = "Cloudflare tunnel exited before becoming ready"
loop.call_soon_threadsafe(self._url_event.set)
def start(self, process, loop):
async def _async_reader(self):
process = self._process
if process is None or process.stdout is None:
if self._url_event:
self._url_event.set()
return
while True:
raw_line = await process.stdout.readline()
if not raw_line:
break
line = self._normalize_line(raw_line)
if not line:
continue
found_url = self._process_log_line(line)
if found_url and self._url_event:
self._url_event.set()
if self._url_event:
if not self._last_error and not self._public_url:
self._last_error = "Cloudflare tunnel exited before becoming ready"
self._url_event.set()
def start(self, process: Any, loop: asyncio.AbstractEventLoop) -> None:
self._process = process
self._loop = loop
self._url_event = asyncio.Event()
self._public_url = None
self._last_error = None
self._recent_logs = []
self._thread = threading.Thread(target=self._reader, daemon=True)
self._thread.start()
self._thread = None
self._task = None
stdout = getattr(process, "stdout", None)
if stdout is None:
return
if asyncio.iscoroutinefunction(getattr(stdout, "readline", None)):
self._task = loop.create_task(self._async_reader())
else:
self._thread = threading.Thread(target=self._reader, daemon=True)
self._thread.start()
async def wait_for_url(self, timeout):
if not self._url_event:
@@ -79,7 +130,11 @@ class ProcessReader:
await asyncio.wait_for(self._url_event.wait(), timeout=timeout)
return self._public_url
def stop(self):
def stop(self) -> None:
if self._task and not self._task.done():
self._task.cancel()
self._task = None
if self._thread and self._thread.is_alive():
self._thread.join(timeout=1)
self._thread = None
+42 -30
View File
@@ -1,17 +1,30 @@
"""Cloudflare tunnel state persistence helpers."""
from dataclasses import dataclass
from typing import Any
from ..config import load_config, save_config
from ..network import normalize_host
def _get_tunnel_config(cfg):
@dataclass(frozen=True)
class TunnelStateUpdate:
status: str | None = None
public_url: str | None = None
pid: int | None = None
log_file: str | None = None
previous_host: str | None = None
master_host: str | None = None
def _get_tunnel_config(cfg: dict[str, Any]) -> dict[str, Any]:
tunnel_cfg = cfg.get("tunnel", {})
if isinstance(tunnel_cfg, dict):
return tunnel_cfg
return {}
def load_tunnel_state():
def load_tunnel_state() -> dict[str, Any]:
cfg = load_config()
tunnel_cfg = _get_tunnel_config(cfg)
master_cfg = cfg.get("master", {}) if isinstance(cfg.get("master", {}), dict) else {}
@@ -25,46 +38,45 @@ def load_tunnel_state():
}
def persist_tunnel_state(
status=None,
public_url=None,
pid=None,
log_file=None,
previous_host=None,
master_host=None,
):
def persist_tunnel_state(update: TunnelStateUpdate) -> None:
cfg = load_config()
tunnel_cfg = _get_tunnel_config(cfg)
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
if update.status is not None:
tunnel_cfg["status"] = update.status
if update.public_url is not None:
tunnel_cfg["public_url"] = update.public_url
if update.pid is not None:
tunnel_cfg["pid"] = update.pid
if update.log_file is not None:
tunnel_cfg["log_file"] = update.log_file
if update.previous_host is not None:
tunnel_cfg["previous_master_host"] = update.previous_host
if update.master_host is not None:
cfg.setdefault("master", {})["host"] = update.master_host
cfg["tunnel"] = tunnel_cfg
save_config(cfg)
def clear_tunnel_state(log_file=None, previous_host=None, master_host=None):
def clear_tunnel_state(
log_file: str | None = None,
previous_host: str | None = None,
master_host: str | None = None,
) -> None:
persist_tunnel_state(
status="stopped",
public_url="",
pid=None,
log_file=log_file,
previous_host=previous_host,
master_host=master_host,
TunnelStateUpdate(
status="stopped",
public_url="",
pid=None,
log_file=log_file,
previous_host=previous_host,
master_host=master_host,
)
)
def resolve_restore_master_host(previous_master_host):
def resolve_restore_master_host(previous_master_host: str | None) -> str | None:
"""Determine whether master host should be restored after tunnel stop."""
cfg = load_config()
tunnel_cfg = _get_tunnel_config(cfg)

Some files were not shown because too many files have changed in this diff Show More