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..0a3a1f8ce7 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,353 @@ 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): + 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("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) + 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 + if check_already_exited: + raise ProcessLookupError + + 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 + async def test_timed_out_aopen_cleans_up_exited_process( + self, + foreground_tunnel: SSHTunnel, + ssh_process: Mock, + kill_process_group: Mock, + monkeypatch: pytest.MonkeyPatch, + ) -> 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) + 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: + 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 + 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, + ) -> None: + 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): + 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): + 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): + 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) + class TestSSHTunnelControlTimeouts: @pytest.fixture 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..f12eddf0da --- /dev/null +++ b/src/tests/_internal/proxy/lib/services/test_service_connection.py @@ -0,0 +1,331 @@ +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 + 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() + + # 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 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_failed_reconnects_back_off_and_then_recover( + self, project, service, pool, tunnel_factory, monitor_clock + ): + tunnel = tunnel_factory.return_value + 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() + + assert 0.5 <= await monitor_clock.tick() <= 1 + assert 1 <= await monitor_clock.tick() <= 2 + + assert tunnel.aopen.await_count == 4 + + 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 monitor_clock.pending.empty() + + async def test_in_server_connection_does_not_monitor( + self, project, service, pool, tunnel_factory + ): + 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 + + 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]) + 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 + tunnel.exits.put_nowait(255) + await _yield_to_event_loop() + assert lifecycle == ["closed", "reconnecting"] + + 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 + with pytest.raises(UnexpectedProxyError, match="closed"): + await connection.open() + assert tunnel.aopen.await_count == 2 + + +@pytest.mark.asyncio +class TestServiceConnectionPool: + async def test_cancelled_removal_finishes_monitor_and_resource_cleanup( + self, project, service, pool, tunnel_factory + ): + 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(2): + 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() + + @pytest.mark.parametrize("error", [RuntimeError, asyncio.CancelledError]) + async def test_initial_open_failure_closes_tunnel_and_removes_connection( + self, project, service, pool, tunnel_factory, error + ): + tunnel = tunnel_factory.return_value + tunnel.aopen.side_effect = error("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() + tunnel.wait_closed.assert_not_awaited() + + async def test_removed_initial_open_succeeding_cannot_resurrect_or_remove_replacement( + 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 opening.wait() + 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(UnexpectedProxyError, match="removed while opening"): + await adding + await removing + + assert await pool.get(replica.id) is replacement + assert not (await replacement.client()).is_closed + 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() + finally: + release_open.set() + adding.cancel() + await asyncio.gather(adding, return_exceptions=True) + if removing is not None: + await removing + + 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() + + async def sleep(self, delay): + 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