Compare commits

..
Author SHA1 Message Date
Robert Wojciechowski 539f94ba39 fix: harden local worker startup preflight and database paths
Handle auto-launch before the master binds, allow TIME_WAIT restarts without accepting active listeners, resolve relative paths in the worker cwd, and round-trip SQLite URLs through SQLAlchemy. Add regressions and remove the README additions.
2026-10-10 22:13:13 +00:00
Robert Wojciechowski fbfb411402 fix: isolate local worker databases and avoid port conflicts 2026-10-10 21:59:18 +00:00
15 changed files with 709 additions and 19 deletions
+14 -1
View File
@@ -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)
+15
View File
@@ -164,11 +164,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)
+34 -1
View 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}]}
+22 -10
View File
@@ -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):
+13 -2
View File
@@ -302,8 +302,8 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
self.assertNotIn("1", result)
self.assertNotIn("2", result)
def test_replaces_dangling_upstream_ref_with_one_empty_image_placeholder(self):
"""Collector must point to exactly one valid placeholder after pruning."""
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", {})
@@ -314,6 +314,10 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
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]
@@ -380,6 +384,13 @@ class PrepareDelegateMasterPromptTests(unittest.TestCase):
self.assertEqual(result["15"], prompt["15"])
self.assertEqual(result["9"]["inputs"]["filename_prefix"], ["15", 0])
def test_does_not_preserve_non_primitive_upstream_for_collector(self):
prompt = _delegate_prompt()
result = pt.prepare_delegate_master_prompt(prompt, ["3"])
self.assertNotIn("2", result)
self.assertNotEqual(result["3"]["inputs"]["images"], ["2", 0])
def test_preserves_load_image_for_switch_alternate_required_input(self):
"""Delegate-only master keeps LoadImage inputs needed by switches."""
prompt = {
+113
View File
@@ -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()
+219 -1
View File
@@ -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()
+10 -1
View File
@@ -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" : "",
+15
View File
@@ -0,0 +1,15 @@
import { describe, expect, it } from "vitest";
import { buildWorkerWebSocketUrl } from "../urlUtils.js";
describe("execution decision helpers", () => {
it("buildWorkerWebSocketUrl converts http/https to ws/wss", () => {
expect(buildWorkerWebSocketUrl("http://worker.local:8188")).toBe(
"ws://worker.local:8188/distributed/worker_ws"
);
expect(buildWorkerWebSocketUrl("https://worker.example.com")).toBe(
"wss://worker.example.com/distributed/worker_ws"
);
});
});
+65
View File
@@ -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);
});
});
+5
View File
@@ -106,6 +106,11 @@ describe("buildWorkerWebSocketUrl", () => {
"wss://worker.example.com/distributed/worker_ws"
);
});
it("always appends /distributed/worker_ws", () => {
const url = buildWorkerWebSocketUrl("http://worker.local:8188");
expect(url.endsWith("/distributed/worker_ws")).toBe(true);
});
});
+25 -3
View File
@@ -3,7 +3,7 @@ import { afterEach, beforeEach, describe, expect, it } from "vitest";
import { getWorkerUrl } from "../workerLifecycle.js";
describe("workerLifecycle URL wiring", () => {
describe("workerLifecycle URL construction", () => {
let originalWindow;
beforeEach(() => {
@@ -22,9 +22,31 @@ describe("workerLifecycle URL wiring", () => {
globalThis.window = originalWindow;
});
// URL variants are covered in urlUtils.test.js; keep the wrapper wiring here.
it("forwards the worker and endpoint using window.location", () => {
it("builds local worker URL with explicit local port", () => {
const worker = { id: "w1", port: 8189, type: "local" };
expect(getWorkerUrl({}, worker, "/prompt")).toBe("http://127.0.0.1:8189/prompt");
});
it("builds remote worker URL with host:port", () => {
const worker = { id: "w2", host: "worker.example.com", port: 9000, type: "remote" };
expect(getWorkerUrl({}, worker, "/prompt")).toBe("http://worker.example.com:9000/prompt");
});
it("builds cloud worker URL as https", () => {
const worker = { id: "w3", host: "cloud.example.com", port: 443, type: "cloud" };
expect(getWorkerUrl({}, worker, "/prompt")).toBe("https://cloud.example.com/prompt");
});
it("rewrites runpod proxy hostname for local worker ports", () => {
globalThis.window = {
location: {
hostname: "podabc.proxy.runpod.net",
protocol: "https:",
port: "",
origin: "https://podabc.proxy.runpod.net",
},
};
const worker = { id: "w4", port: 8189, type: "local" };
expect(getWorkerUrl({}, worker, "/prompt")).toBe("https://podabc-8189.proxy.runpod.net/prompt");
});
});
+97
View File
@@ -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")
+49
View File
@@ -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
+13
View File
@@ -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))