diff --git a/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py index 46d6e9a40cb..2d8a50a2d33 100644 --- a/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py +++ b/tests/integration/providers/test_cache_control_tool_call_marks_owned_proxy.py @@ -1,20 +1,21 @@ import asyncio +import os import re import signal import socket +import threading from collections import Counter -from collections.abc import Sequence +from collections.abc import Callable, Sequence from dataclasses import dataclass from pathlib import Path from typing import Final import httpx -import psutil import pytest from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.process import owned_proxy_process -from integration._support.wire import Request, wire_server +from integration._support.wire import Reply, Request, wire_server from integration.providers._cache_control_marks_support import ( ASK_LABEL, CITIES, @@ -52,7 +53,6 @@ class _Sent: status: int text: str call_id: str - client_port: int def _free_port() -> int: @@ -87,16 +87,9 @@ async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = F surface: Final = _SURFACES[index % len(_SURFACES)] marker: Final = new_marker() path, body = _surface_request(surface, marker) - async with client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {key}"}) as response: - client_port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) - await response.aread() + response: Final = await client.post(path, json=body, headers={"Authorization": f"Bearer {key}"}) return _Sent( - surface, - marker, - response.status_code, - response.text, - response.headers.get("x-litellm-call-id", ""), - client_port, + surface, marker, response.status_code, response.text, response.headers.get("x-litellm-call-id", "") ) async with httpx.AsyncClient(base_url=owned_url, timeout=30, trust_env=False) as client: @@ -108,6 +101,14 @@ async def _fire(owned_url: str, key: str, *, tolerate_transport_errors: bool = F return tuple(result for result in results if isinstance(result, _Sent)) +def _held_anthropic_peer(release: threading.Event) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert release.wait(timeout=120), "Held upstream was never released" + return anthropic_peer(request) + + return respond + + def _by_marker(received: Sequence[Request]) -> dict[str, tuple[Request, ...]]: counted: Final = Counter(marker_of(request) for request in received) return {marker: tuple(request for request in received if marker_of(request) == marker) for marker in counted} @@ -176,7 +177,8 @@ async def test_capped_burst_rides_out_a_provider_outage_and_logs_every_request_o async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_requests( gateway: Gateway, tmp_path: Path ) -> None: - with wire_server(anthropic_peer) as wire: + release: Final = threading.Event() + with wire_server(_held_anthropic_peer(release)) as wire: config: Final = owned_config( tmp_path, [anthropic_deployment(_MODEL, wire.url, cache_control_injection_points=POINTS)] ) @@ -188,26 +190,21 @@ async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_capped_reques ) owned_url: Final = str(owned.gateway.client.base_url) burst: Final = asyncio.create_task(_fire(owned_url, owned.gateway.key, tolerate_transport_errors=True)) - await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30) - victim: Final = psutil.Process(workers[0]) - victim.suspend() - victim_ports: Final = frozenset( - connection.raddr.port for connection in victim.net_connections(kind="tcp") if connection.raddr - ) - victim.send_signal(signal.SIGKILL) + try: + await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= _BURST, 20) + os.kill(workers[0], signal.SIGKILL) + finally: + release.set() served: Final = await burst during: Final = wire.drain() after: Final = await _fire(owned_url, owned.gateway.key) after_received: Final = wire.drain() + assert 0 < len(served) < _BURST, len(served) assert all(len(requests) == 1 for requests in _by_marker(during).values()) - _assert_capped(tuple(item for item in served if item.status == 200), during) + _assert_capped(served, during) _assert_capped(after, after_received) - survivors: Final = tuple(item for item in served if item.client_port not in victim_ports) - assert survivors, [item.client_port for item in served] - for item in (*survivors, *after): + for item in (*served, *after): _single_spend_row(item) - for item in served: - assert len(_spend_rows(item.call_id)) <= 1, item.call_id @pytest.mark.timeout(240)