mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
1f8714c440
commit
3d2bb748f5
5 changed files with 307 additions and 6 deletions
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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=())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue