Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
116 changes: 88 additions & 28 deletions src/dstack/_internal/core/services/ssh/tunnel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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()

Expand All @@ -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:
Expand Down Expand Up @@ -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]:
Expand Down
4 changes: 3 additions & 1 deletion src/dstack/_internal/server/services/gateways/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
114 changes: 97 additions & 17 deletions src/tests/_internal/core/services/ssh/test_tunnel.py
Original file line number Diff line number Diff line change
@@ -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 (
Expand Down Expand Up @@ -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:
Expand Down
Loading