Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a2b341fb99 | ||
|
|
9de7e158e1 |
+1
-14
@@ -16,14 +16,12 @@ 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,
|
||||
@@ -279,13 +277,6 @@ 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)
|
||||
@@ -294,9 +285,7 @@ def _collect_network_info_sync():
|
||||
"all_ips": all_ips,
|
||||
"recommended_ip": recommended_ip,
|
||||
"cuda_device": cuda_device,
|
||||
"cuda_device_count": device_count,
|
||||
"master_port": master_port,
|
||||
"local_worker_ports": worker_ports,
|
||||
"cuda_device_count": physical_device_count if physical_device_count > 0 else cuda_device_count,
|
||||
}
|
||||
|
||||
|
||||
@@ -494,8 +483,6 @@ async def launch_worker_endpoint(request):
|
||||
"message": f"Worker {worker['name']} launched",
|
||||
"log_file": log_file
|
||||
})
|
||||
except ValueError as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 400)
|
||||
except Exception as e:
|
||||
return await handle_api_error(request, f"Failed to launch worker: {str(e)}", 500)
|
||||
|
||||
|
||||
+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.5.0"
|
||||
version = "1.4.7"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = []
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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__ = [str(module_path.parents[1] / "workers")]
|
||||
workers_pkg.__path__ = []
|
||||
workers_pkg.get_worker_manager = lambda: _DummyWorkerManager()
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
@@ -213,7 +213,6 @@ def _load_worker_routes_module():
|
||||
raise RuntimeError("not used in these tests")
|
||||
|
||||
network_module.get_client_session = _get_client_session
|
||||
network_module.get_server_port = lambda: 8189
|
||||
sys.modules[f"{package_name}.utils.network"] = network_module
|
||||
|
||||
constants_module = types.ModuleType(f"{package_name}.utils.constants")
|
||||
@@ -270,38 +269,6 @@ worker_routes = _load_worker_routes_module()
|
||||
|
||||
|
||||
class WorkerRoutesTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_network_info_reports_actual_master_and_available_ports(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 1, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=["127.0.0.1"]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": []}), \
|
||||
patch.object(worker_routes, "get_server_port", return_value=8189), \
|
||||
patch.object(worker_routes, "allocate_worker_ports", return_value=[8191, 8192, 8193]) as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["master_port"], 8189)
|
||||
self.assertEqual(response.payload["local_worker_ports"], [8191, 8192, 8193])
|
||||
self.assertEqual(response.payload["cuda_device_count"], 4)
|
||||
allocate.assert_called_once_with(8189, [], 3)
|
||||
|
||||
async def test_network_info_does_not_reallocate_existing_configuration(self):
|
||||
with patch.object(worker_routes, "_get_cuda_info", return_value=(0, 4, 4)), \
|
||||
patch.object(worker_routes, "get_network_ips", return_value=[]), \
|
||||
patch.object(worker_routes, "load_config", return_value={"workers": [{"id": "existing", "port": 8189}]}), \
|
||||
patch.object(worker_routes, "allocate_worker_ports") as allocate:
|
||||
response = await worker_routes.get_network_info_endpoint(_FakeRequest())
|
||||
self.assertEqual(response.payload["local_worker_ports"], [])
|
||||
allocate.assert_not_called()
|
||||
|
||||
async def test_launch_conflict_is_a_clear_client_error(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8189}]}
|
||||
with patch.object(worker_routes, "get_worker_manager", return_value=manager), \
|
||||
patch.object(worker_routes, "load_config", return_value=config), \
|
||||
patch.object(manager, "launch_worker", side_effect=ValueError("Worker port 8189 conflicts with the master port 8189")):
|
||||
response = await worker_routes.launch_worker_endpoint(_FakeRequest({"worker_id": "worker-a"}))
|
||||
self.assertEqual(response.status, 400)
|
||||
self.assertIn("master port 8189", response.payload["message"])
|
||||
self.assertEqual(manager.processes, {})
|
||||
|
||||
async def test_launch_worker_valid_id_returns_200(self):
|
||||
manager = _DummyWorkerManager()
|
||||
config = {"workers": [{"id": "worker-a", "name": "Worker A", "port": 8188}]}
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
import importlib.util
|
||||
import os
|
||||
import socket
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
def load_ports():
|
||||
path = Path(__file__).resolve().parents[1] / "workers" / "ports.py"
|
||||
spec = importlib.util.spec_from_file_location("worker_ports_test", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
class WorkerPortTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.ports = load_ports()
|
||||
|
||||
def test_allocation_starts_above_actual_master_and_skips_assigned_and_occupied(self):
|
||||
workers = [
|
||||
{"id": "existing", "host": "localhost", "port": 8190, "enabled": False},
|
||||
{"id": "remote", "host": "remote.example", "port": 8192},
|
||||
]
|
||||
with patch.object(self.ports, "is_port_available", side_effect=lambda port: port != 8191):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, workers, 3), [8192, 8193, 8194])
|
||||
|
||||
def test_local_host_forms_reserve_ports(self):
|
||||
for host in ["::1", "[::1]", "http://localhost/", "HTTPS://LOCALHOST", "0.0.0.0", None]:
|
||||
with self.subTest(host=host), patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.assertEqual(self.ports.allocate_worker_ports(8189, [{"id": "local", "host": host, "port": 8190}], 1), [8191])
|
||||
|
||||
def test_exhaustion_is_explicit_not_a_partial_allocation(self):
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
with self.assertRaisesRegex(ValueError, "available.*ports"):
|
||||
self.ports.allocate_worker_ports(65534, [], 2)
|
||||
|
||||
def test_launch_rejects_master_port_without_reassigning_configuration(self):
|
||||
worker = {"id": "manual", "port": 8189}
|
||||
with self.assertRaisesRegex(ValueError, "master.*8189"):
|
||||
self.ports.validate_worker_port(worker, 8189, [])
|
||||
self.assertEqual(worker["port"], 8189)
|
||||
|
||||
def test_launch_rejects_another_local_workers_reserved_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "other", "host": None, "port": 8190, "enabled": False}
|
||||
with self.assertRaisesRegex(ValueError, "other"):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_ignores_itself_and_remote_workers_with_same_port(self):
|
||||
worker = {"id": "manual", "port": 8190}
|
||||
other = {"id": "remote", "host": "example.com", "port": 8190}
|
||||
with patch.object(self.ports, "is_port_available", return_value=True):
|
||||
self.ports.validate_worker_port(worker, 8189, [worker, other])
|
||||
|
||||
def test_launch_rejects_an_occupied_port(self):
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener:
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen()
|
||||
port = listener.getsockname()[1]
|
||||
self.assertFalse(self.ports.is_port_available(port))
|
||||
with self.assertRaisesRegex(ValueError, "already in use"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
@unittest.skipIf(os.name == "nt", "asyncio does not reuse addresses on Windows")
|
||||
def test_recently_closed_connection_does_not_block_worker_restart(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
port = listener.getsockname()[1]
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=2) as client:
|
||||
accepted, _ = listener.accept()
|
||||
with accepted:
|
||||
accepted.settimeout(2)
|
||||
accepted.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(client.recv(1), b"")
|
||||
client.shutdown(socket.SHUT_WR)
|
||||
self.assertEqual(accepted.recv(1), b"")
|
||||
self.assertTrue(self.ports.is_port_available(port))
|
||||
self.ports.validate_worker_port({"id": "restart", "port": port}, 1, [])
|
||||
with socket.socket() as restarted:
|
||||
restarted.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
restarted.bind(("127.0.0.1", port))
|
||||
restarted.listen(1)
|
||||
|
||||
def test_active_listener_with_reuseaddr_is_still_a_conflict(self):
|
||||
with socket.socket() as listener:
|
||||
listener.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
@unittest.skipUnless(socket.has_ipv6, "IPv6 unavailable")
|
||||
def test_ipv6_only_listener_is_still_a_conflict(self):
|
||||
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as listener:
|
||||
listener.setsockopt(socket.IPPROTO_IPV6, socket.IPV6_V6ONLY, 1)
|
||||
try:
|
||||
listener.bind(("::1", 0))
|
||||
except OSError as exc:
|
||||
self.skipTest(f"IPv6 loopback unavailable: {exc}")
|
||||
listener.listen(1)
|
||||
self.assertFalse(self.ports.is_port_available(listener.getsockname()[1]))
|
||||
|
||||
def test_invalid_ports_fail_clearly(self):
|
||||
for port in [0, 65536, "not-a-port", None, True, 8190.5]:
|
||||
with self.subTest(port=port), self.assertRaisesRegex(ValueError, "port"):
|
||||
self.ports.validate_worker_port({"id": "manual", "port": port}, 8189, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,16 +1,11 @@
|
||||
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
|
||||
@@ -26,7 +21,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__ = [str(module_path.parents[1])]
|
||||
workers_pkg.__path__ = []
|
||||
sys.modules[f"{package_name}.workers"] = workers_pkg
|
||||
|
||||
process_pkg = types.ModuleType(f"{package_name}.workers.process")
|
||||
@@ -44,23 +39,8 @@ 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,
|
||||
@@ -90,145 +70,6 @@ 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(
|
||||
@@ -289,64 +130,5 @@ 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()
|
||||
|
||||
+1
-10
@@ -42,15 +42,6 @@ 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;
|
||||
@@ -65,7 +56,7 @@ export async function detectMasterIP(extension) {
|
||||
id: generateUUID(),
|
||||
name: `Worker ${workerNum}`,
|
||||
host: isRunpod ? null : "localhost",
|
||||
port: workerPorts[portOffset],
|
||||
port: 8189 + portOffset,
|
||||
cuda_device: i,
|
||||
enabled: true,
|
||||
extra_args: isRunpod ? "--listen" : "",
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
});
|
||||
});
|
||||
@@ -1,65 +0,0 @@
|
||||
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);
|
||||
});
|
||||
});
|
||||
@@ -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);
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
|
||||
@@ -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");
|
||||
});
|
||||
});
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
"""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,5 +1,4 @@
|
||||
import glob
|
||||
import hashlib
|
||||
import os
|
||||
import shlex
|
||||
import shutil
|
||||
@@ -88,51 +87,6 @@ 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")
|
||||
@@ -187,7 +141,4 @@ 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,9 +7,7 @@ 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
|
||||
@@ -30,17 +28,6 @@ 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