mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
fix(passthrough): await the shielded upstream error report at stream end so a delivered error always lands its row
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b6ba92a641
commit
db834354bb
2 changed files with 82 additions and 29 deletions
|
|
@ -875,20 +875,26 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
upstream: httpx.Response,
|
||||
report: _ReportPreview,
|
||||
log_warning: Callable[..., None],
|
||||
enqueue: Callable[[Coroutine[None, None, None]], None],
|
||||
spawn: Callable[[Coroutine[None, None, None]], asyncio.Future[None]],
|
||||
) -> None:
|
||||
self._upstream: Final = upstream
|
||||
self._report: Final = report
|
||||
self._log_warning: Final = log_warning
|
||||
self._enqueue: Final = enqueue
|
||||
self._spawn: Final = spawn
|
||||
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
|
||||
self._dispatched = False
|
||||
self._pending: asyncio.Future[None] | None = None
|
||||
|
||||
def _dispatch_report(self) -> None:
|
||||
if self._dispatched:
|
||||
return
|
||||
self._dispatched = True
|
||||
self._enqueue(self._report(b"".join(self._collected)))
|
||||
self._pending = self._spawn(self._report(b"".join(self._collected)))
|
||||
|
||||
async def _drain_pending_report(self) -> None:
|
||||
pending: Final = self._pending
|
||||
if pending is not None:
|
||||
await asyncio.shield(pending)
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
total = 0 # rebind-ok: running byte count against the preview budget
|
||||
|
|
@ -901,6 +907,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
self._dispatch_report()
|
||||
yield chunk
|
||||
self._dispatch_report()
|
||||
await self._drain_pending_report()
|
||||
except httpx.HTTPError as err:
|
||||
self._log_warning(
|
||||
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
|
||||
|
|
@ -908,6 +915,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
type(err).__name__,
|
||||
)
|
||||
self._dispatch_report()
|
||||
await self._drain_pending_report()
|
||||
finally:
|
||||
self._dispatch_report()
|
||||
|
||||
|
|
@ -916,6 +924,10 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
await self._upstream.aclose()
|
||||
|
||||
|
||||
def _spawn_report_task(report: Coroutine[None, None, None]) -> asyncio.Future[None]:
|
||||
return asyncio.ensure_future(report)
|
||||
|
||||
|
||||
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:
|
||||
return httpx.Headers(
|
||||
[(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")]
|
||||
|
|
@ -998,7 +1010,7 @@ async def _log_passthrough_upstream_failure(
|
|||
upstream=response,
|
||||
report=report,
|
||||
log_warning=log_warning,
|
||||
enqueue=GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue,
|
||||
spawn=_spawn_report_task,
|
||||
),
|
||||
request=response.request,
|
||||
extensions=response.extensions,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import logging
|
|||
import os
|
||||
import sys
|
||||
import zlib
|
||||
from collections.abc import Callable, Coroutine
|
||||
from collections.abc import Callable
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from io import BytesIO
|
||||
from types import SimpleNamespace
|
||||
|
|
@ -4454,16 +4454,17 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_enqueue(coro):
|
||||
def _recording_spawn(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(coro)
|
||||
enqueued.append(asyncio.ensure_future(coro))
|
||||
return enqueued[-1]
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
|
||||
side_effect=_recording_enqueue,
|
||||
"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(
|
||||
|
|
@ -4523,16 +4524,17 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_enqueue(coro):
|
||||
def _recording_spawn(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(coro)
|
||||
enqueued.append(asyncio.ensure_future(coro))
|
||||
return enqueued[-1]
|
||||
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
|
||||
side_effect=_recording_enqueue,
|
||||
"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(
|
||||
|
|
@ -4615,7 +4617,7 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
|
|
@ -4625,8 +4627,8 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
"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),
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
|
|
@ -4687,7 +4689,7 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
|
||||
with patch(
|
||||
|
|
@ -4697,8 +4699,8 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
"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),
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
|
|
@ -4754,7 +4756,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
enqueued: list[asyncio.Future[None]] = []
|
||||
|
||||
async def _held_hook(**kwargs):
|
||||
await release.wait()
|
||||
|
|
@ -4767,8 +4769,8 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
"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),
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
|
||||
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
|
||||
):
|
||||
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
|
||||
|
|
@ -4793,6 +4795,8 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not release.is_set()
|
||||
assert not enqueued[0].done()
|
||||
release.set()
|
||||
relayed: Final = b"".join(
|
||||
[first]
|
||||
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
|
||||
|
|
@ -4800,11 +4804,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
assert relayed == first_chunk
|
||||
|
||||
assert len(enqueued) == 1, enqueued
|
||||
report_task: Final = asyncio.ensure_future(enqueued[0])
|
||||
await asyncio.sleep(0)
|
||||
assert not report_task.done()
|
||||
release.set()
|
||||
await report_task
|
||||
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 (
|
||||
|
|
@ -4812,6 +4812,47 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
|
|||
), detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_drains_the_report_before_stream_end():
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/plain"},
|
||||
stream=_ChunkedUpstreamErrorBodyStream((b"x" * 512,)),
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
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)
|
||||
assert response.status_code == 500
|
||||
drained: Final = [chunk async for chunk in response.body_iterator]
|
||||
assert drained == [b"x" * 512]
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
|
||||
|
||||
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
|
||||
async def __aiter__(self):
|
||||
yield b'{"error": "half'
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue