mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(passthrough): hand the upstream error report to the logging worker instead of awaiting it in the relay
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
b0d65741a5
commit
e45275a1e6
2 changed files with 163 additions and 54 deletions
|
|
@ -882,44 +882,37 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
|
|||
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
|
||||
self._dispatched = False
|
||||
|
||||
async def _report_once(self) -> None:
|
||||
if self._reported:
|
||||
def _dispatch_report(self) -> None:
|
||||
if self._dispatched:
|
||||
return
|
||||
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())
|
||||
self._dispatched = True
|
||||
self._enqueue(self._report(b"".join(self._collected)))
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[bytes]:
|
||||
total = 0 # rebind-ok: running byte count against the preview budget
|
||||
try:
|
||||
async for chunk in self._upstream.aiter_bytes():
|
||||
if not self._reported:
|
||||
if not self._dispatched:
|
||||
self._collected.append(chunk)
|
||||
total += len(chunk)
|
||||
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
|
||||
await self._report_once()
|
||||
self._dispatch_report()
|
||||
yield chunk
|
||||
await self._report_once()
|
||||
self._dispatch_report()
|
||||
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__,
|
||||
)
|
||||
await self._report_once()
|
||||
self._dispatch_report()
|
||||
finally:
|
||||
self._enqueue_pending_report()
|
||||
self._dispatch_report()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self._enqueue_pending_report()
|
||||
self._dispatch_report()
|
||||
await self._upstream.aclose()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4258,6 +4258,7 @@ async def test_pass_through_request_streaming_upstream_error_body_reaches_client
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
await _poll(lambda: mock_proxy_logging.post_call_failure_hook.called)
|
||||
mock_proxy_logging.post_call_failure_hook.assert_called_once()
|
||||
original_exception: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"]
|
||||
assert "was not found or your project" in original_exception.detail
|
||||
|
|
@ -4453,15 +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"),
|
||||
)
|
||||
|
||||
served_at_warning: list[int] = []
|
||||
real_warning: Final = verbose_proxy_logger.warning
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_warning(*args, **kwargs):
|
||||
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
|
||||
served_at_warning.append(body_stream.served)
|
||||
return real_warning(*args, **kwargs)
|
||||
def _recording_enqueue(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(coro)
|
||||
|
||||
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
|
||||
side_effect=_recording_enqueue,
|
||||
):
|
||||
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"
|
||||
|
|
@ -4495,9 +4498,11 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
assert served_at_warning == [5], (
|
||||
"each raw chunk is yielded as-is; five 1024-byte chunks are the first point the preview budget is exceeded"
|
||||
assert served_at_dispatch == [5], (
|
||||
"the report is dispatched after the fifth 1024-byte chunk, the first point the preview budget is exceeded"
|
||||
)
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
|
||||
assert (
|
||||
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
|
|
@ -4518,15 +4523,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"),
|
||||
)
|
||||
|
||||
served_at_warning: list[int] = []
|
||||
real_warning: Final = verbose_proxy_logger.warning
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
served_at_dispatch: list[int] = []
|
||||
|
||||
def _recording_warning(*args, **kwargs):
|
||||
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
|
||||
served_at_warning.append(body_stream.served)
|
||||
return real_warning(*args, **kwargs)
|
||||
def _recording_enqueue(coro):
|
||||
served_at_dispatch.append(body_stream.served)
|
||||
enqueued.append(coro)
|
||||
|
||||
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
|
||||
with patch(
|
||||
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
|
||||
side_effect=_recording_enqueue,
|
||||
):
|
||||
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"
|
||||
|
|
@ -4560,9 +4567,11 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
assert served_at_warning == [1], (
|
||||
"the rechunked preview is served from the first raw chunk; the second must not be pulled before the warning"
|
||||
assert served_at_dispatch == [1], (
|
||||
"the report is dispatched after the first raw chunk crosses the preview budget; the second is not pulled first"
|
||||
)
|
||||
assert len(enqueued) == 1, enqueued
|
||||
await enqueued[0]
|
||||
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
|
||||
assert (
|
||||
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
|
||||
|
|
@ -4570,6 +4579,13 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
|
|||
)
|
||||
|
||||
|
||||
async def _poll(condition: Callable[[], bool], seconds: float = 5) -> None:
|
||||
deadline: Final = asyncio.get_running_loop().time() + seconds
|
||||
while not condition():
|
||||
assert asyncio.get_running_loop().time() < deadline, "condition not met in time"
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
|
||||
class _GatedUpstreamErrorBodyStream(httpx.AsyncByteStream):
|
||||
def __init__(self, first: bytes, second: bytes) -> None:
|
||||
self._first: Final = first
|
||||
|
|
@ -4599,6 +4615,8 @@ 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]] = []
|
||||
|
||||
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"
|
||||
|
|
@ -4606,29 +4624,34 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
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
|
||||
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)
|
||||
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,
|
||||
)
|
||||
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 == 429
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
assert isinstance(response, StreamingResponse)
|
||||
assert response.status_code == 429
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not body_stream.gate.is_set()
|
||||
mock_proxy_logging.post_call_failure_hook.assert_not_called()
|
||||
body_stream.gate.set()
|
||||
rest: Final = [chunk async for chunk in iterator]
|
||||
relayed: Final = b"".join(
|
||||
|
|
@ -4637,6 +4660,8 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
|
|||
assert relayed == first_chunk + second_chunk
|
||||
await upstream_response.aclose()
|
||||
|
||||
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 (
|
||||
|
|
@ -4712,6 +4737,81 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
|
|||
assert detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_through_request_streaming_upstream_error_yields_over_budget_chunk_before_report_finishes():
|
||||
"""
|
||||
Regression: the failure report is never awaited inside the relay, so a
|
||||
single chunk that crosses the preview budget reaches the client even
|
||||
while the report coroutine is still running.
|
||||
"""
|
||||
first_chunk: Final = b"x" * 6144
|
||||
release: Final = asyncio.Event()
|
||||
body_stream: Final = _ChunkedUpstreamErrorBodyStream((first_chunk,))
|
||||
upstream_response: Final = httpx.Response(
|
||||
status_code=500,
|
||||
headers={"content-type": "text/plain"},
|
||||
stream=body_stream,
|
||||
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
|
||||
)
|
||||
|
||||
enqueued: list[Coroutine[None, None, None]] = []
|
||||
|
||||
async def _held_hook(**kwargs):
|
||||
await release.wait()
|
||||
|
||||
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(side_effect=_held_hook)
|
||||
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
|
||||
iterator: Final = response.body_iterator.__aiter__()
|
||||
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
|
||||
assert not release.is_set()
|
||||
relayed: Final = b"".join(
|
||||
[first]
|
||||
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
|
||||
)
|
||||
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
|
||||
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 == f"Upstream passthrough request failed with status 500: {'x' * 4096}... (truncated at 4096 chars)"
|
||||
), detail
|
||||
|
||||
|
||||
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
|
||||
async def __aiter__(self):
|
||||
yield b'{"error": "half'
|
||||
|
|
@ -4775,6 +4875,11 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k
|
|||
assert streamed_bytes == b'{"error": "half'
|
||||
await upstream_response.aclose()
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
|
||||
)
|
||||
)
|
||||
rendered: Final = [str(args[0]) for args in recorded_warnings]
|
||||
formats: Final = [args[0] for args in recorded_warnings]
|
||||
assert any(
|
||||
|
|
@ -4859,6 +4964,11 @@ async def test_pass_through_request_streaming_upstream_error_gzip_read_failure_r
|
|||
assert streamed_bytes == plaintext
|
||||
await upstream_response.aclose()
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
|
||||
)
|
||||
)
|
||||
rendered: Final = [str(args[0]) for args in recorded_warnings]
|
||||
assert any(
|
||||
args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" and plaintext.decode() in str(args[4])
|
||||
|
|
@ -4910,6 +5020,12 @@ async def test_pass_through_request_streaming_upstream_error_gzip_body_decoded_f
|
|||
)
|
||||
assert streamed_bytes == upstream_content
|
||||
|
||||
await _poll(
|
||||
lambda: any(
|
||||
call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
|
||||
for call in mock_warning.call_args_list
|
||||
)
|
||||
)
|
||||
upstream_warnings: Final = [
|
||||
call
|
||||
for call in mock_warning.call_args_list
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue