Compare 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
154 changed files with 10717 additions and 6642 deletions
+1
View File
@@ -1,3 +1,4 @@
# These are supported funding model platforms
github: robertvoy
buy_me_a_coffee: robertvoy
+1 -1
View File
@@ -19,6 +19,6 @@ jobs:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
uses: Comfy-Org/publish-node-action@v1
with:
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+5
View File
@@ -7,3 +7,8 @@ __pycache__/
node_modules/
npm-debug.log*
AGENTS.md
.claude/
# Desloppify artifacts
.desloppify/
scorecard.png
+59 -49
View File
@@ -15,7 +15,7 @@
---
## Key Features
## Key Features
#### Parallel Workflow Processing
- Run your workflow on multiple GPUs simultaneously with varied seeds, collect results on the master
@@ -27,12 +27,22 @@
- Intelligent distribution
- Handles single images and videos
#### Ease of Use
- Auto-setup local workers; easily add remote/cloud ones
- Convert any workflow to distributed with 2 nodes
- JSON configuration with UI controls
---
#### Ease of Use
- Auto-setup local workers; easily add remote/cloud ones
- Convert any workflow to distributed with 2 nodes
- JSON configuration with UI controls
---
## Current Architecture
- Workflow-level load balancing is controlled by **Distributed Collector** via the `load_balance` toggle.
- There is **no Distributed Queue node** anymore.
- With `load_balance=true`, orchestration selects one least-busy execution participant:
- If master participation is enabled, master is included as a candidate.
- If master is in orchestrator-only mode, only workers are considered.
---
## Worker Types
@@ -51,6 +61,7 @@ ComfyUI Distributed supports three types of workers:
## Requirements
- ComfyUI
> Note: Desktop app not currently supported
- Multiple NVIDIA GPUs
> No additional GPUs? Use [Cloud Workers](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/worker-setup-guides.md#cloud-workers)
- That's it
@@ -81,19 +92,19 @@ Join Runpod with [this link](https://get.runpod.io/0bw29uf3ug0p) and unlock a sp
## Workflow Examples
### Basic Parallel Generation
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
### Basic Parallel Generation
Generate multiple images in the time it takes to generate one. Each worker uses a different seed.
![Clipboard Image (6)](https://github.com/user-attachments/assets/9598c94c-d9b4-4ccf-ab16-a21398220aeb)
> [Download workflow](/workflows/distributed-txt2img.json)
1. Open your ComfyUI workflow
2. Add **Distributed Seed** → connect to sampler's seed
3. Add **Distributed Collector** → after VAE Decode
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
5. Enable workers in the UI
6. Run the workflow!
2. Add **Distributed Seed** → connect to sampler's seed
3. Add **Distributed Collector** → after VAE Decode
4. Optional: enable `load_balance` on Distributed Collector to run on one least-busy participant
5. Enable workers in the UI
6. Run the workflow!
### Parallel WAN Generation
Generate multiple videos in the time it takes to generate one. Each worker uses a different seed.
@@ -111,12 +122,12 @@ Generate multiple videos in the time it takes to generate one. Each worker uses
7. Enable workers in the UI
8. Run the workflow!
### Distributed Image Upscaling
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
### Distributed Image Upscaling
Accelerate Ultimate SD Upscaler by distributing tiles across multiple workers, with speed scaling as you add more GPUs.
![Clipboard Image (3)](https://github.com/user-attachments/assets/ffb57a0d-7b75-4497-96d2-875d60865a1a)
> [Download workflow](/workflows/distributed-upscale.json)
> [Download workflow](/workflows/distributed-upscale.json)
1. Load your image
2. Upscale with ESRGAN or similar
@@ -142,40 +153,40 @@ Accelerate Ultimate SD Upscaler by distributing video tiles across multiple work
---
## Developer API
## Developer API
Control your distributed cluster programmatically without opening the browser.
* **Endpoint:** `POST /distributed/queue`
* **Functionality:** Accepts a ComfyUI API-format prompt, dispatches it to the requested reachable workers, and returns the master `prompt_id`.
* **Documentation:** [See API Examples & Scripts](https://github.com/robertvoy/ComfyUI-Distributed/blob/main/docs/comfyui-distributed-api.md)
* **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).
---
## Distributed Value
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
- The node shows only currently enabled workers.
- If worker enablement changes, worker fields update automatically.
- When disconnected, it resets to default string mode and clears per-worker overrides.
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
---
## Nodes
| Node | Description |
|------|-------------|
| **Distributed Seed** | Generates unique seeds for each worker |
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
| **Distributed Value** | Outputs per-worker override values with fallback to default |
| **Ultimate SD Upscale Distributed** | Distributes upscale tiles across workers |
| **Image Batch Divider** | Splits image batches for multi-GPU output |
| **Audio Segment Divider** | Splits an audio waveform into up to ten sequential time segments |
> **⚠️ Security Warning:** Do not expose your ComfyUI port to the public internet. If you need remote access, run ComfyUI behind a secure proxy (like Cloudflare or a VPN).
---
## Distributed Value
Use **Distributed Value** when you want per-worker overrides (for example, different prompts/models/settings per worker).
- Output type adapts to the connected input where possible (`STRING`, `INT`, `FLOAT`, `COMBO`).
- The node shows only currently enabled workers.
- If worker enablement changes, worker fields update automatically.
- When disconnected, it resets to default string mode and clears per-worker overrides.
- On execution, master uses `default_value`; workers use their mapped override with typed coercion fallback to default.
---
## Nodes
| Node | Description |
|------|-------------|
| **Distributed Seed** | Generates unique seeds for each worker |
| **Distributed Collector** | Collects results (image/video frames and optionally audio) from workers back to the master; `load_balance` can route the run to one least-busy participant |
| **Distributed Value** | Outputs per-worker override values with fallback to default |
| **Ultimate SD Upscale Distributed** | Distributes upscale tiles across workers |
| **Image Batch Divider** | Splits image batches for multi-GPU output |
| **Audio Batch Divider** | Splits audio batches for multi-GPU output |
| **Distributed Model Name** | Passes model paths to workers, enabling workflows to use models not present on the master in orchestrator-only mode |
| **Distributed Empty Image** | Produces an empty IMAGE batch used when the master delegates all work |
@@ -195,7 +206,7 @@ No, it does not speed up the generation of a single image or video. Instead, it
<details>
<summary>Does it work with the ComfyUI desktop app?</summary>
Yes, it does now.
Currently, it is not compatible with the ComfyUI desktop app.
</details>
<details>
@@ -238,4 +249,3 @@ Buy me a coffee at: https://buymeacoffee.com/robertvoy
+7 -18
View File
@@ -1,22 +1,11 @@
"""ComfyUI-Distributed's native V3 extension entrypoint."""
from comfy_api.v0_0_2 import ComfyExtension, io
"""ComfyUI-Distributed package entrypoint."""
from __future__ import annotations
from .nodes.v3 import NODES
from .runtime.bootstrap import initialize
from .bootstrap.entrypoint import build_node_mappings, initialize_runtime
WEB_DIRECTORY = './web'
WEB_DIRECTORY = "./web"
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
initialize_runtime()
class DistributedExtension(ComfyExtension):
async def on_load(self) -> None:
initialize()
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return list(NODES)
async def comfy_entrypoint() -> DistributedExtension:
return DistributedExtension()
__all__ = ['comfy_entrypoint', 'WEB_DIRECTORY']
__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)
+94 -54
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:
@@ -217,7 +243,7 @@ async def distributed_queue_endpoint(request):
return await handle_api_error(request, exc, 400)
try:
prompt_id, prompt_number, worker_count, node_errors = await orchestrate_distributed_execution(
prompt_id, worker_count = await orchestrate_distributed_execution(
payload.prompt,
payload.workflow_meta,
payload.client_id,
@@ -227,17 +253,17 @@ async def distributed_queue_endpoint(request):
)
return web.json_response({
"prompt_id": prompt_id,
"number": prompt_number,
"node_errors": node_errors,
"worker_count": worker_count,
"auto_prepare_supported": True,
})
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")
@@ -253,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")
@@ -271,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:
@@ -286,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():
@@ -295,42 +327,50 @@ async def job_complete_endpoint(request):
errors.append("worker_id: expected non-empty string")
if not isinstance(batch_idx, int) or batch_idx < 0:
errors.append("batch_idx: expected non-negative integer")
image_provided = image_payload is not None
audio_provided = audio_payload is not None
if not image_provided and not audio_provided:
errors.append("expected at least one of image or audio")
if image_provided and (not isinstance(image_payload, str) or not image_payload.strip()):
if not isinstance(image_payload, str) or not image_payload.strip():
errors.append("image: expected non-empty base64 PNG string")
if audio_provided and not isinstance(audio_payload, dict):
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)
try:
tensor = _decode_canonical_png_tensor(image_payload) if image_provided else None
decoded_audio = _decode_audio_payload(audio_payload) if audio_provided else None
except ValueError as exc:
return await handle_api_error(request, exc, 400)
tensor = _decode_canonical_png_tensor(image_payload)
decoded_audio = _decode_audio_payload(audio_payload) if audio_payload is not None else None
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:
+400 -303
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,218 +159,193 @@ def _find_upstream_nodes(prompt_obj, start_ids):
return connected
_DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES = {
"PrimitiveBoolean",
"PrimitiveFloat",
"PrimitiveInt",
"PrimitiveNode",
"PrimitiveString",
}
_DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES = {
"LoadImage",
}
_DELEGATE_MASTER_SAFE_SCALAR_TYPES = {"BOOLEAN", "FLOAT", "INT", "STRING"}
_DELEGATE_MASTER_SAFE_LIST_TYPES = {"LIST"}
# ComfyUI 0.23 exposes CreateList via the newer schema API rather than the
# legacy RETURN_TYPES/INPUT_TYPES attributes. Treat it as a safe config utility
# only after its connected inputs recursively prove safe.
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES = {"CreateList"}
_DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES = {
"CreateList": ("inputs.",),
}
# Test hook. At runtime this stays None and the ComfyUI node registry is loaded lazily.
_DELEGATE_MASTER_NODE_CLASS_MAPPINGS = None
def _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 _get_delegate_master_node_class_mappings():
"""Return ComfyUI node-class mappings when available."""
if _DELEGATE_MASTER_NODE_CLASS_MAPPINGS is not None:
return _DELEGATE_MASTER_NODE_CLASS_MAPPINGS
try:
import nodes as comfy_nodes # type: ignore
except Exception: # pragma: no cover - depends on ComfyUI runtime imports
return {}
return getattr(comfy_nodes, "NODE_CLASS_MAPPINGS", {}) or {}
def _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 _get_delegate_master_node_class(class_type):
mappings = _get_delegate_master_node_class_mappings()
return mappings.get(class_type) if isinstance(mappings, dict) else None
def _normalize_delegate_master_return_type(return_type):
if return_type is None:
return ""
return str(return_type).strip().upper()
def _delegate_master_type_is_safe_scalar(type_name):
return type_name in _DELEGATE_MASTER_SAFE_SCALAR_TYPES
def _delegate_master_type_is_safe_config(type_name):
return _delegate_master_type_is_safe_scalar(type_name) or type_name in _DELEGATE_MASTER_SAFE_LIST_TYPES
def _delegate_master_output_is_safe_scalar(class_type, output_index):
"""Return True when a registered node output is lightweight config data."""
if class_type in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_OUTPUT_CLASSES:
return True
node_class = _get_delegate_master_node_class(class_type)
return_types = getattr(node_class, "RETURN_TYPES", ()) if node_class is not None else ()
try:
output_type = return_types[int(output_index)]
except (IndexError, TypeError, ValueError):
return False
return _delegate_master_type_is_safe_config(_normalize_delegate_master_return_type(output_type))
def _get_delegate_master_input_types(class_type):
node_class = _get_delegate_master_node_class(class_type)
input_types = getattr(node_class, "INPUT_TYPES", None) if node_class is not None else None
if callable(input_types):
try:
input_types = input_types()
except TypeError:
return {}
return input_types if isinstance(input_types, dict) else {}
def _normalize_delegate_master_input_type(input_spec):
if isinstance(input_spec, (list, tuple)) and input_spec:
return _normalize_delegate_master_return_type(input_spec[0])
return _normalize_delegate_master_return_type(input_spec)
def _delegate_master_input_is_safe_scalar(class_type, input_name):
"""Return True when a registered downstream input expects config data."""
for prefix in _DELEGATE_MASTER_SAFE_DYNAMIC_CONFIG_INPUT_PREFIXES.get(class_type, ()):
if input_name.startswith(prefix):
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
input_types = _get_delegate_master_input_types(class_type)
for section_name in ("required", "optional"):
section = input_types.get(section_name, {})
if isinstance(section, dict) and input_name in section:
input_type = _normalize_delegate_master_input_type(section[input_name])
return _delegate_master_type_is_safe_config(input_type)
return False
def _is_delegate_master_always_retained_upstream_node(node):
if not isinstance(node, dict):
return False
class_type = node.get("class_type")
return isinstance(class_type, str) and class_type in _DELEGATE_MASTER_ALWAYS_RETAINED_UPSTREAM_CLASSES
def _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 _is_delegate_master_retained_upstream_node(node, output_index=0):
"""Return True for lightweight upstream nodes safe to keep on the master."""
if not isinstance(node, dict):
return False
class_type = node.get("class_type")
if not isinstance(class_type, str):
return False
return (
class_type in _DELEGATE_MASTER_RETAINED_UPSTREAM_CLASSES
or class_type.startswith("Primitive")
or _delegate_master_output_is_safe_scalar(class_type, output_index)
)
def _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
def _collect_delegate_master_retained_upstream_branch(
prompt_obj,
node_id,
output_index,
memo,
visiting,
):
"""Return safe retained branch nodes, or None when the branch is not safe."""
node_id = str(node_id)
cache_key = (node_id, output_index)
if cache_key in memo:
cached = memo[cache_key]
return None if cached is None else set(cached)
if cache_key in visiting:
memo[cache_key] = None
return None
next_id = _create_numeric_id_generator(prompt_obj)
node = prompt_obj.get(node_id)
if not _is_delegate_master_retained_upstream_node(node, output_index):
memo[cache_key] = None
return None
visiting.add(cache_key)
retained = {node_id}
inputs = node.get("inputs", {}) if isinstance(node, dict) else {}
class_type = node.get("class_type") if isinstance(node, dict) else None
for input_name, value in inputs.items():
if not (isinstance(value, list) and len(value) == 2):
# 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
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
visiting.remove(cache_key)
memo[cache_key] = None
return None
source_id = str(value[0])
branch = _collect_delegate_master_retained_upstream_branch(
prompt_obj,
source_id,
value[1],
memo,
visiting,
)
if branch is None:
visiting.remove(cache_key)
memo[cache_key] = None
return None
retained.update(branch)
visiting.remove(cache_key)
memo[cache_key] = frozenset(retained)
return retained
def _find_delegate_master_retained_upstream_nodes(prompt_obj, start_ids):
"""Return lightweight upstream nodes needed by kept delegate-master nodes."""
connected = set()
memo = {}
for node_id in start_ids:
node = prompt_obj.get(str(node_id)) or {}
inputs = node.get("inputs", {})
class_type = node.get("class_type") if isinstance(node, dict) else None
for input_name, value in inputs.items():
if not (isinstance(value, list) and len(value) == 2):
continue
source_node = prompt_obj.get(str(value[0]))
if _is_delegate_master_always_retained_upstream_node(source_node):
connected.add(str(value[0]))
continue
if not _delegate_master_input_is_safe_scalar(class_type, input_name):
continue
branch = _collect_delegate_master_retained_upstream_branch(
prompt_obj,
value[0],
value[1],
memo,
set(),
)
if branch is not None:
connected.update(branch)
return connected
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_prompt_for_worker(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)
@@ -349,47 +358,30 @@ 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:
original_node = prompt_obj.get(str(dist_id), {})
class_type = original_node.get("class_type")
inputs = original_node.get("inputs", {})
image_connected = class_type != "DistributedCollector" or (
isinstance(inputs.get("images"), list)
and len(inputs["images"]) == 2
)
audio_connected = (
class_type == "DistributedCollector"
and isinstance(inputs.get("audio"), list)
and len(inputs["audio"]) == 2
)
if image_connected:
preview_id = next_id()
pruned_prompt[preview_id] = {
"inputs": {"images": [dist_id, 0]},
"class_type": "PreviewImage",
"_meta": {"title": "Preview Image (auto-added)"},
}
elif audio_connected:
preview_id = next_id()
pruned_prompt[preview_id] = {
"inputs": {"audio": [dist_id, 1]},
"class_type": "PreviewAudio",
"_meta": {"title": "Preview Audio (auto-added)"},
}
preview_id = next_id()
pruned_prompt[preview_id] = {
"inputs": {
"images": [dist_id, 0],
},
"class_type": "PreviewImage",
"_meta": {
"title": "Preview Image (auto-added)",
},
}
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)
nodes_to_keep.update(downstream)
nodes_to_keep.update(
_find_delegate_master_retained_upstream_nodes(prompt_obj, nodes_to_keep)
)
pruned_prompt = {}
for node_id in nodes_to_keep:
@@ -417,9 +409,16 @@ def prepare_delegate_master_prompt(prompt_obj, collector_ids):
collector_entry = pruned_prompt.get(collector_id)
if not collector_entry:
continue
original_inputs = (prompt_obj.get(collector_id) or {}).get("inputs", {})
original_images = original_inputs.get("images")
if not (isinstance(original_images, list) and len(original_images) == 2):
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] = {
@@ -441,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,
@@ -476,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,
@@ -566,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
+245 -165
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
@@ -13,11 +12,13 @@ from ..utils.constants import (
ORCHESTRATION_WORKER_PREP_CONCURRENCY,
)
from ..utils.logging import debug_log, log
from ..utils.network import build_master_url, build_master_callback_url
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,
)
@@ -144,7 +133,6 @@ async def _prepare_worker_payload(
enabled_ids,
job_id_map,
master_url,
config,
delegate_master,
trace_execution_id,
worker_prep_semaphore,
@@ -154,11 +142,6 @@ async def _prepare_worker_payload(
"""Prepare one worker prompt payload with bounded concurrency and media-sync timeout."""
async with worker_prep_semaphore:
worker_prompt = prompt_index.copy_prompt()
worker_master_url = build_master_callback_url(
worker,
config=config,
prompt_server_instance=prompt_server,
)
worker_type = str(worker.get("type") or "local").strip().lower()
is_remote_like = bool(worker.get("host")) and worker_type != "local"
@@ -173,7 +156,7 @@ async def _prepare_worker_payload(
worker["id"],
enabled_ids,
job_id_map,
worker_master_url,
master_url,
delegate_master,
prompt_index,
)
@@ -197,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, int, dict]: (prompt_id, number, worker_count, node_errors)
"""
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,
@@ -262,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",
@@ -287,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(
@@ -296,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,
@@ -310,29 +246,47 @@ async def orchestrate_distributed_execution(
active_workers = []
delegate_master = False
enabled_ids = [worker["id"] for worker in active_workers]
discovery_prefix = f"exec_{int(time.time() * 1000)}_{uuid.uuid4().hex[:6]}"
job_id_map = generate_job_id_map(prompt_index, discovery_prefix)
if not job_id_map:
trace_debug(execution_trace_id, "No distributed nodes detected; queueing prompt on master only.")
queue_result = await queue_prompt_payload(
prompt_obj,
workflow_meta,
client_id,
include_queue_metadata=True,
)
return (
queue_result["prompt_id"],
queue_result["number"],
0,
queue_result.get("node_errors", {}),
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,
)
return active_workers, delegate_master
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,
@@ -344,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(
@@ -364,55 +470,29 @@ 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,
config,
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
]
)
queue_result = await queue_prompt_payload(
master_prompt,
workflow_meta,
client_id,
include_queue_metadata=True,
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,
)
prompt_id = queue_result["prompt_id"]
prompt_number = queue_result["number"]
node_errors = queue_result.get("node_errors", {})
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(
execution_trace_id,
f"Orchestration complete: prompt_id={prompt_id}, dispatched_workers={len(worker_payloads)}, delegate_master={delegate_master}",
)
return prompt_id, prompt_number, len(worker_payloads), node_errors
return prompt_id, len(worker_payloads)
+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})
+109 -61
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
@@ -16,35 +19,33 @@ from ..utils.logging import debug_log, log
from ..utils.network import (
build_worker_url,
get_client_session,
get_server_port,
handle_api_error,
normalize_host,
probe_worker,
)
from ..utils.constants import CHUNK_SIZE
from ..workers import get_worker_manager
from ..workers.ports import allocate_worker_ports
from .schemas import require_fields, validate_worker_id
from .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)
@@ -115,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()
@@ -126,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:
@@ -141,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()
@@ -153,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:
@@ -164,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):
@@ -181,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
@@ -195,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)
@@ -236,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).
@@ -252,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,
@@ -276,16 +298,9 @@ 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()
device_count = physical_device_count if physical_device_count > 0 else cuda_device_count
master_port = get_server_port()
config = load_config()
worker_ports = []
if not config.get("settings", {}).get("has_auto_populated_workers") and not config.get("workers"):
worker_count = device_count - (1 if cuda_device is not None and 0 <= cuda_device < device_count else 0)
worker_ports = allocate_worker_ports(master_port, config.get("workers", []), worker_count)
hostname = socket.gethostname()
all_ips = get_network_ips()
recommended_ip = get_recommended_ip(all_ips)
@@ -294,13 +309,11 @@ def _collect_network_info_sync():
"all_ips": all_ips,
"recommended_ip": recommended_ip,
"cuda_device": cuda_device,
"cuda_device_count": device_count,
"master_port": master_port,
"local_worker_ports": worker_ports,
"cuda_device_count": physical_device_count if physical_device_count > 0 else cuda_device_count,
}
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)
@@ -336,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)
@@ -348,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:
@@ -357,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:
@@ -402,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)
@@ -417,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
@@ -441,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()
@@ -454,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)
@@ -474,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)
@@ -494,8 +526,6 @@ async def launch_worker_endpoint(request):
"message": f"Worker {worker['name']} launched",
"log_file": log_file
})
except ValueError as e:
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 400)
except Exception as e:
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 500)
@@ -504,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()
@@ -515,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)
@@ -534,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({
@@ -547,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 = {}
@@ -617,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']
@@ -660,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()
@@ -684,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 .nodes.v3 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",
+15
View File
@@ -0,0 +1,15 @@
"""Compatibility facade for legacy imports from `distributed`."""
from __future__ import annotations
from .bootstrap.entrypoint import (
build_node_mappings,
initialize_runtime,
)
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = build_node_mappings()
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"initialize_runtime",
]
+9 -25
View File
@@ -19,8 +19,8 @@ This document describes the **public HTTP API** added to ComfyUI-Distributed to
- `POST /distributed/queue` — queues a workflow using the same distributed orchestration rules as the UI:
- Detects distributed nodes in the prompt (`DistributedCollector`, `UltimateSDUpscaleDistributed`).
- Resolves enabled/selected workers.
- Probes and dispatches workers through `/distributed/worker_ws` by default.
- If `settings.websocket_orchestration=false`, probes with `GET /prompt` and dispatches with `POST /prompt` instead.
- Pings workers (`GET /prompt`) to include only reachable ones.
- Dispatches the workflow to workers (`POST /prompt`).
- Queues the master workflow in ComfyUI’s prompt queue.
- If any `DistributedCollector` has `load_balance=true`, selects one least-busy participant for this run.
@@ -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"
}
```
@@ -61,8 +60,7 @@ Queue a workflow for distributed execution.
#### Fields
- `prompt` (required unless `workflow.prompt` is present, object)
- A complete ComfyUI API-format prompt graph, using the same shape as `POST /prompt`.
- This is not the normal visual workflow export from the ComfyUI editor.
- The ComfyUI prompt/workflow graph, same shape as used by `POST /prompt`.
- `workflow` (optional, object)
- Workflow metadata that ComfyUI normally stores in `extra_pnginfo.workflow`.
- If you don’t care about UI metadata, you can omit it.
@@ -75,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>]`.
@@ -106,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
}
```
@@ -124,14 +117,10 @@ $cfg.workers | Select-Object id,name,enabled,host,port,type | Format-Table -Auto
## Worker requirements (important)
For a worker to participate, it must be reachable from the master. By default:
- WebSocket probe and dispatch: `<worker-base>/distributed/worker_ws` must accept the connection.
If `settings.websocket_orchestration=false`:
For a worker to participate, it must be reachable from the master:
- Health check: `GET <worker-base>/prompt` must return HTTP 200.
- Dispatch: `POST <worker-base>/prompt` must accept the prompt.
- Dispatch: `POST <worker-base>/prompt` must accept the workflow.
Also, for collector-based flows:
@@ -206,7 +195,7 @@ Proxy endpoint on master that fetches logs from a configured remote/cloud worker
## Examples
### 1) Minimal `curl` request envelope
### 1) Minimal `curl`
```bash
curl -X POST "http://127.0.0.1:8188/distributed/queue" \
@@ -214,23 +203,18 @@ curl -X POST "http://127.0.0.1:8188/distributed/queue" \
-d @payload.json
```
`payload.json` must contain a complete ComfyUI API-format prompt. The abbreviated envelope below illustrates the request shape but is not directly executable:
Where `payload.json` contains at least:
```json
{
"prompt": {
"<node_id>": {
"class_type": "<node class>",
"inputs": {"<required_input>": "<value or connection>"}
}
"1": {"class_type": "KSampler", "inputs": {} }
},
"enabled_worker_ids": [],
"client_id": "external-client"
}
```
Export or construct a valid API-format prompt with all required node inputs and at least one output node before submitting it.
### 2) Python (`requests`)
```python
-53
View File
@@ -1,53 +0,0 @@
# Native V3 node API
The package registers one `ComfyExtension` through `comfy_entrypoint()` and
imports the versioned `comfy_api.v0_0_2` API. Use a ComfyUI version that provides
that API; there is no V1 registration fallback.
## Compatibility
- All eight node IDs, display names, categories, visible input order, defaults
and output order are retained. Existing execution algorithms remain in the
private collector, utilities and upscale modules.
- The image/audio dividers explicitly declare ten typed outputs. Their existing
frontend extensions still show only the selected number of outputs. This
replaces the V1 `ByPassTypeTuple` indexing workaround without changing saved
IMAGE/AUDIO socket indices or the ten returned values.
- Standard hidden context uses `cls.hidden`. Worker/orchestration metadata keeps
its existing prompt-input names and Python defaults via `accept_all_inputs`;
it does not become visible widgets.
- The collector retains list-input handling. Upscale retains its always-changing
fingerprint and creates a private runtime helper for each V3 execution, so
mutable helper state is not attached to sanitized V3 class clones.
- Routes, distributed state and the existing worker startup/shutdown hooks are
initialized by `runtime/bootstrap.py` during `ComfyExtension.on_load()`.
Worker mode still suppresses automatic worker launch. `distributed.py` is no
longer an entrypoint.
## Verification
Run ordinary unit tests from this repository:
```bash
python -m pytest tests -q -o addopts=
```
Opt into real-framework acceptance using a ComfyUI checkout and its interpreter:
```bash
COMFYUI_SOURCE_ROOT=/path/to/ComfyUI \
/path/to/ComfyUI/.venv/bin/python -m pytest tests -q -o addopts=
```
The acceptance subprocess uses CPU mode and the real ComfyUI loader, input
parser, V3 class preparation, prompt validation and `PromptExecutor`. It compares
all node schemas with the V1 fixture, checks all five bundled workflows' node
IDs and socket/link contracts, exercises injected worker metadata, image/audio
lists, collector aggregation and divider outputs, rejects invalid upscale enums,
and decodes an actual preview PNG referenced by executor history.
The upscale GPU/model boundary is mocked to check argument forwarding and
per-execution helper isolation. Full checkpoint inference, browser canvas
acceptance and multi-host HTTP transport are not exercised. No HTTP listener,
workers or model downloads are started; preview files use a temporary scratch
directory. The test does not install this branch into a live custom-node folder.
+2 -2
View File
@@ -79,7 +79,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
📺 [Watch Tutorial](https://www.youtube.com/watch?v=wxKKWMQhYTk)
**On Runpod:**
> If using your own template, launch ComfyUI with `--listen --enable-cors-header` and clone `ComfyUI-Distributed` into `custom_nodes`. ⚠️ **Required!**
> If using your own template, make sure you launch ComfyUI with the `--enable-cors-header` argument and you `git clone ComfyUI-Distributed` into custom_nodes. ⚠️ **Required!**
1. Register a [Runpod](https://get.runpod.io/0bw29uf3ug0p) account.
2. On Runpod, go to Storage > New Network Volume and create a volume that will store the models you need. Start with 40 GB, you can always add more later. Learn more [about Network Volumes](https://docs.runpod.io/pods/storage/create-network-volumes).
@@ -92,7 +92,7 @@ The master can either contribute GPU work or stay in **orchestrator-only** mode:
- SAGE_ATTENTION: optional optimisation (set to true/false)
5. Deploy your pod.
6. Connect to your pod using JupyterLabs. This gives us access to the pod's file system.
7. Download models into `/workspace/ComfyUI/models/` (these will remain on your network drive even after you terminate the pod). Example commands below:
7. Download models into /workspaces/ComfyUI/models/ (these will remain on your network drive even after you terminate the pod). Example commands below:
```
# Download from CivitAI
comfy model download --url https://civitai.com/api/download/models/1759168 --relative-path /workspace/ComfyUI/models/checkpoints --set-civitai-api-token $CIVITAI_API_TOKEN
+37 -1
View File
@@ -1 +1,37 @@
"""Private execution helpers; public registration lives in nodes.v3."""
from .utilities import (
DistributedSeed,
DistributedModelName,
DistributedValue,
ImageBatchDivider,
AudioBatchDivider,
DistributedEmptyImage,
AnyType,
ByPassTypeTuple,
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,
"ImageBatchDivider": ImageBatchDivider,
"AudioBatchDivider": AudioBatchDivider,
"DistributedEmptyImage": DistributedEmptyImage,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DistributedCollector": "Distributed Collector",
"DistributedBranch": "Distributed Branch",
"DistributedBranchCollector": "Distributed Branch Collector",
"DistributedSeed": "Distributed Seed",
"DistributedModelName": "Distributed Model Name",
"DistributedValue": "Distributed Value",
"ImageBatchDivider": "Image Batch Divider",
"AudioBatchDivider": "Audio Batch Divider",
"DistributedEmptyImage": "Distributed Empty Image",
}
+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)
+338 -361
View File
@@ -1,34 +1,52 @@
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:
INPUT_IS_LIST = True
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",),
"load_balance": (
"BOOLEAN",
{
@@ -37,171 +55,102 @@ class DistributedCollectorNode:
},
),
},
"optional": {
"images": ("IMAGE",),
"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}),
},
"optional": { "audio": ("AUDIO",) },
"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"
@staticmethod
def _unwrap_list_input(value):
"""Unwrap scalar inputs when ComfyUI passes them via INPUT_IS_LIST."""
if isinstance(value, (list, tuple)) and len(value) == 1:
return value[0]
return value
def _normalize_images_input(self, images):
"""Collapse ComfyUI list IMAGE inputs into a normal batched IMAGE tensor."""
if isinstance(images, (list, tuple)):
if not images:
raise ValueError("Collector received an empty image list")
if not all(isinstance(image, torch.Tensor) for image in images):
raise TypeError("Collector expected IMAGE list items to be torch.Tensor instances")
if len(images) == 1:
return ensure_contiguous(images[0])
return ensure_contiguous(torch.cat([ensure_contiguous(image) for image in images], dim=0))
return ensure_contiguous(images)
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
def _normalize_audio_input(self, audio):
"""Collapse ComfyUI list AUDIO inputs into a single AUDIO payload when present."""
if not isinstance(audio, (list, tuple)):
return audio
audio_items = [item for item in audio if item is not None]
if not audio_items:
return None
if len(audio_items) == 1:
return audio_items[0]
waveforms = []
sample_rate = 44100
for item in audio_items:
if not isinstance(item, dict):
raise TypeError("Collector expected AUDIO list items to be dictionaries")
waveform = item.get("waveform")
if waveform is None or waveform.numel() == 0:
continue
waveforms.append(waveform)
if sample_rate == 44100:
sample_rate = item.get("sample_rate", 44100)
if not waveforms:
return None
return {"waveform": torch.cat(waveforms, dim=-1), "sample_rate": sample_rate}
def run(self, images=None, load_balance=False, audio=None, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", pass_through=False, delegate_only=False):
if images is not None:
images = self._normalize_images_input(images)
audio = self._normalize_audio_input(audio)
load_balance = self._unwrap_list_input(load_balance)
multi_job_id = self._unwrap_list_input(multi_job_id)
is_worker = self._unwrap_list_input(is_worker)
master_url = self._unwrap_list_input(master_url)
enabled_worker_ids = self._unwrap_list_input(enabled_worker_ids)
worker_batch_size = self._unwrap_list_input(worker_batch_size)
worker_id = self._unwrap_list_input(worker_id)
pass_through = self._unwrap_list_input(pass_through)
delegate_only = self._unwrap_list_input(delegate_only)
remote_only_master = (
bool(multi_job_id)
and not is_worker
and (delegate_only or is_master_delegate_only())
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,
)
if images is None and audio is None and not remote_only_master:
raise ValueError("DistributedCollector requires at least one image or audio input")
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):
"""Send an image batch, optionally with audio, or an audio-only completion."""
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:
return
encoded_audio = encode_audio_payload(audio)
session = await get_client_session()
url = f"{master_url}/distributed/job_complete"
for batch_idx in range(batch_size):
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
byte_io = io.BytesIO()
img.save(byte_io, format='PNG', compress_level=0)
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
payload = {
"job_id": str(multi_job_id),
"worker_id": str(worker_id),
"batch_idx": int(batch_idx),
"image": f"data:image/png;base64,{encoded_image}",
"is_last": bool(batch_idx == batch_size - 1),
}
if payload["is_last"] and encoded_audio is not None:
payload["audio"] = encoded_audio
payloads = []
batch_size = 0 if image_batch is None else image_batch.shape[0]
if batch_size == 0:
if encoded_audio is None:
raise ValueError("Worker completion requires image or audio data")
payloads.append(
{
"job_id": str(multi_job_id),
"worker_id": str(worker_id),
"batch_idx": 0,
"audio": encoded_audio,
"is_last": True,
}
)
else:
for batch_idx in range(batch_size):
img = tensor_to_pil(image_batch[batch_idx:batch_idx+1], 0)
byte_io = io.BytesIO()
img.save(byte_io, format='PNG', compress_level=0)
encoded_image = base64.b64encode(byte_io.getvalue()).decode('utf-8')
payload = {
"job_id": str(multi_job_id),
"worker_id": str(worker_id),
"batch_idx": int(batch_idx),
"image": f"data:image/png;base64,{encoded_image}",
"is_last": bool(batch_idx == batch_size - 1),
}
if payload["is_last"] and encoded_audio is not None:
payload["audio"] = encoded_audio
payloads.append(payload)
for payload in payloads:
timeout_seconds = 60 if "image" in payload else 600
try:
async with session.post(
url,
json=payload,
timeout=aiohttp.ClientTimeout(total=timeout_seconds),
headers=distributed_auth_headers(load_config()),
timeout=aiohttp.ClientTimeout(total=60),
) as response:
response.raise_for_status()
except Exception as e:
media_type = "image/audio" if "image" in payload else "audio-only"
log(f"Worker - Failed to send canonical {media_type} envelope to master: {e}")
log(f"Worker - Failed to send canonical image envelope to master: {e}")
debug_log(f"Worker - Full error details: URL={url}")
raise
raise # Re-raise to handle at caller level
def _combine_audio(self, master_audio, worker_audio, empty_audio, worker_order=None):
"""Combine audio from master and workers into a single audio output.
@@ -283,8 +232,8 @@ class DistributedCollectorNode:
images_on_cpu,
delegate_mode: bool,
fallback_images,
):
"""Assemble final tensor, or return None when the job contains only audio."""
) -> torch.Tensor:
"""Assemble final tensor: master first, then workers in enabled order."""
ordered_tensors = []
if not delegate_mode and images_on_cpu is not None:
for i in range(master_batch_size):
@@ -315,242 +264,270 @@ class DistributedCollectorNode:
if cpu_tensors:
return ensure_contiguous(torch.cat(cpu_tensors, dim=0))
if fallback_images is not None:
elif fallback_images is not None:
return ensure_contiguous(fallback_images)
return None
async def execute(self, images, audio, load_balance=False, multi_job_id="", is_worker=False, master_url="", enabled_worker_ids="[]", worker_batch_size=1, worker_id="", delegate_only=False):
if is_worker:
# Worker mode: send images and audio to master in a single batch
image_count = 0 if images is None else images.shape[0]
debug_log(f"Worker - Job {multi_job_id} complete. Sending {image_count} image(s) to master")
await self.send_batch_to_master(images, audio, multi_job_id, master_url, worker_id)
return (images, audio if audio is not None else self.EMPTY_AUDIO)
else:
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)
raise ValueError("No image data collected from master or workers")
# 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:
if images is None:
images_on_cpu = None
master_batch_size = 0
else:
images_on_cpu = ensure_contiguous(images.cpu())
master_batch_size = images.shape[0]
master_audio = audio # Keep master's audio for later
debug_log(f"Master - Job {multi_job_id}: Master has {master_batch_size} images, collecting from {num_workers} workers...")
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})")
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)
# 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()
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
# NEW: Initialize progress bar for workers (total = num_workers)
p = ProgressBar(num_workers)
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
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
combined_audio = self._combine_audio(master_audio, worker_audio, self.EMPTY_AUDIO, enabled_workers)
try:
combined = self._reorder_and_combine_tensors(
worker_images, enabled_workers, master_batch_size, images_on_cpu, delegate_mode, images
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}"
)
if combined is None:
debug_log(f"Master - Job {multi_job_id} complete with audio only")
else:
debug_log(f"Master - Job {multi_job_id} complete. Combined {combined.shape[0]} images total "
f"(master: {master_batch_size}, workers: {combined.shape[0] - master_batch_size})")
return
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)
return (combined, combined_audio)
except Exception as e:
log(f"Master - Error combining images: {e}")
# Preserve collected audio even when image assembly fails.
return (images, combined_audio)
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]
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)),
}
+335 -88
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,
@@ -56,22 +48,22 @@ class UltimateSDUpscaleDistributed(
"""
Distributed version of Ultimate SD Upscale (No Upscale).
Supports two currently selected processing modes:
Supports three processing modes:
1. Single GPU: No workers available, process everything locally
2. Distributed tile queue: Workers pull tiles from a shared queue
2. Static Mode: Small batches, distributes tiles across workers (flattened)
3. Dynamic Mode: Large batches, assigns whole images to workers dynamically
Features:
- Tile-based batch handling for video/image upscaling
- Multi-mode batch handling for efficient video/image upscaling
- Tiled VAE support for memory efficiency
- Shared work queue so faster workers can process more tiles
- Dynamic load balancing for large batches
- Backward compatible with single-image workflows
Environment Variables:
- COMFYUI_MAX_BATCH: Chunk size for tile sending (default 20)
- COMFYUI_MAX_PAYLOAD_SIZE: Max API payload bytes (default 50MB)
The hidden dynamic_threshold input is retained for workflow compatibility but
does not affect the current mode-selection policy.
Threshold: dynamic_threshold input controls mode switch (default 8)
"""
def __init__(self):
@@ -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,47 +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
+75 -50
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]
@@ -269,10 +294,10 @@ class ImageBatchDivider:
class AudioBatchDivider:
"""Divides an audio waveform into sequential segments along the time dimension."""
"""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",),
@@ -282,7 +307,7 @@ class AudioBatchDivider:
"max": 10,
"step": 1,
"display": "number",
"tooltip": "Number of sequential time segments to create"
"tooltip": "Number of parts to divide the audio into"
}),
}
}
@@ -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)
-181
View File
@@ -1,181 +0,0 @@
"""Native V3 schemas with the existing execution algorithms kept intact.
Runtime objects are private, per-execution helpers, not registered V1 nodes.
This avoids sharing mutable instance state through V3's sanitized class clones.
"""
import comfy.samplers
from comfy_api.v0_0_2 import io
from . import utilities as _utilities
from .collector import DistributedCollectorNode as _CollectorRuntime
from .distributed_upscale import UltimateSDUpscaleDistributed as _UpscaleRuntime
class DistributedSeed(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='DistributedSeed', display_name='Distributed Seed', category='utils',
inputs=[io.Int.Input('seed', default=1125899906842, min=0,
max=1125899906842624, force_input=False)],
outputs=[io.Int.Output(display_name='seed')],
accept_all_inputs=True,
)
@classmethod
def execute(cls, seed, is_worker=False, worker_id=''):
return io.NodeOutput(*_utilities.DistributedSeed().distribute(seed, is_worker, worker_id))
class DistributedValue(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='DistributedValue', display_name='Distributed Value', category='utils',
inputs=[io.String.Input('default_value', default=''),
io.String.Input('worker_values', default='{}')],
outputs=[io.AnyType.Output(display_name='value')],
accept_all_inputs=True,
)
@classmethod
def execute(cls, default_value, worker_values='{}', is_worker=False, worker_id=''):
return io.NodeOutput(*_utilities.DistributedValue().distribute(
default_value, worker_values, is_worker, worker_id))
class DistributedModelName(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='DistributedModelName', display_name='Distributed Model Name', category='utils',
inputs=[io.String.Input('text', default='')],
outputs=[io.AnyType.Output(display_name='output')],
hidden=[io.Hidden.unique_id, io.Hidden.extra_pnginfo], is_output_node=True,
)
@classmethod
def execute(cls, text):
result = _utilities.DistributedModelName().log_input(
text, unique_id=cls.hidden.unique_id, extra_pnginfo=cls.hidden.extra_pnginfo)
return io.NodeOutput(*result['result'], ui=result['ui'])
class ImageBatchDivider(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='ImageBatchDivider', display_name='Image Batch Divider', category='image',
inputs=[io.Image.Input('images'),
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
display_mode=io.NumberDisplay.number,
tooltip='Number of parts to divide the batch into')],
# The existing frontend still displays only divide_by sockets.
outputs=[io.Image.Output(display_name=f'batch_{index + 1}') for index in range(10)],
is_output_node=True,
)
@classmethod
def execute(cls, images, divide_by):
return io.NodeOutput(*_utilities.ImageBatchDivider().divide_batch(images, divide_by))
class AudioBatchDivider(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='AudioBatchDivider', display_name='Audio Segment Divider', category='audio',
inputs=[io.Audio.Input('audio'),
io.Int.Input('divide_by', default=2, min=1, max=10, step=1,
display_mode=io.NumberDisplay.number,
tooltip='Number of sequential time segments to create')],
outputs=[io.Audio.Output(display_name=f'audio_{index + 1}') for index in range(10)],
is_output_node=True,
)
@classmethod
def execute(cls, audio, divide_by):
return io.NodeOutput(*_utilities.AudioBatchDivider().divide_audio(audio, divide_by))
class DistributedEmptyImage(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='DistributedEmptyImage', display_name='Distributed Empty Image', category='image',
inputs=[io.Int.Input('height', default=64, min=1, max=4096, step=1),
io.Int.Input('width', default=64, min=1, max=4096, step=1),
io.Int.Input('channels', default=3, min=1, max=4, step=1)],
outputs=[io.Image.Output()],
)
@classmethod
def execute(cls, height, width, channels):
return io.NodeOutput(*_utilities.DistributedEmptyImage().create(height, width, channels))
class DistributedCollector(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='DistributedCollector', display_name='Distributed Collector', category='image',
inputs=[io.Boolean.Input('load_balance', default=False,
tooltip='Run this workflow on one least-busy participant (master included when participating).'),
io.Image.Input('images', optional=True), io.Audio.Input('audio', optional=True)],
outputs=[io.Image.Output(display_name='images'), io.Audio.Output(display_name='audio')],
is_input_list=True, accept_all_inputs=True,
)
@classmethod
def execute(cls, images=None, load_balance=False, audio=None, multi_job_id='',
is_worker=False, master_url='', enabled_worker_ids='[]', worker_batch_size=1,
worker_id='', pass_through=False, delegate_only=False):
return io.NodeOutput(*_CollectorRuntime().run(
images=images, load_balance=load_balance, audio=audio, multi_job_id=multi_job_id,
is_worker=is_worker, master_url=master_url, enabled_worker_ids=enabled_worker_ids,
worker_batch_size=worker_batch_size, worker_id=worker_id,
pass_through=pass_through, delegate_only=delegate_only))
class UltimateSDUpscaleDistributed(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id='UltimateSDUpscaleDistributed',
display_name='Ultimate SD Upscale Distributed (No Upscale)', category='image/upscaling',
inputs=[
io.Image.Input('upscaled_image'), io.Model.Input('model'),
io.Conditioning.Input('positive'), io.Conditioning.Input('negative'), io.Vae.Input('vae'),
io.Int.Input('seed', default=0, min=0, max=0xffffffffffffffff),
io.Int.Input('steps', default=20, min=1, max=10000),
io.Float.Input('cfg', default=8.0, min=0.0, max=100.0),
io.Combo.Input('sampler_name', options=comfy.samplers.KSampler.SAMPLERS),
io.Combo.Input('scheduler', options=comfy.samplers.KSampler.SCHEDULERS),
io.Float.Input('denoise', default=0.5, min=0.0, max=1.0, step=0.01),
io.Int.Input('tile_width', default=512, min=64, max=2048, step=8),
io.Int.Input('tile_height', default=512, min=64, max=2048, step=8),
io.Int.Input('padding', default=32, min=0, max=256, step=8),
io.Int.Input('mask_blur', default=8, min=0, max=256),
io.Boolean.Input('force_uniform_tiles', default=True),
io.Boolean.Input('tiled_decode', default=False),
], outputs=[io.Image.Output()], accept_all_inputs=True,
)
@classmethod
def fingerprint_inputs(cls, **kwargs):
return _UpscaleRuntime.IS_CHANGED(**kwargs)
@classmethod
def execute(cls, 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):
return io.NodeOutput(*_UpscaleRuntime().run(
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,
master_url, enabled_worker_ids, worker_id, tile_indices, dynamic_threshold))
NODES = [DistributedCollector, DistributedSeed, DistributedModelName, DistributedValue,
ImageBatchDivider, AudioBatchDivider, DistributedEmptyImage, UltimateSDUpscaleDistributed]
+5 -2
View File
@@ -1,9 +1,12 @@
[project]
name = "ComfyUI-Distributed"
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
version = "1.5.0"
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"
-1
View File
@@ -1 +0,0 @@
"""Internal extension lifecycle support (not a node provider)."""
-35
View File
@@ -1,35 +0,0 @@
"""Initialize routes, distributed state and the existing worker lifecycle."""
import atexit
import os
import server
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
from ..upscale.job_store import ensure_tile_jobs_initialized
_initialized = False
def initialize():
"""Called by ComfyExtension.on_load; initialize once per loaded package."""
global _initialized
if _initialized:
return
ensure_config_exists()
from .. import api # noqa: F401 - registers the existing @routes.* handlers
from ..api.queue_orchestration import ensure_distributed_state
ensure_distributed_state(server.PromptServer.instance)
ensure_tile_jobs_initialized()
if not os.environ.get('COMFYUI_IS_WORKER'):
atexit.register(sync_cleanup)
delayed_auto_launch()
register_async_signals()
_initialized = True
debug_log('Loaded Distributed nodes.')
debug_log(f'Config file: {CONFIG_FILE}')
+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,12 +132,53 @@ 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
queue_orchestration_module = types.ModuleType(f"{package_name}.api.queue_orchestration")
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 7, 1, {}))
queue_orchestration_module.orchestrate_distributed_execution = AsyncMock(return_value=("prompt_dist", 1))
sys.modules[f"{package_name}.api.queue_orchestration"] = queue_orchestration_module
@dataclass(frozen=True)
@@ -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
@@ -221,7 +233,7 @@ job_routes = _load_job_routes_module()
class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
async def test_distributed_queue_happy_path_returns_prompt_metadata(self):
async def test_distributed_queue_happy_path_returns_prompt_id(self):
request = _FakeRequest(
{
"prompt": {"1": {"class_type": "Node"}},
@@ -233,15 +245,12 @@ class DistributedQueueEndpointTests(unittest.IsolatedAsyncioTestCase):
with patch.object(
job_routes,
"orchestrate_distributed_execution",
new=AsyncMock(return_value=("prompt_123", 42, 2, {})),
new=AsyncMock(return_value=("prompt_123", 2)),
):
response = await job_routes.distributed_queue_endpoint(request)
self.assertEqual(response.status, 200)
self.assertEqual(response.payload.get("prompt_id"), "prompt_123")
self.assertEqual(response.payload.get("number"), 42)
self.assertEqual(response.payload.get("node_errors"), {})
self.assertTrue(response.payload.get("auto_prepare_supported"))
async def test_distributed_queue_missing_prompt_returns_400(self):
request = _FakeRequest(
@@ -278,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",
@@ -302,76 +313,27 @@ class JobCompleteAudioPayloadTests(unittest.IsolatedAsyncioTestCase):
self.assertEqual(queued["audio"]["sample_rate"], 44100)
self.assertEqual(tuple(queued["audio"]["waveform"].shape), (1, 2, 4))
async def test_job_complete_accepts_audio_without_image(self):
async def test_job_complete_rejects_worker_not_in_job_allowlist(self):
queue = asyncio.Queue()
job_routes.prompt_server.distributed_jobs_lock = asyncio.Lock()
job_routes.prompt_server.distributed_pending_jobs = {"audio-only-job": 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": "audio-only-job",
"worker_id": "worker-1",
"job_id": "job-allow",
"worker_id": "worker-unexpected",
"batch_idx": 0,
"audio": self._encoded_audio_payload(),
"image": "data:image/png;base64,AAAA",
"is_last": True,
}
)
with patch.object(job_routes, "_decode_canonical_png_tensor") as decode_image:
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, 200)
decode_image.assert_not_called()
queued = await queue.get()
self.assertIsNone(queued["tensor"])
self.assertEqual(queued["audio"]["sample_rate"], 44100)
async def test_job_complete_rejects_payload_without_image_or_audio(self):
request = _FakeRequest(
{
"job_id": "job-1",
"worker_id": "worker-1",
"batch_idx": 0,
"is_last": True,
}
)
response = await job_routes.job_complete_endpoint(request)
self.assertEqual(response.status, 400)
self.assertIn("image or audio", response.payload.get("message", "").lower())
async def test_job_complete_rejects_invalid_image_even_when_audio_is_present(self):
request = _FakeRequest(
{
"job_id": "job-1",
"worker_id": "worker-1",
"batch_idx": 0,
"image": 123,
"audio": self._encoded_audio_payload(),
"is_last": True,
}
)
response = await job_routes.job_complete_endpoint(request)
self.assertEqual(response.status, 400)
self.assertIn("image", response.payload.get("message", "").lower())
async def test_job_complete_returns_400_for_malformed_audio(self):
request = _FakeRequest(
{
"job_id": "job-1",
"worker_id": "worker-1",
"batch_idx": 0,
"audio": {"data": "AAAA", "shape": [1, 2], "dtype": "float32"},
"is_last": True,
}
)
response = await job_routes.job_complete_endpoint(request)
self.assertEqual(response.status, 400)
self.assertIn("audio.shape", response.payload.get("message", "").lower())
self.assertEqual(response.status, 403)
self.assertIn("unauthorized", response.payload.get("message", "").lower())
def test_decode_audio_payload_rejects_bad_shape(self):
bad = {
+47 -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
@@ -164,11 +157,26 @@ class ConvertPathsForPlatformTests(unittest.TestCase):
# ---------------------------------------------------------------------------
class FindMediaReferencesTests(unittest.TestCase):
def test_finds_image_input(self):
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.png"}}}
refs = ms._find_media_references(prompt)
self.assertIn("photo.png", refs)
def test_finds_video_input(self):
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"video": "clip.mp4"}}}
refs = ms._find_media_references(prompt)
self.assertIn("clip.mp4", refs)
def test_finds_file_input_for_load_video(self):
prompt = {"1": {"class_type": "LoadVideo", "inputs": {"file": "1 - Copy.mp4"}}}
refs = ms._find_media_references(prompt)
self.assertIn("1 - Copy.mp4", refs)
def test_finds_audio_input(self):
prompt = {"1": {"class_type": "LoadAudio", "inputs": {"audio": "track.wav"}}}
refs = ms._find_media_references(prompt)
self.assertIn("track.wav", refs)
def test_strips_annotation_suffix(self):
prompt = {"1": {"class_type": "LoadImage", "inputs": {"image": "photo.jpg [abc123]"}}}
refs = ms._find_media_references(prompt)
@@ -245,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])
+34 -111
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,24 +98,10 @@ 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__ = [str(module_path.parents[1] / "workers")]
workers_pkg.__path__ = []
workers_pkg.get_worker_manager = lambda: _DummyWorkerManager()
sys.modules[f"{package_name}.workers"] = workers_pkg
@@ -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:
@@ -213,7 +159,6 @@ def _load_worker_routes_module():
raise RuntimeError("not used in these tests")
network_module.get_client_session = _get_client_session
network_module.get_server_port = lambda: 8189
sys.modules[f"{package_name}.utils.network"] = network_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
@@ -249,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
@@ -270,38 +225,6 @@ worker_routes = _load_worker_routes_module()
class WorkerRoutesTests(unittest.IsolatedAsyncioTestCase):
async def test_network_info_reports_actual_master_and_available_ports(self):
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 1, 4)), \
patch.object(worker_routes, "get_network_ips", return_value=["127.0.0.1"]), \
patch.object(worker_routes, "load_config", return_value={"workers": []}), \
patch.object(worker_routes, "get_server_port", return_value=8189), \
patch.object(worker_routes, "allocate_worker_ports", return_value=[8191, 8192, 8193]) as allocate:
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
self.assertEqual(response.payload["master_port"], 8189)
self.assertEqual(response.payload["local_worker_ports"], [8191, 8192, 8193])
self.assertEqual(response.payload["cuda_device_count"], 4)
allocate.assert_called_once_with(8189, [], 3)
async def test_network_info_does_not_reallocate_existing_configuration(self):
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 4, 4)), \
patch.object(worker_routes, "get_network_ips", return_value=[]), \
patch.object(worker_routes, "load_config", return_value={"workers": [{"id": "existing", "port": 8189}]}), \
patch.object(worker_routes, "allocate_worker_ports") as allocate:
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
self.assertEqual(response.payload["local_worker_ports"], [])
allocate.assert_not_called()
async def test_launch_conflict_is_a_clear_client_error(self):
manager = _DummyWorkerManager()
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8189}]}
with patch.object(worker_routes, "get_worker_manager", return_value=manager), \
patch.object(worker_routes, "load_config", return_value=config), \
patch.object(manager, "launch_worker", side_effect=ValueError("Worker port 8189 conflicts with the master port 8189")):
response = await worker_routes.launch_worker_endpoint(_FakeRequest({"worker_id": "worker-a"}))
self.assertEqual(response.status, 400)
self.assertIn("master port 8189", response.payload["message"])
self.assertEqual(manager.processes, {})
async def test_launch_worker_valid_id_returns_200(self):
manager = _DummyWorkerManager()
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8188}]}
-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.
-691
View File
@@ -1,691 +0,0 @@
{
"DistributedCollector": {
"input": {
"required": {
"load_balance": [
"BOOLEAN",
{
"default": false,
"tooltip": "Run this workflow on one least-busy participant (master included when participating)."
}
]
},
"optional": {
"images": [
"IMAGE"
],
"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
}
]
}
},
"input_order": {
"required": [
"load_balance"
],
"optional": [
"images",
"audio"
],
"hidden": [
"multi_job_id",
"is_worker",
"master_url",
"enabled_worker_ids",
"worker_batch_size",
"worker_id",
"pass_through",
"delegate_only"
]
},
"output": [
"IMAGE",
"AUDIO"
],
"output_name": [
"images",
"audio"
],
"output_is_list": [
false,
false
],
"is_input_list": true,
"output_node": false,
"category": "image",
"display_name": "Distributed Collector"
},
"DistributedSeed": {
"input": {
"required": {
"seed": [
"INT",
{
"default": 1125899906842,
"min": 0,
"max": 1125899906842624,
"forceInput": false
}
]
},
"hidden": {
"is_worker": [
"BOOLEAN",
{
"default": false
}
],
"worker_id": [
"STRING",
{
"default": ""
}
]
}
},
"input_order": {
"required": [
"seed"
],
"hidden": [
"is_worker",
"worker_id"
]
},
"output": [
"INT"
],
"output_name": [
"seed"
],
"output_is_list": [
false
],
"is_input_list": false,
"output_node": false,
"category": "utils",
"display_name": "Distributed Seed"
},
"DistributedModelName": {
"input": {
"required": {
"text": [
"STRING",
{
"default": ""
}
]
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO"
}
},
"input_order": {
"required": [
"text"
],
"hidden": [
"unique_id",
"extra_pnginfo"
]
},
"output": [
"*"
],
"output_name": [
"output"
],
"output_is_list": [
false
],
"is_input_list": false,
"output_node": true,
"category": "utils",
"display_name": "Distributed Model Name"
},
"DistributedValue": {
"input": {
"required": {
"default_value": [
"STRING",
{
"default": ""
}
],
"worker_values": [
"STRING",
{
"default": "{}"
}
]
},
"hidden": {
"is_worker": [
"BOOLEAN",
{
"default": false
}
],
"worker_id": [
"STRING",
{
"default": ""
}
]
}
},
"input_order": {
"required": [
"default_value",
"worker_values"
],
"hidden": [
"is_worker",
"worker_id"
]
},
"output": [
"*"
],
"output_name": [
"value"
],
"output_is_list": [
false
],
"is_input_list": false,
"output_node": false,
"category": "utils",
"display_name": "Distributed Value"
},
"ImageBatchDivider": {
"input": {
"required": {
"images": [
"IMAGE"
],
"divide_by": [
"INT",
{
"default": 2,
"min": 1,
"max": 10,
"step": 1,
"display": "number",
"tooltip": "Number of parts to divide the batch into"
}
]
}
},
"input_order": {
"required": [
"images",
"divide_by"
]
},
"output": [
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*"
],
"output_name": [
"batch_1",
"batch_2",
"batch_3",
"batch_4",
"batch_5",
"batch_6",
"batch_7",
"batch_8",
"batch_9",
"batch_10"
],
"output_is_list": [
false,
false,
false,
false,
false,
false,
false,
false,
false,
false
],
"is_input_list": false,
"output_node": true,
"category": "image",
"display_name": "Image Batch Divider"
},
"AudioBatchDivider": {
"input": {
"required": {
"audio": [
"AUDIO"
],
"divide_by": [
"INT",
{
"default": 2,
"min": 1,
"max": 10,
"step": 1,
"display": "number",
"tooltip": "Number of sequential time segments to create"
}
]
}
},
"input_order": {
"required": [
"audio",
"divide_by"
]
},
"output": [
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*",
"*"
],
"output_name": [
"audio_1",
"audio_2",
"audio_3",
"audio_4",
"audio_5",
"audio_6",
"audio_7",
"audio_8",
"audio_9",
"audio_10"
],
"output_is_list": [
false,
false,
false,
false,
false,
false,
false,
false,
false,
false
],
"is_input_list": false,
"output_node": true,
"category": "audio",
"display_name": "Audio Segment Divider"
},
"DistributedEmptyImage": {
"input": {
"required": {
"height": [
"INT",
{
"default": 64,
"min": 1,
"max": 4096,
"step": 1
}
],
"width": [
"INT",
{
"default": 64,
"min": 1,
"max": 4096,
"step": 1
}
],
"channels": [
"INT",
{
"default": 3,
"min": 1,
"max": 4,
"step": 1
}
]
}
},
"input_order": {
"required": [
"height",
"width",
"channels"
]
},
"output": [
"IMAGE"
],
"output_name": [
"IMAGE"
],
"output_is_list": [
false
],
"is_input_list": false,
"output_node": false,
"category": "image",
"display_name": "Distributed Empty Image"
},
"UltimateSDUpscaleDistributed": {
"input": {
"required": {
"upscaled_image": [
"IMAGE"
],
"model": [
"MODEL"
],
"positive": [
"CONDITIONING"
],
"negative": [
"CONDITIONING"
],
"vae": [
"VAE"
],
"seed": [
"INT",
{
"default": 0,
"min": 0,
"max": 18446744073709551615
}
],
"steps": [
"INT",
{
"default": 20,
"min": 1,
"max": 10000
}
],
"cfg": [
"FLOAT",
{
"default": 8.0,
"min": 0.0,
"max": 100.0
}
],
"sampler_name": [
[
"euler",
"euler_cfg_pp",
"euler_ancestral",
"euler_ancestral_cfg_pp",
"heun",
"heunpp2",
"exp_heun_2_x0",
"exp_heun_2_x0_sde",
"dpm_2",
"dpm_2_ancestral",
"lms",
"dpm_fast",
"dpm_adaptive",
"dpmpp_2s_ancestral",
"dpmpp_2s_ancestral_cfg_pp",
"dpmpp_sde",
"dpmpp_sde_gpu",
"dpmpp_2m",
"dpmpp_2m_cfg_pp",
"dpmpp_2m_sde",
"dpmpp_2m_sde_gpu",
"dpmpp_2m_sde_heun",
"dpmpp_2m_sde_heun_gpu",
"dpmpp_3m_sde",
"dpmpp_3m_sde_gpu",
"ddpm",
"lcm",
"ipndm",
"ipndm_v",
"deis",
"cfgpp_ud10_ab",
"res_multistep",
"res_multistep_cfg_pp",
"res_multistep_ancestral",
"res_multistep_ancestral_cfg_pp",
"gradient_estimation",
"gradient_estimation_cfg_pp",
"er_sde",
"seeds_2",
"seeds_3",
"sa_solver",
"sa_solver_pece",
"ddim",
"uni_pc",
"uni_pc_bh2"
]
],
"scheduler": [
[
"simple",
"sgm_uniform",
"karras",
"exponential",
"ddim_uniform",
"beta",
"normal",
"linear_quadratic",
"kl_optimal"
]
],
"denoise": [
"FLOAT",
{
"default": 0.5,
"min": 0.0,
"max": 1.0,
"step": 0.01
}
],
"tile_width": [
"INT",
{
"default": 512,
"min": 64,
"max": 2048,
"step": 8
}
],
"tile_height": [
"INT",
{
"default": 512,
"min": 64,
"max": 2048,
"step": 8
}
],
"padding": [
"INT",
{
"default": 32,
"min": 0,
"max": 256,
"step": 8
}
],
"mask_blur": [
"INT",
{
"default": 8,
"min": 0,
"max": 256
}
],
"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": ""
}
],
"dynamic_threshold": [
"INT",
{
"default": 8,
"min": 1,
"max": 64
}
]
}
},
"input_order": {
"required": [
"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"
],
"hidden": [
"multi_job_id",
"is_worker",
"master_url",
"enabled_worker_ids",
"worker_id",
"tile_indices",
"dynamic_threshold"
]
},
"output": [
"IMAGE"
],
"output_name": [
"IMAGE"
],
"output_is_list": [
false
],
"is_input_list": false,
"output_node": false,
"category": "image/upscaling",
"display_name": "Ultimate SD Upscale Distributed (No Upscale)"
}
}
-252
View File
@@ -1,252 +0,0 @@
"""Exercise the actual loader, schemas and executor in a fresh CPU process.
No HTTP listener, workers, model downloads or live installation changes.
The original contracts were captured from 32ac027 using the same core.
"""
import asyncio
import inspect
import json
import os
from pathlib import Path
import sys
import tempfile
from unittest.mock import patch
from PIL import Image
ROOT = Path(__file__).resolve().parents[2]
COMFY_ROOT = Path(sys.argv[1]).resolve()
sys.path.insert(0, str(COMFY_ROOT))
os.environ['COMFYUI_IS_WORKER'] = '1'
import comfy.cli_args
comfy.cli_args.args.cpu = True
comfy.cli_args.args.disable_assets = True
from app.assets.manager import default_asset_manager
from comfy_api.v0_0_2 import io
import execution
import nodes
import server
import torch
BASELINE = json.loads((ROOT / 'tests/fixtures/v1_node_contracts.json').read_text())
def normalized_input(value):
kind = value[0]
opts = dict(value[1]) if len(value) > 1 else {}
if kind == 'STRING':
# V3 serializes the same default single-line widget explicitly.
opts.setdefault('multiline', False)
if kind == 'COMBO':
kind = opts['options']
opts = {key: val for key, val in opts.items() if key != 'options'}
if isinstance(kind, list):
# Single-selection is the V1 dropdown default as well.
opts.setdefault('multiselect', False)
return [kind, opts]
async def map_node(cls, values, extra=None):
"""Use the actual input parser and sanitized V3 executor clones."""
prepared, missing, hidden = execution.get_input_data(values, cls, unique_id='probe', extra_data=extra or {})
assert not missing, missing
# Use a worker thread for synchronous nodes: collector/upscale bridge back
# to PromptServer's running loop, just as ComfyUI's prompt worker does.
def run():
return asyncio.run(execution._async_map_node_over_list(
'v3-acceptance', 'probe', cls, prepared, cls.FUNCTION, v3_data=hidden))
return await asyncio.to_thread(run)
def check_saved_workflows(mapping):
checked = 0
occurrences = 0
for path in sorted((ROOT / 'workflows').glob('*.json')):
workflow = json.loads(path.read_text())
local_nodes = {node['id']: node for node in workflow['nodes']}
for saved in workflow['nodes']:
if saved['type'] not in mapping:
continue
occurrences += 1
info = mapping[saved['type']].GET_NODE_INFO_V1()
inputs = {**info['input'].get('required', {}), **info['input'].get('optional', {})}
for socket in saved.get('inputs', []):
assert socket['name'] in inputs, (path.name, saved['id'], socket)
assert socket['type'] == inputs[socket['name']][0], (path.name, socket)
for index, output in enumerate(saved.get('outputs', [])):
assert output['type'] == info['output'][index], (path.name, index)
for link in workflow['links']:
if link[1] == saved['id']:
assert link[2] < len(info['output']), (path.name, link)
target = local_nodes[link[3]]['inputs'][link[4]]
assert target['type'] == info['output'][link[2]], (path.name, link)
checked += 1
assert checked == 5 and occurrences == 10, (checked, occurrences)
print('SAVED_WORKFLOW_CONTRACTS_OK', checked, occurrences)
async def check_execution(mapping, module, prompt_server, asset_manager):
result = await map_node(mapping['DistributedSeed'],
{'seed': 123, 'is_worker': True, 'worker_id': 'worker_2'})
assert result[0].result == (126,)
values = json.dumps({'_type': 'INT', '3': '17'})
result = await map_node(mapping['DistributedValue'],
{'default_value': '4', 'worker_values': values,
'is_worker': True, 'worker_id': 'worker_2'})
assert result[0].result == (17,)
metadata = {'workflow': {'nodes': [{'id': 'probe', 'widgets_values': []}]}}
result = await map_node(mapping['DistributedModelName'], {'text': 'model.ckpt'},
{'extra_pnginfo': metadata})
assert result[0].result == ('model.ckpt',)
assert result[0].ui == {'text': ['model.ckpt']}
assert metadata['workflow']['nodes'][0]['widgets_values'] == [['model.ckpt']]
result = await map_node(mapping['DistributedEmptyImage'],
{'height': 8, 'width': 8, 'channels': 3})
empty = result[0].result[0]
assert empty.shape == (0, 8, 8, 3) and empty.numel() == 0
images = torch.arange(10 * 8 * 8 * 3, dtype=torch.float32).reshape(10, 8, 8, 3)
for count in (1, 3, 10):
result = await map_node(mapping['ImageBatchDivider'], {'images': images, 'divide_by': count})
assert len(result[0].result) == 10
assert torch.equal(torch.cat(result[0].result[:count]), images)
audio = {'waveform': torch.arange(33, dtype=torch.float32).reshape(1, 1, 33), 'sample_rate': 24000}
result = await map_node(mapping['AudioBatchDivider'], {'audio': audio, 'divide_by': 3})
assert len(result[0].result) == 10
assert torch.equal(torch.cat([item['waveform'] for item in result[0].result[:3]], dim=-1), audio['waveform'])
# INPUT_IS_LIST must preserve all images/audio but unwrap transport scalars.
collector = mapping['DistributedCollector']
prepared, missing, hidden = execution.get_input_data(
{'load_balance': False, 'multi_job_id': '', 'pass_through': True}, collector, 'collector-list')
assert not missing
prepared['images'] = [images[:2], images[2:]]
prepared['audio'] = [audio, audio]
result = await execution._async_map_node_over_list(
'v3-acceptance', 'collector-list', collector, prepared, collector.FUNCTION, v3_data=hidden)
assert torch.equal(result[0].result[0], images)
assert result[0].result[1]['waveform'].shape[-1] == 66
# Exercise real async aggregation without an HTTP endpoint or worker.
queue = asyncio.Queue()
await queue.put({'worker_id': 'worker_2', 'tensor': images[2:3], 'image_index': 0, 'is_last': True})
prompt_server.distributed_pending_jobs['v3-aggregate'] = queue
result = await map_node(collector,
{'images': images[:2], 'load_balance': False,
'multi_job_id': 'v3-aggregate', 'enabled_worker_ids': '["worker_2"]'})
assert torch.equal(result[0].result[0], images[:3])
assert 'v3-aggregate' not in prompt_server.distributed_pending_jobs
# Prove V3's per-execution class clones do not leak helper instance state;
# replace only the GPU/model boundary, not the input parser or executor.
upscale_runtime = sys.modules[module.__name__ + '.nodes.distributed_upscale'].UltimateSDUpscaleDistributed
seen = []
def fake_upscale(self, *args):
assert not hasattr(self, 'acceptance_marker')
self.acceptance_marker = True
seen.append(args)
return (args[0],)
inputs = {'upscaled_image': images, 'model': object(), 'positive': [], 'negative': [],
'vae': object(), 'seed': 1, 'steps': 1, 'cfg': 1.0,
'sampler_name': 'euler', 'scheduler': 'normal', 'denoise': 0.5,
'tile_width': 64, 'tile_height': 64, 'padding': 0, 'mask_blur': 0,
'force_uniform_tiles': True, 'tiled_decode': False,
'multi_job_id': 'tile-job', 'is_worker': True, 'master_url': 'http://master.invalid',
'enabled_worker_ids': '["worker_2"]', 'worker_id': 'worker_2',
'tile_indices': '[2,3]', 'dynamic_threshold': 9}
with patch.object(upscale_runtime, 'run', fake_upscale):
for _ in range(2):
result = await map_node(mapping['UltimateSDUpscaleDistributed'], inputs)
assert result[0].result[0] is images
assert len(seen) == 2 and seen[0][-7:] == tuple(inputs[name] for name in (
'multi_job_id', 'is_worker', 'master_url', 'enabled_worker_ids', 'worker_id',
'tile_indices', 'dynamic_threshold')), seen[0]
import math
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id='tile-job'))
assert math.isnan(mapping['UltimateSDUpscaleDistributed'].fingerprint_inputs(multi_job_id=''))
graph = {
'1': {'class_type': 'EmptyImage', 'inputs': {'height': 8, 'width': 8, 'batch_size': 10, 'color': 0}},
'2': {'class_type': 'DistributedCollector', 'inputs': {'images': ['1', 0], 'load_balance': False,
'pass_through': True}},
'3': {'class_type': 'ImageBatchDivider', 'inputs': {'images': ['2', 0], 'divide_by': 10}},
'4': {'class_type': 'PreviewImage', 'inputs': {'images': ['3', 9]}},
}
valid = await execution.validate_prompt('v3-graph', graph, None)
assert valid[0], valid
invalid = {'1': {'class_type': 'UltimateSDUpscaleDistributed',
'inputs': {**inputs, 'sampler_name': 'INVALID_ENUM', 'scheduler': 'INVALID_ENUM'}},
'2': {'class_type': 'PreviewImage', 'inputs': {'images': ['1', 0]}}}
rejected = await execution.validate_prompt('v3-invalid', invalid, None)
assert not rejected[0]
errors = [error for entry in rejected[3].values() for error in entry['errors']]
enum_errors = [error['extra_info']['input_name'] for error in errors if error['type'] == 'value_not_in_list']
assert {'sampler_name', 'scheduler'} <= set(enum_errors), errors
import folder_paths
scratch = Path(os.environ.get('TMPDIR', Path.home() / '.hermes/cache/scratch'))
with tempfile.TemporaryDirectory(prefix='v3-preview-', dir=scratch) as temp:
with patch.object(folder_paths, 'temp_directory', temp):
executor = execution.PromptExecutor(
prompt_server, cache_args={'ram': 0, 'ram_inactive': 0}, asset_manager=asset_manager)
await asyncio.to_thread(executor.execute, graph, 'v3-graph', {}, valid[2])
assert executor.success, executor.status_messages
history = executor.history_result
record = history['outputs']['4']['images'][0]
preview = Path(temp) / record.get('subfolder', '') / record['filename']
with Image.open(preview) as image:
assert image.size == (8, 8) and image.mode == 'RGB'
assert image.getextrema() == ((0, 0), (0, 0), (0, 0))
print('V3_EXECUTION_OK eight nodes; upscale GPU boundary mocked; preview artifact verified')
async def main():
asset_manager = default_asset_manager()
prompt_server = server.PromptServer(asyncio.get_running_loop(), asset_manager)
assert not (ROOT / 'distributed.py').exists(), 'obsolete root bootstrap remains'
assert await nodes.load_custom_node(str(ROOT)), 'ComfyUI loader rejected the pack'
module = sys.modules[str(ROOT).replace('.', '_x_')]
assert not hasattr(module, 'NODE_CLASS_MAPPINGS'), 'V1 map shadows V3 entrypoint'
extension = await module.comfy_entrypoint()
classes = await extension.get_node_list()
mapping = {cls.GET_SCHEMA().node_id: cls for cls in classes}
assert len(classes) == len(mapping) == len(BASELINE) == 8
assert set(mapping) == set(BASELINE)
for node_id, cls in mapping.items():
assert issubclass(cls, io.ComfyNode)
assert nodes.NODE_CLASS_MAPPINGS[node_id] is cls
old = dict(BASELINE[node_id])
if node_id in ('ImageBatchDivider', 'AudioBatchDivider'):
# V1's ByPassTypeTuple advertises '*' when indexed, while its
# underlying tuple and existing frontend declare IMAGE/AUDIO.
# Native V3 declares all ten existing typed sockets explicitly.
assert old['output'] == ['*'] * 10
old['output'] = ['IMAGE' if node_id == 'ImageBatchDivider' else 'AUDIO'] * 10
new = json.loads(json.dumps(cls.GET_NODE_INFO_V1()))
for group in ('required', 'optional'):
old_inputs = old['input'].get(group, {})
new_inputs = new['input'].get(group, {})
assert list(old_inputs) == list(new_inputs), (node_id, group, 'input order')
for name, original in old_inputs.items():
assert normalized_input(original) == normalized_input(new_inputs[name]), (node_id, name, original, new_inputs[name])
for key in ('output', 'output_name', 'output_is_list', 'is_input_list', 'output_node', 'category', 'display_name'):
assert old[key] == new[key], (node_id, key, old[key], new[key])
# Standard context lives in cls.hidden; orchestrator metadata remains
# accepted by its original kwarg name, without creating new widgets.
signature = inspect.signature(cls.execute)
for name, field in old['input'].get('hidden', {}).items():
if isinstance(field, list):
assert cls.GET_SCHEMA().accept_all_inputs, node_id
assert name in signature.parameters, (node_id, name)
assert signature.parameters[name].default == field[1]['default'], (node_id, name)
else:
assert name in new['input']['hidden'], (node_id, name)
expected_hidden = []
if node_id == 'DistributedModelName':
expected_hidden.extend(['unique_id', 'extra_pnginfo'])
if old['output_node']:
expected_hidden.extend(name for name in ['prompt', 'extra_pnginfo'] if name not in expected_hidden)
assert list(new['input'].get('hidden', {})) == expected_hidden, (node_id, new['input'].get('hidden'))
print('SCHEMA_PARITY_OK', len(mapping))
check_saved_workflows(mapping)
await check_execution(mapping, module, prompt_server, asset_manager)
print('V3_ACCEPTANCE_OK')
asyncio.run(main())
-364
View File
@@ -1,364 +0,0 @@
import importlib.util
import sys
import types
import asyncio
from pathlib import Path
import torch
def _load_collector_module():
module_path = Path(__file__).resolve().parents[1] / "nodes" / "collector.py"
package_name = "dist_collector_list_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
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
class _Routes:
def post(self, _path):
return lambda fn: fn
def get(self, _path):
return lambda fn: fn
prompt_server = types.SimpleNamespace(
routes=_Routes(),
distributed_jobs_lock=None,
distributed_pending_jobs={},
)
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
sys.modules["server"] = server_module
comfy_module = types.ModuleType("comfy")
model_management = types.ModuleType("comfy.model_management")
class InterruptProcessingException(Exception):
pass
model_management.InterruptProcessingException = InterruptProcessingException
model_management.throw_exception_if_processing_interrupted = lambda: None
comfy_module.model_management = model_management
comfy_utils = types.ModuleType("comfy.utils")
class ProgressBar:
def __init__(self, _total):
self.total = _total
self.updates = []
def update(self, value):
self.updates.append(value)
comfy_utils.ProgressBar = ProgressBar
comfy_module.utils = comfy_utils
sys.modules["comfy"] = comfy_module
sys.modules["comfy.model_management"] = model_management
sys.modules["comfy.utils"] = comfy_utils
aiohttp_module = types.ModuleType("aiohttp")
aiohttp_module.ClientTimeout = lambda total: types.SimpleNamespace(total=total)
sys.modules["aiohttp"] = aiohttp_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
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.get_worker_timeout_seconds = lambda: 0.1
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() if hasattr(tensor, "contiguous") else tensor
image_module.ensure_contiguous = _ensure_contiguous
image_module.tensor_to_pil = lambda *_args, **_kwargs: None
image_module.pil_to_tensor = lambda value: value
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"
network_module.get_client_session = lambda: None
network_module.probe_worker = lambda *_args, **_kwargs: None
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 audio: audio
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
spec = importlib.util.spec_from_file_location(f"{package_name}.nodes.collector", 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
def test_collector_opts_into_comfyui_list_inputs():
collector = _load_collector_module().DistributedCollectorNode
assert collector.INPUT_IS_LIST is True
def test_collector_exposes_images_as_optional_input():
input_types = _load_collector_module().DistributedCollectorNode.INPUT_TYPES()
assert "images" not in input_types["required"]
assert input_types["optional"]["images"] == ("IMAGE",)
def test_audio_only_pass_through_returns_no_images_and_preserves_audio():
collector = _load_collector_module().DistributedCollectorNode()
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
images, returned_audio = collector.run(images=None, audio=[audio])
assert images is None
assert returned_audio is audio
def test_collector_rejects_missing_images_and_audio():
collector = _load_collector_module().DistributedCollectorNode()
try:
collector.run(images=None, audio=None)
except ValueError as exc:
assert "image or audio" in str(exc).lower()
else:
raise AssertionError("Expected collector to reject a run with no media input")
def test_delegate_only_master_allows_no_local_media_input():
collector = _load_collector_module().DistributedCollectorNode()
images, audio = collector.run(
images=None,
audio=None,
multi_job_id=["delegate-audio-job"],
delegate_only=[True],
enabled_worker_ids=["[]"],
)
assert images is None
assert tuple(audio["waveform"].shape) == (1, 2, 1)
def test_pass_through_collapses_comfyui_image_list_to_batch_and_unwraps_hidden_inputs():
collector = _load_collector_module().DistributedCollectorNode()
first = torch.zeros(1, 2, 2, 3)
second = torch.ones(1, 2, 2, 3)
images, audio = collector.run(
images=[first, second],
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],
)
assert tuple(images.shape) == (2, 2, 2, 3)
assert torch.equal(images[0:1], first)
assert torch.equal(images[1:2], second)
assert tuple(audio["waveform"].shape) == (1, 2, 1)
def test_worker_list_input_sends_one_completion_sequence_with_last_only_on_final_item():
module = _load_collector_module()
collector = module.DistributedCollectorNode()
first = torch.zeros(1, 2, 2, 3)
second = torch.ones(1, 2, 2, 3)
posted_payloads = []
class _FakeImage:
def save(self, fp, format=None, compress_level=None):
fp.write(b"png-bytes")
class _FakeResponse:
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
return False
def raise_for_status(self):
return None
class _FakeSession:
def post(self, url, json, timeout):
posted_payloads.append(json)
return _FakeResponse()
async def _fake_get_client_session():
return _FakeSession()
module.tensor_to_pil = lambda *_args, **_kwargs: _FakeImage()
module.get_client_session = _fake_get_client_session
module.encode_audio_payload = lambda _audio: None
images, audio = collector.run(
images=[first, second],
load_balance=[False],
audio=[None],
multi_job_id=["job-list-1"],
is_worker=[True],
master_url=["http://master"],
enabled_worker_ids=["[]"],
worker_batch_size=[1],
worker_id=["worker-a"],
pass_through=[False],
delegate_only=[False],
)
assert tuple(images.shape) == (2, 2, 2, 3)
assert tuple(audio["waveform"].shape) == (1, 2, 1)
assert len(posted_payloads) == 2
assert [payload["batch_idx"] for payload in posted_payloads] == [0, 1]
assert [payload["is_last"] for payload in posted_payloads] == [False, True]
assert {payload["job_id"] for payload in posted_payloads} == {"job-list-1"}
assert {payload["worker_id"] for payload in posted_payloads} == {"worker-a"}
def test_audio_only_worker_sends_one_completion_without_image():
module = _load_collector_module()
collector = module.DistributedCollectorNode()
audio = {"waveform": torch.ones(1, 2, 4), "sample_rate": 48000}
posted = []
class _FakeResponse:
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
return False
def raise_for_status(self):
return None
class _FakeSession:
def post(self, url, json, timeout):
posted.append((url, json, timeout.total))
return _FakeResponse()
async def _fake_get_client_session():
return _FakeSession()
module.get_client_session = _fake_get_client_session
module.encode_audio_payload = lambda value: {"encoded": value is audio}
images, returned_audio = collector.run(
images=None,
audio=[audio],
multi_job_id=["audio-job"],
is_worker=[True],
master_url=["http://master"],
worker_id=["worker-a"],
)
assert images is None
assert returned_audio is audio
assert len(posted) == 1
assert posted[0][0] == "http://master/distributed/job_complete"
assert posted[0][2] == 600
assert posted[0][1] == {
"job_id": "audio-job",
"worker_id": "worker-a",
"batch_idx": 0,
"audio": {"encoded": True},
"is_last": True,
}
def test_audio_only_master_combines_local_and_worker_audio():
module = _load_collector_module()
collector = module.DistributedCollectorNode()
master_audio = {"waveform": torch.ones(1, 2, 2), "sample_rate": 48000}
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
queue = asyncio.Queue()
queue.put_nowait(
{
"worker_id": "worker-a",
"image_index": 0,
"tensor": None,
"audio": worker_audio,
"is_last": True,
}
)
module.prompt_server.distributed_pending_jobs = {"audio-job": queue}
images, combined_audio = asyncio.run(
collector.execute(
images=None,
audio=master_audio,
multi_job_id="audio-job",
enabled_worker_ids='["worker-a"]',
)
)
assert images is None
assert combined_audio["sample_rate"] == 48000
assert tuple(combined_audio["waveform"].shape) == (1, 2, 5)
assert torch.equal(combined_audio["waveform"][..., :2], master_audio["waveform"])
assert torch.equal(combined_audio["waveform"][..., 2:], worker_audio["waveform"])
def test_delegate_only_audio_collects_worker_audio_without_placeholder_image():
module = _load_collector_module()
collector = module.DistributedCollectorNode()
worker_audio = {"waveform": torch.full((1, 2, 3), 2.0), "sample_rate": 48000}
module.prompt_server.distributed_jobs_lock = asyncio.Lock()
queue = asyncio.Queue()
queue.put_nowait(
{
"worker_id": "worker-a",
"image_index": 0,
"tensor": None,
"audio": worker_audio,
"is_last": True,
}
)
module.prompt_server.distributed_pending_jobs = {"delegate-audio-job": queue}
images, combined_audio = asyncio.run(
collector.execute(
images=None,
audio=None,
multi_job_id="delegate-audio-job",
enabled_worker_ids='["worker-a"]',
delegate_only=True,
)
)
assert images is None
assert combined_audio["sample_rate"] == 48000
assert torch.equal(combined_audio["waveform"], worker_audio["waveform"])
File diff suppressed because it is too large Load Diff
-21
View File
@@ -1,21 +0,0 @@
"""Real-framework acceptance; opt in with COMFYUI_SOURCE_ROOT."""
import os
from pathlib import Path
import subprocess
import sys
import pytest
def test_native_v3_runtime():
comfy_root = os.environ.get('COMFYUI_SOURCE_ROOT')
if not comfy_root:
pytest.skip('Set COMFYUI_SOURCE_ROOT to run native ComfyUI V3 acceptance')
helper = Path(__file__).parent / 'helpers' / 'v3_runtime_check.py'
result = subprocess.run(
[sys.executable, str(helper), str(Path(comfy_root).resolve())],
capture_output=True, text=True, timeout=120,
env={**os.environ, 'COMFYUI_IS_WORKER': '1'},
)
assert result.returncode == 0, result.stdout + '\n' + result.stderr
assert 'V3_ACCEPTANCE_OK' in result.stdout
-113
View File
@@ -1,113 +0,0 @@
import importlib.util
import os
import socket
import unittest
from pathlib import Path
from unittest.mock import patch
def load_ports():
path = Path(__file__).resolve().parents[1] / "workers" / "ports.py"
spec = importlib.util.spec_from_file_location("worker_ports_test", path)
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
class WorkerPortTests(unittest.TestCase):
def setUp(self):
self.ports = load_ports()
def test_allocation_starts_above_actual_master_and_skips_assigned_and_occupied(self):
workers = [
{"id": "existing", "host": "localhost", "port": 8190, "enabled": False},
{"id": "remote", "host": "remote.example", "port": 8192},
]
with patch.object(self.ports, "is_port_available", side_effect=lambda port: port != 8191):
self.assertEqual(self.ports.allocate_worker_ports(8189, workers, 3), [8192, 8193, 8194])
def test_local_host_forms_reserve_ports(self):
for host in ["::1", "[::1]", "http://localhost/", "HTTPS://LOCALHOST", "0.0.0.0", None]:
with self.subTest(host=host), patch.object(self.ports, "is_port_available", return_value=True):
self.assertEqual(self.ports.allocate_worker_ports(8189, [{"id": "local", "host": host, "port": 8190}], 1), [8191])
def test_exhaustion_is_explicit_not_a_partial_allocation(self):
with patch.object(self.ports, "is_port_available", return_value=True):
with self.assertRaisesRegex(ValueError, "available.*ports"):
self.ports.allocate_worker_ports(65534, [], 2)
def test_launch_rejects_master_port_without_reassigning_configuration(self):
worker = {"id": "manual", "port": 8189}
with self.assertRaisesRegex(ValueError, "master.*8189"):
self.ports.validate_worker_port(worker, 8189, [])
self.assertEqual(worker["port"], 8189)
def test_launch_rejects_another_local_workers_reserved_port(self):
worker = {"id": "manual", "port": 8190}
other = {"id": "other", "host": None, "port": 8190, "enabled": False}
with self.assertRaisesRegex(ValueError, "other"):
self.ports.validate_worker_port(worker, 8189, [worker, other])
def test_launch_ignores_itself_and_remote_workers_with_same_port(self):
worker = {"id": "manual", "port": 8190}
other = {"id": "remote", "host": "example.com", "port": 8190}
with patch.object(self.ports, "is_port_available", return_value=True):
self.ports.validate_worker_port(worker, 8189, [worker, other])
def test_launch_rejects_an_occupied_port(self):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener:
listener.bind(("127.0.0.1", 0))
listener.listen()
port = listener.getsockname()[1]
self.assertFalse(self.ports.is_port_available(port))
with self.assertRaisesRegex(ValueError, "already in use"):
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
@unittest.skipIf(os.name == "nt", "asyncio does not reuse addresses on Windows")
def test_recently_closed_connection_does_not_block_worker_restart(self):
with socket.socket() as listener:
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind(("127.0.0.1", 0))
listener.listen(1)
port = listener.getsockname()[1]
with socket.create_connection(("127.0.0.1", port), timeout=2) as client:
accepted, _ = listener.accept()
with accepted:
accepted.settimeout(2)
accepted.shutdown(socket.SHUT_WR)
self.assertEqual(client.recv(1), b"")
client.shutdown(socket.SHUT_WR)
self.assertEqual(accepted.recv(1), b"")
self.assertTrue(self.ports.is_port_available(port))
self.ports.validate_worker_port({"id": "restart", "port": port}, 1, [])
with socket.socket() as restarted:
restarted.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
restarted.bind(("127.0.0.1", port))
restarted.listen(1)
def test_active_listener_with_reuseaddr_is_still_a_conflict(self):
with socket.socket() as listener:
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
listener.bind(("127.0.0.1", 0))
listener.listen(1)
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
@unittest.skipUnless(socket.has_ipv6, "IPv6 unavailable")
def test_ipv6_only_listener_is_still_a_conflict(self):
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as listener:
listener.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
try:
listener.bind(("::1", 0))
except OSError as exc:
self.skipTest(f"IPv6 loopback unavailable: {exc}")
listener.listen(1)
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
def test_invalid_ports_fail_clearly(self):
for port in [0, 65536, "not-a-port", None, True, 8190.5]:
with self.subTest(port=port), self.assertRaisesRegex(ValueError, "port"):
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
if __name__ == "__main__":
unittest.main()
-352
View File
@@ -1,352 +0,0 @@
import importlib.util
import os
import sys
import tempfile
import types
import unittest
from argparse import Namespace
from pathlib import Path
from unittest.mock import patch
from sqlalchemy import create_engine
from sqlalchemy.engine import URL, make_url
def _load_process_module(module_filename: str):
module_path = Path(__file__).resolve().parents[1] / "workers" / "process" / module_filename
package_name = "dist_proc_testpkg"
module_name = module_filename[:-3]
for mod_name in list(sys.modules):
if mod_name == package_name or mod_name.startswith(f"{package_name}."):
del sys.modules[mod_name]
root_pkg = types.ModuleType(package_name)
root_pkg.__path__ = []
sys.modules[package_name] = root_pkg
workers_pkg = types.ModuleType(f"{package_name}.workers")
workers_pkg.__path__ = [str(module_path.parents[1])]
sys.modules[f"{package_name}.workers"] = workers_pkg
process_pkg = types.ModuleType(f"{package_name}.workers.process")
process_pkg.__path__ = []
sys.modules[f"{package_name}.workers.process"] = process_pkg
utils_pkg = types.ModuleType(f"{package_name}.utils")
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
logging_module = types.ModuleType(f"{package_name}.utils.logging")
logging_module.debug_log = lambda *_args, **_kwargs: None
logging_module.log = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.logging"] = logging_module
process_module = types.ModuleType(f"{package_name}.utils.process")
process_module.get_python_executable = lambda: "/usr/bin/test-python"
process_module.is_process_alive = lambda _pid: False
process_module.terminate_process = lambda *_args, **_kwargs: None
sys.modules[f"{package_name}.utils.process"] = process_module
config_module = types.ModuleType(f"{package_name}.utils.config")
config_module.load_config = lambda: {"workers": [], "settings": {"stop_workers_on_master_exit": False}}
config_module.save_config = lambda _config: None
sys.modules[f"{package_name}.utils.config"] = config_module
constants_module = types.ModuleType(f"{package_name}.utils.constants")
constants_module.PROCESS_TERMINATION_TIMEOUT = 1
constants_module.PROCESS_WAIT_TIMEOUT = 1
constants_module.WORKER_CHECK_INTERVAL = 0.01
sys.modules[f"{package_name}.utils.constants"] = constants_module
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.get_server_port = lambda: 8189
sys.modules[f"{package_name}.utils.network"] = network_module
spec = importlib.util.spec_from_file_location(
f"{package_name}.workers.process.{module_name}",
module_path,
)
module = importlib.util.module_from_spec(spec)
assert spec is not None and spec.loader is not None
spec.loader.exec_module(module)
return module
root_discovery_module = _load_process_module("root_discovery.py")
launch_builder_module = _load_process_module("launch_builder.py")
class ComfyRootDiscoveryTests(unittest.TestCase):
def test_prefers_loaded_comfyui_module_path(self):
discovery = root_discovery_module.ComfyRootDiscovery()
server_module = types.SimpleNamespace(__file__="/opt/ComfyUI/server.py")
def fake_exists(path):
return path == "/opt/ComfyUI/main.py"
with patch.dict(sys.modules, {"server": server_module}, clear=False), \
patch.object(root_discovery_module.os.path, "exists", side_effect=fake_exists), \
patch.dict(root_discovery_module.os.environ, {}, clear=True):
self.assertEqual(discovery.find_comfy_root(), "/opt/ComfyUI")
class LaunchCommandBuilderTests(unittest.TestCase):
def build_command(self, root, worker=None, runtime=None):
builder = launch_builder_module.LaunchCommandBuilder()
worker = worker or {"id": "worker-a", "port": 8190}
runtime = runtime if runtime is not None else Namespace(database_url=None)
with patch.object(builder, "_get_runtime_args", return_value=runtime):
return builder.build_launch_command(worker, str(root))
def test_database_is_unique_and_stable_by_worker_id(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
first = self.build_command(root, {"id": "worker-a", "port": 8190})
renamed = self.build_command(root, {
"id": "worker-a", "port": 9001, "name": "Renamed", "cuda_device": 3,
})
second = self.build_command(root, {"id": "worker-b", "port": 8191})
database = first[first.index("--database-url") + 1]
self.assertEqual(database, renamed[renamed.index("--database-url") + 1])
self.assertNotEqual(database, second[second.index("--database-url") + 1])
self.assertTrue(database.startswith("sqlite:///" + root.as_posix() + "/"))
self.assertTrue(Path(database.removeprefix("sqlite:///")).parent.is_dir())
self.assertNotIn("--disable-assets", first)
def test_database_uses_effective_user_or_base_directory(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
for extra_args, expected in [
(f'--user-directory "{root / "custom user"}"', root / "custom user"),
(f'--base-directory "{root / "custom base"}"', root / "custom base" / "user"),
]:
with self.subTest(extra_args=extra_args):
cmd = self.build_command(root, {
"id": "worker-a", "port": 8190, "extra_args": extra_args,
})
database = cmd[cmd.index("--database-url") + 1]
self.assertTrue(database.startswith("sqlite:///" + expected.as_posix() + "/"))
cmd = self.build_command(root, runtime=Namespace(
database_url="sqlite:///master.db", user_directory=str(root / "runtime user"),
))
database = cmd[cmd.index("--database-url") + 1]
self.assertTrue(database.startswith("sqlite:///" + (root / "runtime user").as_posix() + "/"))
self.assertNotEqual(database, "sqlite:///master.db")
def test_relative_database_directories_use_worker_cwd(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory) / "ComfyUI"
root.mkdir()
(root / "main.py").touch()
launcher = Path(directory) / "launcher"
launcher.mkdir()
cases = [
(Namespace(database_url=None, user_directory="profiles"), "", root / "profiles"),
(Namespace(database_url=None, base_directory="data"), "", root / "data/user"),
(Namespace(database_url=None), "--user-directory=profiles", root / "profiles"),
(Namespace(database_url=None), "--base-directory data", root / "data/user"),
]
for runtime, extra, expected in cases:
with self.subTest(runtime=runtime, extra=extra), \
patch.object(os, "getcwd", return_value=str(launcher)):
cmd = self.build_command(root, {
"id": "worker-a", "port": 8190, "extra_args": extra,
}, runtime)
database = make_url(cmd[cmd.index("--database-url") + 1]).database
self.assertEqual(Path(database).parent, expected / "distributed/workers")
self.assertFalse((launcher / "profiles").exists())
self.assertFalse((launcher / "data").exists())
def test_database_url_round_trips_special_characters_without_collisions(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
for name in ["user?profile", "user%3Fprofile"]:
with self.subTest(name=name):
user = root / name
runtime = Namespace(database_url=None, user_directory=str(user))
# SQLAlchemy 2.0 cannot round-trip '?' in a filename. Fail
# clearly on those versions rather than silently sharing a DB.
sample = (user / "test.db").as_posix()
serialized = URL.create("sqlite", database=sample).render_as_string()
if make_url(serialized).database != sample:
with self.assertRaisesRegex(ValueError, "SQLAlchemy.*database path"):
self.build_command(root, runtime=runtime)
self.assertFalse(user.exists())
continue
databases = []
for worker_id in ["worker-a", "worker-b"]:
cmd = self.build_command(root, {"id": worker_id, "port": 8190}, runtime)
url = cmd[cmd.index("--database-url") + 1]
database = make_url(url).database
self.assertEqual(Path(database).parent, user / "distributed/workers")
engine = create_engine(url)
try:
with engine.connect() as connection:
actual = connection.exec_driver_sql("PRAGMA database_list").one()[2]
self.assertEqual(Path(actual), Path(database))
finally:
engine.dispose()
databases.append(database)
self.assertNotEqual(*databases)
def test_explicit_database_and_disabled_assets_are_preserved(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
for extra in ["--database-url sqlite:///explicit.db", "--database-url=sqlite:///explicit.db", "--disable-assets"]:
with self.subTest(extra=extra):
cmd = self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
self.assertEqual(sum(arg.split("=")[0] == "--database-url" for arg in cmd), 0 if extra == "--disable-assets" else 1)
self.assertFalse((root / "user").exists())
def test_older_comfyui_without_database_flag_keeps_existing_command(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
cmd = self.build_command(root, runtime=Namespace())
self.assertNotIn("--database-url", cmd)
self.assertFalse((root / "user").exists())
def test_worker_id_cannot_escape_database_directory(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
databases = []
for worker_id in ["../outside", "a/b", "a_b"]:
cmd = self.build_command(root, {"id": worker_id, "port": 8190})
path = Path(cmd[cmd.index("--database-url") + 1].removeprefix("sqlite:///"))
self.assertTrue(path.is_relative_to(root / "user"))
databases.append(path)
self.assertEqual(len(set(databases)), 3)
def test_extra_args_cannot_silently_override_the_configured_port(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
(root / "main.py").touch()
for extra in ["--port 8189", "--port=8189"]:
with self.subTest(extra=extra), self.assertRaisesRegex(ValueError, "configured worker port"):
self.build_command(root, {"id": "worker-a", "port": 8190, "extra_args": extra})
def test_inherits_runtime_layout_args_for_desktop(self):
builder = launch_builder_module.LaunchCommandBuilder()
runtime_args = Namespace(
listen="127.0.0.1",
base_directory="C:/Users/test/ComfyUI",
temp_directory=None,
input_directory="C:/Users/test/ComfyUI/input",
output_directory="C:/Users/test/ComfyUI/output",
user_directory="C:/Users/test/ComfyUI/user",
front_end_root="C:/Program Files/ComfyUI/web_custom_versions/desktop_app",
extra_model_paths_config=[["C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml"]],
enable_manager=True,
disable_manager_ui=False,
enable_manager_legacy_ui=False,
windows_standalone_build=True,
log_stdout=True,
verbose="INFO",
enable_cors_header="*",
)
comfy_module = types.ModuleType("comfy")
comfy_cli_args = types.ModuleType("comfy.cli_args")
comfy_cli_args.args = runtime_args
worker_config = {
"port": 9001,
"extra_args": "--preview-method auto",
}
def fake_exists(path):
return path == "/desktop/ComfyUI/main.py"
with patch.dict(
sys.modules,
{"comfy": comfy_module, "comfy.cli_args": comfy_cli_args},
clear=False,
), patch.object(launch_builder_module.os.path, "exists", side_effect=fake_exists):
cmd = builder.build_launch_command(worker_config, "/desktop/ComfyUI")
self.assertEqual(cmd[:2], ["/usr/bin/test-python", "/desktop/ComfyUI/main.py"])
self.assertIn("--listen", cmd)
self.assertIn("127.0.0.1", cmd)
self.assertIn("--base-directory", cmd)
self.assertIn("C:/Users/test/ComfyUI", cmd)
self.assertIn("--input-directory", cmd)
self.assertIn("--output-directory", cmd)
self.assertIn("--user-directory", cmd)
self.assertIn("--front-end-root", cmd)
self.assertIn("--extra-model-paths-config", cmd)
self.assertIn("C:/Users/test/AppData/Roaming/ComfyUI/extra_models_config.yaml", cmd)
self.assertIn("--enable-manager", cmd)
self.assertIn("--windows-standalone-build", cmd)
self.assertIn("--log-stdout", cmd)
self.assertIn("--disable-auto-launch", cmd)
self.assertIn("--enable-cors-header", cmd)
self.assertIn("*", cmd)
self.assertIn("--port", cmd)
self.assertIn("9001", cmd)
self.assertNotIn("--auto-launch", cmd)
class ProcessLaunchPortTests(unittest.TestCase):
def test_launch_before_master_listens_uses_configured_master_port(self):
lifecycle_module = _load_process_module("lifecycle.py")
cli_args = types.ModuleType("comfy.cli_args")
cli_args.args = Namespace(port=8189)
for missing in [AttributeError("PromptServer has no port"), None]:
with self.subTest(missing=missing), tempfile.TemporaryDirectory() as directory:
manager = types.SimpleNamespace(
find_comfy_root=lambda: directory,
build_launch_command=lambda _worker, _root: ["python", "main.py"],
processes={}, save_processes=lambda: None,
)
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
patch.object(lifecycle_module, "get_server_port", side_effect=missing if isinstance(missing, Exception) else None, return_value=None), \
patch.object(lifecycle_module, "validate_worker_port") as validate, \
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
spawn.assert_called_once()
self.assertEqual(pid, 1234)
def test_pre_listen_launch_still_rejects_the_configured_master_port(self):
lifecycle_module = _load_process_module("lifecycle.py")
cli_args = types.ModuleType("comfy.cli_args")
cli_args.args = Namespace(port=8189)
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
with patch.dict(sys.modules, {"comfy": types.ModuleType("comfy"), "comfy.cli_args": cli_args}), \
patch.object(lifecycle_module, "get_server_port", side_effect=AttributeError("port")), \
patch.object(lifecycle_module.subprocess, "Popen") as spawn:
with self.assertRaisesRegex(ValueError, "master port 8189"):
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
spawn.assert_not_called()
def test_conflict_is_rejected_before_building_or_spawning(self):
lifecycle_module = _load_process_module("lifecycle.py")
manager = types.SimpleNamespace(find_comfy_root=lambda: "/ComfyUI", processes={})
with patch.object(lifecycle_module.subprocess, "Popen") as spawn:
with self.assertRaisesRegex(ValueError, "master port 8189"):
lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8189})
spawn.assert_not_called()
self.assertEqual(manager.processes, {})
def test_available_port_reaches_existing_launch_path(self):
lifecycle_module = _load_process_module("lifecycle.py")
with tempfile.TemporaryDirectory() as directory:
manager = types.SimpleNamespace(
find_comfy_root=lambda: directory,
build_launch_command=lambda _worker, _root: ["python", "main.py", "--port", "8190"],
processes={}, save_processes=lambda: None,
)
with patch.object(lifecycle_module, "validate_worker_port") as validate, \
patch.object(lifecycle_module.subprocess, "Popen", return_value=types.SimpleNamespace(pid=1234)) as spawn:
pid = lifecycle_module.ProcessLifecycle(manager).launch_worker({"id": "manual", "port": 8190})
validate.assert_called_once_with({"id": "manual", "port": 8190}, 8189, [])
spawn.assert_called_once()
self.assertEqual(pid, 1234)
self.assertEqual(manager.processes["manual"]["pid"], 1234)
if __name__ == "__main__":
unittest.main()
+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."""
@@ -1,3 +1,4 @@
import asyncio
import importlib.util
import sys
import types
@@ -5,7 +6,7 @@ import unittest
from pathlib import Path
class _PromptQueue:
class _FakePromptQueue:
def __init__(self):
self.items = []
@@ -13,10 +14,16 @@ class _PromptQueue:
self.items.append(item)
def _load_async_helpers_module():
module_path = Path(__file__).resolve().parents[1] / "utils" / "async_helpers.py"
package_name = "dist_async_helpers_testpkg"
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]
@@ -29,62 +36,68 @@ def _load_async_helpers_module():
utils_pkg.__path__ = []
sys.modules[f"{package_name}.utils"] = utils_pkg
execution_module = types.ModuleType("execution")
async def _validate_prompt(prompt_id, prompt, partial_execution_targets):
return (True, None, ["9"], {})
execution_module.validate_prompt = _validate_prompt
execution_module.SENSITIVE_EXTRA_DATA_KEYS = []
sys.modules["execution"] = execution_module
prompt_server = types.SimpleNamespace(
trigger_on_prompt=lambda payload: payload,
number=12,
prompt_queue=_PromptQueue(),
)
server_module = types.ModuleType("server")
server_module.PromptServer = types.SimpleNamespace(instance=prompt_server)
sys.modules["server"] = server_module
network_module = types.ModuleType(f"{package_name}.utils.network")
network_module.get_server_loop = lambda: None
network_module.get_server_loop = asyncio.get_event_loop
sys.modules[f"{package_name}.utils.network"] = network_module
spec = importlib.util.spec_from_file_location(f"{package_name}.utils.async_helpers", module_path)
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, prompt_server
return module, fake_prompt_server
async_helpers, prompt_server = _load_async_helpers_module()
class AsyncHelpersQueuePromptPayloadTests(unittest.TestCase):
def test_queue_prompt_payload_includes_create_time_metadata(self):
async_helpers_module, fake_prompt_server = _load_async_helpers_module()
class QueuePromptPayloadTests(unittest.IsolatedAsyncioTestCase):
async def test_queue_prompt_payload_includes_create_time_and_client_metadata(self):
result = await async_helpers.queue_prompt_payload(
{"1": {"class_type": "Node"}},
workflow_meta={"id": "workflow-1"},
client_id="client-1",
include_queue_metadata=True,
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(result["prompt_id"], str)
self.assertTrue(result["prompt_id"])
self.assertEqual(result["number"], 12)
self.assertEqual(result["node_errors"], {})
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)
self.assertEqual(prompt_server.number, 13)
self.assertEqual(len(prompt_server.prompt_queue.items), 1)
queued_item = prompt_server.prompt_queue.items[0]
self.assertEqual(queued_item[0], 12)
extra_data = queued_item[3]
self.assertEqual(extra_data["client_id"], "client-1")
self.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)
self.assertEqual(extra_data["extra_pnginfo"]["workflow"], {"id": "workflow-1"})
if __name__ == "__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)
@@ -90,48 +94,14 @@ class NetworkHelpersTests(unittest.TestCase):
"https://master.example.com",
)
def test_build_master_url_ignores_stale_saved_port_and_uses_runtime_port(self):
cfg = {"master": {"host": "192.168.68.56", "port": 8001}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8188",
)
def test_build_master_url_keeps_explicit_port_in_host(self):
cfg = {"master": {"host": "192.168.68.56:8001"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8188)
self.assertEqual(
network.build_master_url(config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8001",
)
def test_build_master_url_falls_back_to_server_address(self):
cfg = {"master": {"host": "", "port": 8001}}
prompt_server = types.SimpleNamespace(address="0.0.0.0", port=8190)
cfg = {"master": {"host": ""}}
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",
)
def test_build_master_callback_url_uses_loopback_for_local_worker(self):
cfg = {"master": {"host": "192.168.68.56"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
worker = {"id": "w1", "type": "local", "host": "localhost", "port": 8189}
self.assertEqual(
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
"http://127.0.0.1:8001",
)
def test_build_master_callback_url_keeps_public_master_url_for_remote_worker(self):
cfg = {"master": {"host": "192.168.68.56"}}
prompt_server = types.SimpleNamespace(address="127.0.0.1", port=8001)
worker = {"id": "w2", "type": "remote", "host": "192.168.68.99", "port": 8189}
self.assertEqual(
network.build_master_callback_url(worker, config=cfg, prompt_server_instance=prompt_server),
"http://192.168.68.56:8001",
)
if __name__ == "__main__":
unittest.main()
+705
View File
@@ -0,0 +1,705 @@
import importlib.util
import json
import sys
import types
import unittest
from pathlib import Path
def _load_prompt_transform_module():
module_path = Path(__file__).resolve().parents[3] / "api" / "orchestration" / "prompt_transform.py"
package_name = "dist_pt_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
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
spec = importlib.util.spec_from_file_location(
f"{package_name}.api.orchestration.prompt_transform",
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
pt = _load_prompt_transform_module()
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _linear_prompt():
"""1 → 2 → 3 → 4(DistributedCollector) → 5(SaveImage)"""
return {
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
"2": {"class_type": "CLIPTextEncode", "inputs": {"clip": ["1", 1]}},
"3": {"class_type": "KSampler", "inputs": {"model": ["1", 0], "positive": ["2", 0]}},
"4": {"class_type": "DistributedCollector", "inputs": {"images": ["3", 0]}},
"5": {"class_type": "SaveImage", "inputs": {"images": ["4", 0]}},
}
def _collector_only_prompt():
"""1(Checkpoint) → 2(DistributedCollector) [no downstream from 2]"""
return {
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
"2": {"class_type": "DistributedCollector", "inputs": {"images": ["1", 0]}},
}
def _delegate_prompt():
"""1 → 2 → 3(DistributedCollector) → 4(SaveImage)"""
return {
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
"2": {"class_type": "KSampler", "inputs": {"model": ["1", 0]}},
"3": {"class_type": "DistributedCollector", "inputs": {"images": ["2", 0]}},
"4": {"class_type": "SaveImage", "inputs": {"images": ["3", 0]}},
}
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"]
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_copy,
participant_id=participant_id,
enabled_worker_ids=enabled_worker_ids,
job_id_map=job_id_map,
master_url="http://master.example.com",
delegate_master=delegate_master,
prompt_index=idx,
)
# ---------------------------------------------------------------------------
# PromptIndex
# ---------------------------------------------------------------------------
class PromptIndexTests(unittest.TestCase):
def test_nodes_by_class_groups_correctly(self):
prompt = {
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
"2": {"class_type": "DistributedCollector", "inputs": {}},
"3": {"class_type": "DistributedCollector", "inputs": {}},
}
idx = pt.PromptIndex(prompt)
self.assertCountEqual(idx.nodes_for_class("DistributedCollector"), ["2", "3"])
self.assertEqual(idx.nodes_for_class("CheckpointLoaderSimple"), ["1"])
def test_nodes_for_class_unknown_returns_empty(self):
idx = pt.PromptIndex({"1": {"class_type": "KSampler", "inputs": {}}})
self.assertEqual(idx.nodes_for_class("Nonexistent"), [])
def test_nodes_without_class_type_are_indexed_under_none(self):
prompt = {"1": {"inputs": {}}}
idx = pt.PromptIndex(prompt)
# Should not raise; nodes_for_class with None key or missing class_type
self.assertEqual(idx.nodes_for_class("KSampler"), [])
def test_copy_prompt_is_a_deep_copy(self):
prompt = {"1": {"class_type": "KSampler", "inputs": {"seed": 42}}}
idx = pt.PromptIndex(prompt)
copy = idx.copy_prompt()
copy["1"]["inputs"]["seed"] = 999
self.assertEqual(prompt["1"]["inputs"]["seed"], 42)
def test_has_upstream_direct_connection(self):
"""Node 4 reads directly from node 3 (KSampler)."""
idx = pt.PromptIndex(_linear_prompt())
self.assertTrue(idx.has_upstream("4", "KSampler"))
def test_has_upstream_transitive_connection(self):
"""Node 4 → 3 → 2 → 1 (CheckpointLoaderSimple)."""
idx = pt.PromptIndex(_linear_prompt())
self.assertTrue(idx.has_upstream("4", "CheckpointLoaderSimple"))
def test_has_upstream_returns_false_when_no_path(self):
idx = pt.PromptIndex(_linear_prompt())
# CheckpointLoaderSimple has no upstream nodes
self.assertFalse(idx.has_upstream("1", "DistributedCollector"))
def test_has_upstream_result_is_cached(self):
idx = pt.PromptIndex(_linear_prompt())
r1 = idx.has_upstream("4", "KSampler")
r2 = idx.has_upstream("4", "KSampler")
self.assertEqual(r1, r2)
self.assertIn(("4", "KSampler"), idx._upstream_cache)
def test_has_upstream_does_not_infinite_loop_on_cycle(self):
"""Cyclic references in inputs should not cause infinite recursion."""
prompt = {
"1": {"class_type": "A", "inputs": {"x": ["2", 0]}},
"2": {"class_type": "B", "inputs": {"x": ["1", 0]}},
}
idx = pt.PromptIndex(prompt)
# Should terminate without error
result = idx.has_upstream("1", "NonExistent")
self.assertFalse(result)
# ---------------------------------------------------------------------------
# find_nodes_by_class
# ---------------------------------------------------------------------------
class FindNodesByClassTests(unittest.TestCase):
def test_finds_matching_nodes(self):
prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "DistributedCollector", "inputs": {}},
}
result = pt.find_nodes_by_class(prompt, "KSampler")
self.assertEqual(result, ["1"])
def test_returns_empty_when_no_match(self):
prompt = {"1": {"class_type": "KSampler", "inputs": {}}}
self.assertEqual(pt.find_nodes_by_class(prompt, "DistributedCollector"), [])
def test_skips_non_dict_nodes(self):
prompt = {"1": "not a dict", "2": {"class_type": "KSampler", "inputs": {}}}
result = pt.find_nodes_by_class(prompt, "KSampler")
self.assertEqual(result, ["2"])
# ---------------------------------------------------------------------------
# prune_prompt_for_worker
# ---------------------------------------------------------------------------
class PrunePromptForWorkerTests(unittest.TestCase):
def test_no_distributed_nodes_returns_prompt_unchanged(self):
prompt = {
"1": {"class_type": "CheckpointLoaderSimple", "inputs": {}},
"2": {"class_type": "SaveImage", "inputs": {"images": ["1", 0]}},
}
result = pt.prune_prompt_for_worker(prompt)
self.assertCountEqual(result.keys(), ["1", "2"])
def test_keeps_collector_and_upstream(self):
prompt = _linear_prompt()
result = pt.prune_prompt_for_worker(prompt)
for node_id in ("1", "2", "3", "4"):
self.assertIn(node_id, result)
def test_removes_downstream_of_collector(self):
prompt = _linear_prompt()
result = pt.prune_prompt_for_worker(prompt)
self.assertNotIn("5", result)
def test_injects_preview_image_when_downstream_exists(self):
prompt = _linear_prompt()
result = pt.prune_prompt_for_worker(prompt)
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
self.assertEqual(len(preview_nodes), 1)
self.assertEqual(preview_nodes[0]["inputs"]["images"], ["4", 0])
def test_no_preview_image_when_no_downstream(self):
result = pt.prune_prompt_for_worker(_collector_only_prompt())
preview_nodes = [n for n in result.values() if n.get("class_type") == "PreviewImage"]
self.assertEqual(len(preview_nodes), 0)
def test_unrelated_nodes_are_pruned(self):
prompt = {
"1": {"class_type": "DistributedCollector", "inputs": {}},
"2": {"class_type": "UnrelatedNode", "inputs": {}}, # no connection to 1
}
result = pt.prune_prompt_for_worker(prompt)
self.assertIn("1", result)
self.assertNotIn("2", result)
def test_result_is_a_copy_not_same_object(self):
prompt = _linear_prompt()
result = pt.prune_prompt_for_worker(prompt)
# Mutating the result should not affect the original
original_keys = set(prompt.keys())
result["NEW"] = {"class_type": "Test", "inputs": {}}
self.assertEqual(set(prompt.keys()), original_keys)
def test_upscale_node_is_treated_as_distributed(self):
prompt = {
"1": {"class_type": "KSampler", "inputs": {}},
"2": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"image": ["1", 0]}},
"3": {"class_type": "SaveImage", "inputs": {"images": ["2", 0]}},
}
result = pt.prune_prompt_for_worker(prompt)
self.assertIn("1", result)
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
# ---------------------------------------------------------------------------
class PrepareDelegateMasterPromptTests(unittest.TestCase):
def test_keeps_collector_and_downstream(self):
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
self.assertIn("3", result)
self.assertIn("4", result)
self.assertNotIn("1", result)
self.assertNotIn("2", result)
def test_removes_dangling_upstream_refs(self):
"""Collector must not retain dangling refs to pruned upstream nodes."""
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
collector_inputs = result["3"].get("inputs", {})
# Original "images" pointed at node 2, which is pruned.
# It should now point at a newly injected placeholder node.
self.assertIn("images", collector_inputs)
source_id = str(collector_inputs["images"][0])
self.assertNotEqual(source_id, "2")
self.assertIn(source_id, result)
self.assertEqual(result[source_id].get("class_type"), "DistributedEmptyImage")
def test_injects_empty_image_placeholder(self):
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
empty_nodes = [(nid, n) for nid, n in result.items() if n.get("class_type") == "DistributedEmptyImage"]
self.assertEqual(len(empty_nodes), 1)
placeholder_id = empty_nodes[0][0]
self.assertEqual(result["3"]["inputs"]["images"], [placeholder_id, 0])
def test_one_placeholder_per_collector(self):
"""Two collectors → two placeholders."""
prompt = {
"1": {"class_type": "DistributedCollector", "inputs": {}},
"2": {"class_type": "DistributedCollector", "inputs": {}},
"3": {"class_type": "SaveImage", "inputs": {"images": ["1", 0]}},
}
result = pt.prepare_delegate_master_prompt(prompt, ["1", "2"])
empty_nodes = [n for n in result.values() if n.get("class_type") == "DistributedEmptyImage"]
self.assertEqual(len(empty_nodes), 2)
def test_result_is_independent_copy(self):
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
result["3"]["inputs"]["NEW"] = "injected"
# Original should be untouched
self.assertNotIn("NEW", prompt["3"].get("inputs", {}))
# ---------------------------------------------------------------------------
# generate_job_id_map
# ---------------------------------------------------------------------------
class GenerateJobIdMapTests(unittest.TestCase):
def test_maps_collector_nodes(self):
prompt = {
"1": {"class_type": "DistributedCollector", "inputs": {}},
"2": {"class_type": "KSampler", "inputs": {}},
}
idx = pt.PromptIndex(prompt)
job_map = pt.generate_job_id_map(idx, "prefix")
self.assertEqual(job_map["1"], "prefix_1")
self.assertNotIn("2", job_map)
def test_maps_upscale_nodes(self):
prompt = {
"5": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}},
}
idx = pt.PromptIndex(prompt)
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"), {})
def test_stable_ids_across_calls(self):
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {}}}
idx = pt.PromptIndex(prompt)
m1 = pt.generate_job_id_map(idx, "run")
m2 = pt.generate_job_id_map(idx, "run")
self.assertEqual(m1, m2)
# ---------------------------------------------------------------------------
# apply_participant_overrides – DistributedCollector
# ---------------------------------------------------------------------------
class ApplyOverridesCollectorTests(unittest.TestCase):
def _collector_prompt(self):
return {"1": {"class_type": "DistributedCollector", "inputs": {}}}
def test_worker_sets_is_worker_true(self):
result = _apply(self._collector_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
def test_worker_sets_master_url(self):
result = _apply(self._collector_prompt(), "worker-a")
self.assertEqual(result["1"]["inputs"]["master_url"], "http://master.example.com")
def test_worker_sets_worker_id(self):
result = _apply(self._collector_prompt(), "worker-a")
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-a")
def test_worker_sets_delegate_only_false(self):
result = _apply(self._collector_prompt(), "worker-a")
self.assertFalse(result["1"]["inputs"]["delegate_only"])
def test_master_sets_is_worker_false(self):
result = _apply(self._collector_prompt(), "master")
self.assertFalse(result["1"]["inputs"]["is_worker"])
def test_master_clears_stale_master_url(self):
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {"master_url": "stale"}}}
result = _apply(prompt, "master")
self.assertNotIn("master_url", result["1"]["inputs"])
def test_master_clears_stale_worker_id(self):
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {"worker_id": "stale"}}}
result = _apply(prompt, "master")
self.assertNotIn("worker_id", result["1"]["inputs"])
def test_master_with_delegate_master_sets_delegate_only_true(self):
result = _apply(self._collector_prompt(), "master", delegate_master=True)
self.assertTrue(result["1"]["inputs"]["delegate_only"])
def test_master_without_delegate_master_sets_delegate_only_false(self):
result = _apply(self._collector_prompt(), "master", delegate_master=False)
self.assertFalse(result["1"]["inputs"]["delegate_only"])
def test_enabled_worker_ids_serialized_as_json(self):
enabled = ["worker-a", "worker-b"]
result = _apply(self._collector_prompt(), "master", enabled_worker_ids=enabled)
self.assertEqual(result["1"]["inputs"]["enabled_worker_ids"], json.dumps(enabled))
def test_multi_job_id_is_set_from_job_map(self):
prompt = {"1": {"class_type": "DistributedCollector", "inputs": {}}}
idx = pt.PromptIndex(prompt)
job_id_map = {"1": "run_abc_1"}
result = pt.apply_participant_overrides(
prompt,
participant_id="worker-a",
enabled_worker_ids=["worker-a"],
job_id_map=job_id_map,
master_url="http://master",
delegate_master=False,
prompt_index=idx,
)
self.assertEqual(result["1"]["inputs"]["multi_job_id"], "run_abc_1")
# ---------------------------------------------------------------------------
# apply_participant_overrides – DistributedSeed
# ---------------------------------------------------------------------------
class ApplyOverridesSeedTests(unittest.TestCase):
def _seed_prompt(self):
return {"1": {"class_type": "DistributedSeed", "inputs": {}}}
def test_worker_sets_is_worker_true(self):
result = _apply(self._seed_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
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-b")
def test_master_sets_is_worker_false(self):
result = _apply(self._seed_prompt(), "master")
self.assertFalse(result["1"]["inputs"]["is_worker"])
def test_master_sets_empty_worker_id(self):
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
# ---------------------------------------------------------------------------
class ApplyOverridesUpscaleTests(unittest.TestCase):
def _upscale_prompt(self):
return {"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}}}
def test_worker_sets_is_worker_true(self):
result = _apply(self._upscale_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
def test_worker_sets_master_url_and_worker_id(self):
result = _apply(self._upscale_prompt(), "worker-a")
self.assertEqual(result["1"]["inputs"]["master_url"], "http://master.example.com")
self.assertEqual(result["1"]["inputs"]["worker_id"], "worker-a")
def test_master_clears_master_url_and_worker_id(self):
prompt = {"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {"master_url": "x", "worker_id": "y"}}}
result = _apply(prompt, "master")
self.assertNotIn("master_url", result["1"]["inputs"])
self.assertNotIn("worker_id", result["1"]["inputs"])
def test_collector_downstream_of_upscale_gets_pass_through(self):
"""A DistributedCollector that is downstream of UltimateSDUpscaleDistributed → pass_through=True."""
prompt = {
"1": {"class_type": "UltimateSDUpscaleDistributed", "inputs": {}},
"2": {"class_type": "DistributedCollector", "inputs": {"images": ["1", 0]}},
}
result = _apply(prompt, "worker-a", enabled_worker_ids=["worker-a"])
self.assertTrue(result["2"]["inputs"].get("pass_through"))
# ---------------------------------------------------------------------------
# apply_participant_overrides – DistributedValue
# ---------------------------------------------------------------------------
class ApplyOverridesValueTests(unittest.TestCase):
def _value_prompt(self):
return {"1": {"class_type": "DistributedValue", "inputs": {}}}
def test_worker_sets_is_worker_true(self):
result = _apply(self._value_prompt(), "worker-a")
self.assertTrue(result["1"]["inputs"]["is_worker"])
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-b")
def test_master_sets_is_worker_false(self):
result = _apply(self._value_prompt(), "master")
self.assertFalse(result["1"]["inputs"]["is_worker"])
def test_master_sets_empty_worker_id(self):
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
@@ -86,25 +86,27 @@ class ParseTilesFromFormTests(unittest.TestCase):
# --- happy paths ---
def test_single_tile_returns_image_and_metadata(self):
def test_single_tile_returns_one_entry(self):
tiles = pp._parse_tiles_from_form(_make_form(1))
self.assertEqual(len(tiles), 1)
def test_multiple_tiles_all_returned(self):
tiles = pp._parse_tiles_from_form(_make_form(3))
self.assertEqual(len(tiles), 3)
def test_tile_image_is_pil_image(self):
tiles = pp._parse_tiles_from_form(_make_form(1))
self.assertIsInstance(tiles[0]["image"], PILImage.Image)
def test_tile_metadata_fields_are_parsed(self):
tiles = pp._parse_tiles_from_form(_make_form(1))
tile = tiles[0]
self.assertIsInstance(tile["image"], PILImage.Image)
self.assertEqual(tile["tile_idx"], 0)
self.assertEqual(tile["x"], 0)
self.assertEqual(tile["y"], 0)
self.assertEqual(tile["extracted_width"], 64)
self.assertEqual(tile["extracted_height"], 64)
def test_multiple_tiles_preserve_count_order_and_coordinates(self):
tiles = pp._parse_tiles_from_form(_make_form(3))
self.assertEqual(len(tiles), 3)
for i, tile in enumerate(tiles):
self.assertEqual(tile["tile_idx"], i)
self.assertEqual(tiles[1]["x"], 64)
self.assertEqual(tiles[2]["x"], 128)
def test_padding_is_parsed_from_form(self):
tiles = pp._parse_tiles_from_form(_make_form(1, padding=16))
self.assertEqual(tiles[0]["padding"], 16)
@@ -136,6 +138,16 @@ class ParseTilesFromFormTests(unittest.TestCase):
self.assertNotIn("batch_idx", tiles[0])
self.assertNotIn("global_idx", tiles[0])
def test_tile_indices_match_metadata_order(self):
tiles = pp._parse_tiles_from_form(_make_form(3))
for i, tile in enumerate(tiles):
self.assertEqual(tile["tile_idx"], i)
def test_x_coordinates_reflect_metadata(self):
tiles = pp._parse_tiles_from_form(_make_form(3))
self.assertEqual(tiles[1]["x"], 64)
self.assertEqual(tiles[2]["x"], 128)
# --- error cases ---
def test_missing_tiles_metadata_raises_value_error(self):
+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]

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