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
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
111 changes: 102 additions & 9 deletions src/dstack/_internal/core/services/ssh/tunnel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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)
Expand All @@ -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)
Expand All @@ -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,
]
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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(
Expand All @@ -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"
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
101 changes: 93 additions & 8 deletions src/dstack/_internal/proxy/lib/services/service_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand All @@ -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,
Expand Down Expand Up @@ -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)),
Expand All @@ -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)
Expand All @@ -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

Expand Down
Loading
Loading