Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6f85129e45 | ||
|
|
a53d88f7fb | ||
|
|
9be342f5fb | ||
|
|
a1faf7e916 |
+15
-22
@@ -1,29 +1,22 @@
|
||||
# Import everything needed from the main module
|
||||
from .distributed import (
|
||||
NODE_CLASS_MAPPINGS as DISTRIBUTED_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as DISTRIBUTED_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
"""ComfyUI-Distributed's native V3 extension entrypoint."""
|
||||
from comfy_api.v0_0_2 import ComfyExtension, io
|
||||
|
||||
# Import utilities
|
||||
from .utils.config import ensure_config_exists, CONFIG_FILE
|
||||
from .utils.logging import debug_log
|
||||
from .nodes.v3 import NODES
|
||||
from .runtime.bootstrap import initialize
|
||||
|
||||
# Import distributed upscale nodes
|
||||
from .nodes.distributed_upscale import (
|
||||
NODE_CLASS_MAPPINGS as UPSCALE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS as UPSCALE_DISPLAY_NAME_MAPPINGS
|
||||
)
|
||||
WEB_DIRECTORY = './web'
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
|
||||
ensure_config_exists()
|
||||
class DistributedExtension(ComfyExtension):
|
||||
async def on_load(self) -> None:
|
||||
initialize()
|
||||
|
||||
# Merge node mappings
|
||||
NODE_CLASS_MAPPINGS = {**DISTRIBUTED_CLASS_MAPPINGS, **UPSCALE_CLASS_MAPPINGS}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {**DISTRIBUTED_DISPLAY_NAME_MAPPINGS, **UPSCALE_DISPLAY_NAME_MAPPINGS}
|
||||
async def get_node_list(self) -> list[type[io.ComfyNode]]:
|
||||
return list(NODES)
|
||||
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
debug_log("Loaded Distributed nodes.")
|
||||
debug_log(f"Config file: {CONFIG_FILE}")
|
||||
debug_log(f"Available nodes: {list(NODE_CLASS_MAPPINGS.keys())}")
|
||||
async def comfy_entrypoint() -> DistributedExtension:
|
||||
return DistributedExtension()
|
||||
|
||||
|
||||
__all__ = ['comfy_entrypoint', 'WEB_DIRECTORY']
|
||||
|
||||
+14
-1
@@ -16,12 +16,14 @@ 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 ..workers.detection import (
|
||||
get_machine_id,
|
||||
@@ -277,6 +279,13 @@ def _get_cuda_info():
|
||||
def _collect_network_info_sync():
|
||||
"""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)
|
||||
@@ -285,7 +294,9 @@ def _collect_network_info_sync():
|
||||
"all_ips": all_ips,
|
||||
"recommended_ip": recommended_ip,
|
||||
"cuda_device": cuda_device,
|
||||
"cuda_device_count": physical_device_count if physical_device_count > 0 else cuda_device_count,
|
||||
"cuda_device_count": device_count,
|
||||
"master_port": master_port,
|
||||
"local_worker_ports": worker_ports,
|
||||
}
|
||||
|
||||
|
||||
@@ -483,6 +494,8 @@ 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)
|
||||
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
# conftest.py — project-level pytest configuration.
|
||||
#
|
||||
# Problem: custom_nodes/ComfyUI-Distributed/__init__.py uses relative imports
|
||||
# (from .distributed import ...) that fail when pytest tries to import it as a
|
||||
# (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.
|
||||
|
||||
@@ -1,51 +0,0 @@
|
||||
"""
|
||||
ComfyUI-Distributed: thin entry point.
|
||||
All implementation lives in workers/, nodes/, api/.
|
||||
"""
|
||||
import atexit
|
||||
import os
|
||||
|
||||
import server
|
||||
|
||||
from .utils.config import ensure_config_exists
|
||||
from .utils.logging import debug_log
|
||||
from .utils.network import cleanup_client_session
|
||||
from .workers import get_worker_manager
|
||||
from .workers.startup import delayed_auto_launch, register_async_signals, sync_cleanup
|
||||
from .upscale.job_store import ensure_tile_jobs_initialized
|
||||
from .nodes import (
|
||||
NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS,
|
||||
ImageBatchDivider,
|
||||
DistributedCollectorNode,
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from . import api # noqa: F401 - triggers all @routes.* registrations
|
||||
from .api.queue_orchestration import ensure_distributed_state
|
||||
|
||||
ensure_config_exists()
|
||||
|
||||
# Aiohttp session cleanup
|
||||
async def _cleanup_session():
|
||||
await cleanup_client_session()
|
||||
|
||||
|
||||
atexit.register(lambda: None) # placeholder; real cleanup in sync_cleanup
|
||||
|
||||
# Initialize distributed job state on prompt_server
|
||||
prompt_server = server.PromptServer.instance
|
||||
ensure_distributed_state(prompt_server)
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# Worker startup
|
||||
if not os.environ.get('COMFYUI_IS_WORKER'):
|
||||
atexit.register(sync_cleanup)
|
||||
delayed_auto_launch()
|
||||
register_async_signals()
|
||||
@@ -0,0 +1,53 @@
|
||||
# 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.
|
||||
+1
-31
@@ -1,31 +1 @@
|
||||
from .utilities import (
|
||||
DistributedSeed,
|
||||
DistributedModelName,
|
||||
DistributedValue,
|
||||
ImageBatchDivider,
|
||||
AudioBatchDivider,
|
||||
DistributedEmptyImage,
|
||||
AnyType,
|
||||
ByPassTypeTuple,
|
||||
any_type,
|
||||
)
|
||||
from .collector import DistributedCollectorNode
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DistributedCollector": DistributedCollectorNode,
|
||||
"DistributedSeed": DistributedSeed,
|
||||
"DistributedModelName": DistributedModelName,
|
||||
"DistributedValue": DistributedValue,
|
||||
"ImageBatchDivider": ImageBatchDivider,
|
||||
"AudioBatchDivider": AudioBatchDivider,
|
||||
"DistributedEmptyImage": DistributedEmptyImage,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DistributedCollector": "Distributed Collector",
|
||||
"DistributedSeed": "Distributed Seed",
|
||||
"DistributedModelName": "Distributed Model Name",
|
||||
"DistributedValue": "Distributed Value",
|
||||
"ImageBatchDivider": "Image Batch Divider",
|
||||
"AudioBatchDivider": "Audio Segment Divider",
|
||||
"DistributedEmptyImage": "Distributed Empty Image",
|
||||
}
|
||||
"""Private execution helpers; public registration lives in nodes.v3."""
|
||||
|
||||
@@ -268,12 +268,3 @@ class UltimateSDUpscaleDistributed(
|
||||
|
||||
# Ensure initialization before registering routes
|
||||
ensure_tile_jobs_initialized()
|
||||
|
||||
# Node registration
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": UltimateSDUpscaleDistributed,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"UltimateSDUpscaleDistributed": "Ultimate SD Upscale Distributed (No Upscale)",
|
||||
}
|
||||
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
"""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]
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "ComfyUI-Distributed"
|
||||
description = "ComfyUI extension that enables multi-GPU processing locally, remotely and in the cloud"
|
||||
version = "1.4.7"
|
||||
version = "1.5.0"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = []
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Internal extension lifecycle support (not a node provider)."""
|
||||
@@ -0,0 +1,35 @@
|
||||
"""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}')
|
||||
@@ -106,7 +106,7 @@ def _load_worker_routes_module():
|
||||
sys.modules[f"{package_name}.utils"] = utils_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
workers_pkg.__path__ = [str(module_path.parents[1] / "workers")]
|
||||
workers_pkg.get_worker_manager = lambda: _DummyWorkerManager()
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
@@ -213,6 +213,7 @@ 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")
|
||||
@@ -269,6 +270,38 @@ 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}]}
|
||||
|
||||
+691
@@ -0,0 +1,691 @@
|
||||
{
|
||||
"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)"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,252 @@
|
||||
"""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())
|
||||
@@ -0,0 +1,21 @@
|
||||
"""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
|
||||
@@ -0,0 +1,113 @@
|
||||
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()
|
||||
@@ -1,11 +1,16 @@
|
||||
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
|
||||
@@ -21,7 +26,7 @@ def _load_process_module(module_filename: str):
|
||||
sys.modules[package_name] = root_pkg
|
||||
|
||||
workers_pkg = types.ModuleType(f"{package_name}.workers")
|
||||
workers_pkg.__path__ = []
|
||||
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")
|
||||
@@ -39,8 +44,23 @@ def _load_process_module(module_filename: str):
|
||||
|
||||
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,
|
||||
@@ -70,6 +90,145 @@ class ComfyRootDiscoveryTests(unittest.TestCase):
|
||||
|
||||
|
||||
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(
|
||||
@@ -130,5 +289,64 @@ class LaunchCommandBuilderTests(unittest.TestCase):
|
||||
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()
|
||||
|
||||
+9
-12
@@ -2,7 +2,6 @@ import { app } from "/scripts/app.js";
|
||||
import { ENDPOINTS } from "./constants.js";
|
||||
|
||||
const NODE_CLASS = "DistributedValue";
|
||||
const CONVERTED_WIDGET = "converted-widget";
|
||||
const DYNAMIC_DEFAULT_WIDGET = "_dv_default";
|
||||
const DYNAMIC_WORKER_WIDGET_PREFIX = "_dv_worker_";
|
||||
const WORKERS_CHANGED_EVENT = "distributed:workers-changed";
|
||||
@@ -42,15 +41,13 @@ function getDynamicWorkerWidgets(node) {
|
||||
return (node.widgets || []).filter((w) => w.name.startsWith(DYNAMIC_WORKER_WIDGET_PREFIX));
|
||||
}
|
||||
|
||||
function hideWidgetForGood(node, widget, suffix = "") {
|
||||
function hideWidgetForGood(widget) {
|
||||
if (!widget) return;
|
||||
if (typeof widget.type === "string" && widget.type.startsWith(CONVERTED_WIDGET)) return;
|
||||
|
||||
widget.origType = widget.type;
|
||||
widget.origComputeSize = widget.computeSize;
|
||||
widget.origSerializeValue = widget.serializeValue;
|
||||
widget.computeSize = () => [0, -4];
|
||||
widget.type = `${CONVERTED_WIDGET}${suffix}`;
|
||||
// These are serialized backing fields, not widgets converted into sockets.
|
||||
// Keep their native type/serializer and mutate the shared options in place
|
||||
// so both the classic canvas and Nodes 2.0 suppress the whole widget row.
|
||||
widget.options ??= {};
|
||||
widget.options.hidden = true;
|
||||
|
||||
// Hide any attached DOM element (multiline widgets).
|
||||
if (widget.element) widget.element.style.display = "none";
|
||||
@@ -58,14 +55,14 @@ function hideWidgetForGood(node, widget, suffix = "") {
|
||||
|
||||
if (widget.linkedWidgets) {
|
||||
for (const linked of widget.linkedWidgets) {
|
||||
hideWidgetForGood(node, linked, `:${widget.name}`);
|
||||
hideWidgetForGood(linked);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function hideRawWidgets(node) {
|
||||
hideWidgetForGood(node, getRawDefaultWidget(node), ":default_value");
|
||||
hideWidgetForGood(node, getRawWorkerValuesWidget(node), ":worker_values");
|
||||
hideWidgetForGood(getRawDefaultWidget(node));
|
||||
hideWidgetForGood(getRawWorkerValuesWidget(node));
|
||||
}
|
||||
|
||||
function removeDynamicDefaultWidget(node) {
|
||||
|
||||
+10
-1
@@ -42,6 +42,15 @@ export async function detectMasterIP(extension) {
|
||||
if (shouldAutoPopulate) {
|
||||
extension.log(`Auto-populating workers based on ${extension.cudaDeviceCount} CUDA devices (excluding master on CUDA ${extension.masterCudaDevice})`, "info");
|
||||
|
||||
const workerCount = extension.cudaDeviceCount - (
|
||||
Number.isInteger(extension.masterCudaDevice) &&
|
||||
extension.masterCudaDevice >= 0 &&
|
||||
extension.masterCudaDevice < extension.cudaDeviceCount ? 1 : 0
|
||||
);
|
||||
const workerPorts = data.local_worker_ports;
|
||||
if (!Array.isArray(workerPorts) || workerPorts.length !== workerCount) {
|
||||
throw new Error("Could not allocate enough available local worker ports; check the server log");
|
||||
}
|
||||
const newWorkers = [];
|
||||
let workerNum = 1;
|
||||
let portOffset = 0;
|
||||
@@ -56,7 +65,7 @@ export async function detectMasterIP(extension) {
|
||||
id: generateUUID(),
|
||||
name: `Worker ${workerNum}`,
|
||||
host: isRunpod ? null : "localhost",
|
||||
port: 8189 + portOffset,
|
||||
port: workerPorts[portOffset],
|
||||
cuda_device: i,
|
||||
enabled: true,
|
||||
extra_args: isRunpod ? "--listen" : "",
|
||||
|
||||
@@ -0,0 +1,120 @@
|
||||
import { readFileSync } from "node:fs";
|
||||
import vm from "node:vm";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
// Evaluate this repository's browser entrypoint without needing a ComfyUI server.
|
||||
// The real-renderer acceptance additionally checks the input rows in Chromium.
|
||||
const source = readFileSync(new URL("../distributedValue.js", import.meta.url), "utf8")
|
||||
.replace(/^import .*;\r?\n/gm, "");
|
||||
|
||||
function setup({ connected = false, targetType = "STRING" } = {}) {
|
||||
let extension;
|
||||
const listeners = new Map();
|
||||
const tasks = [];
|
||||
const workers = [
|
||||
{ id: "gpu-a", name: "GPU A", enabled: true },
|
||||
{ id: "gpu-b", name: "GPU B", enabled: true },
|
||||
{ id: "off", name: "Disabled", enabled: false },
|
||||
];
|
||||
const target = {
|
||||
inputs: [{ name: "value" }],
|
||||
widgets: [{ name: "value", type: targetType === "INT" ? "number" : "string", options: { step: 1, precision: 0 } }],
|
||||
};
|
||||
const graph = { links: { 1: { target_id: 2, target_slot: 0 } }, getNodeById: () => target };
|
||||
const raw = ["default_value", "worker_values"].map((name, index) => ({
|
||||
name,
|
||||
type: "text",
|
||||
value: index === 0 ? "saved default" : '{"1":"saved worker"}',
|
||||
options: { multiline: false },
|
||||
computeSize: vi.fn(() => [200, 20]),
|
||||
serializeValue: vi.fn(function () { return this.value; }),
|
||||
}));
|
||||
const node = {
|
||||
comfyClass: "DistributedValue", graph, widgets: [...raw], size: [200, 100],
|
||||
inputs: raw.map(w => ({ name: w.name, widget: { name: w.name }, link: null })),
|
||||
outputs: [{ links: connected ? [1] : [] }],
|
||||
computeSize: () => [200, 100], setSize: vi.fn(), setDirtyCanvas: vi.fn(),
|
||||
addWidget(type, name, value, callback, options) {
|
||||
const widget = { type, name, value, callback, options };
|
||||
this.widgets.push(widget);
|
||||
return widget;
|
||||
},
|
||||
configure(data) {
|
||||
data.widgets_values.forEach((value, i) => { this.widgets[i].value = value; });
|
||||
return "configured";
|
||||
},
|
||||
};
|
||||
vm.runInNewContext(source, {
|
||||
app: { graph, registerExtension: ext => { extension = ext; } },
|
||||
ENDPOINTS: { CONFIG: "/distributed/config" },
|
||||
fetch: async () => ({ ok: true, json: async () => ({ workers }) }),
|
||||
window: { addEventListener: (name, handler) => listeners.set(name, handler) },
|
||||
setTimeout: callback => { tasks.push(callback); }, console,
|
||||
}, { filename: "distributedValue.js" });
|
||||
return { node, raw, extension, listeners, tasks, flush: () => { while (tasks.length) tasks.shift()(); } };
|
||||
}
|
||||
|
||||
describe("Distributed Value internal widget visibility", () => {
|
||||
it("hides raw fields natively without converting their types or replacing their options", async () => {
|
||||
const h = setup();
|
||||
const options = h.raw.map(w => w.options);
|
||||
const sizes = h.raw.map(w => w.computeSize);
|
||||
const serializers = h.raw.map(w => w.serializeValue);
|
||||
const slots = h.node.inputs;
|
||||
await h.extension.nodeCreated(h.node);
|
||||
|
||||
for (const [i, widget] of h.raw.entries()) {
|
||||
expect(widget.type).toBe("text");
|
||||
expect(widget.options).toBe(options[i]);
|
||||
expect(widget.options.hidden).toBe(true);
|
||||
expect(widget.computeSize).toBe(sizes[i]);
|
||||
expect(widget.serializeValue).toBe(serializers[i]);
|
||||
expect(h.node.widgets[i]).toBe(widget);
|
||||
}
|
||||
expect(h.node.inputs).toBe(slots);
|
||||
expect(h.node.widgets.slice(2).map(w => w.label)).toEqual(["default_value", "GPU A", "GPU B"]);
|
||||
expect(h.node.widgets.slice(2).every(w => !w.options.hidden)).toBe(true);
|
||||
});
|
||||
|
||||
it("keeps typed edits in the original serialized fields across rebuild/configure", async () => {
|
||||
const h = setup({ connected: true, targetType: "INT" });
|
||||
await h.extension.nodeCreated(h.node);
|
||||
h.node.widgets.find(w => w.name === "_dv_default").callback(42);
|
||||
h.node.widgets.find(w => w.name === "_dv_worker_1").callback(7);
|
||||
h.node.widgets.find(w => w.name === "_dv_worker_2").callback(9);
|
||||
const values = h.raw.map(w => w.serializeValue());
|
||||
expect(values[0]).toBe(42);
|
||||
expect(JSON.parse(values[1])).toMatchObject({
|
||||
_type: "INT", 1: "7", 2: "9", _by_worker_id: { "gpu-a": "7", "gpu-b": "9" },
|
||||
});
|
||||
expect(h.node.configure({ widgets_values: values })).toBe("configured");
|
||||
h.flush();
|
||||
expect(h.raw.map(w => w.serializeValue())).toEqual(values);
|
||||
expect(h.raw.every(w => w.options.hidden && w.type === "text")).toBe(true);
|
||||
expect(h.node.widgets.slice(2).map(w => w.value)).toEqual([42, 7, 9]);
|
||||
});
|
||||
|
||||
it("reapplies visibility on worker refresh without duplicating raw or dynamic fields", async () => {
|
||||
const h = setup();
|
||||
await h.extension.nodeCreated(h.node);
|
||||
h.raw[0].options.hidden = false;
|
||||
await h.listeners.get("distributed:workers-changed")({ detail: { workers: [
|
||||
{ id: "gpu-b", name: "Renamed B", enabled: true },
|
||||
] } });
|
||||
expect(h.raw.every(w => w.options.hidden)).toBe(true);
|
||||
expect(h.node.widgets.map(w => w.name)).toEqual(["default_value", "worker_values", "_dv_default", "_dv_worker_1"]);
|
||||
expect(h.node.widgets[3].label).toBe("Renamed B");
|
||||
});
|
||||
|
||||
it("hides linked backing widgets and handles absent options without altering serializers", async () => {
|
||||
const h = setup();
|
||||
delete h.raw[1].options;
|
||||
const linked = { type: "text", options: {}, serializeValue: vi.fn(() => "linked") };
|
||||
h.raw[0].linkedWidgets = [linked];
|
||||
await h.extension.nodeCreated(h.node);
|
||||
expect(h.raw[1].options.hidden).toBe(true);
|
||||
expect(linked.options.hidden).toBe(true);
|
||||
expect(linked.type).toBe("text");
|
||||
expect(linked.serializeValue()).toBe("linked");
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,65 @@
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { detectMasterIP } from "../masterDetection.js";
|
||||
|
||||
describe("automatic local worker ports", () => {
|
||||
let originalWindow;
|
||||
beforeEach(() => {
|
||||
originalWindow = globalThis.window;
|
||||
globalThis.window = { location: { hostname: "localhost", port: "443" } };
|
||||
});
|
||||
afterEach(() => { globalThis.window = originalWindow; });
|
||||
|
||||
function makeExtension(overrides = {}) {
|
||||
return {
|
||||
config: { workers: [], settings: {}, master: {} },
|
||||
log: vi.fn(),
|
||||
api: {
|
||||
getNetworkInfo: vi.fn().mockResolvedValue({
|
||||
cuda_device: 0, cuda_device_count: 3,
|
||||
master_port: 8189, local_worker_ports: [8191, 8193],
|
||||
...overrides,
|
||||
}),
|
||||
updateMaster: vi.fn().mockResolvedValue({}),
|
||||
updateWorker: vi.fn().mockResolvedValue({}),
|
||||
updateSetting: vi.fn().mockResolvedValue({}),
|
||||
},
|
||||
ui: { updateMasterDisplay: vi.fn() },
|
||||
app: {},
|
||||
loadConfig: vi.fn().mockResolvedValue(),
|
||||
};
|
||||
}
|
||||
|
||||
it("uses server-allocated ports, not hardcoded ports or the browser proxy port", async () => {
|
||||
const extension = makeExtension();
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers.map(worker => worker.port)).toEqual([8191, 8193]);
|
||||
expect(extension.config.workers.map(worker => worker.cuda_device)).toEqual([1, 2]);
|
||||
expect(extension.api.updateWorker).toHaveBeenCalledTimes(2);
|
||||
});
|
||||
|
||||
it("does not overwrite an existing manually configured worker", async () => {
|
||||
const extension = makeExtension();
|
||||
const worker = { id: "manual", host: "localhost", port: 8189 };
|
||||
extension.config.workers = [worker];
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers).toEqual([worker]);
|
||||
expect(extension.api.updateWorker).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("fails without partially saving workers if not enough ports are available", async () => {
|
||||
const extension = makeExtension({ local_worker_ports: [8191] });
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers).toEqual([]);
|
||||
expect(extension.api.updateWorker).not.toHaveBeenCalled();
|
||||
expect(extension.api.updateSetting).not.toHaveBeenCalled();
|
||||
expect(extension.log).toHaveBeenCalledWith(expect.stringContaining("port"), "error");
|
||||
});
|
||||
|
||||
it("preserves Runpod worker host and launch settings", async () => {
|
||||
globalThis.window.location.hostname = "pod-8189.proxy.runpod.net";
|
||||
const extension = makeExtension();
|
||||
await detectMasterIP(extension);
|
||||
expect(extension.config.workers.map(worker => worker.port)).toEqual([8191, 8193]);
|
||||
expect(extension.config.workers.every(worker => worker.host === null && worker.extra_args === "--listen")).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Host-local worker port selection and launch preflight (no ComfyUI imports)."""
|
||||
import errno
|
||||
import os
|
||||
import socket
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
|
||||
def _port_number(value):
|
||||
try:
|
||||
port = int(value)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError("Worker port must be an integer between 1 and 65535") from None
|
||||
if isinstance(value, (bool, float)) or not 1 <= port <= 65535:
|
||||
raise ValueError("Worker port must be an integer between 1 and 65535")
|
||||
return port
|
||||
|
||||
|
||||
def _local_config(worker):
|
||||
local_hosts = ("localhost", "127.0.0.1", "0.0.0.0", "::1", "::")
|
||||
host = (worker.get("host") or "localhost").strip().lower()
|
||||
if worker.get("type") == "local" or host in local_hosts:
|
||||
return True
|
||||
try:
|
||||
host = urlsplit(host if "://" in host else "//" + host).hostname
|
||||
except ValueError:
|
||||
return False
|
||||
return host in local_hosts
|
||||
|
||||
|
||||
def _reserved_ports(workers, exclude_id=None):
|
||||
reserved = {}
|
||||
for worker in workers:
|
||||
if exclude_id is not None and str(worker.get("id")) == str(exclude_id):
|
||||
continue
|
||||
if not _local_config(worker):
|
||||
continue
|
||||
try:
|
||||
port = _port_number(worker.get("port"))
|
||||
except ValueError:
|
||||
continue
|
||||
reserved[port] = worker.get("name") or worker.get("id") or "another local worker"
|
||||
return reserved
|
||||
|
||||
|
||||
def is_port_available(port):
|
||||
"""Match asyncio's address reuse while rejecting active IPv4/IPv6 listeners."""
|
||||
addresses = [(socket.AF_INET, "0.0.0.0")]
|
||||
if socket.has_ipv6:
|
||||
addresses.append((socket.AF_INET6, "::"))
|
||||
for family, address in addresses:
|
||||
try:
|
||||
with socket.socket(family, socket.SOCK_STREAM) as probe:
|
||||
if os.name != "nt":
|
||||
# Like asyncio, allow a restart while old connections are
|
||||
# in TIME_WAIT. Do not enable SO_REUSEPORT.
|
||||
probe.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
elif hasattr(socket, "SO_EXCLUSIVEADDRUSE"):
|
||||
# Windows SO_REUSEADDR can steal an active listener's port.
|
||||
probe.setsockopt(socket.SOL_SOCKET, socket.SO_EXCLUSIVEADDRUSE, 1)
|
||||
probe.bind((address, port))
|
||||
probe.listen(1)
|
||||
except OSError as exc:
|
||||
if family == socket.AF_INET6 and exc.errno in (errno.EAFNOSUPPORT, errno.EADDRNOTAVAIL):
|
||||
continue
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def allocate_worker_ports(master_port, workers, count):
|
||||
"""Select ports above the actual master, skipping reservations/listeners.
|
||||
|
||||
This is a snapshot, not a reservation. Launch must validate again because
|
||||
another process can claim a suggested port before the user starts a worker.
|
||||
"""
|
||||
master_port = _port_number(master_port)
|
||||
reserved = _reserved_ports(workers)
|
||||
selected = []
|
||||
if count == 0:
|
||||
return selected
|
||||
for port in range(master_port + 1, 65536):
|
||||
if port not in reserved and is_port_available(port):
|
||||
selected.append(port)
|
||||
if len(selected) == count:
|
||||
return selected
|
||||
raise ValueError(f"Not enough available worker ports above master port {master_port}; need {count}")
|
||||
|
||||
|
||||
def validate_worker_port(worker, master_port, workers):
|
||||
"""Reject conflicts without silently changing a configured worker port."""
|
||||
port = _port_number(worker.get("port"))
|
||||
if port == _port_number(master_port):
|
||||
raise ValueError(f"Worker port {port} conflicts with the master port {master_port}; choose another worker port")
|
||||
reserved = _reserved_ports(workers, exclude_id=worker.get("id"))
|
||||
if port in reserved:
|
||||
raise ValueError(f"Worker port {port} is assigned to local worker {reserved[port]}; choose another worker port")
|
||||
if not is_port_available(port):
|
||||
raise ValueError(f"Worker port {port} is already in use on this host; choose another worker port")
|
||||
@@ -1,4 +1,5 @@
|
||||
import glob
|
||||
import hashlib
|
||||
import os
|
||||
import shlex
|
||||
import shutil
|
||||
@@ -87,6 +88,51 @@ class LaunchCommandBuilder:
|
||||
return wt_path
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _arg_value(cmd, flag):
|
||||
"""Return the last value, matching argparse's override semantics."""
|
||||
value = None
|
||||
for index, arg in enumerate(cmd):
|
||||
if arg == flag and index + 1 < len(cmd):
|
||||
value = cmd[index + 1]
|
||||
elif arg.startswith(flag + "="):
|
||||
value = arg.split("=", 1)[1]
|
||||
return value
|
||||
|
||||
def _add_worker_database(self, cmd, worker_config, comfy_root):
|
||||
args = self._get_runtime_args()
|
||||
# The parsed namespace advertises support without launching ComfyUI
|
||||
# (or importing another checkout's argument parser).
|
||||
if args is None or not hasattr(args, "database_url"):
|
||||
return
|
||||
if any(arg.split("=", 1)[0] in ("--database-url", "--disable-assets") for arg in cmd):
|
||||
return
|
||||
|
||||
user_directory = self._arg_value(cmd, "--user-directory")
|
||||
base_directory = self._arg_value(cmd, "--base-directory")
|
||||
if not user_directory:
|
||||
user_directory = os.path.join(base_directory or comfy_root, "user")
|
||||
if not os.path.isabs(user_directory):
|
||||
# Popen uses comfy_root as cwd, which can differ from the master's.
|
||||
user_directory = os.path.join(comfy_root, user_directory)
|
||||
database_directory = os.path.abspath(os.path.join(user_directory, "distributed", "workers"))
|
||||
# Hash the immutable ID: names/GPU/ports may change, and IDs must not
|
||||
# become filesystem paths or collide after filename sanitization.
|
||||
worker_id = str(worker_config["id"])
|
||||
filename = hashlib.sha256(worker_id.encode("utf-8")).hexdigest() + ".db"
|
||||
database_path = os.path.join(database_directory, filename).replace(os.sep, "/")
|
||||
# Use the same URL codec as ComfyUI; raw '?' and '%xx' can otherwise
|
||||
# truncate/decode the path and make distinct workers share a database.
|
||||
from sqlalchemy.engine import URL, make_url
|
||||
database_url = URL.create("sqlite", database=database_path).render_as_string()
|
||||
if make_url(database_url).database != database_path:
|
||||
raise ValueError(
|
||||
"Installed SQLAlchemy cannot safely encode the worker database path; "
|
||||
"use a user directory without '?' or an explicit --database-url with a safe path"
|
||||
)
|
||||
os.makedirs(database_directory, exist_ok=True)
|
||||
cmd.extend(["--database-url", database_url])
|
||||
|
||||
def build_launch_command(self, worker_config, comfy_root):
|
||||
"""Build the command to launch a worker."""
|
||||
main_py = os.path.join(comfy_root, "main.py")
|
||||
@@ -141,4 +187,7 @@ class LaunchCommandBuilder:
|
||||
raise ValueError(f"Invalid characters in extra_args: {arg}. Forbidden: {forbidden}")
|
||||
cmd.extend(extra_args_list)
|
||||
|
||||
if int(self._arg_value(cmd, "--port")) != int(worker_config["port"]):
|
||||
raise ValueError("extra_args must not override the configured worker port; set the port in the UI/config instead")
|
||||
self._add_worker_database(cmd, worker_config, comfy_root)
|
||||
return cmd
|
||||
|
||||
@@ -7,7 +7,9 @@ import time
|
||||
from ...utils.config import load_config, save_config
|
||||
from ...utils.constants import PROCESS_TERMINATION_TIMEOUT, PROCESS_WAIT_TIMEOUT, WORKER_CHECK_INTERVAL
|
||||
from ...utils.logging import debug_log, log
|
||||
from ...utils.network import get_server_port
|
||||
from ...utils.process import get_python_executable, is_process_alive, terminate_process
|
||||
from ..ports import validate_worker_port
|
||||
|
||||
try:
|
||||
import psutil
|
||||
@@ -28,6 +30,17 @@ class ProcessLifecycle:
|
||||
"""Launch a worker process with logging."""
|
||||
_ = show_window # Kept for API compatibility.
|
||||
comfy_root = self._manager.find_comfy_root()
|
||||
config = load_config()
|
||||
try:
|
||||
master_port = get_server_port()
|
||||
except AttributeError:
|
||||
master_port = None
|
||||
if master_port is None:
|
||||
# The auto-launch timer can fire during custom-node imports, before
|
||||
# PromptServer assigns its port. Still reserve the configured port.
|
||||
from comfy.cli_args import args
|
||||
master_port = args.port
|
||||
validate_worker_port(worker_config, master_port, config.get("workers", []))
|
||||
|
||||
env = os.environ.copy()
|
||||
env["CUDA_VISIBLE_DEVICES"] = str(worker_config.get("cuda_device", 0))
|
||||
|
||||
Reference in New Issue
Block a user