fix(passthrough): finish the upstream error report when the client disconnects mid-relay

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-24 21:50:52 +00:00
parent ecc2a94d2c
commit b0d65741a5
3 changed files with 139 additions and 5 deletions

View file

@ -5,7 +5,7 @@ import json
import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import count, groupby
@ -875,12 +875,15 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
upstream: httpx.Response,
report: _ReportPreview,
log_warning: Callable[..., None],
enqueue: Callable[[Coroutine[None, None, None]], None],
) -> None:
self._upstream: Final = upstream
self._report: Final = report
self._log_warning: Final = log_warning
self._enqueue: Final = enqueue
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
self._reported = False
self._enqueued = False
async def _report_once(self) -> None:
if self._reported:
@ -888,6 +891,12 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
self._reported = True
await self._report(b"".join(self._collected))
def _enqueue_pending_report(self) -> None:
if self._reported or self._enqueued:
return
self._enqueued = True
self._enqueue(self._report_once())
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
try:
@ -898,17 +907,19 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
await self._report_once()
yield chunk
await self._report_once()
except httpx.HTTPError as err:
self._log_warning(
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
sum(len(part) for part in self._collected),
type(err).__name__,
)
finally:
await self._report_once()
finally:
self._enqueue_pending_report()
async def aclose(self) -> None:
await self._report_once()
self._enqueue_pending_report()
await self._upstream.aclose()
@ -990,7 +1001,12 @@ async def _log_passthrough_upstream_failure(
return httpx.Response(
status_code=response.status_code,
headers=_headers_without_body_framing(response.headers),
stream=_PreviewReportingStream(upstream=response, report=report, log_warning=log_warning),
stream=_PreviewReportingStream(
upstream=response,
report=report,
log_warning=log_warning,
enqueue=GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue,
),
request=response.request,
extensions=response.extensions,
)

View file

@ -100,6 +100,18 @@ def _spend_error_information(call_id: str) -> dict[str, JsonValue]:
return object_value(parsed["error_information"])
def _spend_error_information_or_none(call_id: str, seconds: float = 20) -> dict[str, JsonValue] | None:
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=seconds,
)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
error_information: Final = parsed.get("error_information")
return None if error_information is None else object_value(error_information)
def _spend_status(call_id: str) -> str:
rows: Final = eventually(
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
@ -671,3 +683,40 @@ def test_gemini_passthrough_quota_wording_in_upstream_body_keeps_passthrough_nor
warning: Final = _upstream_warning(owned.log)
assert "exceeded your current quota" in warning, warning
assert _LEAKED_UPSTREAM_KEY not in warning, warning
def test_gemini_passthrough_streaming_429_client_disconnect_still_logs_failure(
gateway: Gateway, tmp_path: Path
) -> None:
gate: Final = threading.Event()
frames: Final = (b'data: {"error":"rate limited"}\n\n', b"data: [DONE]\n\n")
def respond(request: Request) -> Reply:
return Reply(status=429, content_type="text/event-stream", chunks=frames, gate_after_first=gate)
path: Final = tmp_path / "gemini-stream-429-disconnect.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
try:
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
timeout=httpx.Timeout(3, connect=5),
) as response:
assert response.status_code == 429, response.text
first: Final = next(response.iter_bytes())
assert first.startswith(b'data: {"error":"rate limited"}'), first
call_id: Final = response.headers["x-litellm-call-id"]
finally:
gate.set()
error_information: Final = _spend_error_information_or_none(call_id)
assert error_information is not None, f"no spend row for {call_id} after client disconnect"
assert error_information["error_code"] == "429", error_information
assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information
warnings: Final = _upstream_warnings(owned.log, "returned 429")
assert len(warnings) == 1, warnings

View file

@ -1,11 +1,12 @@
import asyncio
import gc
import gzip
import json
import logging
import os
import sys
import zlib
from collections.abc import Callable
from collections.abc import Callable, Coroutine
from contextlib import ExitStack, contextmanager
from io import BytesIO
from types import SimpleNamespace
@ -4643,6 +4644,74 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
), detail
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_client_disconnect_enqueues_failure_report():
"""
Regression: when the client disconnects mid-relay the response task is
cancelled, so the preview report cannot be awaited inline; it must be
handed to the logging worker, which then fires the failure hook once
with the chunks already relayed.
"""
first_chunk: Final = b'data: {"error":"rate limited"}\n\n'
second_chunk: Final = b"data: [DONE]\n\n"
body_stream: Final = _GatedUpstreamErrorBodyStream(first_chunk, second_chunk)
upstream_response: Final = httpx.Response(
status_code=429,
headers={"content-type": "text/event-stream"},
stream=body_stream,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
iterator = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert first_chunk in (first if isinstance(first, bytes) else first.encode())
await iterator.aclose()
del iterator
gc.collect()
for _ in range(40):
if enqueued:
break
await asyncio.sleep(0.05)
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_called_once()
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
async def __aiter__(self):
yield b'{"error": "half'