From a6bf68103f9e2644bf86cc1bbb5987889eea5aa8 Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Thu, 8 Oct 2026 10:47:28 +0200 Subject: [PATCH 1/3] Reconnect gateway replica SSH tunnels after process exit --- pyproject.toml | 2 + .../_internal/core/services/ssh/tunnel.py | 111 +++- .../proxy/gateway/resources/systemd/start.sh | 2 +- .../proxy/lib/services/service_connection.py | 101 +++- .../core/services/ssh/test_tunnel.py | 513 +++++++++++++++++- .../_internal/proxy/lib/services/__init__.py | 0 .../lib/services/test_service_connection.py | 497 +++++++++++++++++ 7 files changed, 1207 insertions(+), 19 deletions(-) create mode 100644 src/tests/_internal/proxy/lib/services/__init__.py create mode 100644 src/tests/_internal/proxy/lib/services/test_service_connection.py diff --git a/pyproject.toml b/pyproject.toml index bba8433701..47758c2fef 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -217,6 +217,8 @@ gateway = [ "fastapi", "starlette>=0.26.0", "uvicorn", + # Supervise long-lived SSH children without a waiting thread for each replica. + "uvloop>=0.18.0; sys_platform != 'win32' and platform_python_implementation != 'PyPy'", "aiorwlock", "aiocache", "httpx>=0.28.0", diff --git a/src/dstack/_internal/core/services/ssh/tunnel.py b/src/dstack/_internal/core/services/ssh/tunnel.py index 9dff46c6e0..313eeea6e1 100644 --- a/src/dstack/_internal/core/services/ssh/tunnel.py +++ b/src/dstack/_internal/core/services/ssh/tunnel.py @@ -2,11 +2,13 @@ import asyncio import os import shlex +import signal import subprocess import tempfile from dataclasses import dataclass from typing import Dict, Iterable, List, Literal, NoReturn, Optional, Union +from dstack._internal.compat import IS_WINDOWS from dstack._internal.core.errors import SSHError from dstack._internal.core.models.instances import SSHConnectionParams from dstack._internal.core.services.ssh import get_ssh_error @@ -71,6 +73,7 @@ def __init__( port: Optional[int] = None, ssh_proxies: Iterable[tuple[SSHConnectionParams, Optional[FilePathOrContent]]] = (), batch_mode: bool = False, + background: bool = True, ): """ :param forwarded_sockets: Connections to the specified local sockets will be @@ -91,6 +94,10 @@ def __init__( Control commands (`check`, `close`, `exec`) always run in batch mode, since they only talk to the local master and must not prompt if ssh falls back to a direct connection. + :param background: If False, own a foreground SSH process instead of a daemon. + Use only the async methods and a dedicated control socket. `aopen()` waits + for startup readiness; `wait_closed()` waits for exit without polling; + `aclose()` cleans up the process and its ProxyCommand children. """ self.destination = destination self.forwarded_sockets = list(forwarded_sockets) @@ -114,6 +121,8 @@ def __init__( ) self.ssh_proxies.append((proxy_params, proxy_identity_path)) self.batch_mode = batch_mode + self.background = background + self._process: Optional[asyncio.subprocess.Process] = None self.log_path = normalize_path(os.path.join(temp_dir.name, "tunnel.log")) self.ssh_client_info = get_ssh_client_info() self.ssh_exec_path = str(self.ssh_client_info.path) @@ -136,19 +145,21 @@ def open_command(self) -> List[str]: self.log_path, "-N", # do not run commands on remote ] - if self.ssh_client_info.supports_background_mode: + if self.background: + if not self.ssh_client_info.supports_background_mode: + raise SSHError("Unsupported SSH client") command += ["-f"] # go to background after successful authentication else: - raise SSHError("Unsupported SSH client") + command += ["-o", "ForkAfterAuthentication=no", "-o", "ControlPersist=no"] if self.ssh_client_info.supports_control_socket: # It's safe to use ControlMaster even if the ssh client does not support multiplexing # as long as we don't allow more than one tunnel to the specific host to be running. # We use this feature for control only (see :meth:`close_command`). command += [ - # Not `-M`, which means `ControlMaster=yes`, to avoid spawning uncontrollable - # ssh instances if more than one tunnel is started (precaution). + # Background connections may reuse a master. Foreground connections must + # own their process, rather than attach to another master and exit. "-o", - "ControlMaster=auto", + "ControlMaster=auto" if self.background else "ControlMaster=yes", "-S", self.control_sock_path, ] @@ -185,6 +196,8 @@ def exec_command(self) -> List[str]: return [*self._control_command_prefix(), self.destination] def open(self) -> None: + if not self.background: + raise SSHError("Foreground SSH tunnels require aopen()") # We cannot use `stderr=subprocess.PIPE` here since the forked process (daemon) does not # close standard streams if ProxyJump is used, therefore we will wait EOF from the pipe # as long as the daemon exists. @@ -206,6 +219,9 @@ def open(self) -> None: self._raise_ssh_error_from_log_output(log_output) async def aopen(self) -> None: + if not self.background: + await self._aopen_foreground() + return await run_async(self._remove_log_file) proc = await asyncio.create_subprocess_exec( *self.open_command(), stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL @@ -222,6 +238,12 @@ async def aopen(self) -> None: log_output = await run_async(self._read_log_file) self._raise_ssh_error_from_log_output(log_output) + async def wait_closed(self) -> int: + """Wait for the foreground process to exit; retain it for `aclose()` cleanup.""" + if self._process is None: + raise SSHError("No foreground SSH process to wait for") + return await self._process.wait() + def close(self) -> None: if not os.path.exists(self.control_sock_path): logger.debug( @@ -247,6 +269,16 @@ def close(self) -> None: ) async def aclose(self) -> None: + if not self.background: + if self._process is None: + return + cleanup = asyncio.create_task(self._close_foreground(self._process)) + try: + await asyncio.shield(cleanup) + except asyncio.CancelledError: + await cleanup + raise + return if not os.path.exists(self.control_sock_path): logger.debug( "Control socket does not exist, it seems that ssh process has already exited" @@ -330,6 +362,64 @@ def _control_command_prefix(self) -> List[str]: self.control_sock_path, ] + async def _aopen_foreground(self) -> None: + if self._process is not None: + raise SSHError("Close the previous foreground SSH process before opening") + await run_async(self._remove_log_file) + creation = asyncio.create_task( + asyncio.create_subprocess_exec( + *self.open_command(), + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=not IS_WINDOWS, + ) + ) + try: + # Retain the process handle even if cancellation arrives during its creation. + proc = await asyncio.shield(creation) + self._process = proc + await asyncio.wait_for(self._wait_until_ready(proc), SSH_TIMEOUT) + except BaseException as e: + if self._process is None: + try: + self._process = await creation + except Exception: + # Preserve cancellation if creation also failed. + pass + await self.aclose() + if isinstance(e, asyncio.TimeoutError): + raise SSHError( + f"SSH tunnel to {self.destination} did not open in {SSH_TIMEOUT} seconds" + ) from e + raise + + async def _wait_until_ready(self, proc: asyncio.subprocess.Process) -> None: + while proc.returncode is None: + if os.path.exists(self.control_sock_path) and await self.acheck(): + if proc.returncode is None: + return + break + await asyncio.sleep(0.1) + log_output = await run_async(self._read_log_file) + self._raise_ssh_error_from_log_output(log_output) + + async def _close_foreground(self, proc: asyncio.subprocess.Process) -> None: + try: + if IS_WINDOWS: + proc.kill() + else: + # The launcher may have exited while a ProxyCommand child remains alive. + os.killpg(proc.pid, signal.SIGKILL) # pyright: ignore[reportAttributeAccessIssue] + except ProcessLookupError: + pass + await proc.wait() + self._process = None + # SIGKILL leaves the control socket behind. This path is owned by this tunnel. + try: + os.remove(self.control_sock_path) + except FileNotFoundError: + pass + def _get_proxy_command(self) -> Optional[str]: proxy_command: Optional[str] = None for params, identity_path in self.ssh_proxies: @@ -410,16 +500,19 @@ def _get_identity_path(self, identity: FilePathOrContent, tmp_filename: str) -> async def _arun(command: List[str], timeout: float) -> tuple[int, bytes, bytes]: """ Runs `command` with stdin redirected from /dev/null and returns its exit status, stdout, - and stderr. Kills the process and raises `asyncio.TimeoutError` if it does not exit in - `timeout` seconds. + and stderr. Kills and reaps the process on cancellation or if it does not exit in + `timeout` seconds, preserving the cancellation or `asyncio.TimeoutError`. """ proc = await asyncio.create_subprocess_exec( *command, stdin=subprocess.DEVNULL, stdout=subprocess.PIPE, stderr=subprocess.PIPE ) try: stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout) - except asyncio.TimeoutError: - proc.kill() + except (asyncio.CancelledError, asyncio.TimeoutError): + try: + proc.kill() + except ProcessLookupError: + pass await proc.wait() raise assert proc.returncode is not None diff --git a/src/dstack/_internal/proxy/gateway/resources/systemd/start.sh b/src/dstack/_internal/proxy/gateway/resources/systemd/start.sh index 932740acc5..bcd39d2159 100644 --- a/src/dstack/_internal/proxy/gateway/resources/systemd/start.sh +++ b/src/dstack/_internal/proxy/gateway/resources/systemd/start.sh @@ -8,4 +8,4 @@ else version="blue" echo "$version" > "$root/version" fi -"$root/$version/bin/uvicorn" dstack._internal.proxy.gateway.main:app +"$root/$version/bin/uvicorn" dstack._internal.proxy.gateway.main:app --loop uvloop diff --git a/src/dstack/_internal/proxy/lib/services/service_connection.py b/src/dstack/_internal/proxy/lib/services/service_connection.py index c8229ad53a..030b6cb4fd 100644 --- a/src/dstack/_internal/proxy/lib/services/service_connection.py +++ b/src/dstack/_internal/proxy/lib/services/service_connection.py @@ -24,6 +24,9 @@ logger = get_logger(__name__) OPEN_TUNNEL_TIMEOUT = 10 +# Bound SSH startup work during a shared outage, without limiting service requests. +MAX_CONCURRENT_TUNNEL_RECONNECTS = 8 +MAX_TUNNEL_RECONNECT_DELAY = 30 class ServiceClient(httpx.AsyncClient): @@ -33,7 +36,19 @@ def build_request(self, *args, **kwargs) -> httpx.Request: class ServiceConnection: - def __init__(self, project: Project, service: Service, replica: Replica) -> None: + """Forward a replica's HTTP traffic over SSH to a stable local Unix socket. + + Gateways supervise and reconnect the SSH process so Nginx can keep + using the same socket path. The in-server proxy opens connections on demand. + """ + + def __init__( + self, + project: Project, + service: Service, + replica: Replica, + reconnect_semaphore: asyncio.Semaphore, + ) -> None: self._temp_dir = TemporaryDirectory() options = { **SSH_DEFAULT_OPTIONS, @@ -66,6 +81,7 @@ def __init__(self, project: Project, service: Service, replica: Replica) -> None ), ], options=options, + background=service.domain is None, ) self._client = ServiceClient( transport=AsyncHTTPTransport(uds=str(self._app_socket_path)), @@ -75,29 +91,93 @@ def __init__(self, project: Project, service: Service, replica: Replica) -> None timeout=service.read_timeout, ) self._is_open = asyncio.locks.Event() + self._lifecycle_lock = asyncio.Lock() + self._closed = False + self._monitor_task: Optional[asyncio.Task] = None + self._auto_reconnect = service.domain is not None + self._replica_id = replica.id + self._reconnect_semaphore = reconnect_semaphore @property def app_socket_path(self) -> Path: return self._app_socket_path async def open(self) -> None: - await self._tunnel.aopen() - self._is_open.set() + async with self._lifecycle_lock: + if self._closed: + raise UnexpectedProxyError("Cannot open a closed service connection") + if self._is_open.is_set(): + return + await self._tunnel.aopen() + if self._closed: + # Removal may have started while SSH was connecting. + raise UnexpectedProxyError("Service connection was removed while opening") + self._is_open.set() + if self._auto_reconnect: + self._monitor_task = asyncio.create_task(self._monitor_tunnel()) async def close(self) -> None: - self._is_open.clear() - await self._client.aclose() - await self._tunnel.aclose() + self._closed = True + # Removal must finish cleaning up even if its caller is cancelled. + cleanup = asyncio.create_task(self._close()) + cancelled = None + while not cleanup.done(): + try: + await asyncio.shield(cleanup) + except asyncio.CancelledError as e: + cancelled = e + cleanup.result() + if cancelled is not None: + raise cancelled async def client(self) -> ServiceClient: await asyncio.wait_for(self._is_open.wait(), timeout=OPEN_TUNNEL_TIMEOUT) return self._client + async def _close(self) -> None: + async with self._lifecycle_lock: + if self._monitor_task is not None: + self._monitor_task.cancel() + await asyncio.gather(self._monitor_task, return_exceptions=True) + self._monitor_task = None + self._is_open.clear() + try: + await self._client.aclose() + finally: + await self._tunnel.aclose() + + async def _monitor_tunnel(self) -> None: + loop = asyncio.get_running_loop() + retry_delay = 0 + while True: + started_at = loop.time() + await self._tunnel.wait_closed() + if loop.time() - started_at >= MAX_TUNNEL_RECONNECT_DELAY: + retry_delay = 0 + logger.warning("SSH tunnel to replica %s exited, reconnecting", self._replica_id) + # Reap any surviving ProxyCommand children before opening a replacement. + await self._tunnel.aclose() + while True: + if retry_delay: + await asyncio.sleep(random.uniform(retry_delay / 2, retry_delay)) + # Back off failed starts and tunnels that repeatedly exit just after startup. + retry_delay = min(max(1, retry_delay * 2), MAX_TUNNEL_RECONNECT_DELAY) + try: + async with self._reconnect_semaphore: + # Keep the socket path configured in Nginx. SSH replaces stale sockets. + await self._tunnel.aopen() + except Exception as e: + logger.warning("Could not reconnect to replica %s: %s", self._replica_id, e) + else: + logger.info("SSH tunnel to replica %s reconnected", self._replica_id) + break + class ServiceConnectionPool: def __init__(self) -> None: # TODO(#2238): remove connections to stopped replicas in-server self.connections: Dict[str, ServiceConnection] = {} + self._reconnect_semaphore = asyncio.Semaphore(MAX_CONCURRENT_TUNNEL_RECONNECTS) async def get(self, replica_id: str) -> Optional[ServiceConnection]: return self.connections.get(replica_id) @@ -108,12 +188,17 @@ async def get_or_add( connection = self.connections.get(replica.id) if connection is not None: return connection - connection = ServiceConnection(project, service, replica) + connection = ServiceConnection(project, service, replica, self._reconnect_semaphore) self.connections[replica.id] = connection try: await connection.open() except BaseException: - self.connections.pop(replica.id, None) + if self.connections.get(replica.id) is connection: + self.connections.pop(replica.id) + try: + await connection.close() + except Exception: + logger.exception("Error closing failed connection to replica %s", replica.id) raise return connection diff --git a/src/tests/_internal/core/services/ssh/test_tunnel.py b/src/tests/_internal/core/services/ssh/test_tunnel.py index da61536bd5..b1bca4b85f 100644 --- a/src/tests/_internal/core/services/ssh/test_tunnel.py +++ b/src/tests/_internal/core/services/ssh/test_tunnel.py @@ -1,11 +1,13 @@ import asyncio +import signal import subprocess from pathlib import Path from typing import NoReturn, Optional -from unittest.mock import Mock +from unittest.mock import AsyncMock, Mock import pytest +from dstack._internal.compat import IS_WINDOWS from dstack._internal.core.errors import SSHError from dstack._internal.core.models.instances import SSHConnectionParams from dstack._internal.core.services.ssh.client import SSHClientInfo @@ -44,6 +46,33 @@ def sample_tunnel_with_all_params(self, ssh_client_info: SSHClientInfo) -> SSHTu reverse_forwarded_sockets=[SocketPair(UnixSocket("/1"), UnixSocket("/2"))], ) + @pytest.fixture + def foreground_tunnel(self, ssh_client_info: SSHClientInfo) -> SSHTunnel: + return SSHTunnel( + destination="ubuntu@my-server", + identity=FilePath("/home/user/.ssh/id_rsa"), + background=False, + ) + + @pytest.fixture + def kill_process_group(self, monkeypatch: pytest.MonkeyPatch) -> Mock: + kill_process_group = Mock() + monkeypatch.setattr( + "dstack._internal.core.services.ssh.tunnel.os.killpg", + kill_process_group, + raising=False, + ) + return kill_process_group + + @pytest.fixture + def ssh_process(self, monkeypatch: pytest.MonkeyPatch, kill_process_group: Mock) -> Mock: + process = Mock(spec=asyncio.subprocess.Process) + process.pid = 12345 + process.returncode = None + process.communicate.return_value = (b"", b"") + monkeypatch.setattr(asyncio, "create_subprocess_exec", AsyncMock(return_value=process)) + return process + @pytest.mark.usefixtures("ssh_client_info") def test_open_command_basic(self) -> None: tunnel = SSHTunnel( @@ -242,6 +271,431 @@ def test_exec_command(self, sample_tunnel_with_all_params: SSHTunnel) -> None: "/usr/bin/ssh -F none -o BatchMode=yes -S /tmp/control.sock ubuntu@my-server" ) + def test_foreground_command_owns_process(self, foreground_tunnel: SSHTunnel) -> None: + foreground_tunnel.options = { + "ForkAfterAuthentication": "yes", + "ControlPersist": "yes", + "ControlMaster": "auto", + } + command = foreground_tunnel.open_command() + + assert "-f" not in command + # OpenSSH honors the first value, so user/config options cannot detach the process. + assert command.index("ForkAfterAuthentication=no") < command.index( + "ForkAfterAuthentication=yes" + ) + assert command.index("ControlPersist=no") < command.index("ControlPersist=yes") + assert command.index("ControlMaster=yes") < command.index("ControlMaster=auto") + + def test_foreground_requires_async_open(self, foreground_tunnel: SSHTunnel) -> None: + with pytest.raises(SSHError, match="require aopen"): + foreground_tunnel.open() + + @pytest.mark.asyncio + async def test_background_aopen_keeps_existing_behavior( + self, + sample_tunnel_with_all_params: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + create_process = AsyncMock(return_value=ssh_process) + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_process) + ssh_process.returncode = 0 + + await sample_tunnel_with_all_params.aopen() + + assert "-f" in create_process.call_args.args + assert "start_new_session" not in create_process.call_args.kwargs + ssh_process.communicate.assert_awaited_once_with() + ssh_process.wait.assert_not_called() + ssh_process.kill.assert_not_called() + kill_process_group.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "is_windows", + [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], + ) + async def test_foreground_waits_for_readiness_then_process_exit( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + is_windows: bool, + ) -> None: + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) + create_process = AsyncMock(return_value=ssh_process) + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_process) + Path(foreground_tunnel.control_sock_path).touch() + checking = asyncio.Event() + ready = asyncio.Event() + waiting_for_exit = asyncio.Event() + exited = asyncio.Event() + + async def check(): + checking.set() + await ready.wait() + return True + + async def wait(): + waiting_for_exit.set() + await exited.wait() + return ssh_process.returncode + + check_mock = AsyncMock(side_effect=check) + monkeypatch.setattr(foreground_tunnel, "acheck", check_mock) + ssh_process.wait.side_effect = wait + opening = asyncio.create_task(foreground_tunnel.aopen()) + await checking.wait() + assert not opening.done() + ready.set() + await opening + + assert create_process.call_args.kwargs["start_new_session"] is not is_windows + ssh_process.communicate.assert_not_called() + ssh_process.wait.assert_not_called() + waiting = asyncio.create_task(foreground_tunnel.wait_closed()) + await waiting_for_exit.wait() + assert not waiting.done() + ssh_process.returncode = 255 + exited.set() + assert await waiting == 255 + check_mock.assert_awaited_once_with() + + # Keep the handle after exit: the process group can still contain proxy children. + await foreground_tunnel.aclose() + if is_windows: + ssh_process.kill.assert_called_once_with() + kill_process_group.assert_not_called() + else: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.kill.assert_not_called() + assert ssh_process.wait.await_count == 2 + assert not Path(foreground_tunnel.control_sock_path).exists() + await foreground_tunnel.aclose() + assert ssh_process.wait.await_count == 2 + with pytest.raises(SSHError, match="No foreground SSH process"): + await foreground_tunnel.wait_closed() + + @pytest.mark.asyncio + async def test_foreground_checks_only_after_control_socket_exists( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + polling = asyncio.Event() + create_socket = asyncio.Event() + check = AsyncMock(return_value=True) + monkeypatch.setattr(foreground_tunnel, "acheck", check) + + async def startup_poll(interval): + assert interval == 0.1 + polling.set() + await create_socket.wait() + Path(foreground_tunnel.control_sock_path).touch() + + monkeypatch.setattr(asyncio, "sleep", startup_poll) + opening = asyncio.create_task(foreground_tunnel.aopen()) + await polling.wait() + check.assert_not_called() + create_socket.set() + await opening + check.assert_awaited_once_with() + + with pytest.raises(SSHError, match="Close the previous foreground SSH process"): + await foreground_tunnel.aopen() + await foreground_tunnel.aclose() + + @pytest.mark.asyncio + @pytest.mark.parametrize("during_check", [False, True]) + async def test_foreground_early_exit_fails_startup_and_cleans_up( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + during_check: bool, + ) -> None: + Path(foreground_tunnel.control_sock_path).touch() + + async def check(): + ssh_process.returncode = 255 + return True + + check_mock = AsyncMock(side_effect=check) + monkeypatch.setattr(foreground_tunnel, "acheck", check_mock) + if not during_check: + ssh_process.returncode = 255 + + with pytest.raises(SSHError): + await foreground_tunnel.aopen() + + assert check_mock.await_count == int(during_check) + if IS_WINDOWS: + ssh_process.kill.assert_called_once_with() + else: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.wait.assert_awaited_once_with() + assert not Path(foreground_tunnel.control_sock_path).exists() + with pytest.raises(SSHError, match="No foreground SSH process"): + await foreground_tunnel.wait_closed() + + @pytest.mark.asyncio + @pytest.mark.parametrize("method", ["aopen", "acheck"]) + @pytest.mark.parametrize("already_exited", [False, True]) + @pytest.mark.parametrize( + "is_windows", + [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], + ) + async def test_cancelled_command_kills_and_reaps_process( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + method: str, + already_exited: bool, + is_windows: bool, + ) -> None: + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) + waiting = asyncio.Event() + + async def wait_for_process(*args): + if ssh_process.returncode is None: + waiting.set() + await asyncio.Future() + return ssh_process.returncode + + def kill(*args) -> None: + ssh_process.returncode = 0 if already_exited else -9 + if already_exited: + raise ProcessLookupError + + monkeypatch.setattr( + foreground_tunnel, "_wait_until_ready", AsyncMock(side_effect=wait_for_process) + ) + ssh_process.communicate.side_effect = wait_for_process + ssh_process.wait.side_effect = wait_for_process + ssh_process.kill.side_effect = kill + kill_process_group.side_effect = kill + task = asyncio.create_task(getattr(foreground_tunnel, method)()) + await waiting.wait() + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + if method == "aopen" and not is_windows: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.kill.assert_not_called() + else: + ssh_process.kill.assert_called_once_with() + kill_process_group.assert_not_called() + ssh_process.wait.assert_awaited_once_with() + assert task.cancelled() + + @pytest.mark.asyncio + async def test_cancelled_foreground_readiness_cleans_check_and_master( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + Path(foreground_tunnel.control_sock_path).touch() + check_process = Mock(spec=asyncio.subprocess.Process) + check_process.returncode = None + waiting = asyncio.Event() + + async def wait(): + if check_process.returncode is None: + waiting.set() + await asyncio.Future() + return check_process.returncode + + def kill(): + check_process.returncode = -9 + + check_process.communicate.side_effect = wait + check_process.kill.side_effect = kill + create_process = AsyncMock(side_effect=[ssh_process, check_process]) + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_process) + opening = asyncio.create_task(foreground_tunnel.aopen()) + await waiting.wait() + opening.cancel() + with pytest.raises(asyncio.CancelledError): + await opening + + check_process.kill.assert_called_once_with() + check_process.communicate.assert_awaited_once_with() + check_process.wait.assert_awaited_once_with() + if IS_WINDOWS: + ssh_process.kill.assert_called_once_with() + else: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.wait.assert_awaited_once_with() + assert foreground_tunnel._process is None + + @pytest.mark.asyncio + @pytest.mark.parametrize("already_exited", [False, True]) + @pytest.mark.parametrize( + "is_windows", + [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], + ) + async def test_timed_out_aopen_kills_and_reaps_process( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + already_exited: bool, + is_windows: bool, + ) -> None: + # A zero timeout exercises wait_for's cleanup without waiting on real time. + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.SSH_TIMEOUT", 0) + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) + if already_exited: + ssh_process.kill.side_effect = ProcessLookupError + kill_process_group.side_effect = ProcessLookupError + + with pytest.raises(SSHError, match="in 0 seconds") as exc_info: + await foreground_tunnel.aopen() + + if not is_windows: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.kill.assert_not_called() + else: + ssh_process.kill.assert_called_once_with() + kill_process_group.assert_not_called() + ssh_process.wait.assert_awaited_once_with() + assert isinstance(exc_info.value.__cause__, asyncio.TimeoutError) + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "is_windows", + [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], + ) + async def test_cancelled_aopen_during_creation_kills_and_reaps_process( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + is_windows: bool, + ) -> None: + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) + spawned = asyncio.Event() + return_process = asyncio.Event() + returned = asyncio.Event() + + async def create_process(*args, **kwargs): + # Creation has spawned SSH but has not yet returned its handle to aopen(). + spawned.set() + await return_process.wait() + returned.set() + return ssh_process + + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_process) + task = asyncio.create_task(foreground_tunnel.aopen()) + await spawned.wait() + task.cancel("cancel while creating SSH") + return_process.set() + with pytest.raises(asyncio.CancelledError, match="cancel while creating SSH"): + await task + + assert returned.is_set() + if is_windows: + ssh_process.kill.assert_called_once_with() + kill_process_group.assert_not_called() + else: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + ssh_process.kill.assert_not_called() + ssh_process.wait.assert_awaited_once_with() + ssh_process.communicate.assert_not_called() + + @pytest.mark.asyncio + @pytest.mark.parametrize("cancelled", [False, True]) + async def test_aopen_process_creation_failure_preserves_cancellation( + self, + foreground_tunnel: SSHTunnel, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + cancelled: bool, + ) -> None: + creating = asyncio.Event() + fail_creation = asyncio.Event() + failure = OSError("SSH could not start") + + async def create_process(*args, **kwargs): + creating.set() + await fail_creation.wait() + raise failure + + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_process) + task = asyncio.create_task(foreground_tunnel.aopen()) + await creating.wait() + if cancelled: + task.cancel("cancel while creating SSH") + fail_creation.set() + + if cancelled: + with pytest.raises(asyncio.CancelledError, match="cancel while creating SSH"): + await task + else: + with pytest.raises(OSError) as exc_info: + await task + assert exc_info.value is failure + kill_process_group.assert_not_called() + + @pytest.mark.asyncio + async def test_cancelled_aclose_waits_for_cleanup( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + ) -> None: + foreground_tunnel._process = ssh_process + Path(foreground_tunnel.control_sock_path).touch() + reaping = asyncio.Event() + reaped = asyncio.Event() + + async def wait(): + reaping.set() + await reaped.wait() + return -9 + + ssh_process.wait.side_effect = wait + closing = asyncio.create_task(foreground_tunnel.aclose()) + await reaping.wait() + closing.cancel("cancel during cleanup") + reaped.set() + with pytest.raises(asyncio.CancelledError, match="cancel during cleanup"): + await closing + + ssh_process.wait.assert_awaited_once_with() + assert foreground_tunnel._process is None + assert not Path(foreground_tunnel.control_sock_path).exists() + if IS_WINDOWS: + ssh_process.kill.assert_called_once_with() + else: + kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) + + @pytest.mark.asyncio + @pytest.mark.parametrize("returncode", [0, 255]) + async def test_acheck_returns_process_status( + self, + sample_tunnel_with_all_params: SSHTunnel, + ssh_process: Mock, + returncode: int, + ) -> None: + ssh_process.returncode = returncode + + assert await sample_tunnel_with_all_params.acheck() is (returncode == 0) + + ssh_process.kill.assert_not_called() + class TestSSHTunnelControlTimeouts: @pytest.fixture @@ -307,6 +761,63 @@ async def test_aexec_kills_process_and_raises_on_timeout( await tunnel.aexec("true", timeout=0) assert hanging_process.killed + @pytest.mark.asyncio + @pytest.mark.parametrize("method", ["acheck", "aclose", "aexec"]) + @pytest.mark.parametrize("already_exited", [False, True]) + async def test_cancelled_control_command_kills_and_reaps_process( + self, + tunnel: SSHTunnel, + monkeypatch: pytest.MonkeyPatch, + method: str, + already_exited: bool, + ) -> None: + process = Mock(spec=asyncio.subprocess.Process) + communicating = asyncio.Event() + + async def communicate(): + communicating.set() + await asyncio.Future() + + process.communicate.side_effect = communicate + if already_exited: + process.kill.side_effect = ProcessLookupError + monkeypatch.setattr(asyncio, "create_subprocess_exec", AsyncMock(return_value=process)) + command = getattr(tunnel, method)(*(["true"] if method == "aexec" else [])) + task = asyncio.create_task(command) + await communicating.wait() + task.cancel("cancel control command") + + with pytest.raises(asyncio.CancelledError, match="cancel control command"): + await task + + process.kill.assert_called_once_with() + process.wait.assert_awaited_once_with() + + @pytest.mark.asyncio + @pytest.mark.parametrize("method", ["acheck", "aclose", "aexec"]) + async def test_timeout_preserves_behavior_when_process_has_already_exited( + self, + tunnel: SSHTunnel, + hanging_process: "_HangingProcess", + monkeypatch: pytest.MonkeyPatch, + method: str, + ) -> None: + def kill(): + hanging_process.returncode = 0 + raise ProcessLookupError + + reaped = AsyncMock(wraps=hanging_process.wait) + monkeypatch.setattr(hanging_process, "kill", kill) + monkeypatch.setattr(hanging_process, "wait", reaped) + if method == "aexec": + with pytest.raises(SSHError, match="did not complete") as exc_info: + await tunnel.aexec("true", timeout=0) + assert isinstance(exc_info.value.__cause__, asyncio.TimeoutError) + else: + result = await getattr(tunnel, method)() + assert result is (False if method == "acheck" else None) + reaped.assert_awaited_once_with() + class _HangingProcess: def __init__(self) -> None: diff --git a/src/tests/_internal/proxy/lib/services/__init__.py b/src/tests/_internal/proxy/lib/services/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/src/tests/_internal/proxy/lib/services/test_service_connection.py b/src/tests/_internal/proxy/lib/services/test_service_connection.py new file mode 100644 index 0000000000..4026572794 --- /dev/null +++ b/src/tests/_internal/proxy/lib/services/test_service_connection.py @@ -0,0 +1,497 @@ +import asyncio +from typing import AsyncIterator +from unittest.mock import create_autospec + +import pytest +import pytest_asyncio + +from dstack._internal.core.errors import SSHError +from dstack._internal.core.services.ssh.tunnel import SSHTunnel +from dstack._internal.proxy.lib.errors import UnexpectedProxyError +from dstack._internal.proxy.lib.services import service_connection +from dstack._internal.proxy.lib.services.service_connection import ServiceConnectionPool +from dstack._internal.proxy.lib.testing.common import make_project, make_service + + +@pytest.fixture +def project(): + return make_project("test-project") + + +@pytest.fixture +def service(): + return make_service("test-project", "test-service", domain="service.gateway.test") + + +@pytest.fixture +def tunnel_factory(mocker): + factory = mocker.patch.object(service_connection, "SSHTunnel", autospec=True) + factory.return_value = _make_tunnel() + return factory + + +@pytest.fixture +def monitor_clock(monkeypatch): + clock = _MonitorClock() + monkeypatch.setattr(service_connection.asyncio, "sleep", clock.sleep) + return clock + + +@pytest_asyncio.fixture +async def pool() -> AsyncIterator[ServiceConnectionPool]: + pool = ServiceConnectionPool() + try: + yield pool + finally: + await pool.remove_all() + + +@pytest.mark.asyncio +class TestServiceConnection: + async def test_gateway_recovers_the_same_tunnel_and_socket_without_polling( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + connection = await pool.get_or_add(project, service, service.replicas[0]) + socket_path = connection.app_socket_path + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + + # Recovery runs without an HTTP request and keeps Nginx's configured socket. + tunnel_factory.assert_called_once() + assert tunnel.aopen.await_count == 2 + tunnel.aclose.assert_awaited_once() + tunnel.acheck.assert_not_awaited() + assert not monitor_clock.tasks + assert tunnel_factory.call_args.kwargs["background"] is False + assert connection.app_socket_path == socket_path + assert tunnel_factory.call_args.kwargs["forwarded_sockets"][0].local.path == socket_path + + async def test_healthy_gateway_does_not_reopen_or_start_duplicate_monitors( + self, project, service, pool, tunnel_factory, monitor_clock + ): + connection = await pool.get_or_add(project, service, service.replicas[0]) + await connection.open() + + await _yield_to_event_loop() + + tunnel_factory.return_value.aopen.assert_awaited_once() + tunnel_factory.return_value.wait_closed.assert_awaited_once() + tunnel_factory.return_value.acheck.assert_not_awaited() + assert not monitor_clock.tasks + + async def test_failed_reconnects_back_off_and_then_recover( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + connection = await pool.get_or_add(project, service, service.replicas[0]) + original_client = await connection.client() + tunnel.aopen.side_effect = [SSHError("offline"), SSHError("offline"), None] + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + + assert 0.5 <= await monitor_clock.tick() <= 1 + assert 1 <= await monitor_clock.tick() <= 2 + + assert tunnel.aopen.await_count == 4 + assert await asyncio.wait_for(connection.client(), timeout=1) is original_client + assert not original_client.is_closed + + async def test_repeated_early_exits_back_off_to_a_bounded_delay( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + await pool.get_or_add(project, service, service.replicas[0]) + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + for maximum in (1, 2, 4, 8, 16, 30, 30): + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + assert maximum / 2 <= await monitor_clock.tick() <= maximum + + async def test_stable_tunnel_resets_retry_delay( + self, project, service, pool, tunnel_factory, monitor_clock, mocker + ): + clock = mocker.patch.object(asyncio.get_running_loop(), "time", return_value=0) + tunnel = tunnel_factory.return_value + await pool.get_or_add(project, service, service.replicas[0]) + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + clock.return_value = service_connection.MAX_TUNNEL_RECONNECT_DELAY + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + + assert tunnel.aopen.await_count == 3 + assert not monitor_clock.tasks + + async def test_in_server_connection_does_not_monitor( + self, project, service, pool, tunnel_factory, monitor_clock + ): + service = service.model_copy(update={"domain": None}) + connection = await pool.get_or_add(project, service, service.replicas[0]) + await _yield_to_event_loop() + + assert not (await connection.client()).is_closed + tunnel_factory.return_value.aopen.assert_awaited_once() + tunnel_factory.return_value.wait_closed.assert_not_awaited() + assert tunnel_factory.call_args.kwargs["background"] is True + assert not monitor_clock.tasks + + async def test_client_does_not_wait_for_a_blocked_reconnect( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + connection = await pool.get_or_add(project, service, service.replicas[0]) + original_client = await connection.client() + reconnecting = asyncio.Event() + release_reconnect = asyncio.Event() + + async def reconnect(): + reconnecting.set() + await release_reconnect.wait() + + tunnel.aopen.side_effect = reconnect + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + assert reconnecting.is_set() + + # The timeout is only a deadlock guard; the reconnect stays blocked throughout. + client = await asyncio.wait_for(connection.client(), timeout=1) + assert client is original_client + assert not release_reconnect.is_set() + + @pytest.mark.parametrize("remove_from_pool", [False, True], ids=["close", "remove"]) + async def test_close_cancels_reconnect_before_final_tunnel_cleanup( + self, project, service, pool, tunnel_factory, monitor_clock, remove_from_pool + ): + tunnel = tunnel_factory.return_value + connection = await pool.get_or_add(project, service, service.replicas[0]) + client = await connection.client() + release_reconnect = asyncio.Event() + lifecycle = [] + + async def reconnect(): + lifecycle.append("reconnecting") + try: + await release_reconnect.wait() + except asyncio.CancelledError: + lifecycle.append("cancelled") + raise + + async def close_tunnel(): + lifecycle.append("closed") + + tunnel.aopen.side_effect = reconnect + tunnel.aclose.side_effect = close_tunnel + monitor_task = connection._monitor_task + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + assert lifecycle == ["closed", "reconnecting"] + + if remove_from_pool: + await pool.remove(service.replicas[0].id) + assert await pool.get(service.replicas[0].id) is None + else: + await connection.close() + + release_reconnect.set() + await _yield_to_event_loop() + assert lifecycle == ["closed", "reconnecting", "cancelled", "closed"] + assert client.is_closed + assert monitor_task.cancelled() + with pytest.raises(UnexpectedProxyError, match="closed"): + await connection.open() + assert tunnel.aopen.await_count == 2 + + +@pytest.mark.asyncio +class TestServiceConnectionPool: + @pytest.mark.parametrize("cancel_count", [1, 2]) + async def test_cancelled_removal_finishes_monitor_and_resource_cleanup( + self, project, service, pool, tunnel_factory, cancel_count + ): + tunnel = tunnel_factory.return_value + waiting = asyncio.Event() + monitor_cancelled = asyncio.Event() + release_monitor = asyncio.Event() + + async def wait_closed(): + waiting.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + monitor_cancelled.set() + await release_monitor.wait() + raise + + tunnel.wait_closed.side_effect = wait_closed + replica = service.replicas[0] + connection = await pool.get_or_add(project, service, replica) + client = await connection.client() + await waiting.wait() + removing = asyncio.create_task(pool.remove(replica.id)) + try: + await monitor_cancelled.wait() + for _ in range(cancel_count): + removing.cancel() + await _yield_to_event_loop() + assert not removing.done() + finally: + release_monitor.set() + with pytest.raises(asyncio.CancelledError): + await removing + + assert await pool.get(replica.id) is None + assert client.is_closed + assert connection._monitor_task is None + tunnel.aclose.assert_awaited_once() + + async def test_initial_open_failure_closes_tunnel_and_removes_connection( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + tunnel.aopen.side_effect = RuntimeError("cannot reach replica") + + with pytest.raises(RuntimeError, match="cannot reach replica"): + await pool.get_or_add(project, service, service.replicas[0]) + + assert await pool.get(service.replicas[0].id) is None + tunnel.aclose.assert_awaited_once() + assert not monitor_clock.tasks + + async def test_cancelled_initial_open_closes_tunnel_and_removes_connection( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + opening = asyncio.Event() + + async def open_tunnel(): + opening.set() + await asyncio.Event().wait() + + tunnel.aopen.side_effect = open_tunnel + adding = asyncio.create_task(pool.get_or_add(project, service, service.replicas[0])) + try: + await _yield_to_event_loop() + assert opening.is_set() + adding.cancel() + with pytest.raises(asyncio.CancelledError): + await adding + finally: + adding.cancel() + await asyncio.gather(adding, return_exceptions=True) + + assert await pool.get(service.replicas[0].id) is None + tunnel.aclose.assert_awaited_once() + assert not monitor_clock.tasks + + async def test_remove_finishing_after_readd_keeps_new_connection( + self, project, service, pool, tunnel_factory, monitor_clock + ): + old_tunnel = tunnel_factory.return_value + new_tunnel = _make_tunnel() + tunnel_factory.side_effect = [old_tunnel, new_tunnel] + replica = service.replicas[0] + await pool.get_or_add(project, service, replica) + closing = asyncio.Event() + release_close = asyncio.Event() + + async def close_tunnel(): + closing.set() + await release_close.wait() + + old_tunnel.aclose.side_effect = close_tunnel + removing = asyncio.create_task(pool.remove(replica.id)) + try: + await asyncio.wait_for(closing.wait(), timeout=1) + assert await pool.get(replica.id) is None + replacement = await pool.get_or_add(project, service, replica) + finally: + release_close.set() + await removing + + assert closing.is_set() + assert await pool.get(replica.id) is replacement + assert not (await replacement.client()).is_closed + new_tunnel.aclose.assert_not_awaited() + + async def test_failed_open_cleanup_after_readd_keeps_new_connection( + self, project, service, pool, tunnel_factory, monitor_clock + ): + old_tunnel = tunnel_factory.return_value + new_tunnel = _make_tunnel() + tunnel_factory.side_effect = [old_tunnel, new_tunnel] + replica = service.replicas[0] + release_open = asyncio.Event() + + async def open_tunnel(): + await release_open.wait() + raise RuntimeError("old connection failed") + + old_tunnel.aopen.side_effect = open_tunnel + adding = asyncio.create_task(pool.get_or_add(project, service, replica)) + removing = None + try: + await _yield_to_event_loop() + removing = asyncio.create_task(pool.remove(replica.id)) + await _yield_to_event_loop() + replacement = await pool.get_or_add(project, service, replica) + release_open.set() + with pytest.raises(RuntimeError, match="old connection failed"): + await adding + await removing + + assert await pool.get(replica.id) is replacement + assert not (await replacement.client()).is_closed + old_tunnel.aclose.assert_awaited() + new_tunnel.aclose.assert_not_awaited() + finally: + release_open.set() + adding.cancel() + await asyncio.gather(adding, return_exceptions=True) + if removing is not None: + await removing + + async def test_removed_initial_open_succeeding_cannot_resurrect_or_remove_replacement( + self, project, service, pool, tunnel_factory, monitor_clock + ): + old_tunnel = tunnel_factory.return_value + new_tunnel = _make_tunnel() + tunnel_factory.side_effect = [old_tunnel, new_tunnel] + replica = service.replicas[0] + release_open = asyncio.Event() + + async def open_tunnel(): + await release_open.wait() + + old_tunnel.aopen.side_effect = open_tunnel + adding = asyncio.create_task(pool.get_or_add(project, service, replica)) + removing = None + try: + await _yield_to_event_loop() + original = await pool.get(replica.id) + assert original is not None + removing = asyncio.create_task(pool.remove(replica.id)) + await _yield_to_event_loop() + replacement = await pool.get_or_add(project, service, replica) + await _yield_to_event_loop() + replacement_monitor = replacement._monitor_task + assert replacement_monitor is not None + + release_open.set() + with pytest.raises(UnexpectedProxyError, match="removed while opening"): + await adding + await removing + + assert await pool.get(replica.id) is replacement + assert not (await replacement.client()).is_closed + assert replacement._monitor_task is replacement_monitor + assert not replacement_monitor.done() + old_tunnel.aopen.assert_awaited_once() + old_tunnel.wait_closed.assert_not_awaited() + old_tunnel.aclose.assert_awaited() + new_tunnel.aclose.assert_not_awaited() + with pytest.raises(UnexpectedProxyError, match="closed"): + await original.open() + finally: + release_open.set() + adding.cancel() + await asyncio.gather(adding, return_exceptions=True) + if removing is not None: + await removing + + async def test_remove_all_cancels_every_monitor_even_if_one_close_fails( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnels = [_make_tunnel() for _ in range(2)] + tunnel_factory.side_effect = tunnels + clients = [] + monitors = [] + for index in range(2): + replica = service.replicas[0].model_copy(update={"id": f"replica-{index}"}) + connection = await pool.get_or_add(project, service, replica) + clients.append(await connection.client()) + monitors.append(connection._monitor_task) + await _yield_to_event_loop() + tunnels[0].aclose.side_effect = RuntimeError("already gone") + + await pool.remove_all() + + assert not pool.connections + assert all(task.cancelled() for task in monitors) + assert all(client.is_closed for client in clients) + for tunnel in tunnels: + tunnel.aclose.assert_awaited_once() + + async def test_reconnect_limit_is_shared_and_removal_cancels_queued_reconnects( + self, project, service, pool, tunnel_factory + ): + count = service_connection.MAX_CONCURRENT_TUNNEL_RECONNECTS + 3 + tunnels = [_make_tunnel() for _ in range(count)] + tunnel_factory.side_effect = tunnels + active = 0 + peak = 0 + started = asyncio.Event() + release = asyncio.Event() + + async def reconnect(): + nonlocal active, peak + active += 1 + peak = max(active, peak) + if active == service_connection.MAX_CONCURRENT_TUNNEL_RECONNECTS: + started.set() + try: + await release.wait() + finally: + active -= 1 + + for index, tunnel in enumerate(tunnels): + replica = service.replicas[0].model_copy(update={"id": f"replica-{index}"}) + await pool.get_or_add(project, service, replica) + tunnel.aopen.side_effect = reconnect + tunnel.exits.put_nowait(255) + await asyncio.wait_for(started.wait(), timeout=1) + await _yield_to_event_loop() + + # All tunnels have exited, but only a bounded number are starting SSH processes. + assert peak == service_connection.MAX_CONCURRENT_TUNNEL_RECONNECTS + await pool.remove_all() + release.set() + await _yield_to_event_loop() + assert active == 0 + assert not pool.connections + assert sum(t.aopen.await_count - 1 for t in tunnels) == peak + + +class _MonitorClock: + """Wake monitor timers explicitly, without waiting for wall-clock time.""" + + def __init__(self): + self.pending = asyncio.Queue() + self.tasks = set() + + async def sleep(self, delay): + self.tasks.add(asyncio.current_task()) + wakeup = asyncio.get_running_loop().create_future() + self.pending.put_nowait((delay, wakeup)) + await wakeup + + async def tick(self): + await _yield_to_event_loop() + delay, wakeup = self.pending.get_nowait() + wakeup.set_result(None) + await _yield_to_event_loop() + return delay + + +def _make_tunnel(): + tunnel = create_autospec(SSHTunnel, instance=True) + tunnel.exits = asyncio.Queue() + tunnel.wait_closed.side_effect = tunnel.exits.get + return tunnel + + +async def _yield_to_event_loop(): + ready = asyncio.get_running_loop().create_future() + asyncio.get_running_loop().call_soon(ready.set_result, None) + await ready From be7efe05ce032b8ae9dc20ebabe03871981f847a Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Thu, 8 Oct 2026 10:58:08 +0200 Subject: [PATCH 2/3] Make SSH cancellation tests compatible with Python 3.10 --- src/tests/_internal/core/services/ssh/test_tunnel.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/src/tests/_internal/core/services/ssh/test_tunnel.py b/src/tests/_internal/core/services/ssh/test_tunnel.py index b1bca4b85f..292e36e96c 100644 --- a/src/tests/_internal/core/services/ssh/test_tunnel.py +++ b/src/tests/_internal/core/services/ssh/test_tunnel.py @@ -602,7 +602,7 @@ async def create_process(*args, **kwargs): await spawned.wait() task.cancel("cancel while creating SSH") return_process.set() - with pytest.raises(asyncio.CancelledError, match="cancel while creating SSH"): + with pytest.raises(asyncio.CancelledError): await task assert returned.is_set() @@ -641,7 +641,7 @@ async def create_process(*args, **kwargs): fail_creation.set() if cancelled: - with pytest.raises(asyncio.CancelledError, match="cancel while creating SSH"): + with pytest.raises(asyncio.CancelledError): await task else: with pytest.raises(OSError) as exc_info: @@ -671,7 +671,7 @@ async def wait(): await reaping.wait() closing.cancel("cancel during cleanup") reaped.set() - with pytest.raises(asyncio.CancelledError, match="cancel during cleanup"): + with pytest.raises(asyncio.CancelledError): await closing ssh_process.wait.assert_awaited_once_with() @@ -787,7 +787,7 @@ async def communicate(): await communicating.wait() task.cancel("cancel control command") - with pytest.raises(asyncio.CancelledError, match="cancel control command"): + with pytest.raises(asyncio.CancelledError): await task process.kill.assert_called_once_with() From 57a77b1b063b921e715a3d1060d55d3cf28fb474 Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Thu, 8 Oct 2026 12:14:51 +0200 Subject: [PATCH 3/3] Trim SSH tunnel recovery regression tests --- .../core/services/ssh/test_tunnel.py | 157 +------------ .../lib/services/test_service_connection.py | 216 ++---------------- 2 files changed, 36 insertions(+), 337 deletions(-) diff --git a/src/tests/_internal/core/services/ssh/test_tunnel.py b/src/tests/_internal/core/services/ssh/test_tunnel.py index 292e36e96c..0a3a1f8ce7 100644 --- a/src/tests/_internal/core/services/ssh/test_tunnel.py +++ b/src/tests/_internal/core/services/ssh/test_tunnel.py @@ -391,8 +391,7 @@ async def test_foreground_checks_only_after_control_socket_exists( check = AsyncMock(return_value=True) monkeypatch.setattr(foreground_tunnel, "acheck", check) - async def startup_poll(interval): - assert interval == 0.1 + async def startup_poll(_interval): polling.set() await create_socket.wait() Path(foreground_tunnel.control_sock_path).touch() @@ -444,65 +443,14 @@ async def check(): await foreground_tunnel.wait_closed() @pytest.mark.asyncio - @pytest.mark.parametrize("method", ["aopen", "acheck"]) - @pytest.mark.parametrize("already_exited", [False, True]) - @pytest.mark.parametrize( - "is_windows", - [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], - ) - async def test_cancelled_command_kills_and_reaps_process( - self, - foreground_tunnel: SSHTunnel, - ssh_process: Mock, - kill_process_group: Mock, - monkeypatch: pytest.MonkeyPatch, - method: str, - already_exited: bool, - is_windows: bool, - ) -> None: - monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) - waiting = asyncio.Event() - - async def wait_for_process(*args): - if ssh_process.returncode is None: - waiting.set() - await asyncio.Future() - return ssh_process.returncode - - def kill(*args) -> None: - ssh_process.returncode = 0 if already_exited else -9 - if already_exited: - raise ProcessLookupError - - monkeypatch.setattr( - foreground_tunnel, "_wait_until_ready", AsyncMock(side_effect=wait_for_process) - ) - ssh_process.communicate.side_effect = wait_for_process - ssh_process.wait.side_effect = wait_for_process - ssh_process.kill.side_effect = kill - kill_process_group.side_effect = kill - task = asyncio.create_task(getattr(foreground_tunnel, method)()) - await waiting.wait() - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task - - if method == "aopen" and not is_windows: - kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) - ssh_process.kill.assert_not_called() - else: - ssh_process.kill.assert_called_once_with() - kill_process_group.assert_not_called() - ssh_process.wait.assert_awaited_once_with() - assert task.cancelled() - - @pytest.mark.asyncio + @pytest.mark.parametrize("check_already_exited", [False, True]) async def test_cancelled_foreground_readiness_cleans_check_and_master( self, foreground_tunnel: SSHTunnel, ssh_process: Mock, kill_process_group: Mock, monkeypatch: pytest.MonkeyPatch, + check_already_exited: bool, ) -> None: Path(foreground_tunnel.control_sock_path).touch() check_process = Mock(spec=asyncio.subprocess.Process) @@ -517,6 +465,8 @@ async def wait(): def kill(): check_process.returncode = -9 + if check_already_exited: + raise ProcessLookupError check_process.communicate.side_effect = wait check_process.kill.side_effect = kill @@ -539,31 +489,23 @@ def kill(): assert foreground_tunnel._process is None @pytest.mark.asyncio - @pytest.mark.parametrize("already_exited", [False, True]) - @pytest.mark.parametrize( - "is_windows", - [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], - ) - async def test_timed_out_aopen_kills_and_reaps_process( + async def test_timed_out_aopen_cleans_up_exited_process( self, foreground_tunnel: SSHTunnel, ssh_process: Mock, kill_process_group: Mock, monkeypatch: pytest.MonkeyPatch, - already_exited: bool, - is_windows: bool, ) -> None: # A zero timeout exercises wait_for's cleanup without waiting on real time. monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.SSH_TIMEOUT", 0) - monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) - if already_exited: - ssh_process.kill.side_effect = ProcessLookupError - kill_process_group.side_effect = ProcessLookupError + ssh_process.returncode = 255 + ssh_process.kill.side_effect = ProcessLookupError + kill_process_group.side_effect = ProcessLookupError with pytest.raises(SSHError, match="in 0 seconds") as exc_info: await foreground_tunnel.aopen() - if not is_windows: + if not IS_WINDOWS: kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) ssh_process.kill.assert_not_called() else: @@ -573,19 +515,13 @@ async def test_timed_out_aopen_kills_and_reaps_process( assert isinstance(exc_info.value.__cause__, asyncio.TimeoutError) @pytest.mark.asyncio - @pytest.mark.parametrize( - "is_windows", - [True, pytest.param(False, marks=pytest.mark.skipif(IS_WINDOWS, reason="POSIX signals"))], - ) async def test_cancelled_aopen_during_creation_kills_and_reaps_process( self, foreground_tunnel: SSHTunnel, ssh_process: Mock, kill_process_group: Mock, monkeypatch: pytest.MonkeyPatch, - is_windows: bool, ) -> None: - monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.IS_WINDOWS", is_windows) spawned = asyncio.Event() return_process = asyncio.Event() returned = asyncio.Event() @@ -606,7 +542,7 @@ async def create_process(*args, **kwargs): await task assert returned.is_set() - if is_windows: + if IS_WINDOWS: ssh_process.kill.assert_called_once_with() kill_process_group.assert_not_called() else: @@ -682,20 +618,6 @@ async def wait(): else: kill_process_group.assert_called_once_with(ssh_process.pid, signal.SIGKILL) - @pytest.mark.asyncio - @pytest.mark.parametrize("returncode", [0, 255]) - async def test_acheck_returns_process_status( - self, - sample_tunnel_with_all_params: SSHTunnel, - ssh_process: Mock, - returncode: int, - ) -> None: - ssh_process.returncode = returncode - - assert await sample_tunnel_with_all_params.acheck() is (returncode == 0) - - ssh_process.kill.assert_not_called() - class TestSSHTunnelControlTimeouts: @pytest.fixture @@ -761,63 +683,6 @@ async def test_aexec_kills_process_and_raises_on_timeout( await tunnel.aexec("true", timeout=0) assert hanging_process.killed - @pytest.mark.asyncio - @pytest.mark.parametrize("method", ["acheck", "aclose", "aexec"]) - @pytest.mark.parametrize("already_exited", [False, True]) - async def test_cancelled_control_command_kills_and_reaps_process( - self, - tunnel: SSHTunnel, - monkeypatch: pytest.MonkeyPatch, - method: str, - already_exited: bool, - ) -> None: - process = Mock(spec=asyncio.subprocess.Process) - communicating = asyncio.Event() - - async def communicate(): - communicating.set() - await asyncio.Future() - - process.communicate.side_effect = communicate - if already_exited: - process.kill.side_effect = ProcessLookupError - monkeypatch.setattr(asyncio, "create_subprocess_exec", AsyncMock(return_value=process)) - command = getattr(tunnel, method)(*(["true"] if method == "aexec" else [])) - task = asyncio.create_task(command) - await communicating.wait() - task.cancel("cancel control command") - - with pytest.raises(asyncio.CancelledError): - await task - - process.kill.assert_called_once_with() - process.wait.assert_awaited_once_with() - - @pytest.mark.asyncio - @pytest.mark.parametrize("method", ["acheck", "aclose", "aexec"]) - async def test_timeout_preserves_behavior_when_process_has_already_exited( - self, - tunnel: SSHTunnel, - hanging_process: "_HangingProcess", - monkeypatch: pytest.MonkeyPatch, - method: str, - ) -> None: - def kill(): - hanging_process.returncode = 0 - raise ProcessLookupError - - reaped = AsyncMock(wraps=hanging_process.wait) - monkeypatch.setattr(hanging_process, "kill", kill) - monkeypatch.setattr(hanging_process, "wait", reaped) - if method == "aexec": - with pytest.raises(SSHError, match="did not complete") as exc_info: - await tunnel.aexec("true", timeout=0) - assert isinstance(exc_info.value.__cause__, asyncio.TimeoutError) - else: - result = await getattr(tunnel, method)() - assert result is (False if method == "acheck" else None) - reaped.assert_awaited_once_with() - class _HangingProcess: def __init__(self) -> None: diff --git a/src/tests/_internal/proxy/lib/services/test_service_connection.py b/src/tests/_internal/proxy/lib/services/test_service_connection.py index 4026572794..f12eddf0da 100644 --- a/src/tests/_internal/proxy/lib/services/test_service_connection.py +++ b/src/tests/_internal/proxy/lib/services/test_service_connection.py @@ -54,6 +54,11 @@ async def test_gateway_recovers_the_same_tunnel_and_socket_without_polling( tunnel = tunnel_factory.return_value connection = await pool.get_or_add(project, service, service.replicas[0]) socket_path = connection.app_socket_path + await connection.open() + await _yield_to_event_loop() + tunnel.aopen.assert_awaited_once() + tunnel.wait_closed.assert_awaited_once() + tunnel.exits.put_nowait(255) await _yield_to_event_loop() @@ -62,30 +67,16 @@ async def test_gateway_recovers_the_same_tunnel_and_socket_without_polling( assert tunnel.aopen.await_count == 2 tunnel.aclose.assert_awaited_once() tunnel.acheck.assert_not_awaited() - assert not monitor_clock.tasks + assert monitor_clock.pending.empty() assert tunnel_factory.call_args.kwargs["background"] is False assert connection.app_socket_path == socket_path assert tunnel_factory.call_args.kwargs["forwarded_sockets"][0].local.path == socket_path - async def test_healthy_gateway_does_not_reopen_or_start_duplicate_monitors( - self, project, service, pool, tunnel_factory, monitor_clock - ): - connection = await pool.get_or_add(project, service, service.replicas[0]) - await connection.open() - - await _yield_to_event_loop() - - tunnel_factory.return_value.aopen.assert_awaited_once() - tunnel_factory.return_value.wait_closed.assert_awaited_once() - tunnel_factory.return_value.acheck.assert_not_awaited() - assert not monitor_clock.tasks - async def test_failed_reconnects_back_off_and_then_recover( self, project, service, pool, tunnel_factory, monitor_clock ): tunnel = tunnel_factory.return_value - connection = await pool.get_or_add(project, service, service.replicas[0]) - original_client = await connection.client() + await pool.get_or_add(project, service, service.replicas[0]) tunnel.aopen.side_effect = [SSHError("offline"), SSHError("offline"), None] tunnel.exits.put_nowait(255) await _yield_to_event_loop() @@ -94,8 +85,6 @@ async def test_failed_reconnects_back_off_and_then_recover( assert 1 <= await monitor_clock.tick() <= 2 assert tunnel.aopen.await_count == 4 - assert await asyncio.wait_for(connection.client(), timeout=1) is original_client - assert not original_client.is_closed async def test_repeated_early_exits_back_off_to_a_bounded_delay( self, project, service, pool, tunnel_factory, monitor_clock @@ -122,10 +111,10 @@ async def test_stable_tunnel_resets_retry_delay( await _yield_to_event_loop() assert tunnel.aopen.await_count == 3 - assert not monitor_clock.tasks + assert monitor_clock.pending.empty() async def test_in_server_connection_does_not_monitor( - self, project, service, pool, tunnel_factory, monitor_clock + self, project, service, pool, tunnel_factory ): service = service.model_copy(update={"domain": None}) connection = await pool.get_or_add(project, service, service.replicas[0]) @@ -135,34 +124,9 @@ async def test_in_server_connection_does_not_monitor( tunnel_factory.return_value.aopen.assert_awaited_once() tunnel_factory.return_value.wait_closed.assert_not_awaited() assert tunnel_factory.call_args.kwargs["background"] is True - assert not monitor_clock.tasks - - async def test_client_does_not_wait_for_a_blocked_reconnect( - self, project, service, pool, tunnel_factory, monitor_clock - ): - tunnel = tunnel_factory.return_value - connection = await pool.get_or_add(project, service, service.replicas[0]) - original_client = await connection.client() - reconnecting = asyncio.Event() - release_reconnect = asyncio.Event() - async def reconnect(): - reconnecting.set() - await release_reconnect.wait() - - tunnel.aopen.side_effect = reconnect - tunnel.exits.put_nowait(255) - await _yield_to_event_loop() - assert reconnecting.is_set() - - # The timeout is only a deadlock guard; the reconnect stays blocked throughout. - client = await asyncio.wait_for(connection.client(), timeout=1) - assert client is original_client - assert not release_reconnect.is_set() - - @pytest.mark.parametrize("remove_from_pool", [False, True], ids=["close", "remove"]) - async def test_close_cancels_reconnect_before_final_tunnel_cleanup( - self, project, service, pool, tunnel_factory, monitor_clock, remove_from_pool + async def test_removal_cancels_reconnect_before_final_tunnel_cleanup( + self, project, service, pool, tunnel_factory ): tunnel = tunnel_factory.return_value connection = await pool.get_or_add(project, service, service.replicas[0]) @@ -183,22 +147,17 @@ async def close_tunnel(): tunnel.aopen.side_effect = reconnect tunnel.aclose.side_effect = close_tunnel - monitor_task = connection._monitor_task tunnel.exits.put_nowait(255) await _yield_to_event_loop() assert lifecycle == ["closed", "reconnecting"] - if remove_from_pool: - await pool.remove(service.replicas[0].id) - assert await pool.get(service.replicas[0].id) is None - else: - await connection.close() + await pool.remove(service.replicas[0].id) release_reconnect.set() await _yield_to_event_loop() + assert await pool.get(service.replicas[0].id) is None assert lifecycle == ["closed", "reconnecting", "cancelled", "closed"] assert client.is_closed - assert monitor_task.cancelled() with pytest.raises(UnexpectedProxyError, match="closed"): await connection.open() assert tunnel.aopen.await_count == 2 @@ -206,9 +165,8 @@ async def close_tunnel(): @pytest.mark.asyncio class TestServiceConnectionPool: - @pytest.mark.parametrize("cancel_count", [1, 2]) async def test_cancelled_removal_finishes_monitor_and_resource_cleanup( - self, project, service, pool, tunnel_factory, cancel_count + self, project, service, pool, tunnel_factory ): tunnel = tunnel_factory.return_value waiting = asyncio.Event() @@ -232,7 +190,7 @@ async def wait_closed(): removing = asyncio.create_task(pool.remove(replica.id)) try: await monitor_cancelled.wait() - for _ in range(cancel_count): + for _ in range(2): removing.cancel() await _yield_to_event_loop() assert not removing.done() @@ -246,137 +204,42 @@ async def wait_closed(): assert connection._monitor_task is None tunnel.aclose.assert_awaited_once() + @pytest.mark.parametrize("error", [RuntimeError, asyncio.CancelledError]) async def test_initial_open_failure_closes_tunnel_and_removes_connection( - self, project, service, pool, tunnel_factory, monitor_clock + self, project, service, pool, tunnel_factory, error ): tunnel = tunnel_factory.return_value - tunnel.aopen.side_effect = RuntimeError("cannot reach replica") + tunnel.aopen.side_effect = error("cannot reach replica") - with pytest.raises(RuntimeError, match="cannot reach replica"): + with pytest.raises(error): await pool.get_or_add(project, service, service.replicas[0]) assert await pool.get(service.replicas[0].id) is None tunnel.aclose.assert_awaited_once() - assert not monitor_clock.tasks - - async def test_cancelled_initial_open_closes_tunnel_and_removes_connection( - self, project, service, pool, tunnel_factory, monitor_clock - ): - tunnel = tunnel_factory.return_value - opening = asyncio.Event() - - async def open_tunnel(): - opening.set() - await asyncio.Event().wait() - - tunnel.aopen.side_effect = open_tunnel - adding = asyncio.create_task(pool.get_or_add(project, service, service.replicas[0])) - try: - await _yield_to_event_loop() - assert opening.is_set() - adding.cancel() - with pytest.raises(asyncio.CancelledError): - await adding - finally: - adding.cancel() - await asyncio.gather(adding, return_exceptions=True) - - assert await pool.get(service.replicas[0].id) is None - tunnel.aclose.assert_awaited_once() - assert not monitor_clock.tasks - - async def test_remove_finishing_after_readd_keeps_new_connection( - self, project, service, pool, tunnel_factory, monitor_clock - ): - old_tunnel = tunnel_factory.return_value - new_tunnel = _make_tunnel() - tunnel_factory.side_effect = [old_tunnel, new_tunnel] - replica = service.replicas[0] - await pool.get_or_add(project, service, replica) - closing = asyncio.Event() - release_close = asyncio.Event() - - async def close_tunnel(): - closing.set() - await release_close.wait() - - old_tunnel.aclose.side_effect = close_tunnel - removing = asyncio.create_task(pool.remove(replica.id)) - try: - await asyncio.wait_for(closing.wait(), timeout=1) - assert await pool.get(replica.id) is None - replacement = await pool.get_or_add(project, service, replica) - finally: - release_close.set() - await removing - - assert closing.is_set() - assert await pool.get(replica.id) is replacement - assert not (await replacement.client()).is_closed - new_tunnel.aclose.assert_not_awaited() - - async def test_failed_open_cleanup_after_readd_keeps_new_connection( - self, project, service, pool, tunnel_factory, monitor_clock - ): - old_tunnel = tunnel_factory.return_value - new_tunnel = _make_tunnel() - tunnel_factory.side_effect = [old_tunnel, new_tunnel] - replica = service.replicas[0] - release_open = asyncio.Event() - - async def open_tunnel(): - await release_open.wait() - raise RuntimeError("old connection failed") - - old_tunnel.aopen.side_effect = open_tunnel - adding = asyncio.create_task(pool.get_or_add(project, service, replica)) - removing = None - try: - await _yield_to_event_loop() - removing = asyncio.create_task(pool.remove(replica.id)) - await _yield_to_event_loop() - replacement = await pool.get_or_add(project, service, replica) - release_open.set() - with pytest.raises(RuntimeError, match="old connection failed"): - await adding - await removing - - assert await pool.get(replica.id) is replacement - assert not (await replacement.client()).is_closed - old_tunnel.aclose.assert_awaited() - new_tunnel.aclose.assert_not_awaited() - finally: - release_open.set() - adding.cancel() - await asyncio.gather(adding, return_exceptions=True) - if removing is not None: - await removing + tunnel.wait_closed.assert_not_awaited() async def test_removed_initial_open_succeeding_cannot_resurrect_or_remove_replacement( - self, project, service, pool, tunnel_factory, monitor_clock + self, project, service, pool, tunnel_factory ): old_tunnel = tunnel_factory.return_value new_tunnel = _make_tunnel() tunnel_factory.side_effect = [old_tunnel, new_tunnel] replica = service.replicas[0] + opening = asyncio.Event() release_open = asyncio.Event() async def open_tunnel(): + opening.set() await release_open.wait() old_tunnel.aopen.side_effect = open_tunnel adding = asyncio.create_task(pool.get_or_add(project, service, replica)) removing = None try: - await _yield_to_event_loop() - original = await pool.get(replica.id) - assert original is not None + await opening.wait() removing = asyncio.create_task(pool.remove(replica.id)) await _yield_to_event_loop() replacement = await pool.get_or_add(project, service, replica) - await _yield_to_event_loop() - replacement_monitor = replacement._monitor_task - assert replacement_monitor is not None release_open.set() with pytest.raises(UnexpectedProxyError, match="removed while opening"): @@ -385,14 +248,10 @@ async def open_tunnel(): assert await pool.get(replica.id) is replacement assert not (await replacement.client()).is_closed - assert replacement._monitor_task is replacement_monitor - assert not replacement_monitor.done() - old_tunnel.aopen.assert_awaited_once() old_tunnel.wait_closed.assert_not_awaited() old_tunnel.aclose.assert_awaited() + new_tunnel.wait_closed.assert_awaited_once() new_tunnel.aclose.assert_not_awaited() - with pytest.raises(UnexpectedProxyError, match="closed"): - await original.open() finally: release_open.set() adding.cancel() @@ -400,29 +259,6 @@ async def open_tunnel(): if removing is not None: await removing - async def test_remove_all_cancels_every_monitor_even_if_one_close_fails( - self, project, service, pool, tunnel_factory, monitor_clock - ): - tunnels = [_make_tunnel() for _ in range(2)] - tunnel_factory.side_effect = tunnels - clients = [] - monitors = [] - for index in range(2): - replica = service.replicas[0].model_copy(update={"id": f"replica-{index}"}) - connection = await pool.get_or_add(project, service, replica) - clients.append(await connection.client()) - monitors.append(connection._monitor_task) - await _yield_to_event_loop() - tunnels[0].aclose.side_effect = RuntimeError("already gone") - - await pool.remove_all() - - assert not pool.connections - assert all(task.cancelled() for task in monitors) - assert all(client.is_closed for client in clients) - for tunnel in tunnels: - tunnel.aclose.assert_awaited_once() - async def test_reconnect_limit_is_shared_and_removal_cancels_queued_reconnects( self, project, service, pool, tunnel_factory ): @@ -468,10 +304,8 @@ class _MonitorClock: def __init__(self): self.pending = asyncio.Queue() - self.tasks = set() async def sleep(self, delay): - self.tasks.add(asyncio.current_task()) wakeup = asyncio.get_running_loop().create_future() self.pending.put_nowait((delay, wakeup)) await wakeup