From 5895ab32600e3c8bc43c0789f8ef1848d6a64f57 Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Wed, 7 Oct 2026 14:30:48 +0200 Subject: [PATCH 1/2] Fix gateway recovery after high request load --- pyproject.toml | 2 + src/dstack/_internal/proxy/gateway/auth.py | 5 +- .../gateway/resources/licenses/httpcore.txt | 27 ++ .../resources/systemd/dstack.gateway.service | 1 + .../proxy/gateway/services/server_client.py | 156 ++++++- .../gateway/services/test_server_client.py | 440 ++++++++++++++++++ .../services/gateways/test_gateway_update.py | 149 ++++++ 7 files changed, 773 insertions(+), 7 deletions(-) create mode 100644 src/dstack/_internal/proxy/gateway/resources/licenses/httpcore.txt create mode 100644 src/tests/_internal/proxy/gateway/services/test_server_client.py create mode 100644 src/tests/_internal/server/services/gateways/test_gateway_update.py diff --git a/pyproject.toml b/pyproject.toml index bba8433701..6b1c852aa8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -220,6 +220,8 @@ gateway = [ "aiorwlock", "aiocache", "httpx>=0.28.0", + # Gateway pool workarounds use private HTTPcore APIs; review before upgrading. + "httpcore>=1.0.5,<1.1", "jinja2", ] server = [ diff --git a/src/dstack/_internal/proxy/gateway/auth.py b/src/dstack/_internal/proxy/gateway/auth.py index 2523e9ff59..fff1d2f313 100644 --- a/src/dstack/_internal/proxy/gateway/auth.py +++ b/src/dstack/_internal/proxy/gateway/auth.py @@ -4,12 +4,15 @@ 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/licenses/httpcore.txt b/src/dstack/_internal/proxy/gateway/resources/licenses/httpcore.txt new file mode 100644 index 0000000000..311b2b56c5 --- /dev/null +++ b/src/dstack/_internal/proxy/gateway/resources/licenses/httpcore.txt @@ -0,0 +1,27 @@ +Copyright © 2020, [Encode OSS Ltd](https://www.encode.io/). +All rights reserved. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that the following conditions are met: + +* Redistributions of source code must retain the above copyright notice, this + list of conditions and the following disclaimer. + +* Redistributions in binary form must reproduce the above copyright notice, + this list of conditions and the following disclaimer in the documentation + and/or other materials provided with the distribution. + +* Neither the name of the copyright holder nor the names of its + contributors may be used to endorse or promote products derived from + this software without specific prior written permission. + +THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. 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..6300e8bb2b 100644 --- a/src/dstack/_internal/proxy/gateway/services/server_client.py +++ b/src/dstack/_internal/proxy/gateway/services/server_client.py @@ -2,9 +2,11 @@ import logging import random from dataclasses import dataclass, field +from itertools import chain from pathlib import Path from typing import Container, Dict, Generator, List +import httpcore import httpx logger = logging.getLogger(__name__) @@ -13,22 +15,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 +45,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 +75,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: @@ -81,7 +88,7 @@ def _iter_clients_rand(self) -> Generator[CachedClientInfo, None, None]: @staticmethod def _make_client(socket: Path) -> CachedClientInfo: client = httpx.AsyncClient( - transport=httpx.AsyncHTTPTransport(uds=str(socket.absolute())), + transport=_ServerTransport(uds=str(socket.absolute())), base_url=BASE_URL, ) return CachedClientInfo( @@ -93,3 +100,140 @@ def _evict_clients(self, stems_to_keep: Container[str]) -> None: self._clients_cache = { stem: client for stem, client in self._clients_cache.items() if stem in stems_to_keep } + + +class _ServerTransport(httpx.AsyncHTTPTransport): + """HTTPMultiClient's transport for one dstack-server Unix socket. + + Preserve HTTPX's connection limits and use the request's HTTPX timeouts. + _ServerConnectionPool provides recovery after overload. + """ + + def __init__(self, uds: str): + super().__init__(uds=uds) + # Preserve HTTPX's existing default limits. + self._pool = _ServerConnectionPool( + uds=uds, + max_connections=100, + max_keepalive_connections=20, + keepalive_expiry=5.0, + ) + + +class _ServerConnectionPool(httpcore.AsyncConnectionPool): + """Authorization connection pool used by _ServerTransport, with overload recovery. + + Preserve HTTPcore's connection limits, reuse, ordering and timeout behavior. + Failed or cancelled requests must not permanently consume capacity or prevent + later requests from proceeding once connections are available. + See https://github.com/dstackai/dstack/issues/4338. + """ + + async def handle_async_request(self, request: httpcore.Request) -> httpcore.Response: + try: + return await super().handle_async_request(request) + except BaseException: + # HTTPcore removes the failed/cancelled request, but can leave its newly + # assigned connection behind before it starts. That slot is never reused. + # Remove this orphan cleanup once https://github.com/encode/httpcore/pull/1099 ships. + with self._optional_thread_lock: + owners = {id(request.connection) for request in self._requests} + orphaned = [ + connection + for connection in self._connections + if id(connection) not in owners + and not connection.is_closed() + and not connection.is_available() + and not connection.is_idle() + ] + if not orphaned: + raise + for connection in orphaned: + self._connections.remove(connection) + # Wake queued requests now that slots are free. Response streams still + # retain their owners, and HTTPcore shields connection closing below. + closing = orphaned + self._assign_requests_to_connections() + await self._close_connections(closing) + raise + + def _assign_requests_to_connections(self) -> List[httpcore.AsyncConnectionInterface]: + # Adapted from HTTPcore; see resources/licenses/httpcore.txt. + if ( + len(self._requests) > len(self._connections) + and len(self._connections) >= self._max_connections + and all( + not connection.is_closed() + and not connection.has_expired() + and not connection.is_idle() + and not connection.is_available() + for connection in self._connections + ) + ): + # A full, busy pool cannot assign or close anything. Scanning every queued + # request here makes a backlog of timeouts quadratic to drain. + return [] + closing = [] + for connection in list(self._connections): + if connection.is_closed(): + self._connections.remove(connection) + elif connection.has_expired(): + self._connections.remove(connection) + closing.append(connection) + elif connection.is_idle() and len(self._connections) > self._max_keepalive_connections: + # Preserve HTTPcore's existing total-connection keepalive check. + self._connections.remove(connection) + closing.append(connection) + + if not self._requests: + return closing + if len(self._requests) == 1: + request = self._requests[0] + if not request.is_queued(): + return closing + # Reuse a connection without building lists for a single request. + origin = request.request.url.origin + for connection in self._connections: + if connection.can_handle_request(origin) and connection.is_available(): + request.assign_to_connection(connection) + return closing + + # Skip assignment scans when no request is waiting for a connection. + queued_requests = (request for request in self._requests if request.is_queued()) + first_request = next(queued_requests, None) + if first_request is None: + return closing + + # No connection state changes while this synchronous pass runs. Inspect + # availability once, rather than scanning every connection for every waiter. + available = [connection for connection in self._connections if connection.is_available()] + idle = [connection for connection in self._connections if connection.is_idle()] + + def add_connection(origin: httpcore.Origin) -> httpcore.AsyncConnectionInterface: + connection = self.create_connection(origin) + self._connections.append(connection) + if connection.is_available(): + available.append(connection) + if connection.is_idle(): + idle.append(connection) + return connection + + for request in chain((first_request,), queued_requests): + if len(self._connections) >= self._max_connections and not available and not idle: + break + origin = request.request.url.origin + connection = next( + (connection for connection in available if connection.can_handle_request(origin)), + None, + ) + if connection is not None: + request.assign_to_connection(connection) + elif len(self._connections) < self._max_connections: + request.assign_to_connection(add_connection(origin)) + elif idle: + connection = idle.pop(0) + self._connections.remove(connection) + if connection in available: + available.remove(connection) + closing.append(connection) + request.assign_to_connection(add_connection(origin)) + return closing 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..adbcd2fad7 --- /dev/null +++ b/src/tests/_internal/proxy/gateway/services/test_server_client.py @@ -0,0 +1,440 @@ +import asyncio +from pathlib import Path + +import httpcore +import pytest +from httpcore._async.connection_pool import AsyncPoolRequest + +from dstack._internal.proxy.gateway.services.server_client import ( + HTTPMultiClient, + _ServerConnectionPool, + _ServerTransport, +) + +RESPONSE = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok" + + +class TestServerConnectionPool: + @pytest.mark.asyncio + @pytest.mark.parametrize("failure", ["cancel", "timeout"]) + async def test_reclaims_connection_assigned_to_cancelled_waiter(self, monkeypatch, failure): + backend = httpcore.AsyncMockBackend([RESPONSE]) + 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": + # The deadline fires while a connection is assigned, before the + # waiter resumes. Inject the outcome without waiting on a clock. + raise httpcore.PoolTimeout() + return connection + + monkeypatch.setattr(AsyncPoolRequest, "wait_for_connection", wait) + async with _ServerConnectionPool(max_connections=1, network_backend=backend) as pool: + first = await pool.handle_async_request( + httpcore.Request("GET", "http://first/", headers={"Host": "first"}) + ) + second = asyncio.create_task(pool.request("GET", "http://second/")) + await queued.wait() + # Releasing the only slot assigns a new, unstarted connection to second. + await first.aclose() + assert len(pool.connections) == 1 + assert pool.connections[0].info() == "CONNECTING" + if failure == "cancel": + second.cancel() + with pytest.raises( + asyncio.CancelledError if failure == "cancel" else httpcore.PoolTimeout + ): + await second + + assert not pool._requests + assert pool.connections == [] + response = await pool.request("GET", "http://next/") + assert response.status == 200 + assert response.content == b"ok" + + @pytest.mark.asyncio + async def test_assigns_reclaimed_slot_to_queued_waiter(self, monkeypatch): + backend = httpcore.AsyncMockBackend([RESPONSE]) + queued = {host: asyncio.Event() for host in (b"cancelled", b"survivor")} + waiting = {} + wait_for_connection = AsyncPoolRequest.wait_for_connection + + async def wait(self, timeout=None): + host = self.request.url.host + if self.connection is None and host in queued: + waiting[host] = self + queued[host].set() + return await wait_for_connection(self, timeout) + + monkeypatch.setattr(AsyncPoolRequest, "wait_for_connection", wait) + async with _ServerConnectionPool(max_connections=1, network_backend=backend) as pool: + first = await pool.handle_async_request( + httpcore.Request("GET", "http://first/", headers={"Host": "first"}) + ) + cancelled = asyncio.create_task(pool.request("GET", "http://cancelled/")) + survivor = asyncio.create_task(pool.request("GET", "http://survivor/")) + try: + await queued[b"cancelled"].wait() + await queued[b"survivor"].wait() + await first.aclose() + cancelled.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled + # Reclaiming the abandoned slot must wake the existing waiter without + # requiring another incoming request to trigger pool maintenance. + assert waiting[b"survivor"].connection is not None + response = await survivor + assert response.status == 200 + assert response.content == b"ok" + finally: + cancelled.cancel() + survivor.cancel() + await asyncio.gather(cancelled, survivor, return_exceptions=True) + + @pytest.mark.asyncio + async def test_preserves_connections_being_established(self, monkeypatch): + backend = httpcore.AsyncMockBackend([RESPONSE]) + connecting = asyncio.Event() + release_connection = asyncio.Event() + connect_tcp = backend.connect_tcp + + async def connect(host, *args, **kwargs): + if host == "first": + connecting.set() + await release_connection.wait() + else: + raise httpcore.ConnectError("simulated connection failure") + return await connect_tcp(host, *args, **kwargs) + + monkeypatch.setattr(backend, "connect_tcp", connect) + async with _ServerConnectionPool(max_connections=2, network_backend=backend) as pool: + first = asyncio.create_task(pool.request("GET", "http://first/")) + try: + await connecting.wait() + connection = pool.connections[0] + assert connection.info() == "CONNECTING" + # A failed request triggers cleanup while the first still owns its slot. + with pytest.raises(httpcore.ConnectError): + await pool.request("GET", "http://second/") + assert connection in pool.connections + assert not first.done() + release_connection.set() + response = await first + assert response.status == 200 + assert response.content == b"ok" + finally: + first.cancel() + await asyncio.gather(first, return_exceptions=True) + + @pytest.mark.asyncio + async def test_preserves_connections_with_unread_response_streams(self, monkeypatch): + backend = httpcore.AsyncMockBackend([RESPONSE]) + connect_tcp = backend.connect_tcp + + async def connect(host, *args, **kwargs): + if host == "second": + raise httpcore.ConnectError("simulated connection failure") + return await connect_tcp(host, *args, **kwargs) + + monkeypatch.setattr(backend, "connect_tcp", connect) + async with _ServerConnectionPool(max_connections=2, network_backend=backend) as pool: + first = await pool.handle_async_request( + httpcore.Request("GET", "http://first/", headers={"Host": "first"}) + ) + with pytest.raises(httpcore.ConnectError): + await pool.request("GET", "http://second/") + # Failed-request cleanup must not close the first response's connection. + assert await first.aread() == b"ok" + await first.aclose() + + @pytest.mark.asyncio + async def test_reuses_idle_connections(self): + backend = httpcore.AsyncMockBackend([RESPONSE, RESPONSE]) + async with _ServerConnectionPool(network_backend=backend) as pool: + await pool.request("GET", "http://dstack/") + connection = pool.connections[0] + response = await pool.request("GET", "http://dstack/") + assert response.content == b"ok" + assert pool.connections == [connection] + + @pytest.mark.asyncio + @pytest.mark.parametrize("http2", [False, True]) + async def test_single_waiter_reuses_first_match_without_scanning_later_connections( + self, monkeypatch, http2 + ): + async with _ServerConnectionPool( + max_connections=3, + max_keepalive_connections=3, + http2=http2, + network_backend=httpcore.AsyncMockBackend([RESPONSE]), + ) as pool: + scheme = "https" if http2 else "http" + requests = [ + httpcore.Request("GET", f"{scheme}://{host}/", headers={"Host": host}) + for host in ("other", "dstack", "dstack") + ] + if http2: + # HTTPS HTTP/2 connections can be shared while still connecting; + # reusing them must not depend on the connection being idle. + pool._connections = [ + pool.create_connection(request.url.origin) for request in requests + ] + assert all(not connection.is_idle() for connection in pool.connections) + else: + responses = [await pool.handle_async_request(request) for request in requests] + for response in responses: + assert await response.aread() == b"ok" + await response.aclose() + assert all(connection.is_idle() for connection in pool.connections) + different_origin, first_match, later_match = pool.connections + assert not different_origin.can_handle_request(requests[1].url.origin) + + def unexpected_scan(*args): + pytest.fail("A single waiter must stop scanning after its first reusable match") + + monkeypatch.setattr(later_match, "can_handle_request", unexpected_scan) + monkeypatch.setattr(later_match, "is_available", unexpected_scan) + waiter = AsyncPoolRequest(requests[1]) + pool._requests.append(waiter) + + assert pool._assign_requests_to_connections() == [] + assert waiter.connection is first_match + assert pool.connections == [different_origin, first_match, later_match] + + @pytest.mark.asyncio + @pytest.mark.parametrize("active_request", [False, True]) + @pytest.mark.parametrize("state", ["idle", "excess_idle", "expired", "closed"]) + async def test_response_completion_without_waiters_skips_availability_scan( + self, monkeypatch, active_request, state + ): + async with _ServerConnectionPool( + max_connections=2, + max_keepalive_connections=0 if state == "excess_idle" else 2, + keepalive_expiry=-float("inf") if state == "expired" else 5.0, + network_backend=httpcore.AsyncMockBackend([RESPONSE]), + ) as pool: + responses = [ + await pool.handle_async_request( + httpcore.Request("GET", "http://dstack/", headers={"Host": "dstack"}) + ) + for _ in range(1 + active_request) + ] + released, *busy = pool.connections + + def is_available(): + pytest.fail( + "Response completion must not scan availability without queued requests" + ) + + with monkeypatch.context() as patch: + for connection in pool.connections: + patch.setattr(connection, "is_available", is_available) + # Closing an unread HTTP/1.1 response closes its connection. Reading + # it first returns it idle, allowing expiry/keepalive cleanup instead. + if state != "closed": + assert await responses[0].aread() == b"ok" + await responses[0].aclose() + + if state == "idle": + assert pool.connections == [released, *busy] + assert released.is_idle() + else: + assert pool.connections == busy + assert released.is_closed() + assert len(pool._requests) == active_request + if active_request: + assert pool._requests[0].connection is busy[0] + assert not pool._requests[0].is_queued() + assert not busy[0].is_closed() + assert await responses[1].aread() == b"ok" + await responses[1].aclose() + + @pytest.mark.asyncio + async def test_saturated_pool_waits_without_scanning_unavailable_connections( + self, monkeypatch + ): + backend = httpcore.AsyncMockBackend([RESPONSE, RESPONSE]) + queued = asyncio.Event() + wait_for_connection = AsyncPoolRequest.wait_for_connection + + async def wait(self, timeout=None): + if self.connection is None: + queued.set() + return await wait_for_connection(self, timeout) + + monkeypatch.setattr(AsyncPoolRequest, "wait_for_connection", wait) + async with _ServerConnectionPool(max_connections=1, network_backend=backend) as pool: + first = await pool.handle_async_request( + httpcore.Request("GET", "http://dstack/", headers={"Host": "dstack"}) + ) + connection = pool.connections[0] + can_handle_request = connection.can_handle_request + scans = 0 + + def can_handle(origin): + nonlocal scans + scans += 1 + return can_handle_request(origin) + + monkeypatch.setattr(connection, "can_handle_request", can_handle) + second = asyncio.create_task(pool.request("GET", "http://dstack/")) + try: + await queued.wait() + assert scans == 0 + assert not second.done() + await first.aread() + await first.aclose() + response = await second + assert response.status == 200 + assert response.content == b"ok" + finally: + second.cancel() + await asyncio.gather(second, return_exceptions=True) + + @pytest.mark.asyncio + @pytest.mark.parametrize("reuse", [True, False]) + async def test_response_completion_does_not_rescan_busy_connections(self, monkeypatch, reuse): + async with _ServerConnectionPool( + max_connections=3, + max_keepalive_connections=3 if reuse else 1, + network_backend=httpcore.AsyncMockBackend([RESPONSE]), + ) as pool: + responses = [ + await pool.handle_async_request( + httpcore.Request("GET", "http://dstack/", headers={"Host": "dstack"}) + ) + for _ in range(3) + ] + released, *busy = pool.connections + scans = 0 + + def can_handle(origin): + nonlocal scans + scans += 1 + return True + + for connection in busy: + monkeypatch.setattr(connection, "can_handle_request", can_handle) + queued = [ + AsyncPoolRequest(httpcore.Request("GET", "http://dstack/")) for _ in range(100) + ] + pool._requests.extend(queued) + + await responses[0].aread() + await responses[0].aclose() + + # A completed response bypasses the all-busy shortcut. Pool maintenance + # must not rescan busy connections for each of the remaining waiters. + assert scans <= len(busy) + if reuse: + assert all(request.connection is released for request in queued) + else: + assert released.is_closed() + assert queued[0].connection is not None + assert queued[0].connection is not released + assert all(request.connection is None for request in queued[1:]) + assert len(pool.connections) == 3 + + @pytest.mark.asyncio + async def test_evicted_connection_is_not_reused_by_later_waiter(self): + async with _ServerConnectionPool( + max_connections=1, + network_backend=httpcore.AsyncMockBackend([RESPONSE]), + ) as pool: + await pool.request("GET", "http://first/") + original = pool.connections[0] + queued = [ + AsyncPoolRequest(httpcore.Request("GET", f"http://{host}/")) + for host in ("second", "first") + ] + pool._requests.extend(queued) + + closing = pool._assign_requests_to_connections() + + assert closing == [original] + assert queued[0].connection is pool.connections[0] + assert queued[0].connection is not original + assert queued[1].connection is None + await pool._close_connections(closing) + + @pytest.mark.asyncio + @pytest.mark.parametrize("state", ["spare_capacity", "idle", "closed", "expired"]) + async def test_assigns_when_capacity_or_connections_are_available(self, state): + async with _ServerConnectionPool( + max_connections=2 if state == "spare_capacity" else 1, + max_keepalive_connections=1, + network_backend=httpcore.AsyncMockBackend([RESPONSE]), + ) as pool: + response = await pool.handle_async_request( + httpcore.Request("GET", "http://dstack/", headers={"Host": "dstack"}) + ) + connection = pool.connections[0] + assert isinstance(connection, httpcore.AsyncHTTPConnection) + if state != "spare_capacity": + await response.aread() + await response.aclose() + if state == "closed": + await connection.aclose() + elif state == "expired": + assert isinstance(connection._connection, httpcore.AsyncHTTP11Connection) + connection._connection._expire_at = -float("inf") + queued = [ + AsyncPoolRequest(httpcore.Request("GET", "http://dstack/")) for _ in range(2) + ] + pool._requests.extend(queued) + closing = pool._assign_requests_to_connections() + assert queued[0].connection is not None + if state == "idle": + assert queued[0].connection is connection + else: + assert queued[0].connection is not connection + if state == "expired": + assert connection in closing + await pool._close_connections(closing) + + @pytest.mark.asyncio + @pytest.mark.parametrize("existing_connection", [True, False]) + async def test_assigns_available_http2_connection_to_multiple_requests( + self, existing_connection + ): + async with _ServerConnectionPool(max_connections=1, http2=True) as pool: + request = httpcore.Request("GET", "https://dstack/") + if existing_connection: + connection = pool.create_connection(request.url.origin) + pool._connections.append(connection) + assert not connection.is_idle() + assert connection.is_available() + queued = [AsyncPoolRequest(request) for _ in range(2)] + pool._requests.extend(queued) + assert pool._assign_requests_to_connections() == [] + assert len(pool.connections) == 1 + assert all(waiter.connection is pool.connections[0] for waiter in queued) + + +class TestHTTPMultiClient: + @pytest.mark.asyncio + async def test_server_transport_uses_fixed_pool(self, tmp_path: Path): + socket = tmp_path / "server.sock" + socket.touch() + client = HTTPMultiClient(tmp_path) + cached = next(client._iter_clients_rand()) + transport = cached.client._transport + assert isinstance(transport, _ServerTransport) + pool = transport._pool + assert isinstance(pool, _ServerConnectionPool) + assert pool._uds == str(socket) + assert (pool._max_connections, pool._max_keepalive_connections) == (100, 20) + assert pool._keepalive_expiry == 5.0 + pool._network_backend = httpcore.AsyncMockBackend([RESPONSE]) + try: + response = await client.post("/api/projects/test/get") + assert response.status_code == 200 + assert response.text == "ok" + finally: + await cached.client.aclose() + await client.aclose() 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..2a15bf3615 --- /dev/null +++ b/src/tests/_internal/server/services/gateways/test_gateway_update.py @@ -0,0 +1,149 @@ +import shlex +import subprocess +from dataclasses import dataclass +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: "_GatewayUpdateSandbox", + target_version: str, + ): + sandbox = 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=sandbox.connection), + ) + monkeypatch.setattr( + gateways.gateway_connections_pool, + "all", + AsyncMock(return_value=[sandbox.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_with(sandbox.connection, attempts=7) + events = sandbox.events.read_text().splitlines() + if target_version == "0.22.2": + assert events == ["blue pip show dstack"] + assert (sandbox.root / "version").read_text().strip() == "blue" + assert replica.app_updated_at == previously_updated_at + return + + assert (sandbox.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) + ) + + +@dataclass +class _GatewayUpdateSandbox: + root: Path + events: Path + connection: Mock + + +@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) -> str: + return subprocess.run( + shlex.split(command), + check=True, + capture_output=True, + text=True, + 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 _GatewayUpdateSandbox(root, events, connection) + + +def _write_executable(path: Path, script: str): + path.write_text("#!/bin/sh\nset -e\n" + script) + path.chmod(0o755) From 02cd4392eb97f3da1f077801266f67c10da4df1b Mon Sep 17 00:00:00 2001 From: Andrey Cheptsov Date: Thu, 8 Oct 2026 12:14:27 +0200 Subject: [PATCH 2/2] Trim gateway authorization regression tests --- .../gateway/services/test_server_client.py | 196 +++--------------- .../_internal/proxy/gateway/test_auth.py | 47 +---- .../services/gateways/test_gateway_update.py | 26 +-- 3 files changed, 48 insertions(+), 221 deletions(-) diff --git a/src/tests/_internal/proxy/gateway/services/test_server_client.py b/src/tests/_internal/proxy/gateway/services/test_server_client.py index e6f0309d32..f387270c75 100644 --- a/src/tests/_internal/proxy/gateway/services/test_server_client.py +++ b/src/tests/_internal/proxy/gateway/services/test_server_client.py @@ -1,5 +1,4 @@ import asyncio -from contextlib import asynccontextmanager from pathlib import Path from tempfile import TemporaryDirectory @@ -13,13 +12,20 @@ RESPONSE = b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok" -class TestHTTPMultiClientPool: - @pytest.mark.asyncio +@pytest.mark.asyncio +class TestHTTPMultiClient: @pytest.mark.parametrize("failure", ["cancel", "timeout"]) - async def test_reclaims_connection_assigned_to_cancelled_waiter( - self, tmp_path, monkeypatch, failure - ): - backend = httpcore.AsyncMockBackend([RESPONSE]) + 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 @@ -29,172 +35,41 @@ async def wait(self, timeout=None): queued.set() connection = await wait_for_connection(self, timeout) if was_queued and failure == "timeout": - # The deadline fires while a connection is assigned, before the - # waiter resumes. Inject the outcome without waiting on a clock. + # 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) - async with _auth_pool(tmp_path, max_connections=1, network_backend=backend) as pool: - first = await pool.handle_async_request( - httpcore.Request("GET", "http://first/", headers={"Host": "first"}) + second = None + try: + first = await client.send( + client.build_request("POST", "/api/projects/test/get"), stream=True ) - second = asyncio.create_task(pool.request("GET", "http://second/")) - await queued.wait() - # Releasing the only slot assigns a new, unstarted connection to second. + second = asyncio.create_task(client.post("/api/projects/test/get")) + await asyncio.wait_for(queued.wait(), timeout=1) await first.aclose() - assert len(pool.connections) == 1 - assert not pool.connections[0].is_connected() if failure == "cancel": second.cancel() with pytest.raises( - asyncio.CancelledError if failure == "cancel" else httpcore.PoolTimeout + asyncio.CancelledError if failure == "cancel" else httpx.RequestError ): await second - assert not pool._requests + # The abandoned connection must not permanently consume the only slot. assert pool.connections == [] - response = await pool.request("GET", "http://next/") - assert response.status == 200 - assert response.content == b"ok" - - @pytest.mark.asyncio - async def test_assigns_reclaimed_slot_to_queued_waiter(self, tmp_path, monkeypatch): - backend = httpcore.AsyncMockBackend([RESPONSE]) - queued = {host: asyncio.Event() for host in (b"cancelled", b"survivor")} - waiting = {} - wait_for_connection = AsyncPoolRequest.wait_for_connection - - async def wait(self, timeout=None): - host = self.request.url.host - if self.connection is None and host in queued: - waiting[host] = self - queued[host].set() - return await wait_for_connection(self, timeout) - - monkeypatch.setattr(AsyncPoolRequest, "wait_for_connection", wait) - async with _auth_pool(tmp_path, max_connections=1, network_backend=backend) as pool: - first = await pool.handle_async_request( - httpcore.Request("GET", "http://first/", headers={"Host": "first"}) - ) - cancelled = asyncio.create_task(pool.request("GET", "http://cancelled/")) - survivor = asyncio.create_task(pool.request("GET", "http://survivor/")) - try: - await queued[b"cancelled"].wait() - await queued[b"survivor"].wait() - await first.aclose() - cancelled.cancel() - with pytest.raises(asyncio.CancelledError): - await cancelled - # Reclaiming the abandoned slot must wake the existing waiter without - # requiring another incoming request to trigger pool maintenance. - assert waiting[b"survivor"].connection is not None - response = await survivor - assert response.status == 200 - assert response.content == b"ok" - finally: - cancelled.cancel() - survivor.cancel() - await asyncio.gather(cancelled, survivor, return_exceptions=True) - - @pytest.mark.asyncio - async def test_preserves_connections_being_established(self, tmp_path, monkeypatch): - backend = httpcore.AsyncMockBackend([RESPONSE]) - connecting = asyncio.Event() - release_connection = asyncio.Event() - connect_unix_socket = backend.connect_unix_socket - - async def connect(*args, **kwargs): - if not connecting.is_set(): - connecting.set() - await release_connection.wait() - else: - raise httpcore.ConnectError("simulated connection failure") - return await connect_unix_socket(*args, **kwargs) - - monkeypatch.setattr(backend, "connect_unix_socket", connect) - async with _auth_pool(tmp_path, max_connections=2, network_backend=backend) as pool: - first = asyncio.create_task(pool.request("GET", "http://first/")) - try: - await connecting.wait() - connection = pool.connections[0] - assert not connection.is_connected() - # A failed request triggers cleanup while the first still owns its slot. - with pytest.raises(httpcore.ConnectError): - await pool.request("GET", "http://second/") - assert connection in pool.connections - assert not first.done() - release_connection.set() - response = await first - assert response.status == 200 - assert response.content == b"ok" - finally: - first.cancel() - await asyncio.gather(first, return_exceptions=True) - - @pytest.mark.asyncio - async def test_preserves_connections_with_unread_response_streams(self, tmp_path, monkeypatch): - backend = httpcore.AsyncMockBackend([RESPONSE]) - connect_unix_socket = backend.connect_unix_socket - - connected = False - - async def connect(*args, **kwargs): - nonlocal connected - if connected: - raise httpcore.ConnectError("simulated connection failure") - connected = True - return await connect_unix_socket(*args, **kwargs) - - monkeypatch.setattr(backend, "connect_unix_socket", connect) - async with _auth_pool(tmp_path, max_connections=2, network_backend=backend) as pool: - first = await pool.handle_async_request( - httpcore.Request("GET", "http://first/", headers={"Host": "first"}) - ) - with pytest.raises(httpcore.ConnectError): - await pool.request("GET", "http://second/") - # Failed-request cleanup must not close the first response's connection. - assert await first.aread() == b"ok" - await first.aclose() - - @pytest.mark.asyncio - async def test_reuses_idle_connections(self, tmp_path): - backend = httpcore.AsyncMockBackend([RESPONSE, RESPONSE]) - async with _auth_pool(tmp_path, network_backend=backend) as pool: - await pool.request("GET", "http://dstack/") - connection = pool.connections[0] - response = await pool.request("GET", "http://dstack/") - assert response.content == b"ok" - assert pool.connections == [connection] - - -class TestHTTPMultiClient: - @pytest.mark.asyncio - async def test_default_limits_timeouts_and_request(self, tmp_path): - socket = tmp_path / "server.sock" - socket.touch() - client = HTTPMultiClient(tmp_path) - cached = next(client._iter_clients_rand()) - pool = cached.client._transport._pool - assert isinstance(pool, httpcore.AsyncConnectionPool) - assert pool._uds == str(socket) - assert (pool._max_connections, pool._max_keepalive_connections) == (100, 20) - assert pool._keepalive_expiry == 5.0 - assert cached.client.timeout == httpx.Timeout(5.0) - pool._network_backend = httpcore.AsyncMockBackend([RESPONSE]) - try: 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.asyncio @pytest.mark.parametrize("outcome", ["connect_error", "timeout", "forbidden"]) - async def test_failover_uses_httpx2_errors_and_preserves_http_responses( - self, tmp_path, monkeypatch, outcome - ): + 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()) @@ -226,10 +101,8 @@ def second(request): await info.client.aclose() await client.aclose() - @pytest.mark.asyncio async def test_sends_authorization_over_unix_socket(self): - # Real local I/O verifies the migrated transport's UDS support and headers; - # the deterministic pool tests above use mock I/O to force cancellation races. + # Real local I/O verifies the migrated transport's UDS support and headers. headers = [] async def handle(reader, writer): @@ -259,16 +132,3 @@ async def handle(reader, writer): await client.aclose() server.close() await server.wait_closed() - - -@asynccontextmanager -async def _auth_pool(tmp_path, **settings): - # Exercise the pool created by the actual gateway transport. Lower capacity - # and mocked I/O make the production cancellation race deterministic and fast. - info = HTTPMultiClient._make_client(tmp_path / "server.sock") - async with info.client: - pool = info.client._transport._pool - assert isinstance(pool, httpcore.AsyncConnectionPool) - for key, value in settings.items(): - setattr(pool, f"_{key}", value) - yield pool diff --git a/src/tests/_internal/proxy/gateway/test_auth.py b/src/tests/_internal/proxy/gateway/test_auth.py index f0504d08b4..1c5e2a5ec2 100644 --- a/src/tests/_internal/proxy/gateway/test_auth.py +++ b/src/tests/_internal/proxy/gateway/test_auth.py @@ -17,8 +17,8 @@ async def clear_auth_cache(): await cache.clear() +@pytest.mark.asyncio class TestGatewayProxyAuthProvider: - @pytest.mark.asyncio @pytest.mark.parametrize("status", [200, 403]) async def test_caches_project_token_decision_for_sixty_seconds(self, monkeypatch, status): requests = [] @@ -28,8 +28,8 @@ def handle(request): return httpx.Response(status) cache = GatewayProxyAuthProvider.is_project_member.cache - set_value = AsyncMock(wraps=cache._set) - monkeypatch.setattr(cache, "_set", set_value) + 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: @@ -47,20 +47,14 @@ def handle(request): await provider.is_project_member("second", "token") assert len(requests) == 3 - # Eviction must cause a fresh check, without waiting for the real clock. - await cache.clear() - await provider.is_project_member("first", "token") - assert len(requests) == 4 - - @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [500, 429, httpx.ConnectError, httpx.ReadTimeout]) + @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 > 2: + if requests > 1: return httpx.Response(200) if isinstance(failure, int): return httpx.Response(failure) @@ -70,35 +64,16 @@ def handle(request): transport=httpx.MockTransport(handle), base_url="http://dstack/" ) as client: provider = GatewayProxyAuthProvider(client) - for _ in range(2): - 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") + 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 == 3 - - @pytest.mark.asyncio - @pytest.mark.parametrize("failure", [None, httpx.ConnectError, httpx.PoolTimeout]) - async def test_all_server_failures_reach_auth_error_boundary(self, tmp_path, failure): - client = HTTPMultiClient(tmp_path) - if failure is not None: - (tmp_path / "server.sock").touch() - cached = list(client._iter_clients_rand()) - - def fail(request): - raise failure("server unavailable", request=request) + assert requests == 2 - for info in cached: - await info.client.aclose() - info.client = httpx.AsyncClient(transport=httpx.MockTransport(fail)) - try: + 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" - finally: - for info in cached: - await info.client.aclose() - await client.aclose() diff --git a/src/tests/_internal/server/services/gateways/test_gateway_update.py b/src/tests/_internal/server/services/gateways/test_gateway_update.py index e7e29215ee..0d11d3de96 100644 --- a/src/tests/_internal/server/services/gateways/test_gateway_update.py +++ b/src/tests/_internal/server/services/gateways/test_gateway_update.py @@ -1,6 +1,5 @@ import shlex import subprocess -from dataclasses import dataclass from datetime import datetime, timedelta, timezone from pathlib import Path from unittest.mock import AsyncMock, Mock @@ -24,10 +23,10 @@ async def test_release_version_controls_existing_gateway_update( self, session: AsyncSession, monkeypatch: pytest.MonkeyPatch, - gateway_update_sandbox: "_GatewayUpdateSandbox", + gateway_update_sandbox: tuple[Path, Path, Mock], target_version: str, ): - sandbox = gateway_update_sandbox + 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) @@ -37,12 +36,12 @@ async def test_release_version_controls_existing_gateway_update( monkeypatch.setattr( gateways.gateway_connections_pool, "get_or_add", - AsyncMock(return_value=sandbox.connection), + AsyncMock(return_value=connection), ) monkeypatch.setattr( gateways.gateway_connections_pool, "all", - AsyncMock(return_value=[sandbox.connection]), + AsyncMock(return_value=[connection]), ) configure = AsyncMock() monkeypatch.setattr(gateways, "configure_gateway_replica", configure) @@ -54,15 +53,15 @@ async def test_release_version_controls_existing_gateway_update( await gateways.init_gateways(session) - configure.assert_awaited_once_with(sandbox.connection, attempts=7) - events = sandbox.events.read_text().splitlines() + configure.assert_awaited_once() + events = events_file.read_text().splitlines() if target_version == "0.22.2": assert events == ["blue pip show dstack"] - assert (sandbox.root / "version").read_text().strip() == "blue" + assert (root / "version").read_text().strip() == "blue" assert replica.app_updated_at == previously_updated_at return - assert (sandbox.root / "version").read_text().strip() == "green" + 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" @@ -76,13 +75,6 @@ async def test_release_version_controls_existing_gateway_update( ) -@dataclass -class _GatewayUpdateSandbox: - root: Path - events: Path - connection: Mock - - @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 @@ -141,7 +133,7 @@ def execute(command: str, timeout: float) -> str: connection = Mock(ip_address="1.1.1.1") connection.tunnel.aexec = AsyncMock(side_effect=execute) - return _GatewayUpdateSandbox(root, events, connection) + return root, events, connection def _write_executable(path: Path, script: str):