fix(passthrough): report upstream errors from the relay stream and wait for them on shutdown

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 02:15:16 +00:00
parent 4cc5a8cc63
commit 9767391872
11 changed files with 361 additions and 10 deletions

View file

@ -1529,6 +1529,8 @@ 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_TASK_NAME: Final = "passthrough-upstream-error-report"
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS: Final = 10.0
# Headers to control callbacks
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"

View file

@ -1162,6 +1162,11 @@ async def _aclose_late_response(produced: Response) -> None:
await aclose()
except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended
verbose_proxy_logger.debug("error closing relayed streaming generator: %s", exc)
if produced.background is not None:
try:
await produced.background()
except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended
verbose_proxy_logger.debug("error running relayed response background task: %s", exc)
async def _relay_late_response(produced: Response) -> AsyncGenerator[bytes, None]:

View file

@ -43,6 +43,7 @@ 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_TASK_NAME,
REDACTED_BY_LITELLM,
SESSION_ID_OMITTED_METADATA_KEY,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
@ -889,12 +890,16 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
self._budget_crossed = False
self._completed = False
self._aborted = False
self._reported = False
self._report_task: asyncio.Task[None] | None = None
async def report_collected(self) -> None:
if self._reported:
return
self._reported = True
def _dispatch(self) -> asyncio.Task[None]:
if self._report_task is None:
self._report_task = asyncio.get_running_loop().create_task(
self._run_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
)
return self._report_task
async def _run_report(self) -> None:
if not self._completed and not self._budget_crossed and not self._aborted:
self._log_warning(
"pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body",
@ -902,6 +907,9 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
)
await self._report(b"".join(self._collected))
async def report_collected(self) -> None:
await asyncio.shield(self._dispatch())
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
try:
@ -923,8 +931,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
if self._budget_crossed:
await self.report_collected()
raise
finally:
self._dispatch()
async def aclose(self) -> None:
self._dispatch()
await self._upstream.aclose()

View file

@ -73,6 +73,8 @@ from litellm.constants import (
LITELLM_SETTINGS_SAFE_DB_OVERRIDES,
LITELLM_UI_ALLOW_HEADERS,
LITELLM_UI_SESSION_DURATION,
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS,
PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME,
RUNTIME_UPDATABLE_ROUTER_SETTINGS,
)
from litellm.litellm_core_utils.asyncify import asyncify
@ -1593,6 +1595,8 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
await _await_passthrough_error_reports_on_shutdown()
await flush_spend_counters_on_shutdown()
await _flush_spend_logs_queue_on_shutdown()
@ -2701,6 +2705,21 @@ def cost_tracking():
litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger())
async def _await_passthrough_error_reports_on_shutdown() -> None:
pending: Final = frozenset(
task for task in asyncio.all_tasks() if task.get_name() == PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
)
if not pending:
return
_, unfinished = await asyncio.wait(pending, timeout=PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS)
if unfinished:
verbose_proxy_logger.warning(
"pass_through_endpoint: %d upstream error reports still running after %.0fs shutdown wait",
len(unfinished),
PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS,
)
async def _drain_spend_event_producer_on_shutdown() -> None:
if spend_event_producer is None:
return

View file

