Skip to content
Closed
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
8 changes: 7 additions & 1 deletion src/dstack/_internal/proxy/gateway/app.py
Original file line number Diff line number Diff line change
@@ -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

Expand Down Expand Up @@ -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__

Expand All @@ -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()


Expand Down
71 changes: 67 additions & 4 deletions src/dstack/_internal/proxy/lib/services/service_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand All @@ -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)
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Empty file.
108 changes: 108 additions & 0 deletions src/tests/_internal/proxy/lib/services/test_service_connection.py
Original file line number Diff line number Diff line change
@@ -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