fix(passthrough): re-raise post-preview upstream aborts and own error report tasks

Once the preview budget is crossed and the report is dispatched, an upstream
httpx.HTTPError is re-raised so the client still sees the truncated framing
instead of a clean terminator. Report tasks now live in a module registry
behind a semaphore and are drained with a timeout during proxy shutdown, and
a client disconnect before upstream completion logs the preview byte count

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-25 17:48:49 +00:00
parent 1f8714c440
commit 3d2bb748f5
5 changed files with 307 additions and 6 deletions

View file

@ -1520,6 +1520,12 @@ 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 = 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")
)
# Headers to control callbacks
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"

View file

@ -5,7 +5,7 @@ import json
import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import count, groupby
@ -42,6 +42,8 @@ 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,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
@ -875,7 +877,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
upstream: httpx.Response,
report: _ReportPreview,
log_warning: Callable[..., None],
spawn: Callable[[Coroutine[None, None, None]], asyncio.Future[None]],
spawn: Callable[[Awaitable[None]], asyncio.Future[None]],
) -> None:
self._upstream: Final = upstream
self._report: Final = report
@ -884,6 +886,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
self._dispatched = False
self._pending: asyncio.Future[None] | None = None
self._completed = False
def _dispatch_report(self) -> None:
if self._dispatched:
@ -896,6 +899,14 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
if pending is not None:
await asyncio.shield(pending)
def _dispatch_disconnect_report(self) -> None:
if not self._completed and not self._dispatched:
self._log_warning(
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body",
sum(len(part) for part in self._collected),
)
self._dispatch_report()
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
try:
@ -906,9 +917,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
self._dispatch_report()
yield chunk
self._completed = True
self._dispatch_report()
await self._drain_pending_report()
except httpx.HTTPError as err:
dispatched_before: Final = self._dispatched
self._log_warning(
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
sum(len(part) for part in self._collected),
@ -916,16 +929,38 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
)
self._dispatch_report()
await self._drain_pending_report()
if dispatched_before:
raise
finally:
self._dispatch_report()
self._dispatch_disconnect_report()
async def aclose(self) -> None:
self._dispatch_report()
self._dispatch_disconnect_report()
await self._upstream.aclose()
def _spawn_report_task(report: Coroutine[None, None, None]) -> asyncio.Future[None]:
return asyncio.ensure_future(report)
_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
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)
return task
async def drain_passthrough_upstream_error_reports(
timeout: float = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS,
) -> None:
pending: Final = tuple(_REPORT_TASKS)
if pending:
await asyncio.wait(pending, timeout=timeout)
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:

View file

