diff --git a/pyproject.toml b/pyproject.toml index bba8433701..32b6b394fc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -220,6 +220,8 @@ gateway = [ "aiorwlock", "aiocache", "httpx>=0.28.0", + # Gateway authorization needs the pool cancellation and recovery fixes in httpx2 (#4338). + "httpx2>=2.13.1", "jinja2", ] server = [ diff --git a/src/dstack/_internal/proxy/gateway/auth.py b/src/dstack/_internal/proxy/gateway/auth.py index 2523e9ff59..f71f9845d1 100644 --- a/src/dstack/_internal/proxy/gateway/auth.py +++ b/src/dstack/_internal/proxy/gateway/auth.py @@ -1,15 +1,18 @@ -import httpx +import httpx2 as httpx from aiocache import cached from dstack._internal.proxy.lib.auth import BaseProxyAuthProvider from dstack._internal.proxy.lib.errors import UnexpectedProxyError +_AUTH_CACHE_TTL_SECONDS = 60 +"""Seconds to cache allowed or denied access for each project/token pair.""" + class GatewayProxyAuthProvider(BaseProxyAuthProvider): def __init__(self, server_client: httpx.AsyncClient) -> None: self._server_client = server_client - @cached(ttl=60, noself=True, skip_cache_func=lambda r: r is None) + @cached(ttl=_AUTH_CACHE_TTL_SECONDS, noself=True, skip_cache_func=lambda r: r is None) async def is_project_member(self, project_name: str, token: str) -> bool: try: resp = await self._server_client.post( diff --git a/src/dstack/_internal/proxy/gateway/resources/systemd/dstack.gateway.service b/src/dstack/_internal/proxy/gateway/resources/systemd/dstack.gateway.service index 409e51fe3f..a8027f451a 100644 --- a/src/dstack/_internal/proxy/gateway/resources/systemd/dstack.gateway.service +++ b/src/dstack/_internal/proxy/gateway/resources/systemd/dstack.gateway.service @@ -9,6 +9,7 @@ User=ubuntu Group=ubuntu Restart=always RestartSec=5 +LimitNOFILE=65535 [Install] WantedBy=default.target diff --git a/src/dstack/_internal/proxy/gateway/services/server_client.py b/src/dstack/_internal/proxy/gateway/services/server_client.py index cdcf3aa046..06c3b781ec 100644 --- a/src/dstack/_internal/proxy/gateway/services/server_client.py +++ b/src/dstack/_internal/proxy/gateway/services/server_client.py @@ -5,7 +5,7 @@ from pathlib import Path from typing import Container, Dict, Generator, List -import httpx +import httpx2 as httpx logger = logging.getLogger(__name__) BASE_URL = "http://dstack/" # any hostname will work @@ -13,22 +13,25 @@ @dataclass class CachedClientInfo: + """HTTP client and connection-failure history for one server, owned by HTTPMultiClient.""" + client: httpx.AsyncClient socket: Path connect_errors: List[datetime.datetime] = field(default_factory=lambda: []) def seems_disconnected(self) -> bool: + # Only ConnectError is recorded; a returned HTTP response clears this history. if len(self.connect_errors) < 2: return False return self.connect_errors[-1] - self.connect_errors[0] >= datetime.timedelta(minutes=2) class HTTPMultiClient(httpx.AsyncClient): - """ - An HTTP client that sends requests to randomly chosen Unix sockets from a specified - directory. This allows to balance the load between multiple HTTP server replicas. - Automatically deletes sockets that stop responding. - Used for requesting random dstack-server replicas from the gateway. + """Gateway HTTP client for uncached project-access checks against dstack-server. + + GatewayProxyAuthProvider uses this client and owns authorization decisions and + their cache. Return the server's HTTP response, including error statuses. Raise + httpx.RequestError when no connected server can complete the request. """ def __init__(self, sockets_dir: Path): @@ -40,6 +43,7 @@ async def send(self, request: httpx.Request, *args, **kwargs) -> httpx.Response: errors: List[httpx.RequestError] = [] clients_count = 0 + # Try another replica after request errors; HTTP error responses are returned. for clients_count, client in enumerate(self._iter_clients_rand(), start=1): try: resp = await client.client.send(request, *args, **kwargs) @@ -69,6 +73,7 @@ async def send(self, request: httpx.Request, *args, **kwargs) -> httpx.Response: def _iter_clients_rand(self) -> Generator[CachedClientInfo, None, None]: sockets = list(self._sockets_dir.glob("*.sock")) self._evict_clients(stems_to_keep={s.stem for s in sockets}) + # Each socket forwards to a server replica; shuffle to distribute auth checks. random.shuffle(sockets) for socket in sockets: diff --git a/src/tests/_internal/proxy/gateway/services/test_server_client.py b/src/tests/_internal/proxy/gateway/services/test_server_client.py new file mode 100644 index 0000000000..f387270c75 --- /dev/null +++ b/src/tests/_internal/proxy/gateway/services/test_server_client.py @@ -0,0 +1,134 @@ +import asyncio +from pathlib import Path +from tempfile import TemporaryDirectory + +import httpcore2 as httpcore +import httpx2 as httpx +import pytest +from httpcore2._async.connection_pool import AsyncPoolRequest + +from dstack._internal.proxy.gateway.services.server_client import HTTPMultiClient + +RESPONSE = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok" + + +@pytest.mark.asyncio +class TestHTTPMultiClient: + @pytest.mark.parametrize("failure", ["cancel", "timeout"]) + async def test_recovers_after_cancelled_pool_waiter(self, tmp_path, monkeypatch, failure): + (tmp_path / "server.sock").touch() + client = HTTPMultiClient(tmp_path) + cached = next(client._iter_clients_rand()) + pool = cached.client._transport._pool + assert isinstance(pool, httpcore.AsyncConnectionPool) + pool._max_connections = 1 + # Closing the first response assigns a fresh connection to the queued request. + pool._network_backend = httpcore.AsyncMockBackend( + [RESPONSE.replace(b"Content-Length", b"Connection: close\r\nContent-Length")] + ) + queued = asyncio.Event() + wait_for_connection = AsyncPoolRequest.wait_for_connection + + async def wait(self, timeout=None): + was_queued = self.connection is None + if was_queued: + queued.set() + connection = await wait_for_connection(self, timeout) + if was_queued and failure == "timeout": + # Force the original race without waiting on a real deadline: the + # waiter expires after assignment, before starting its connection. + raise httpcore.PoolTimeout() + return connection + + monkeypatch.setattr(AsyncPoolRequest, "wait_for_connection", wait) + second = None + try: + first = await client.send( + client.build_request("POST", "/api/projects/test/get"), stream=True + ) + second = asyncio.create_task(client.post("/api/projects/test/get")) + await asyncio.wait_for(queued.wait(), timeout=1) + await first.aclose() + if failure == "cancel": + second.cancel() + with pytest.raises( + asyncio.CancelledError if failure == "cancel" else httpx.RequestError + ): + await second + + # The abandoned connection must not permanently consume the only slot. + assert pool.connections == [] + response = await client.post("/api/projects/test/get") + assert response.status_code == 200 + assert response.text == "ok" + finally: + if second is not None: + second.cancel() + await asyncio.gather(second, return_exceptions=True) + await cached.client.aclose() + await client.aclose() + + @pytest.mark.parametrize("outcome", ["connect_error", "timeout", "forbidden"]) + async def test_failover_preserves_http_responses(self, tmp_path, monkeypatch, outcome): + for name in ("a.sock", "b.sock"): + (tmp_path / name).touch() + monkeypatch.setattr("random.shuffle", lambda sockets: sockets.sort()) + client = HTTPMultiClient(tmp_path) + cached = list(client._iter_clients_rand()) + requests = [] + + def first(request): + requests.append("first") + if outcome == "connect_error": + raise httpx.ConnectError("disconnected", request=request) + if outcome == "timeout": + raise httpx.PoolTimeout("busy", request=request) + return httpx.Response(403) + + def second(request): + requests.append("second") + return httpx.Response(200) + + for info, handler in zip(cached, (first, second)): + await info.client.aclose() + info.client = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + try: + response = await client.post("/api/projects/test/get") + assert response.status_code == (403 if outcome == "forbidden" else 200) + assert requests == (["first"] if outcome == "forbidden" else ["first", "second"]) + finally: + for info in cached: + await info.client.aclose() + await client.aclose() + + async def test_sends_authorization_over_unix_socket(self): + # Real local I/O verifies the migrated transport's UDS support and headers. + headers = [] + + async def handle(reader, writer): + try: + headers.append(await reader.readuntil(b"\r\n\r\n")) + writer.write(RESPONSE) + await writer.drain() + finally: + writer.close() + await writer.wait_closed() + + # A short path is needed for macOS's Unix-socket pathname limit. + with TemporaryDirectory() as directory: + socket = Path(directory) / "server.sock" + server = await asyncio.start_unix_server(handle, path=socket) + client = HTTPMultiClient(Path(directory)) + try: + response = await client.post( + "/api/projects/test/get", headers={"Authorization": "Bearer test-token"} + ) + assert response.status_code == 200 + assert b"POST /api/projects/test/get HTTP/1.1\r\n" in headers[0] + assert b"Authorization: Bearer test-token\r\n" in headers[0] + finally: + for info in client._clients_cache.values(): + await info.client.aclose() + await client.aclose() + server.close() + await server.wait_closed() diff --git a/src/tests/_internal/proxy/gateway/test_auth.py b/src/tests/_internal/proxy/gateway/test_auth.py new file mode 100644 index 0000000000..1c5e2a5ec2 --- /dev/null +++ b/src/tests/_internal/proxy/gateway/test_auth.py @@ -0,0 +1,79 @@ +from unittest.mock import AsyncMock + +import httpx2 as httpx +import pytest +import pytest_asyncio + +from dstack._internal.proxy.gateway.auth import GatewayProxyAuthProvider +from dstack._internal.proxy.gateway.services.server_client import HTTPMultiClient +from dstack._internal.proxy.lib.errors import UnexpectedProxyError + + +@pytest_asyncio.fixture(autouse=True) +async def clear_auth_cache(): + cache = GatewayProxyAuthProvider.is_project_member.cache + await cache.clear() + yield + await cache.clear() + + +@pytest.mark.asyncio +class TestGatewayProxyAuthProvider: + @pytest.mark.parametrize("status", [200, 403]) + async def test_caches_project_token_decision_for_sixty_seconds(self, monkeypatch, status): + requests = [] + + def handle(request): + requests.append(request) + return httpx.Response(status) + + cache = GatewayProxyAuthProvider.is_project_member.cache + set_value = AsyncMock(wraps=cache.set) + monkeypatch.setattr(cache, "set", set_value) + async with httpx.AsyncClient( + transport=httpx.MockTransport(handle), base_url="http://dstack/" + ) as client: + provider = GatewayProxyAuthProvider(client) + for _ in range(2): + assert await provider.is_project_member("first", "token") is (status == 200) + assert len(requests) == 1 + assert requests[0].method == "POST" + assert requests[0].url.path == "/api/projects/first/get" + assert requests[0].headers["Authorization"] == "Bearer token" + assert set_value.call_args.kwargs["ttl"] == 60 + + # Different tokens and projects must not inherit another cached decision. + await provider.is_project_member("first", "other-token") + await provider.is_project_member("second", "token") + assert len(requests) == 3 + + @pytest.mark.parametrize("failure", [500, httpx.ReadTimeout]) + async def test_httpx2_failures_are_wrapped_and_not_cached(self, failure): + requests = 0 + + def handle(request): + nonlocal requests + requests += 1 + if requests > 1: + return httpx.Response(200) + if isinstance(failure, int): + return httpx.Response(failure) + raise failure("temporary failure", request=request) + + async with httpx.AsyncClient( + transport=httpx.MockTransport(handle), base_url="http://dstack/" + ) as client: + provider = GatewayProxyAuthProvider(client) + with pytest.raises(UnexpectedProxyError) as exc: + await provider.is_project_member("test", "token") + assert isinstance(exc.value.__cause__, httpx.HTTPError) + assert await provider.is_project_member("test", "token") + assert requests == 2 + + async def test_no_servers_reaches_auth_error_boundary(self, tmp_path): + async with HTTPMultiClient(tmp_path) as client: + provider = GatewayProxyAuthProvider(client) + with pytest.raises(UnexpectedProxyError) as exc: + await provider.is_project_member("test", "token") + assert isinstance(exc.value.__cause__, httpx.RequestError) + assert exc.value.__cause__.request.url.path == "/api/projects/test/get" diff --git a/src/tests/_internal/server/services/gateways/test_gateway_update.py b/src/tests/_internal/server/services/gateways/test_gateway_update.py new file mode 100644 index 0000000000..0d11d3de96 --- /dev/null +++ b/src/tests/_internal/server/services/gateways/test_gateway_update.py @@ -0,0 +1,141 @@ +import shlex +import subprocess +from datetime import datetime, timedelta, timezone +from pathlib import Path +from unittest.mock import AsyncMock, Mock + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from dstack._internal import settings as core_settings +from dstack._internal.server.services import gateways +from dstack._internal.server.testing.common import ( + create_backend, + create_gateway_replica, + create_project, +) + + +@pytest.mark.asyncio +class TestInitGateways: + @pytest.mark.parametrize("target_version", ["0.22.2", "0.22.3"]) + async def test_release_version_controls_existing_gateway_update( + self, + session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, + gateway_update_sandbox: tuple[Path, Path, Mock], + target_version: str, + ): + root, events_file, connection = gateway_update_sandbox + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + previously_updated_at = now - timedelta(minutes=2) + monkeypatch.setattr(core_settings, "DSTACK_GATEWAY_PACKAGE_URL", None) + monkeypatch.setattr(core_settings, "DSTACK_VERSION", target_version) + monkeypatch.setattr(gateways.settings, "SKIP_GATEWAY_UPDATE", False) + monkeypatch.setattr(gateways, "get_current_datetime", lambda: now) + monkeypatch.setattr( + gateways.gateway_connections_pool, + "get_or_add", + AsyncMock(return_value=connection), + ) + monkeypatch.setattr( + gateways.gateway_connections_pool, + "all", + AsyncMock(return_value=[connection]), + ) + configure = AsyncMock() + monkeypatch.setattr(gateways, "configure_gateway_replica", configure) + project = await create_project(session=session) + backend = await create_backend(session=session, project_id=project.id) + replica = await create_gateway_replica(session=session, backend=backend) + replica.app_updated_at = previously_updated_at + await session.commit() + + await gateways.init_gateways(session) + + configure.assert_awaited_once() + events = events_file.read_text().splitlines() + if target_version == "0.22.2": + assert events == ["blue pip show dstack"] + assert (root / "version").read_text().strip() == "blue" + assert replica.app_updated_at == previously_updated_at + return + + assert (root / "version").read_text().strip() == "green" + assert replica.app_updated_at == now + install_package = f"green pip install dstack[gateway]=={target_version}" + install_service = "green python -m dstack._internal.proxy.gateway.systemd install" + reload_service = "systemctl daemon-reload active=green" + restart_service = "systemctl restart dstack.gateway active=green" + assert ( + events.index(install_package) + < events.index(install_service) + < events.index(reload_service) + < events.index(restart_service) + ) + + +@pytest.fixture +def gateway_update_sandbox(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + # Execute the real shell: mocking aexec would miss quoting, version gating, and the order + # of service installation / venv switching / restart. Only the remote commands are stubbed; + # no SSH, package installation, sudo, or real systemd operations are performed. + root = tmp_path / "gateway root" + root.mkdir() + (root / "version").write_text("blue\n") + events = tmp_path / "events" + events.touch() + for color in ("blue", "green"): + venv_bin = root / color / "bin" + venv_bin.mkdir(parents=True) + for command in ("pip", "python"): + _write_executable( + venv_bin / command, + f'printf "%s\\n" "{color} {command} $*" >> "$TEST_EVENTS"\n' + + ( + 'if [ "$1" = show ]; then echo "Version: 0.22.2"; fi\n' + if command == "pip" + else "" + ), + ) + command_bin = tmp_path / "commands" + command_bin.mkdir() + _write_executable( + command_bin / "sudo", + 'if [ "$1" = systemctl ]; then\n' + ' printf "%s\\n" "$* active=$(cat "$TEST_ROOT/version")" >> "$TEST_EVENTS"\n' + "else\n" + ' exec "$@"\n' + "fi\n", + ) + # The only script substitution redirects its hard-coded remote root to this temporary dir. + remote_root = "root=/home/ubuntu/dstack" + assert gateways._GATEWAY_UPDATE_SCRIPT.count(remote_root) == 1 + monkeypatch.setattr( + gateways, + "_GATEWAY_UPDATE_SCRIPT", + gateways._GATEWAY_UPDATE_SCRIPT.replace(remote_root, f"root={shlex.quote(str(root))}"), + ) + + def execute(command: str, timeout: float) -> str: + return subprocess.run( + shlex.split(command), + check=True, + capture_output=True, + text=True, + timeout=min(timeout, 5), + env={ + "PATH": f"{command_bin}:/usr/bin:/bin", + "TEST_EVENTS": str(events), + "TEST_ROOT": str(root), + }, + ).stdout + + connection = Mock(ip_address="1.1.1.1") + connection.tunnel.aexec = AsyncMock(side_effect=execute) + return root, events, connection + + +def _write_executable(path: Path, script: str): + path.write_text("#!/bin/sh\nset -e\n" + script) + path.chmod(0o755)