test(integration): hold the upstream so the worker kill lands mid-burst

This commit is contained in:
mateo-berri 2026-09-28 22:31:13 -07:00
parent e341ab4649
commit 4ef93e4cde

View file

@ -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)