@ -730,6 +730,7 @@ from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import (
router as openai_passthrough_router,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
drain_passthrough_upstream_error_reports,
initialize_pass_through_endpoints,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
@ -1584,6 +1585,11 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
await flush_spend_counters_on_shutdown()
try:
await drain_passthrough_upstream_error_reports()
except Exception as e: # noqa: BLE001 # shutdown must continue when a report drain fails
verbose_proxy_logger.error("Error draining passthrough upstream error reports: %s", e)
await _flush_spend_logs_queue_on_shutdown()
await proxy_config.stop_config_sync_subscriber()

View file

@ -1000,6 +1000,50 @@ def test_gemini_passthrough_streaming_429_upstream_abort_after_first_frame_still
assert follow_up.status_code == 200, follow_up.text
def test_upstream_abort_after_preview_budget_reaches_client_as_truncated(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = tuple(b"d" * 1000 for _ in range(5)) + (b"data: tail\n\n",)
def respond(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(status=500, content_type="text/event-stream", chunks=frames, abort_after=5)
return Reply(status=200, body=json.dumps({"ok": True}).encode())
path: Final = tmp_path / "gemini-stream-500-abort-past-preview.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
received: Final = bytearray()
def consume_error_stream() -> None:
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
timeout=httpx.Timeout(15, connect=5),
) as response:
assert response.status_code == 500, response.text
for chunk in response.iter_bytes():
received.extend(chunk)
with pytest.raises(httpx.HTTPError):
consume_error_stream()
assert bytes(received) == b"d" * 5000, bytes(received)[-64:]
eventually(
lambda: _upstream_warnings(owned.log),
lambda lines: (
any("returned 500" in line for line in lines) and any("read failed" in line for line in lines)
),
seconds=30,
)
returned: Final = tuple(line for line in _upstream_warnings(owned.log) if "returned 500" in line)
read_failures: Final = tuple(line for line in _upstream_warnings(owned.log) if "read failed" in line)
assert len(returned) == 1, returned
assert len(read_failures) == 1, read_failures
def test_gemini_passthrough_empty_streaming_429_still_logged(gateway: Gateway, tmp_path: Path) -> None:
def respond(request: Request) -> Reply:
return Reply(status=429, content_type="text/event-stream", chunks=())

View file

@ -27,15 +27,20 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_REPORT_TASKS,
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
HttpPassThroughEndpointHelpers,
InitPassThroughEndpointHelpers,
_bounded_report,
_PreviewReportingStream,
_registered_pass_through_routes,
_spawn_report_task,
_truncate_upstream_error_body,
_with_trace_context,
chat_completion_pass_through_endpoint,
create_pass_through_route,
drain_passthrough_upstream_error_reports,
initialize_pass_through_endpoints,
pass_through_request,
resolve_llm_passthrough_timeout,
@ -4935,6 +4940,211 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k
), rendered
class _UpstreamErrorBodyStreamHeld(httpx.AsyncByteStream):
def __init__(self, chunks: tuple[bytes, ...], hold: asyncio.Event) -> None:
self._chunks: Final = chunks
self._hold: Final = hold
async def __aiter__(self):
for chunk in self._chunks:
yield chunk
await self._hold.wait()
class _UpstreamErrorBodyStreamAbortingAfter(httpx.AsyncByteStream):
def __init__(self, chunks: tuple[bytes, ...]) -> None:
self._chunks: Final = chunks
async def __aiter__(self):
for chunk in self._chunks:
yield chunk
raise httpx.RemoteProtocolError("peer closed connection without sending complete message body")
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_abort_after_preview_budget_reraises_to_client():
chunks: Final = tuple(b"d" * 1000 for _ in range(5))
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStreamAbortingAfter(chunks),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[asyncio.Future[None]] = []
def _recording_spawn(coro):
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
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:
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)
received: list[bytes] = []
async def consume_response() -> None:
async for chunk in response.body_iterator:
received.append(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8"))
with pytest.raises(httpx.RemoteProtocolError):
await consume_response()
assert b"".join(received) == b"d" * 5000
assert len(enqueued) == 1, enqueued
await enqueued[0]
mock_proxy_logging.post_call_failure_hook.assert_called_once()
expected_body: Final = f"{'d' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
== f"Upstream passthrough request failed with status 500: {expected_body}"
)
@pytest.mark.asyncio
async def test_shutdown_drain_finishes_report_after_consumer_cancelled():
upstream_hold: Final = asyncio.Event()
chunks: Final = (b"first",)
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
release: Final = asyncio.Event()
hook_done: Final = asyncio.Event()
log_warning: Final = MagicMock()
async def report(preview: bytes) -> None:
await release.wait()
hook_done.set()
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=log_warning,
spawn=_spawn_report_task,
)
async def consume() -> None:
async for _ in relay.__aiter__():
pass
consumer: Final = asyncio.ensure_future(consume())
for _ in range(20):
await asyncio.sleep(0)
consumer.cancel()
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS)
assert len(pending) == 1, pending
assert not pending[0].done()
assert not hook_done.is_set()
log_warning.assert_called_once_with(
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", 5
)
release.set()
await drain_passthrough_upstream_error_reports(timeout=5)
assert hook_done.is_set()
assert pending[0].done()
assert not _REPORT_TASKS
@pytest.mark.asyncio
async def test_shutdown_drain_returns_with_report_still_pending_on_timeout():
upstream_hold: Final = asyncio.Event()
chunks: Final = (b"first",)
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
release: Final = asyncio.Event()
async def report(preview: bytes) -> None:
await release.wait()
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
spawn=_spawn_report_task,
)
async def consume() -> None:
async for _ in relay.__aiter__():
pass
consumer: Final = asyncio.ensure_future(consume())
for _ in range(20):
await asyncio.sleep(0)
consumer.cancel()
with pytest.raises(asyncio.CancelledError):
await consumer
pending: Final = tuple(_REPORT_TASKS)
assert len(pending) == 1, pending
await drain_passthrough_upstream_error_reports(timeout=0.05)
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] = []
def make_report(index: int):
async def report() -> None:
entered.append(index)
await release.wait()
finished.append(index)
return report
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 asyncio.gather(*tasks)
assert len(entered) == 4, entered
assert len(finished) == 4, finished
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):
def __init__(self, flushed_prefix: bytes) -> None:
self._flushed_prefix: Final = flushed_prefix