From 3fa23e1910b25b2f021065a66fb2aa6394ad0ec3 Mon Sep 17 00:00:00 2001 From: Herman Brunborg Date: Sun, 31 May 2026 12:14:36 -0700 Subject: [PATCH] Add Cloudflare quick tunnel support for manager URLs - Introduce a manager connection abstraction with direct and cloudflared-backed implementations - Advertise a public trycloudflare URL to workers while using loopback locally - Update Slurm backend defaults and add coverage for tunnel startup and shutdown --- src/furu/execution/connection.py | 205 ++++++++++++++++++++++ src/furu/execution/manager.py | 3 + src/furu/execution/server.py | 83 +++++++-- src/furu/worker/backends/slurm/backend.py | 39 +++- tests/test_manager_connection.py | 146 +++++++++++++++ tests/test_slurm_backend.py | 67 +++++++ tests/test_worker_manager.py | 164 ++++++++++++++++- 7 files changed, 681 insertions(+), 26 deletions(-) create mode 100644 src/furu/execution/connection.py create mode 100644 tests/test_manager_connection.py diff --git a/src/furu/execution/connection.py b/src/furu/execution/connection.py new file mode 100644 index 00000000..ec9f9cb8 --- /dev/null +++ b/src/furu/execution/connection.py @@ -0,0 +1,205 @@ +from __future__ import annotations + +import re +import shutil +import subprocess +import threading +import time +from collections import deque +from collections.abc import Iterator +from contextlib import AbstractContextManager, contextmanager, nullcontext +from dataclasses import dataclass, field +from queue import Empty, SimpleQueue +from typing import IO, Protocol + + +class ManagerConnection(Protocol): + def connect(self, *, local_url: str) -> AbstractContextManager[str]: ... + + +@dataclass(frozen=True, slots=True) +class DirectManagerConnection: + def connect(self, *, local_url: str) -> AbstractContextManager[str]: + return nullcontext(local_url) + + +@dataclass(frozen=True, slots=True) +class CloudflareQuickTunnel: + command: tuple[str, ...] = ("cloudflared",) + startup_timeout: float = 30.0 + extra_args: tuple[str, ...] = field(default_factory=tuple) + + def connect(self, *, local_url: str) -> AbstractContextManager[str]: + return _cloudflare_quick_tunnel( + command=self.command, + startup_timeout=self.startup_timeout, + extra_args=self.extra_args, + local_url=local_url, + ) + + +_TRYCLOUDFLARE_URL_RE = re.compile(r"https://[-a-zA-Z0-9.]+\.trycloudflare\.com") + + +@contextmanager +def _cloudflare_quick_tunnel( + *, + command: tuple[str, ...], + startup_timeout: float, + extra_args: tuple[str, ...], + local_url: str, +) -> Iterator[str]: + if not command: + raise ValueError("cloudflared command must not be empty") + + executable = command[0] + if shutil.which(executable) is None: + raise RuntimeError( + f"could not find {executable!r} on PATH; install cloudflared or configure " + "the manager to use a different connection method" + ) + + args = [ + *command, + "tunnel", + *extra_args, + "--url", + local_url, + ] + process = subprocess.Popen( + args, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + encoding="utf-8", + errors="replace", + bufsize=1, + ) + output = _CapturedOutput() + output_reader = _start_output_reader(process.stdout, output) + + try: + public_url = _wait_for_trycloudflare_url( + process=process, + output=output, + output_reader=output_reader, + startup_timeout=startup_timeout, + ) + except BaseException: + _terminate_process(process, output_reader=output_reader, timeout=5.0) + raise + + try: + yield public_url + finally: + _terminate_process(process, output_reader=output_reader, timeout=5.0) + + +class _CapturedOutput: + def __init__(self) -> None: + self._lines: deque[str] = deque(maxlen=50) + self._queue: SimpleQueue[str] = SimpleQueue() + self._lock = threading.Lock() + + def append(self, line: str) -> None: + with self._lock: + self._lines.append(line) + self._queue.put(line) + + def get(self, *, timeout: float) -> str: + return self._queue.get(timeout=timeout) + + def get_nowait(self) -> str: + return self._queue.get_nowait() + + def recent_text(self) -> str: + with self._lock: + text = "".join(self._lines).strip() + if text: + return text + return "" + + +def _start_output_reader( + stream: IO[str] | None, + output: _CapturedOutput, +) -> threading.Thread: + if stream is None: + raise RuntimeError("cloudflared stdout pipe was not created") + + def read_output() -> None: + with stream: + for line in stream: + output.append(line) + + thread = threading.Thread( + target=read_output, + name="furu-cloudflared-output-reader", + daemon=True, + ) + thread.start() + return thread + + +def _wait_for_trycloudflare_url( + *, + process: subprocess.Popen[str], + output: _CapturedOutput, + output_reader: threading.Thread, + startup_timeout: float, +) -> str: + deadline = time.monotonic() + startup_timeout + + while True: + while True: + try: + line = output.get_nowait() + except Empty: + break + if match := _TRYCLOUDFLARE_URL_RE.search(line): + return match.group(0).rstrip("/") + + returncode = process.poll() + if returncode is not None: + output_reader.join(timeout=1.0) + while True: + try: + line = output.get_nowait() + except Empty: + break + if match := _TRYCLOUDFLARE_URL_RE.search(line): + return match.group(0).rstrip("/") + raise RuntimeError( + "cloudflared exited before printing a trycloudflare URL: " + f"{output.recent_text()}" + ) + + remaining = deadline - time.monotonic() + if remaining <= 0: + raise TimeoutError( + "cloudflared did not print a trycloudflare URL within " + f"{startup_timeout:g} seconds; recent output: {output.recent_text()}" + ) + + try: + line = output.get(timeout=min(0.05, remaining)) + except Empty: + continue + if match := _TRYCLOUDFLARE_URL_RE.search(line): + return match.group(0).rstrip("/") + + +def _terminate_process( + process: subprocess.Popen[str], + *, + output_reader: threading.Thread, + timeout: float, +) -> None: + if process.poll() is None: + process.terminate() + try: + process.wait(timeout=timeout) + except subprocess.TimeoutExpired: + process.kill() + process.wait() + output_reader.join(timeout=1.0) diff --git a/src/furu/execution/manager.py b/src/furu/execution/manager.py index c3e5be80..278b1fda 100644 --- a/src/furu/execution/manager.py +++ b/src/furu/execution/manager.py @@ -26,6 +26,7 @@ ) if TYPE_CHECKING: + from furu.execution.connection import ManagerConnection from furu.worker.backends import WorkerBackend @@ -74,6 +75,7 @@ def run( *, worker_backends: tuple[WorkerBackend, ...], port: int = 0, + manager_connection: ManagerConnection | None = None, ) -> None: from furu.execution.server import _run_until_done @@ -81,6 +83,7 @@ def run( self, worker_backends=worker_backends, port=port, + manager_connection=manager_connection, ) @contextmanager diff --git a/src/furu/execution/server.py b/src/furu/execution/server.py index 29c912bf..3a9bcac7 100644 --- a/src/furu/execution/server.py +++ b/src/furu/execution/server.py @@ -3,15 +3,17 @@ import socket import threading import time -from collections.abc import Iterator +from collections.abc import Callable, Iterator from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass +from typing import cast from secrets import token_urlsafe import uvicorn from furu.execution.api import create_manager_api_app +from furu.execution.connection import ManagerConnection from furu.execution.manager import Manager from furu.logging import get_logger from furu.worker.backends import WorkerBackend, WorkerPool @@ -29,6 +31,13 @@ class ManagerServer: def server_url(self) -> str: return f"http://{self.bound_host}:{self.bound_port}" + @property + def local_origin_url(self) -> str: + host = self.bound_host + if host == "0.0.0.0": + host = "127.0.0.1" + return f"http://{host}:{self.bound_port}" + @contextmanager def manager_server( @@ -87,8 +96,14 @@ def _run_until_done( *, worker_backends: tuple[WorkerBackend, ...], port: int, + manager_connection: ManagerConnection | None = None, ) -> None: (bind_host,) = {backend.manager_listen_host for backend in worker_backends} + connection = ( + manager_connection + if manager_connection is not None + else _select_manager_connection(worker_backends) + ) with manager.log_context(): logger.info( @@ -99,22 +114,54 @@ def _run_until_done( len(manager.blocked), ) with manager_server(manager, bind_host=bind_host, port=port) as server: - logger.info( - "manager server listening: server_url=%s", - server.server_url, - ) - pools: list[WorkerPool] = [] - for backend in worker_backends: - pool = backend.start_pool( - server_url=server.server_url, - auth_token=server.auth_token, - executor_dir=manager.executor_dir, + with connection.connect( + local_url=server.local_origin_url + ) as advertised_url: + logger.info( + "manager server listening: local_url=%s advertised_url=%s", + server.local_origin_url, + advertised_url, ) - pools.append(pool) - logger.info("worker pool started: backend=%s", type(backend).__name__) - manager.done.wait() - - with ThreadPoolExecutor(max_workers=len(pools)) as executor: - for pool in pools: - executor.submit(pool.stop, timeout=5) + pools: list[WorkerPool] = [] + for backend in worker_backends: + pool = backend.start_pool( + server_url=advertised_url, + auth_token=server.auth_token, + executor_dir=manager.executor_dir, + ) + pools.append(pool) + logger.info( + "worker pool started: backend=%s", type(backend).__name__ + ) + manager.done.wait() + + with ThreadPoolExecutor(max_workers=len(pools)) as executor: + for pool in pools: + executor.submit(pool.stop, timeout=5) manager.raise_for_failure() + + +def _select_manager_connection( + worker_backends: tuple[WorkerBackend, ...], +) -> ManagerConnection: + from furu.execution.connection import DirectManagerConnection + + connections: list[ManagerConnection] = [] + for backend in worker_backends: + get_connection = cast( + Callable[[], ManagerConnection | None] | None, + getattr(backend, "manager_connection", None), + ) + if get_connection is None: + continue + connection = get_connection() + if connection is not None: + connections.append(connection) + + if not connections: + return DirectManagerConnection() + + first = connections[0] + if any(connection != first for connection in connections[1:]): + raise ValueError("worker backends requested conflicting manager connections") + return first diff --git a/src/furu/worker/backends/slurm/backend.py b/src/furu/worker/backends/slurm/backend.py index f2f7f5bb..7ed886b3 100644 --- a/src/furu/worker/backends/slurm/backend.py +++ b/src/furu/worker/backends/slurm/backend.py @@ -6,9 +6,11 @@ import threading from dataclasses import dataclass, field from pathlib import Path +from urllib.parse import urlsplit, urlunsplit from furu.config import _WORKER_JSON_CONFIG_FILE_ENV_VAR, get_config from furu.execution.api import PoolApiClient +from furu.execution.connection import CloudflareQuickTunnel, ManagerConnection from furu.resources import ResourceRequest from furu.utils import write_private_file from furu.worker.backends.slurm.pool import SlurmWorkerPool @@ -19,17 +21,31 @@ class SlurmWorkerBackend: max_workers: int resources: SlurmResources - worker_connect_host: str + worker_connect_host: str | None = None max_failed_restarts: int = field( default_factory=lambda: get_config().worker.max_failed_restarts ) - manager_listen_host: str = "0.0.0.0" + manager_listen_host: str = "" job_name: str = "furu-worker" poll_interval: float = 10.0 worker_idle_timeout: float = field( default_factory=lambda: get_config().worker.idle_timeout_seconds ) + def __post_init__(self) -> None: + if self.manager_listen_host: + return + + if self.worker_connect_host is None: + object.__setattr__(self, "manager_listen_host", "127.0.0.1") + else: + object.__setattr__(self, "manager_listen_host", "0.0.0.0") + + def manager_connection(self) -> ManagerConnection | None: + if self.worker_connect_host is None: + return CloudflareQuickTunnel() + return None + def start_pool( self, *, @@ -37,10 +53,7 @@ def start_pool( auth_token: str, executor_dir: Path, ) -> SlurmWorkerPool: - scheme, rest = server_url.split("://", maxsplit=1) - server_url = ( - f"{scheme}://{self.worker_connect_host}:{rest.rsplit(':', maxsplit=1)[1]}" - ) + server_url = self._worker_server_url(server_url) chdir = Path.cwd().resolve() worker_dir = executor_dir.resolve() / "workers" @@ -118,3 +131,17 @@ def start_pool( pool_holder.append(pool) pool._scale_thread.start() return pool + + def _worker_server_url(self, manager_server_url: str) -> str: + if self.worker_connect_host is None: + return manager_server_url + + parts = urlsplit(manager_server_url) + if parts.port is None: + netloc = self.worker_connect_host + else: + netloc = f"{self.worker_connect_host}:{parts.port}" + + return urlunsplit( + (parts.scheme, netloc, parts.path, parts.query, parts.fragment) + ) diff --git a/tests/test_manager_connection.py b/tests/test_manager_connection.py new file mode 100644 index 00000000..000a5cfa --- /dev/null +++ b/tests/test_manager_connection.py @@ -0,0 +1,146 @@ +from __future__ import annotations + +import json +import sys +import textwrap +from pathlib import Path + +import pytest + +from furu.execution.connection import DirectManagerConnection, CloudflareQuickTunnel + + +def test_direct_manager_connection_returns_local_url() -> None: + with DirectManagerConnection().connect(local_url="http://127.0.0.1:1234") as url: + assert url == "http://127.0.0.1:1234" + + +def test_cloudflare_quick_tunnel_command_must_not_be_empty() -> None: + with pytest.raises(ValueError, match="must not be empty"): + with CloudflareQuickTunnel(command=()).connect(local_url="http://127.0.0.1:1"): + pass + + +def test_cloudflare_quick_tunnel_missing_command_gives_clear_error() -> None: + with pytest.raises(RuntimeError, match="could not find"): + with CloudflareQuickTunnel( + command=("definitely-not-cloudflared-furu-test",) + ).connect(local_url="http://127.0.0.1:1"): + pass + + +def test_cloudflare_quick_tunnel_parses_url_from_output(tmp_path: Path) -> None: + fake_script = tmp_path / "fake_cloudflared.py" + argv_file = tmp_path / "argv.json" + terminated_file = tmp_path / "terminated" + fake_script.write_text( + textwrap.dedent( + f""" + import json + import signal + import sys + import time + from pathlib import Path + + argv_file = Path({str(argv_file)!r}) + terminated_file = Path({str(terminated_file)!r}) + + def handle_sigterm(signum, frame): + terminated_file.write_text("terminated", encoding="utf-8") + raise SystemExit(0) + + signal.signal(signal.SIGTERM, handle_sigterm) + argv_file.write_text(json.dumps(sys.argv[1:]), encoding="utf-8") + print("starting fake cloudflared", flush=True) + print("https://example-furu.trycloudflare.com", flush=True) + + while True: + time.sleep(1) + """ + ).lstrip(), + encoding="utf-8", + ) + + tunnel = CloudflareQuickTunnel( + command=(sys.executable, str(fake_script)), + startup_timeout=5, + ) + + with tunnel.connect(local_url="http://127.0.0.1:1234") as url: + assert url == "https://example-furu.trycloudflare.com" + + assert json.loads(argv_file.read_text(encoding="utf-8")) == [ + "tunnel", + "--url", + "http://127.0.0.1:1234", + ] + assert terminated_file.exists() + + +def test_cloudflare_quick_tunnel_timeout_stops_process_and_includes_output( + tmp_path: Path, +) -> None: + fake_script = tmp_path / "fake_cloudflared_timeout.py" + terminated_file = tmp_path / "terminated" + fake_script.write_text( + textwrap.dedent( + f""" + import signal + import time + from pathlib import Path + + terminated_file = Path({str(terminated_file)!r}) + + def handle_sigterm(signum, frame): + terminated_file.write_text("terminated", encoding="utf-8") + raise SystemExit(0) + + signal.signal(signal.SIGTERM, handle_sigterm) + print("still starting without a URL", flush=True) + + while True: + time.sleep(1) + """ + ).lstrip(), + encoding="utf-8", + ) + + tunnel = CloudflareQuickTunnel( + command=(sys.executable, str(fake_script)), + startup_timeout=0.1, + ) + + with pytest.raises( + TimeoutError, + match="did not print a trycloudflare URL.*still starting without a URL", + ): + with tunnel.connect(local_url="http://127.0.0.1:1234"): + pass + + assert terminated_file.exists() + + +def test_cloudflare_quick_tunnel_early_exit_includes_captured_output( + tmp_path: Path, +) -> None: + fake_script = tmp_path / "fake_cloudflared_exit.py" + fake_script.write_text( + textwrap.dedent( + """ + import sys + + print("config error", flush=True) + raise SystemExit(2) + """ + ).lstrip(), + encoding="utf-8", + ) + + tunnel = CloudflareQuickTunnel( + command=(sys.executable, str(fake_script)), + startup_timeout=5, + ) + + with pytest.raises(RuntimeError, match="config error"): + with tunnel.connect(local_url="http://127.0.0.1:1234"): + pass diff --git a/tests/test_slurm_backend.py b/tests/test_slurm_backend.py index ebb9670c..decff971 100644 --- a/tests/test_slurm_backend.py +++ b/tests/test_slurm_backend.py @@ -13,6 +13,7 @@ import furu.worker.backends.slurm.backend as slurm_backend_module from furu.config import _FuruConfig, _WORKER_JSON_CONFIG_FILE_ENV_VAR, get_config from furu.execution.api import PoolApiClient +from furu.execution.connection import CloudflareQuickTunnel from furu.resources import ResourceRequest from furu.worker import _cli from furu.worker.backends.slurm.backend import SlurmWorkerBackend @@ -451,6 +452,71 @@ def test_slurm_resources_emit_gpu_option(gpus: int, expected_args: list[str]) -> ] +def test_slurm_backend_defaults_to_cloudflare_connection() -> None: + backend = SlurmWorkerBackend( + max_workers=1, + resources=SlurmResources(cpus_per_worker=1), + ) + + assert backend.worker_connect_host is None + assert backend.manager_listen_host == "127.0.0.1" + assert isinstance(backend.manager_connection(), CloudflareQuickTunnel) + + +def test_slurm_backend_worker_connect_host_preserves_legacy_listen_host() -> None: + backend = SlurmWorkerBackend( + max_workers=1, + resources=SlurmResources(cpus_per_worker=1), + worker_connect_host="manager.cluster", + ) + + assert backend.manager_listen_host == "0.0.0.0" + assert backend.manager_connection() is None + + +def test_slurm_backend_explicit_manager_listen_host_wins() -> None: + backend = SlurmWorkerBackend( + max_workers=1, + resources=SlurmResources(cpus_per_worker=1), + manager_listen_host="0.0.0.0", + ) + + assert backend.manager_listen_host == "0.0.0.0" + + +def test_slurm_backend_keeps_cloudflare_advertised_url( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _disable_slurm_pool_scale_thread(monkeypatch) + record_file, _active_file = _install_fake_slurm(tmp_path, monkeypatch) + _stub_count_satisfiable_jobs(monkeypatch, 1) + backend = SlurmWorkerBackend( + max_workers=1, + resources=SlurmResources(cpus_per_worker=1), + ) + + pool = backend.start_pool( + server_url="https://furu-test.trycloudflare.com", + auth_token="secret-token", + executor_dir=tmp_path / "executor", + ) + pool._scale_once() + + assert pool._job_ids == ["100"] + assert pool._server_url == "https://furu-test.trycloudflare.com" + + records = _read_records(record_file) + sbatch_records = [record for record in records if record["executable"] == "sbatch"] + assert len(sbatch_records) == 1 + + script_path = Path(sbatch_records[0]["argv"][-1]) + script = script_path.read_text() + assert "--server-url https://furu-test.trycloudflare.com" in script + assert "127.0.0.1" not in script + assert "0.0.0.0" not in script + + def test_slurm_backend_rewrites_manager_url_to_worker_connect_host( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -472,6 +538,7 @@ def test_slurm_backend_rewrites_manager_url_to_worker_connect_host( pool._scale_once() assert pool._job_ids == ["100"] + assert pool._server_url == "http://manager.cluster:4321" records = _read_records(record_file) sbatch_records = [record for record in records if record["executable"] == "sbatch"] diff --git a/tests/test_worker_manager.py b/tests/test_worker_manager.py index 571fe9e4..c3e93ad2 100644 --- a/tests/test_worker_manager.py +++ b/tests/test_worker_manager.py @@ -1,4 +1,6 @@ import threading +from collections.abc import Iterator +from contextlib import contextmanager from pathlib import Path from typing import Any, ClassVar, cast from uuid import UUID, uuid4 @@ -18,7 +20,13 @@ Manager, RunningJob, ) -from furu.execution.server import _run_until_done, manager_server +from furu.execution.connection import CloudflareQuickTunnel, ManagerConnection +from furu.execution.server import ( + ManagerServer, + _run_until_done, + _select_manager_connection, + manager_server, +) from furu.metadata import ArtifactSpec from furu.resources import ResourceRequest, ResourceRequirements from furu._storage_layout import manager_log_path_in @@ -37,6 +45,20 @@ ANY_RESOURCES = ResourceRequest() +class RecordingConnection: + def __init__( + self, + advertised_url: str = "https://furu-test.trycloudflare.com", + ) -> None: + self.advertised_url = advertised_url + self.local_urls: list[str] = [] + + @contextmanager + def connect(self, *, local_url: str) -> Iterator[str]: + self.local_urls.append(local_url) + yield self.advertised_url + + def _new_local_pool( *, server_url: str = "http://manager.test", @@ -653,7 +675,7 @@ def start_pool( assert leaf.status() == "completed" assert leaf.load_or_create() == 11 assert len(backend.server_urls) == 1 - assert backend.server_urls[0].startswith("http://0.0.0.0:") + assert backend.server_urls[0].startswith("http://127.0.0.1:") assert len(backend.auth_tokens) == 1 assert backend.auth_tokens[0] @@ -761,6 +783,91 @@ def start_pool( assert pool.stop_timeouts == [5] +def test_run_until_done_passes_advertised_manager_url() -> None: + class RecordingDone: + def wait(self, timeout: float | None = None) -> bool: + return True + + class RecordingPool: + def stop(self, *, timeout: float) -> None: + pass + + class RecordingBackend: + manager_listen_host = "127.0.0.1" + + def __init__(self) -> None: + self.server_urls: list[str] = [] + + def start_pool( + self, + *, + server_url: str, + auth_token: str, + executor_dir: Path, + ) -> RecordingPool: + self.server_urls.append(server_url) + return RecordingPool() + + manager = Manager([ManagerLeaf(value=15)]) + backend = RecordingBackend() + connection = RecordingConnection() + cast(Any, manager).done = RecordingDone() + + _run_until_done( + manager, + worker_backends=(backend,), + port=0, + manager_connection=connection, + ) + + assert backend.server_urls == ["https://furu-test.trycloudflare.com"] + assert len(connection.local_urls) == 1 + assert connection.local_urls[0].startswith("http://127.0.0.1:") + + +def test_run_until_done_uses_backend_requested_manager_connection() -> None: + class RecordingDone: + def wait(self, timeout: float | None = None) -> bool: + return True + + class RecordingPool: + def stop(self, *, timeout: float) -> None: + pass + + class RecordingBackend: + manager_listen_host = "127.0.0.1" + + def __init__(self, connection: ManagerConnection) -> None: + self.connection = connection + self.server_urls: list[str] = [] + + def manager_connection(self) -> ManagerConnection | None: + return self.connection + + def start_pool( + self, + *, + server_url: str, + auth_token: str, + executor_dir: Path, + ) -> RecordingPool: + self.server_urls.append(server_url) + return RecordingPool() + + manager = Manager([ManagerLeaf(value=15)]) + connection = RecordingConnection( + advertised_url="https://backend-requested.trycloudflare.com" + ) + backend = RecordingBackend(connection) + cast(Any, manager).done = RecordingDone() + + _run_until_done(manager, worker_backends=(backend,), port=0) + + assert backend.server_urls == ["https://backend-requested.trycloudflare.com"] + assert len(connection.local_urls) == 1 + assert connection.local_urls[0].startswith("http://127.0.0.1:") + + def test_run_until_done_uses_worker_backend_manager_listen_host() -> None: class RecordingDone: def wait(self, timeout: float | None = None) -> bool: @@ -796,6 +903,59 @@ def start_pool( assert backend.server_urls[0].startswith("http://127.0.0.1:") +def test_select_manager_connection_rejects_conflicting_backend_requests() -> None: + class RecordingPool: + def stop(self, *, timeout: float) -> None: + pass + + class ConnectionBackend: + manager_listen_host = "127.0.0.1" + + def __init__(self, connection: ManagerConnection) -> None: + self.connection = connection + + def manager_connection(self) -> ManagerConnection | None: + return self.connection + + def start_pool( + self, + *, + server_url: str, + auth_token: str, + executor_dir: Path, + ) -> RecordingPool: + raise AssertionError("start_pool should not be called") + + with pytest.raises( + ValueError, + match="worker backends requested conflicting manager connections", + ): + _select_manager_connection( + worker_backends=( + ConnectionBackend( + CloudflareQuickTunnel(extra_args=("--edge-ip-version", "4")) + ), + ConnectionBackend( + CloudflareQuickTunnel(extra_args=("--edge-ip-version", "6")) + ), + ) + ) + + +def test_manager_server_local_origin_url_uses_loopback_for_wildcard_bind() -> None: + server = ManagerServer(bound_host="0.0.0.0", bound_port=1234, auth_token="x") + + assert server.server_url == "http://0.0.0.0:1234" + assert server.local_origin_url == "http://127.0.0.1:1234" + + +def test_manager_server_local_origin_url_preserves_concrete_bind_host() -> None: + server = ManagerServer(bound_host="127.0.0.1", bound_port=1234, auth_token="x") + + assert server.server_url == "http://127.0.0.1:1234" + assert server.local_origin_url == "http://127.0.0.1:1234" + + def test_manager_server_exposes_bound_host_and_port() -> None: manager = Manager([ManagerLeaf(value=12)])