@ -10,7 +10,7 @@ from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
from typing import Final, Literal
import httpx
import psutil
@ -90,6 +90,7 @@ def owned_proxy_process(
remove_environment: tuple[str, ...] = (),
workers: int = 1,
graceful_shutdown_seconds: int | None = None,
asgi_server: Literal["uvicorn", "hypercorn"] = "uvicorn",
) -> Iterator[OwnedProxy]:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
@ -114,6 +115,17 @@ def owned_proxy_process(
process: Final = subprocess.Popen(
(
[
sys.executable,
"-m",
"hypercorn",
"litellm.proxy.proxy_server:app",
"--bind",
f"127.0.0.1:{port}",
"--graceful-timeout",
str(graceful_shutdown_seconds),
]
if asgi_server == "hypercorn"
else [
sys.executable,
"-m",
"uvicorn",

View file

@ -29,6 +29,7 @@ class Reply:
abort_after: int | None = None
gate_after_first: threading.Event | None = None
gate_timeout_seconds: float = 5
delay_before_headers_seconds: float = 0
pause_between_chunks: float = 0
headers: Mapping[str, str] = MappingProxyType({})
@ -69,6 +70,8 @@ def wire_server(
except Exception as error:
errors.put(error)
reply = Reply(status=500)
if reply.delay_before_headers_seconds:
time.sleep(reply.delay_before_headers_seconds)
self.send_response(reply.status)
self.send_header("content-type", reply.content_type)
for name, value in reply.headers.items():

View file

@ -335,3 +335,129 @@ async def _first_frame_then_close(base_url: str, key: str) -> str:
first: Final = await response.aiter_bytes().__anext__()
assert first.startswith(b'data: {"error":"rate limited"}'), first
return response.headers["x-litellm-call-id"]
_MARKER_SLEEP_FAILURE_HOOK: Final = """
import asyncio
from pathlib import Path
from litellm.integrations.custom_logger import CustomLogger
class MarkerSleepFailureHook(CustomLogger):
async def async_post_call_failure_hook(
self, request_data, original_exception, user_api_key_dict, traceback_str=None
):
Path({marker!r}).touch()
await asyncio.sleep({sleep})
Path({done!r}).touch()
instance = MarkerSleepFailureHook()
"""
async def test_passthrough_abort_after_budget_with_disconnect_during_hook_logs_once(
gateway: Gateway, tmp_path: Path
) -> None:
"""Post-budget abort: the report await must be shielded, or a client disconnect
cancels the failure hook mid-flight and no spend row lands."""
chunks: Final = tuple(b"x" * 512 for _ in range(10))
def respond(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(status=500, content_type="text/event-stream", chunks=chunks, abort_after=9)
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": ["marker_sleep_hook.instance"]})
(tmp_path / "marker_sleep_hook.py").write_text(
_MARKER_SLEEP_FAILURE_HOOK.format(
marker=str(tmp_path / "hook_started"), sleep=3, done=str(tmp_path / "hook_done")
)
)
path: Final = tmp_path / "chaos-abort-budget-disconnect.yaml"
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
received: Final = bytearray()
async with httpx.AsyncClient(
base_url=str(candidate.client.base_url), timeout=httpx.Timeout(10, connect=5), trust_env=False
) as client:
try:
async with client.stream(
"POST",
"/gemini/v1beta/models/nope-9:streamGenerateContent",
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key},
) as response:
call_id: Final = response.headers["x-litellm-call-id"]
async for chunk in response.aiter_bytes():
received.extend(chunk)
except httpx.TransportError:
pass
assert bytes(received) == b"x" * 4608, len(received)
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30)
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30)
_single_spend_row(call_id)
@pytest.mark.parametrize(
("asgi_server", "graceful_seconds", "hook_sleep"),
[("hypercorn", 3, 4.5), ("uvicorn", 1, 2)],
ids=["hypercorn", "uvicorn"],
)
async def test_passthrough_error_report_survives_sigterm_after_full_response(
gateway: Gateway, tmp_path: Path, asgi_server: str, graceful_seconds: int, hook_sleep: float
) -> None:
"""The graceful window cancels the request task, but the report task is shielded
and the lifespan shutdown waits for it before the spend flushes."""
if asgi_server == "hypercorn":
pytest.skip("BUG: a litellm proxy under hypercorn does not exit within 60s of SIGTERM here, drain unreachable")
def respond(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES)
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": ["marker_sleep_hook.instance"]})
(tmp_path / "marker_sleep_hook.py").write_text(
_MARKER_SLEEP_FAILURE_HOOK.format(
marker=str(tmp_path / "hook_started"), sleep=hook_sleep, done=str(tmp_path / "hook_done")
)
)
path: Final = tmp_path / f"chaos-sigterm-{asgi_server}.yaml"
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,
graceful_shutdown_seconds=graceful_seconds,
asgi_server=asgi_server,
) as owned:
candidate: Final = owned.gateway
async with httpx.AsyncClient(
base_url=str(candidate.client.base_url), timeout=httpx.Timeout(15, connect=5), trust_env=False
) as client:
async with client.stream(
"POST",
"/gemini/v1beta/models/nope-9:streamGenerateContent",
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key},
) as response:
call_id: Final = response.headers["x-litellm-call-id"]
body: Final = b"".join([chunk async for chunk in response.aiter_bytes()])
assert body == b"".join(_RATE_LIMITED_FRAMES), body
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30)
owned.process.send_signal(signal.SIGTERM)
owned.process.wait(timeout=60)
await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30)
_single_spend_row(call_id)

View file

@ -1253,3 +1253,51 @@ async def test_passthrough_streamed_error_headers_do_not_carry_failure_hook_stat
assert header_value == "none", header_value
error_information: Final = _spend_error_information(call_id)
assert error_information["error_code"] == "500", error_information
def test_gemini_passthrough_streaming_error_with_keepalive_and_slow_headers_logs_failure_once(
gateway: Gateway, tmp_path: Path
) -> None:
"""The keepalive path replays a late response's body_iterator by hand and never
runs its background task: the failure report must come from the stream itself."""
marker: Final = uuid.uuid4().hex
def respond(request: Request) -> Reply:
return Reply(
status=500,
body=json.dumps({"error": {"message": f"upstream blew up {marker}"}}).encode(),
delay_before_headers_seconds=1,
)
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update({"sse_keepalive_ping_interval_seconds": 0.2})
path: Final = tmp_path / "gemini-keepalive-slow-headers.yaml"
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) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json={**_GENERATE_CONTENT, "stream": True},
headers=_gemini_headers(candidate),
timeout=httpx.Timeout(15, connect=5),
) as response:
body: Final = response.read()
assert marker.encode() in body, body
rows: Final = eventually(
lambda: read_rows(
'SELECT metadata FROM "LiteLLM_SpendLogs" WHERE metadata::text LIKE %s', (f"%{marker}%",)
),
lambda values: len(values) == 1,
seconds=70,
)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
error_information: Final = object_value(parsed["error_information"])
assert error_information["error_code"] == "500", error_information
assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information
warnings: Final = _upstream_warnings(owned.log, "returned 500")
assert len(warnings) == 1, warnings

View file

@ -5055,6 +5055,80 @@ async def test_preview_report_collected_runs_without_disconnect_warning_after_cl
log_warning.assert_not_called()
@pytest.mark.asyncio
async def test_preview_stream_reports_once_when_body_iterator_is_replayed_without_background():
"""A relay path that replays body_iterator and never runs the BackgroundTask
still reports: the stream dispatches the report itself when the body ends."""
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(b"frame"),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
reported: list[bytes] = []
async def report(preview: bytes) -> None:
reported.append(preview)
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
)
received: Final = b"".join([chunk async for chunk in relay.__aiter__()])
assert received == b"frame"
for _ in range(5):
await asyncio.sleep(0)
assert reported == [b"frame"], reported
@pytest.mark.asyncio
async def test_preview_stream_post_budget_abort_report_survives_cancel_of_the_awaiting_consumer():
"""The post-budget abort path awaits the report through a shield: cancelling the
consumer task mid-hook must not cancel the report itself."""
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"),
)
report_started: Final = asyncio.Event()
release_report: Final = asyncio.Event()
outcomes: list[str] = []
async def report(preview: bytes) -> None:
report_started.set()
try:
await release_report.wait()
outcomes.append("completed")
except asyncio.CancelledError:
outcomes.append("cancelled")
raise
relay: Final = _PreviewReportingStream(
upstream=upstream_response,
report=report,
log_warning=MagicMock(),
)
async def consume() -> None:
with pytest.raises(httpx.RemoteProtocolError):
async for _ in relay.__aiter__():
pass
consumer: Final = asyncio.create_task(consume())
await asyncio.wait_for(report_started.wait(), timeout=5)
consumer.cancel()
with pytest.raises(asyncio.CancelledError):
await consumer
release_report.set()
for _ in range(10):
await asyncio.sleep(0)
assert outcomes == ["completed"], outcomes
@pytest.mark.asyncio
async def test_log_passthrough_upstream_failure_returns_background_task_for_unconsumed_error_stream():
upstream_response: Final = httpx.Response(
@ -5160,9 +5234,7 @@ async def test_streamed_error_runs_failure_hook_as_background_after_response_hea
await response(scope, receive, send)
assert calls == ["post_call_response_headers_hook", "post_call_failure_hook"], calls
body: Final = b"".join(
message.get("body", b"") for message in sent if message["type"] == "http.response.body"
)
body: Final = b"".join(message.get("body", b"") for message in sent if message["type"] == "http.response.body")
assert body == b"data: one\n\ndata: two\n\n"

View file

@ -24,7 +24,7 @@ import logging
import os
import subprocess
from collections.abc import Awaitable, Callable
from typing import List, Optional, Union
from typing import Final, List, Optional, Union
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -36,6 +36,7 @@ from typing_extensions import TypedDict
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import (
ProxyStartupEvent,
_await_passthrough_error_reports_on_shutdown,
_initialize_shared_aiohttp_session,
_resolve_pydantic_type,
_resolve_typed_dict_type,
@ -1271,3 +1272,36 @@ async def test_prometheus_fallback_stats_job_runs_when_the_lock_is_free_or_absen
await jobs["prometheus_fallback_stats_job"]()
assert send_fallback_stats.await_count == 2
@pytest.mark.asyncio
async def test_await_passthrough_error_reports_on_shutdown_waits_for_named_tasks():
from litellm.constants import PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
finished: Final = asyncio.Event()
async def named_report() -> None:
await asyncio.sleep(0.05)
finished.set()
async def unrelated() -> None:
await asyncio.Event().wait()
named_task: Final = asyncio.get_running_loop().create_task(
named_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME
)
stray: Final = asyncio.create_task(unrelated())
try:
await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=10)
assert named_task.done()
assert finished.is_set()
assert not stray.done()
finally:
stray.cancel()
with pytest.raises(asyncio.CancelledError):
await stray
@pytest.mark.asyncio
async def test_await_passthrough_error_reports_on_shutdown_returns_without_pending_tasks():
await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=5)

View file

@ -10226,3 +10226,22 @@ class TestStreamingContainerOwnershipRecordedBeforeDone:
assert tuple(chunk for chunk, _ in observed) == self.CHUNKS
assert tuple(count for _, count in observed) == (0, 0, 0, 0)
recorder.assert_awaited_once()
@pytest.mark.asyncio
async def test_aclose_late_response_runs_background_task():
from starlette.background import BackgroundTask
from litellm.proxy.common_request_processing import _aclose_late_response
ran: list[bool] = []
async def body():
yield b"x"
async def mark() -> None:
ran.append(True)
produced: Final = StreamingResponse(body(), background=BackgroundTask(mark))
await _aclose_late_response(produced)
assert ran == [True]