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