mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
ecc2a94d2c
commit
b0d65741a5
3 changed files with 139 additions and 5 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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'
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue