From f36ed0105a28fd61200839799852d8e2a0830bc7 Mon Sep 17 00:00:00 2001 From: Alvils Sture Date: Thu, 8 Oct 2026 05:48:58 +0300 Subject: [PATCH] Reopen gateway replica SSH tunnels whose ssh process exited The gateway opens one SSH tunnel per service replica when the replica is registered and reuses it for every request, but never checks it again. If the tunnel's ssh process exits (the replica's sshd drops a gateway that missed keepalives under load, a network blip, or the process is killed), nginx keeps forwarding to the abandoned Unix socket and the service returns 502 until the gateway is restarted, while probes, which use their own tunnels, stay green. Check every replica tunnel periodically (DSTACK_PROXY_TUNNEL_CHECK_INTERVAL, default 15 s) through its control socket and reopen the ones whose ssh process has exited. The local app socket path is kept, so nginx and the HTTP client need no reconfiguration. A stale control socket left by a killed ssh process is removed first, otherwise the new ssh would run without one and fail every check. A per-connection lock keeps a reopen from resurrecting a connection that is being closed. Fixes #4352 --- src/dstack/_internal/proxy/gateway/app.py | 8 +- .../proxy/lib/services/service_connection.py | 71 +++++++++++- .../_internal/proxy/lib/services/__init__.py | 0 .../lib/services/test_service_connection.py | 108 ++++++++++++++++++ 4 files changed, 182 insertions(+), 5 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/src/dstack/_internal/proxy/gateway/app.py b/src/dstack/_internal/proxy/gateway/app.py index 6627e9e3f3..313a29edcd 100644 --- a/src/dstack/_internal/proxy/gateway/app.py +++ b/src/dstack/_internal/proxy/gateway/app.py @@ -1,6 +1,7 @@ """FastAPI app running on a gateway.""" -from contextlib import asynccontextmanager +import asyncio +from contextlib import asynccontextmanager, suppress from pathlib import Path from typing import Optional @@ -29,6 +30,7 @@ from dstack._internal.proxy.gateway.services.server_client import HTTPMultiClient from dstack._internal.proxy.gateway.services.stats import StatsCollector from dstack._internal.proxy.lib.routers.model_proxy import router as model_proxy_router +from dstack._internal.proxy.lib.services.service_connection import maintain_service_connections from dstack._internal.utils.common import run_async from dstack.version import __version__ @@ -45,9 +47,13 @@ async def lifespan(app: FastAPI): service_conn_pool = await injector.get_service_connection_pool() await run_async(nginx.write_global_conf) await apply_all(repo, nginx, service_conn_pool) + maintenance = asyncio.create_task(maintain_service_connections(service_conn_pool)) yield + maintenance.cancel() + with suppress(asyncio.CancelledError): + await maintenance await service_conn_pool.remove_all() diff --git a/src/dstack/_internal/proxy/lib/services/service_connection.py b/src/dstack/_internal/proxy/lib/services/service_connection.py index c8229ad53a..9b8c0160ae 100644 --- a/src/dstack/_internal/proxy/lib/services/service_connection.py +++ b/src/dstack/_internal/proxy/lib/services/service_connection.py @@ -18,12 +18,15 @@ from dstack._internal.proxy.lib.errors import UnexpectedProxyError from dstack._internal.proxy.lib.models import Project, Replica, Service from dstack._internal.proxy.lib.repo import BaseProxyRepo -from dstack._internal.utils.common import get_or_error +from dstack._internal.utils.common import get_or_error, run_async +from dstack._internal.utils.env import environ from dstack._internal.utils.logging import get_logger from dstack._internal.utils.path import FileContent logger = get_logger(__name__) OPEN_TUNNEL_TIMEOUT = 10 +TUNNEL_CHECK_INTERVAL = environ.get_int("DSTACK_PROXY_TUNNEL_CHECK_INTERVAL", default=15) +"""Seconds between checks that reopen replica SSH tunnels whose ssh process has exited.""" class ServiceClient(httpx.AsyncClient): @@ -75,6 +78,9 @@ def __init__(self, project: Project, service: Service, replica: Replica) -> None timeout=service.read_timeout, ) self._is_open = asyncio.locks.Event() + self._replica_id = replica.id + self._lock = asyncio.Lock() + self._closed = False @property def app_socket_path(self) -> Path: @@ -85,9 +91,33 @@ async def open(self) -> None: self._is_open.set() async def close(self) -> None: - self._is_open.clear() - await self._client.aclose() - await self._tunnel.aclose() + async with self._lock: + self._closed = True + self._is_open.clear() + await self._client.aclose() + await self._tunnel.aclose() + + async def reopen_if_exited(self) -> bool: + """ + Reopen the SSH tunnel if its ssh process has exited, e.g. after the replica's SSH server + dropped the connection because the gateway missed keepalives. The local app socket path + is kept, so nginx and the HTTP client keep working without reconfiguration. + + Returns `True` if the tunnel was reopened. + """ + if self._closed or not self._is_open.is_set(): + return False + if await self._tunnel.acheck(): + return False + async with self._lock: + if self._closed or await self._tunnel.acheck(): + return False + logger.warning("SSH tunnel to service replica %s exited, reopening", self._replica_id) + # The control socket of a killed ssh process stays on disk. A new ssh with + # `ControlMaster=auto` would then run without a control socket and fail every check. + await run_async(_remove_file, Path(self._tunnel.control_sock_path)) + await self._tunnel.aopen() + return True async def client(self) -> ServiceClient: await asyncio.wait_for(self._is_open.wait(), timeout=OPEN_TUNNEL_TIMEOUT) @@ -122,6 +152,20 @@ async def remove(self, replica_id: str) -> None: if connection is not None: await connection.close() + async def reopen_exited(self) -> None: + connections = list(self.connections.items()) + results = await asyncio.gather( + *(connection.reopen_if_exited() for _, connection in connections), + return_exceptions=True, + ) + for (replica_id, _), result in zip(connections, results): + if isinstance(result, BaseException): + logger.warning( + "Failed to reopen SSH tunnel to service replica %s: %r", replica_id, result + ) + elif result: + logger.info("Reopened SSH tunnel to service replica %s", replica_id) + async def remove_all(self) -> None: replica_ids = list(self.connections) results = await asyncio.gather( @@ -158,3 +202,22 @@ async def get_service_replica_client( ) connection = await service_conn_pool.get_or_add(project, service, replica) return await connection.client() + + +async def maintain_service_connections( + service_conn_pool: ServiceConnectionPool, interval: float = TUNNEL_CHECK_INTERVAL +) -> None: + """Periodically reopen replica SSH tunnels whose ssh process has exited.""" + while True: + await asyncio.sleep(interval) + try: + await service_conn_pool.reopen_exited() + except Exception: + logger.exception("Failed to check service replica SSH tunnels") + + +def _remove_file(path: Path) -> None: + try: + path.unlink() + except FileNotFoundError: + pass 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..7175896310 --- /dev/null +++ b/src/tests/_internal/proxy/lib/services/test_service_connection.py @@ -0,0 +1,108 @@ +import asyncio +from pathlib import Path +from unittest.mock import AsyncMock + +import pytest + +from dstack._internal.proxy.lib.services.service_connection import ( + ServiceConnection, + ServiceConnectionPool, + maintain_service_connections, +) +from dstack._internal.proxy.lib.testing.common import make_project, make_service + + +def make_connection(alive: bool = True) -> ServiceConnection: + service = make_service("test-proj", "test-run") + connection = ServiceConnection(make_project("test-proj"), service, service.replicas[0]) + connection._tunnel.aopen = AsyncMock() + connection._tunnel.aclose = AsyncMock() + connection._tunnel.acheck = AsyncMock(return_value=alive) + return connection + + +@pytest.mark.asyncio +async def test_reopen_if_exited_keeps_live_tunnel() -> None: + connection = make_connection(alive=True) + await connection.open() + connection._tunnel.aopen.reset_mock() + + assert not await connection.reopen_if_exited() + connection._tunnel.aopen.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_reopen_if_exited_reopens_dead_tunnel_and_removes_stale_control_socket() -> None: + connection = make_connection(alive=False) + await connection.open() + connection._tunnel.aopen.reset_mock() + control_sock = Path(connection._tunnel.control_sock_path) + control_sock.touch() + + assert await connection.reopen_if_exited() + connection._tunnel.aopen.assert_awaited_once() + assert not control_sock.exists() + assert (await connection.client()) is connection._client + + +@pytest.mark.asyncio +async def test_reopen_if_exited_skips_unopened_and_closed_connections() -> None: + unopened = make_connection(alive=False) + assert not await unopened.reopen_if_exited() + unopened._tunnel.acheck.assert_not_awaited() + + closed = make_connection(alive=False) + await closed.open() + await closed.close() + closed._tunnel.aopen.reset_mock() + assert not await closed.reopen_if_exited() + closed._tunnel.aopen.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_reopen_if_exited_does_not_reopen_tunnel_closed_meanwhile() -> None: + connection = make_connection(alive=False) + await connection.open() + connection._tunnel.aopen.reset_mock() + + async def close_during_check() -> bool: + connection._closed = True + return False + + connection._tunnel.acheck = AsyncMock(side_effect=close_during_check) + assert not await connection.reopen_if_exited() + connection._tunnel.aopen.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_pool_reopen_exited_continues_after_failure() -> None: + pool = ServiceConnectionPool() + failing = make_connection(alive=False) + dead = make_connection(alive=False) + live = make_connection(alive=True) + for connection in (failing, dead, live): + await connection.open() + connection._tunnel.aopen.reset_mock() + failing._tunnel.aopen.side_effect = RuntimeError("ssh failed") + pool.connections = {"failing": failing, "dead": dead, "live": live} + + await pool.reopen_exited() + + failing._tunnel.aopen.assert_awaited_once() + dead._tunnel.aopen.assert_awaited_once() + live._tunnel.aopen.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_maintain_service_connections_checks_periodically() -> None: + pool = ServiceConnectionPool() + pool.reopen_exited = AsyncMock(side_effect=[RuntimeError("check failed"), None, None]) + task = asyncio.create_task(maintain_service_connections(pool, interval=0.01)) + for _ in range(100): + if pool.reopen_exited.await_count >= 3: + break + await asyncio.sleep(0.01) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert pool.reopen_exited.await_count >= 3