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 @@ -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 = [
Expand Down
7 changes: 5 additions & 2 deletions src/dstack/_internal/proxy/gateway/auth.py
Original file line number Diff line number Diff line change
@@ -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(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ User=ubuntu
Group=ubuntu
Restart=always
RestartSec=5
LimitNOFILE=65535

[Install]
WantedBy=default.target
17 changes: 11 additions & 6 deletions src/dstack/_internal/proxy/gateway/services/server_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,30 +5,33 @@
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


@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):
Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down
134 changes: 134 additions & 0 deletions src/tests/_internal/proxy/gateway/services/test_server_client.py
Original file line number Diff line number Diff line change
@@ -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()
79 changes: 79 additions & 0 deletions src/tests/_internal/proxy/gateway/test_auth.py
Original file line number Diff line number Diff line change
@@ -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"
Loading
Loading