fix(passthrough): await own report without a shared limiter, register reports per loop and drain without a timeout by default

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-25 19:24:22 +00:00
parent bf803cde01
commit 391d8f0f93
5 changed files with 191 additions and 47 deletions

View file

@ -1529,11 +1529,10 @@ CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTER
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY: Final = max(
1, int(os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "64"))
)
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final = int(
os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", "10")
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final[float | None] = (
max(0.0, float(_raw_drain))
if (_raw_drain := os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS"))
else None
)
# Headers to control callbacks

View file

@ -4,6 +4,7 @@ import copy
import json
import posixpath
import traceback
import weakref
from base64 import b64encode
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
@ -42,7 +43,6 @@ from litellm._uuid import uuid
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS,
PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY,
PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
REDACTED_BY_LITELLM,
SESSION_ID_OMITTED_METADATA_KEY,
@ -944,28 +944,35 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
await self._upstream.aclose()
_REPORT_CONCURRENCY: Final = asyncio.Semaphore(PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY)
_REPORT_TASKS: Final[set[asyncio.Future[None]]] = set() # mutable-ok: in-flight report registry drained at shutdown
async def _bounded_report(report: Awaitable[None], limiter: asyncio.Semaphore) -> None:
async with limiter:
await report
_REPORT_TASKS: Final[ # mutable-ok: in-flight report registry per loop, drained at shutdown
weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, set[asyncio.Future[None]]]
] = weakref.WeakKeyDictionary()
def _spawn_report_task(report: Awaitable[None]) -> asyncio.Future[None]:
task: Final = asyncio.ensure_future(_bounded_report(report, _REPORT_CONCURRENCY))
_REPORT_TASKS.add(task)
task.add_done_callback(_REPORT_TASKS.discard)
task: Final = asyncio.ensure_future(report)
registry: Final = _REPORT_TASKS.setdefault(
asyncio.get_running_loop(),
set(), # mutable-ok: per-loop task set, tasks discard themselves on completion
)
registry.add(task)
task.add_done_callback(registry.discard)
return task
async def drain_passthrough_upstream_error_reports(
timeout: float = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
timeout: float | None = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
log_warning: Callable[..., None] = verbose_proxy_logger.warning,
) -> None:
pending: Final = tuple(_REPORT_TASKS)
if pending:
await asyncio.wait(pending, timeout=timeout)
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
if not pending:
return
_, still_pending = await asyncio.wait(pending, timeout=timeout)
if still_pending:
log_warning(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending",
len(still_pending),
)
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:

View file

@ -193,6 +193,56 @@ async def test_passthrough_disconnect_burst_logs_every_failure_once(gateway: Gat
assert follow_up.status_code == 200, follow_up.text
_SLOW_FAILURE_HOOK: Final = """
import asyncio
from litellm.integrations.custom_logger import CustomLogger
class SlowFailureHook(CustomLogger):
async def async_post_call_failure_hook(
self, request_data, original_exception, user_api_key_dict, traceback_str=None
):
await asyncio.sleep(15)
instance = SlowFailureHook()
"""
async def test_passthrough_sigterm_drains_reports_parked_on_a_slow_failure_hook(
gateway: Gateway, tmp_path: Path
) -> None:
gate: Final = threading.Event()
def respond(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(
status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES, gate_after_first=gate
)
return Reply(status=200, body=json.dumps({"ok": True}).encode())
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"callbacks": ["slow_hook.instance"]})
(tmp_path / "slow_hook.py").write_text(_SLOW_FAILURE_HOOK)
path: Final = tmp_path / "chaos-sigterm-drain.yaml"
path.write_text(yaml.safe_dump(config))
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=1) as owned:
candidate: Final = owned.gateway
call_ids: Final = await asyncio.gather(
*(_first_frame_then_close(str(candidate.client.base_url), candidate.key) for _ in range(20))
)
assert len(set(call_ids)) == 20, call_ids
gate.set()
owned.process.send_signal(signal.SIGTERM)
owned.process.wait(timeout=90)
for call_id in call_ids:
_single_spend_row(call_id)
async def _first_frame_then_close(base_url: str, key: str) -> str:
async with httpx.AsyncClient(base_url=base_url, timeout=httpx.Timeout(5, connect=5), trust_env=False) as client:
async with client.stream(

View file

@ -5,6 +5,7 @@ import json
import logging
import os
import sys
import time
import zlib
from collections.abc import Callable, Mapping
from contextlib import ExitStack, contextmanager
@ -33,7 +34,6 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
HttpPassThroughEndpointHelpers,
InitPassThroughEndpointHelpers,
_bounded_report,
_PreviewReportingStream,
_registered_pass_through_routes,
_spawn_report_task,
@ -5066,7 +5066,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS)
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
assert len(pending) == 1, pending
assert not pending[0].done()
assert not hook_done.is_set()
@ -5078,7 +5078,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
await drain_passthrough_upstream_error_reports(timeout=5)
assert hook_done.is_set()
assert pending[0].done()
assert not _REPORT_TASKS
assert not _REPORT_TASKS.get(asyncio.get_running_loop())
@pytest.mark.asyncio
@ -5114,38 +5114,115 @@ async def test_shutdown_drain_returns_with_report_still_pending_on_timeout():
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS)
pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ()))
assert len(pending) == 1, pending
await drain_passthrough_upstream_error_reports(timeout=0.05)
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
warnings.assert_called_once_with(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
)
assert not pending[0].done()
pending[0].cancel()
await asyncio.gather(pending[0], return_exceptions=True)
@pytest.mark.asyncio
async def test_report_concurrency_is_bounded_by_the_semaphore():
limiter: Final = asyncio.Semaphore(2)
release: Final = asyncio.Event()
entered: list[int] = []
finished: list[int] = []
async def test_parked_reports_do_not_block_the_error_stream():
"""A pile of still-running reports must not sit on the response path: a stream whose
own report finishes immediately delivers its body without waiting on the others."""
hold: Final = asyncio.Event()
def make_report(index: int):
async def report() -> None:
entered.append(index)
async def parked() -> None:
await hold.wait()
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
assert len(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) == 70
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(b"d" * 5000),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
hook_done: Final = asyncio.Event()
async def report(preview: bytes) -> None:
hook_done.set()
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
spawn=_spawn_report_task,
)
started: Final = time.monotonic()
received: Final = b"".join([chunk async for chunk in relay.__aiter__()])
elapsed: Final = time.monotonic() - started
assert received == b"d" * 5000
assert hook_done.is_set()
assert elapsed < 0.5, elapsed
for task in spawned:
task.cancel()
await asyncio.gather(*spawned, return_exceptions=True)
def test_report_registry_is_scoped_to_each_event_loop():
"""Two consecutive asyncio.run calls: each loop's reports register and drain on their
own loop, so a shared module-level primitive bound to the first loop never breaks the second."""
async def run_once() -> None:
loop: Final = asyncio.get_running_loop()
release: Final = asyncio.Event()
async def parked() -> None:
await release.wait()
finished.append(index)
return report
spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70))
assert len(_REPORT_TASKS.get(loop, ())) == 70
tasks: Final = tuple(asyncio.ensure_future(_bounded_report(make_report(index)(), limiter)) for index in range(4))
for _ in range(20):
await asyncio.sleep(0)
assert len(entered) == 2, entered
release.set()
await drain_passthrough_upstream_error_reports()
assert all(task.done() for task in spawned)
assert not _REPORT_TASKS.get(loop)
release.set()
await asyncio.gather(*tasks)
assert len(entered) == 4, entered
assert len(finished) == 4, finished
asyncio.run(run_once())
asyncio.run(run_once())
assert all(len(pending) == 0 for pending in _REPORT_TASKS.values())
@pytest.mark.asyncio
async def test_shutdown_drain_waits_without_a_timeout():
finished: Final = asyncio.Event()
async def report() -> None:
await asyncio.sleep(0.3)
finished.set()
task: Final = _spawn_report_task(report())
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=None, log_warning=warnings)
assert task.done()
assert finished.is_set()
warnings.assert_not_called()
@pytest.mark.asyncio
async def test_shutdown_drain_timeout_warns_with_the_pending_count():
hold: Final = asyncio.Event()
async def parked() -> None:
await hold.wait()
task: Final = _spawn_report_task(parked())
warnings: Final = MagicMock()
await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings)
warnings.assert_called_once_with(
"pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1
)
task.cancel()
await asyncio.gather(task, return_exceptions=True)
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):

View file

@ -70,12 +70,23 @@ def _build_constant_env_var_map() -> dict[str, str]:
return env_var_map
def test_passthrough_error_report_concurrency_env_zero_clamps_to_one(monkeypatch):
"""A zero/negative concurrency would deadlock every report; the constant clamps to 1."""
monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "0")
@pytest.mark.parametrize(
"env_value, expected",
[
(None, None),
("2.5", 2.5),
("-1", 0.0),
],
)
def test_passthrough_error_report_drain_seconds_env_parsing(monkeypatch, env_value, expected):
"""Unset waits for every report; a value clamps to >= 0 seconds."""
if env_value is None:
monkeypatch.delenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", raising=False)
else:
monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", env_value)
try:
reloaded = importlib.reload(constants)
assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY == 1
assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS == expected
finally:
monkeypatch.undo()
importlib.reload(constants)