mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
test(integration): hold the upstream so the worker kill lands mid-burst
This commit is contained in:
parent
e341ab4649
commit
4ef93e4cde
1 changed files with 24 additions and 27 deletions
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue