diff --git a/src/dstack/_internal/core/services/ssh/tunnel.py b/src/dstack/_internal/core/services/ssh/tunnel.py index f4d6a17f70..9dff46c6e0 100644 --- a/src/dstack/_internal/core/services/ssh/tunnel.py +++ b/src/dstack/_internal/core/services/ssh/tunnel.py @@ -88,6 +88,9 @@ def __init__( configured `destination` with `ProxyJump` in the `ssh_config_path` config, the proxy jump connection will ignore this option -- in that case, you should replace `ProxyJump` with explicit `ProxyCommand=ssh [...] -o BatchMode=yes` in your config. + 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. """ self.destination = destination self.forwarded_sockets = list(forwarded_sockets) @@ -173,13 +176,13 @@ def open_command(self) -> List[str]: return command def close_command(self) -> List[str]: - return [self.ssh_exec_path, "-S", self.control_sock_path, "-O", "exit", self.destination] + return [*self._control_command_prefix(), "-O", "exit", self.destination] def check_command(self) -> List[str]: - return [self.ssh_exec_path, "-S", self.control_sock_path, "-O", "check", self.destination] + return [*self._control_command_prefix(), "-O", "check", self.destination] def exec_command(self) -> List[str]: - return [self.ssh_exec_path, "-S", self.control_sock_path, self.destination] + return [*self._control_command_prefix(), self.destination] def open(self) -> None: # We cannot use `stderr=subprocess.PIPE` here since the forked process (daemon) does not @@ -225,9 +228,17 @@ def close(self) -> None: "Control socket does not exist, it seems that ssh process has already exited" ) return - proc = subprocess.run( - self.close_command(), stdout=subprocess.PIPE, stderr=subprocess.STDOUT - ) + try: + proc = subprocess.run( + self.close_command(), + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + timeout=SSH_TIMEOUT, + ) + except subprocess.TimeoutExpired: + logger.error("Failed to close SSH tunnel in %d seconds", SSH_TIMEOUT) + return if proc.returncode: logger.error( "Failed to close SSH tunnel, exit status: %d, output: %s", @@ -241,37 +252,52 @@ async def aclose(self) -> None: "Control socket does not exist, it seems that ssh process has already exited" ) return - proc = await asyncio.create_subprocess_exec( - *self.close_command(), stdout=subprocess.PIPE, stderr=subprocess.STDOUT - ) - await proc.wait() - if proc.returncode: + try: + returncode, stdout, stderr = await _arun(self.close_command(), SSH_TIMEOUT) + except asyncio.TimeoutError: + logger.error("Failed to close SSH tunnel in %d seconds", SSH_TIMEOUT) + return + if returncode: logger.error( "Failed to close SSH tunnel, exit status: %d, output: %s", - proc.returncode, - proc.stdout, + returncode, + stdout + stderr, ) def check(self) -> bool: - proc = subprocess.run( - self.check_command(), stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL - ) + try: + proc = subprocess.run( + self.check_command(), + stdin=subprocess.DEVNULL, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + timeout=SSH_TIMEOUT, + ) + except subprocess.TimeoutExpired: + logger.debug("SSH tunnel check did not complete in %d seconds", SSH_TIMEOUT) + return False return proc.returncode == 0 async def acheck(self) -> bool: - proc = await asyncio.create_subprocess_exec( - *self.check_command(), stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL - ) - await proc.wait() - ok = proc.returncode == 0 - return ok + try: + returncode, _, _ = await _arun(self.check_command(), SSH_TIMEOUT) + except asyncio.TimeoutError: + logger.debug("SSH tunnel check did not complete in %d seconds", SSH_TIMEOUT) + return False + return returncode == 0 - async def aexec(self, command: str) -> str: - proc = await asyncio.create_subprocess_exec( - *self.exec_command(), command, stdout=subprocess.PIPE, stderr=subprocess.PIPE - ) - stdout, stderr = await proc.communicate() - if proc.returncode != 0: + async def aexec(self, command: str, timeout: float = SSH_TIMEOUT) -> str: + """ + Runs `command` on the remote host over the open tunnel. + + :param timeout: Seconds to wait for `command` to complete before killing it + and raising `SSHError`. + """ + try: + returncode, stdout, stderr = await _arun([*self.exec_command(), command], timeout) + except asyncio.TimeoutError as e: + raise SSHError(f"Command did not complete in {timeout} seconds") from e + if returncode != 0: raise SSHError(stderr.decode()) return stdout.decode() @@ -289,6 +315,21 @@ async def __aenter__(self): async def __aexit__(self, exc_type, exc_val, exc_tb): await self.aclose() + def _control_command_prefix(self) -> List[str]: + # If the master does not complete the initial exchange (or, for exec, the control socket + # is missing), OpenSSH falls back to connecting to `destination` directly, even for `-O` + # commands. Ignore the user's ssh config and disable prompts so that such a connection + # fails instead of waiting for input on the terminal. + return [ + self.ssh_exec_path, + "-F", + "none", + "-o", + "BatchMode=yes", + "-S", + self.control_sock_path, + ] + def _get_proxy_command(self) -> Optional[str]: proxy_command: Optional[str] = None for params, identity_path in self.ssh_proxies: @@ -366,6 +407,25 @@ def _get_identity_path(self, identity: FilePathOrContent, tmp_filename: str) -> return identity_path +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. + """ + 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() + await proc.wait() + raise + assert proc.returncode is not None + return proc.returncode, stdout, stderr + + def ports_to_forwarded_sockets( ports: Dict[int, int], bind_local: str = "localhost" ) -> List[SocketPair]: diff --git a/src/dstack/_internal/server/services/gateways/__init__.py b/src/dstack/_internal/server/services/gateways/__init__.py index 100a955a37..793654ded1 100644 --- a/src/dstack/_internal/server/services/gateways/__init__.py +++ b/src/dstack/_internal/server/services/gateways/__init__.py @@ -832,13 +832,15 @@ async def _update_gateway_replica( + " " + shlex.quote(target_version or "") ) - stdout = await connection.tunnel.aexec(command) + stdout = await connection.tunnel.aexec(command, timeout=_GATEWAY_UPDATE_TIMEOUT) if "Update successfully completed" in stdout: logger.info("Gateway replica %s updated", connection.ip_address) return True return False +_GATEWAY_UPDATE_TIMEOUT = 600 + # Blue/green: install the new build into the currently inactive venv and flip to it _GATEWAY_UPDATE_SCRIPT = """\ set -e diff --git a/src/tests/_internal/core/services/ssh/test_tunnel.py b/src/tests/_internal/core/services/ssh/test_tunnel.py index b97c86a1e6..da61536bd5 100644 --- a/src/tests/_internal/core/services/ssh/test_tunnel.py +++ b/src/tests/_internal/core/services/ssh/test_tunnel.py @@ -1,7 +1,12 @@ +import asyncio +import subprocess from pathlib import Path +from typing import NoReturn, Optional +from unittest.mock import Mock import pytest +from dstack._internal.core.errors import SSHError from dstack._internal.core.models.instances import SSHConnectionParams from dstack._internal.core.services.ssh.client import SSHClientInfo from dstack._internal.core.services.ssh.tunnel import ( @@ -221,29 +226,104 @@ def test_open_command_with_forwarding(self) -> None: def test_check_command(self, sample_tunnel_with_all_params: SSHTunnel) -> None: command = sample_tunnel_with_all_params.check_command() - assert command == [ - "/usr/bin/ssh", - "-S", - "/tmp/control.sock", - "-O", - "check", - "ubuntu@my-server", - ] + assert " ".join(command) == ( + "/usr/bin/ssh -F none -o BatchMode=yes -S /tmp/control.sock -O check ubuntu@my-server" + ) def test_close_command(self, sample_tunnel_with_all_params: SSHTunnel) -> None: command = sample_tunnel_with_all_params.close_command() - assert command == [ - "/usr/bin/ssh", - "-S", - "/tmp/control.sock", - "-O", - "exit", - "ubuntu@my-server", - ] + assert " ".join(command) == ( + "/usr/bin/ssh -F none -o BatchMode=yes -S /tmp/control.sock -O exit ubuntu@my-server" + ) def test_exec_command(self, sample_tunnel_with_all_params: SSHTunnel) -> None: command = sample_tunnel_with_all_params.exec_command() - assert command == ["/usr/bin/ssh", "-S", "/tmp/control.sock", "ubuntu@my-server"] + assert " ".join(command) == ( + "/usr/bin/ssh -F none -o BatchMode=yes -S /tmp/control.sock ubuntu@my-server" + ) + + +class TestSSHTunnelControlTimeouts: + @pytest.fixture + def tunnel(self, monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> SSHTunnel: + monkeypatch.setattr( + "dstack._internal.core.services.ssh.client._ssh_client_info", + SSHClientInfo.from_raw_version("OpenSSH_9.7p1", Path("/usr/bin/ssh")), + ) + control_sock_path = tmp_path / "control.sock" + control_sock_path.touch() + return SSHTunnel( + destination="ubuntu@my-server", + identity=FilePath("/home/user/.ssh/id_rsa"), + control_sock_path=control_sock_path, + ) + + @pytest.fixture + def hanging_process(self, monkeypatch: pytest.MonkeyPatch) -> "_HangingProcess": + process = _HangingProcess() + + async def create_subprocess_exec(*args, **kwargs): + return process + + monkeypatch.setattr(asyncio, "create_subprocess_exec", create_subprocess_exec) + monkeypatch.setattr("dstack._internal.core.services.ssh.tunnel.SSH_TIMEOUT", 0) + return process + + @pytest.fixture + def timing_out_run(self, monkeypatch: pytest.MonkeyPatch) -> Mock: + run = Mock(side_effect=subprocess.TimeoutExpired(cmd="ssh", timeout=0)) + monkeypatch.setattr(subprocess, "run", run) + return run + + def test_check_returns_false_on_timeout(self, tunnel: SSHTunnel, timing_out_run: Mock) -> None: + assert tunnel.check() is False + assert timing_out_run.call_args.kwargs["stdin"] == subprocess.DEVNULL + + def test_close_does_not_raise_on_timeout( + self, tunnel: SSHTunnel, timing_out_run: Mock + ) -> None: + tunnel.close() + assert timing_out_run.call_args.kwargs["stdin"] == subprocess.DEVNULL + + @pytest.mark.asyncio + async def test_acheck_kills_process_and_returns_false_on_timeout( + self, tunnel: SSHTunnel, hanging_process: "_HangingProcess" + ) -> None: + assert await tunnel.acheck() is False + assert hanging_process.killed + + @pytest.mark.asyncio + async def test_aclose_kills_process_on_timeout( + self, tunnel: SSHTunnel, hanging_process: "_HangingProcess" + ) -> None: + await tunnel.aclose() + assert hanging_process.killed + + @pytest.mark.asyncio + async def test_aexec_kills_process_and_raises_on_timeout( + self, tunnel: SSHTunnel, hanging_process: "_HangingProcess" + ) -> None: + with pytest.raises(SSHError, match="did not complete"): + await tunnel.aexec("true", timeout=0) + assert hanging_process.killed + + +class _HangingProcess: + def __init__(self) -> None: + self.returncode: Optional[int] = None + self.killed = False + + async def communicate(self) -> NoReturn: + await asyncio.Event().wait() + raise AssertionError("unreachable") + + def kill(self) -> None: + self.killed = True + self.returncode = -9 + + async def wait(self) -> int: + assert self.returncode is not None + return self.returncode def test_ports_to_forwarded_sockets() -> None: