mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #38606 from BerriAI/litellm_bedrock_messages_midstream_fallback
fix(router): fall over on raised mid-stream errors in /v1/messages streams
This commit is contained in:
commit
98c52339d4
2 changed files with 281 additions and 8 deletions
|
|
@ -421,6 +421,36 @@ def _anthropic_stream_should_decline_fallback(has_generated_content: bool, error
|
||||||
return has_generated_content or not error.is_pre_first_chunk
|
return has_generated_content or not error.is_pre_first_chunk
|
||||||
|
|
||||||
|
|
||||||
|
def _anthropic_stream_raised_error_status(error: Exception) -> int | None:
|
||||||
|
raw_status: Final = getattr(error, "status_code", None)
|
||||||
|
if isinstance(raw_status, int):
|
||||||
|
return raw_status
|
||||||
|
if isinstance(raw_status, str) and raw_status.isdigit():
|
||||||
|
return int(raw_status)
|
||||||
|
response_status: Final = getattr(getattr(error, "response", None), "status_code", None)
|
||||||
|
return response_status if isinstance(response_status, int) else None
|
||||||
|
|
||||||
|
|
||||||
|
def _anthropic_stream_fallback_error_for_raised(
|
||||||
|
error: Exception, model: str, has_generated_content: bool
|
||||||
|
) -> "MidStreamFallbackError | None":
|
||||||
|
"""Same gate as a detected SSE error event; None means the raise propagates unchanged."""
|
||||||
|
from litellm.exceptions import MidStreamFallbackError
|
||||||
|
|
||||||
|
if has_generated_content:
|
||||||
|
return None
|
||||||
|
status_code: Final = _anthropic_stream_raised_error_status(error)
|
||||||
|
if status_code is not None and not _is_retriable_anthropic_status(status_code):
|
||||||
|
return None
|
||||||
|
return MidStreamFallbackError(
|
||||||
|
message=str(error),
|
||||||
|
model=model,
|
||||||
|
llm_provider="anthropic",
|
||||||
|
original_exception=error,
|
||||||
|
is_pre_first_chunk=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, buffered_chunk_count: int) -> bool:
|
def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, buffered_chunk_count: int) -> bool:
|
||||||
"""
|
"""
|
||||||
Whether `chunk` should make Router._aanthropic_messages_streaming_iterator
|
Whether `chunk` should make Router._aanthropic_messages_streaming_iterator
|
||||||
|
|
@ -5065,14 +5095,15 @@ class Router:
|
||||||
yield chunk
|
yield chunk
|
||||||
for buffered_chunk in buffered_lifecycle_chunks:
|
for buffered_chunk in buffered_lifecycle_chunks:
|
||||||
yield buffered_chunk
|
yield buffered_chunk
|
||||||
except MidStreamFallbackError as e:
|
except Exception as stream_error: # noqa: BLE001 # any raised provider error must reach the fallback gate
|
||||||
if _anthropic_stream_should_decline_fallback(has_generated_content, e):
|
async for item in self._aanthropic_messages_recover_stream_error(
|
||||||
for buffered_chunk in buffered_lifecycle_chunks:
|
stream_error,
|
||||||
yield buffered_chunk
|
has_generated_content,
|
||||||
if e.original_exception is not None:
|
buffered_lifecycle_chunks,
|
||||||
raise e.original_exception from e
|
model,
|
||||||
raise
|
initial_kwargs,
|
||||||
async for item in self._aanthropic_messages_fallback_attempt(e, initial_kwargs, wrapper):
|
wrapper,
|
||||||
|
):
|
||||||
yield item
|
yield item
|
||||||
finally:
|
finally:
|
||||||
with anyio.CancelScope(shield=True), contextlib.suppress(BaseException):
|
with anyio.CancelScope(shield=True), contextlib.suppress(BaseException):
|
||||||
|
|
@ -5084,6 +5115,36 @@ class Router:
|
||||||
wrapper: Final = FallbackAwareAnthropicMessagesStream(stream_with_fallbacks(), source_iterator)
|
wrapper: Final = FallbackAwareAnthropicMessagesStream(stream_with_fallbacks(), source_iterator)
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
async def _aanthropic_messages_recover_stream_error(
|
||||||
|
self,
|
||||||
|
stream_error: Exception,
|
||||||
|
has_generated_content: bool,
|
||||||
|
buffered_lifecycle_chunks: tuple[bytes, ...],
|
||||||
|
model: str,
|
||||||
|
initial_kwargs: dict[str, Any], # mutable-ok: handed to _aanthropic_messages_fallback_attempt, which mutates it
|
||||||
|
wrapper: "FallbackAwareAnthropicMessagesStream",
|
||||||
|
) -> AsyncGenerator[bytes, None]:
|
||||||
|
"""Turns a source-iterator failure into a fallback attempt or the error reaching the caller."""
|
||||||
|
from litellm.exceptions import MidStreamFallbackError
|
||||||
|
|
||||||
|
if isinstance(stream_error, MidStreamFallbackError) and _anthropic_stream_should_decline_fallback(
|
||||||
|
has_generated_content, stream_error
|
||||||
|
):
|
||||||
|
for buffered_chunk in buffered_lifecycle_chunks:
|
||||||
|
yield buffered_chunk
|
||||||
|
if stream_error.original_exception is not None:
|
||||||
|
raise stream_error.original_exception from stream_error
|
||||||
|
raise stream_error
|
||||||
|
fallback_error: Final = (
|
||||||
|
stream_error
|
||||||
|
if isinstance(stream_error, MidStreamFallbackError)
|
||||||
|
else _anthropic_stream_fallback_error_for_raised(stream_error, model, has_generated_content)
|
||||||
|
)
|
||||||
|
if fallback_error is None:
|
||||||
|
raise stream_error
|
||||||
|
async for item in self._aanthropic_messages_fallback_attempt(fallback_error, initial_kwargs, wrapper):
|
||||||
|
yield item
|
||||||
|
|
||||||
async def _aanthropic_messages_fallback_attempt(
|
async def _aanthropic_messages_fallback_attempt(
|
||||||
self,
|
self,
|
||||||
e: "MidStreamFallbackError",
|
e: "MidStreamFallbackError",
|
||||||
|
|
|
||||||
|
|
@ -5,6 +5,7 @@ import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import threading
|
import threading
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
@ -17,6 +18,7 @@ import litellm
|
||||||
from litellm import Router
|
from litellm import Router
|
||||||
from litellm.exceptions import MidStreamFallbackError
|
from litellm.exceptions import MidStreamFallbackError
|
||||||
from litellm.integrations.custom_logger import CustomLogger
|
from litellm.integrations.custom_logger import CustomLogger
|
||||||
|
from litellm.llms.bedrock.common_utils import BedrockError
|
||||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||||
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
|
SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES,
|
||||||
)
|
)
|
||||||
|
|
@ -24,6 +26,8 @@ from litellm.router import (
|
||||||
MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS,
|
MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS,
|
||||||
FallbackAwareAnthropicMessagesStream,
|
FallbackAwareAnthropicMessagesStream,
|
||||||
_anthropic_stream_commits_now,
|
_anthropic_stream_commits_now,
|
||||||
|
_anthropic_stream_fallback_error_for_raised,
|
||||||
|
_anthropic_stream_raised_error_status,
|
||||||
_anthropic_stream_should_decline_fallback,
|
_anthropic_stream_should_decline_fallback,
|
||||||
_anthropic_stream_error_is_gateway_verdict,
|
_anthropic_stream_error_is_gateway_verdict,
|
||||||
_anthropic_stream_forwards_ping_live,
|
_anthropic_stream_forwards_ping_live,
|
||||||
|
|
@ -10224,6 +10228,214 @@ async def test_anthropic_messages_fallback_also_catches_raised_midstream_error()
|
||||||
assert mock_fallback.await_args.kwargs["e"] is raised_error
|
assert mock_fallback.await_args.kwargs["e"] is raised_error
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"raised_error",
|
||||||
|
[
|
||||||
|
BedrockError(status_code=503, message='serviceUnavailableException {"message": "Service unavailable"}'),
|
||||||
|
BedrockError(status_code=500, message='internalServerException {"message": "Internal error"}'),
|
||||||
|
BedrockError(status_code=429, message='throttlingException {"message": "Too many requests"}'),
|
||||||
|
httpx.ReadError("connection reset by upstream"),
|
||||||
|
],
|
||||||
|
ids=["503", "500", "429", "transport-drop"],
|
||||||
|
)
|
||||||
|
async def test_anthropic_messages_raised_provider_error_before_content_triggers_fallback(raised_error):
|
||||||
|
"""A retriable raise before content falls over exactly like a detected SSE error event."""
|
||||||
|
router = _anthropic_messages_make_router()
|
||||||
|
source = _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], raised_error)
|
||||||
|
fallback_stream = _AnthropicMessagesFallbackByteStream([_anthropic_messages_content_chunk("fallback answer")])
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
router,
|
||||||
|
"async_function_with_fallbacks_common_utils",
|
||||||
|
new=AsyncMock(return_value=fallback_stream),
|
||||||
|
) as mock_fallback:
|
||||||
|
wrapped = await router._aanthropic_messages_streaming_iterator(
|
||||||
|
response=source,
|
||||||
|
initial_kwargs={"model": "primary"},
|
||||||
|
)
|
||||||
|
collected = [chunk async for chunk in wrapped]
|
||||||
|
|
||||||
|
assert collected == [_anthropic_messages_content_chunk("fallback answer")]
|
||||||
|
mock_fallback.assert_awaited_once()
|
||||||
|
converted = mock_fallback.await_args.kwargs["e"]
|
||||||
|
assert isinstance(converted, MidStreamFallbackError)
|
||||||
|
assert converted.original_exception is raised_error
|
||||||
|
assert converted.is_pre_first_chunk is True
|
||||||
|
assert source.closed is True
|
||||||
|
|
||||||
|
|
||||||
|
class _AnthropicMessagesStringStatusError(Exception):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("bad request")
|
||||||
|
self.status_code = "400"
|
||||||
|
|
||||||
|
|
||||||
|
class _AnthropicMessagesResponseOnlyStatusError(Exception):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__("bad request")
|
||||||
|
self.response = SimpleNamespace(status_code=400)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"raised_error",
|
||||||
|
[
|
||||||
|
BedrockError(status_code=400, message='validationException {"message": "Malformed input"}'),
|
||||||
|
BedrockError(status_code=424, message='modelStreamErrorException {"message": "Model stream error"}'),
|
||||||
|
_AnthropicMessagesStringStatusError(),
|
||||||
|
_AnthropicMessagesResponseOnlyStatusError(),
|
||||||
|
],
|
||||||
|
ids=["400", "424", "str-400", "response-only-400"],
|
||||||
|
)
|
||||||
|
async def test_anthropic_messages_raised_non_retriable_provider_error_propagates_unchanged(raised_error):
|
||||||
|
"""A raised client error reaches the caller as the same exception, nothing flushed, no fallback."""
|
||||||
|
router = _anthropic_messages_make_router()
|
||||||
|
source = _AnthropicMessagesRaisingByteStream([_anthropic_messages_message_start_chunk()], raised_error)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
router,
|
||||||
|
"async_function_with_fallbacks_common_utils",
|
||||||
|
new=AsyncMock(),
|
||||||
|
) as mock_fallback:
|
||||||
|
wrapped = await router._aanthropic_messages_streaming_iterator(
|
||||||
|
response=source,
|
||||||
|
initial_kwargs={"model": "primary"},
|
||||||
|
)
|
||||||
|
collected = []
|
||||||
|
|
||||||
|
async def _consume():
|
||||||
|
async for chunk in wrapped:
|
||||||
|
collected.append(chunk)
|
||||||
|
|
||||||
|
with pytest.raises(type(raised_error)) as exc_info:
|
||||||
|
await _consume()
|
||||||
|
|
||||||
|
assert collected == []
|
||||||
|
assert exc_info.value is raised_error
|
||||||
|
mock_fallback.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_anthropic_messages_raised_provider_error_after_content_propagates_unchanged():
|
||||||
|
"""A raise after content propagates unchanged even when its status is retriable."""
|
||||||
|
router = _anthropic_messages_make_router()
|
||||||
|
content = _anthropic_messages_content_chunk("partial answer")
|
||||||
|
raised_error = BedrockError(
|
||||||
|
status_code=503, message='serviceUnavailableException {"message": "Service unavailable"}'
|
||||||
|
)
|
||||||
|
source = _AnthropicMessagesRaisingByteStream([content], raised_error)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
router,
|
||||||
|
"async_function_with_fallbacks_common_utils",
|
||||||
|
new=AsyncMock(),
|
||||||
|
) as mock_fallback:
|
||||||
|
wrapped = await router._aanthropic_messages_streaming_iterator(
|
||||||
|
response=source,
|
||||||
|
initial_kwargs={"model": "primary"},
|
||||||
|
)
|
||||||
|
collected = []
|
||||||
|
|
||||||
|
async def _consume():
|
||||||
|
async for chunk in wrapped:
|
||||||
|
collected.append(chunk)
|
||||||
|
|
||||||
|
with pytest.raises(BedrockError) as exc_info:
|
||||||
|
await _consume()
|
||||||
|
|
||||||
|
assert collected == [content]
|
||||||
|
assert exc_info.value is raised_error
|
||||||
|
mock_fallback.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"error, expected_status",
|
||||||
|
[
|
||||||
|
(BedrockError(status_code=503, message="unavailable"), 503),
|
||||||
|
(_AnthropicMessagesStringStatusError(), 400),
|
||||||
|
(_AnthropicMessagesResponseOnlyStatusError(), 400),
|
||||||
|
(httpx.ReadError("connection reset by upstream"), None),
|
||||||
|
],
|
||||||
|
ids=["int", "digit-str", "response-only", "none"],
|
||||||
|
)
|
||||||
|
def test_anthropic_stream_raised_error_status_reads_every_status_shape(error, expected_status):
|
||||||
|
assert _anthropic_stream_raised_error_status(error) == expected_status
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"error, has_generated_content, converts",
|
||||||
|
[
|
||||||
|
(BedrockError(status_code=503, message="unavailable"), False, True),
|
||||||
|
(httpx.ReadError("connection reset by upstream"), False, True),
|
||||||
|
(BedrockError(status_code=400, message="malformed"), False, False),
|
||||||
|
(BedrockError(status_code=503, message="unavailable"), True, False),
|
||||||
|
],
|
||||||
|
ids=["retriable", "no-status", "client-error", "after-content"],
|
||||||
|
)
|
||||||
|
def test_anthropic_stream_fallback_error_for_raised_gates_like_a_detected_error_event(
|
||||||
|
error, has_generated_content, converts
|
||||||
|
):
|
||||||
|
converted = _anthropic_stream_fallback_error_for_raised(error, "primary", has_generated_content)
|
||||||
|
if not converts:
|
||||||
|
assert converted is None
|
||||||
|
return
|
||||||
|
assert isinstance(converted, MidStreamFallbackError)
|
||||||
|
assert converted.original_exception is error
|
||||||
|
assert converted.is_pre_first_chunk is True
|
||||||
|
assert converted.llm_provider == "anthropic"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_aanthropic_messages_recover_stream_error_flushes_buffered_frames_before_declining():
|
||||||
|
router = _anthropic_messages_make_router()
|
||||||
|
original = BedrockError(status_code=503, message="unavailable")
|
||||||
|
declined = MidStreamFallbackError(
|
||||||
|
message="unavailable",
|
||||||
|
model="primary",
|
||||||
|
llm_provider="anthropic",
|
||||||
|
original_exception=original,
|
||||||
|
is_pre_first_chunk=False,
|
||||||
|
)
|
||||||
|
buffered = (_anthropic_messages_message_start_chunk(),)
|
||||||
|
flushed = []
|
||||||
|
|
||||||
|
async def drain(recovery) -> None:
|
||||||
|
async for chunk in recovery:
|
||||||
|
flushed.append(chunk)
|
||||||
|
|
||||||
|
with patch.object(router, "_aanthropic_messages_fallback_attempt") as mock_attempt:
|
||||||
|
recovery = router._aanthropic_messages_recover_stream_error(
|
||||||
|
declined, True, buffered, "primary", {"model": "primary"}, _anthropic_messages_make_wrapper()
|
||||||
|
)
|
||||||
|
with pytest.raises(BedrockError) as exc_info:
|
||||||
|
await drain(recovery)
|
||||||
|
assert flushed == list(buffered)
|
||||||
|
assert exc_info.value is original
|
||||||
|
mock_attempt.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_aanthropic_messages_recover_stream_error_hands_converted_raise_to_fallback_attempt():
|
||||||
|
router = _anthropic_messages_make_router()
|
||||||
|
raised = BedrockError(status_code=503, message="unavailable")
|
||||||
|
handed_over = []
|
||||||
|
|
||||||
|
async def fake_attempt(fallback_error, initial_kwargs, wrapper):
|
||||||
|
handed_over.append(fallback_error)
|
||||||
|
yield b"fallback"
|
||||||
|
|
||||||
|
with patch.object(router, "_aanthropic_messages_fallback_attempt", new=fake_attempt):
|
||||||
|
recovery = router._aanthropic_messages_recover_stream_error(
|
||||||
|
raised, False, (), "primary", {"model": "primary"}, _anthropic_messages_make_wrapper()
|
||||||
|
)
|
||||||
|
collected = [chunk async for chunk in recovery]
|
||||||
|
assert collected == [b"fallback"]
|
||||||
|
assert len(handed_over) == 1
|
||||||
|
assert isinstance(handed_over[0], MidStreamFallbackError)
|
||||||
|
assert handed_over[0].original_exception is raised
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_anthropic_messages_non_retriable_client_error_skips_fallback():
|
async def test_anthropic_messages_non_retriable_client_error_skips_fallback():
|
||||||
"""A 4xx (non-429) error type (e.g. invalid_request_error) is a client
|
"""A 4xx (non-429) error type (e.g. invalid_request_error) is a client
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue