mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
refactor(router): drop dead provider derivation in raised-stream fallback
This commit is contained in:
parent
9f290d8b99
commit
406db3fccf
2 changed files with 20 additions and 9 deletions
|
|
@ -431,7 +431,7 @@ def _anthropic_stream_raised_error_status(error: Exception) -> int | None:
|
||||||
|
|
||||||
|
|
||||||
def _anthropic_stream_fallback_error_for_raised(
|
def _anthropic_stream_fallback_error_for_raised(
|
||||||
error: Exception, model: str, llm_provider: str, has_generated_content: bool
|
error: Exception, model: str, has_generated_content: bool
|
||||||
) -> "MidStreamFallbackError | None":
|
) -> "MidStreamFallbackError | None":
|
||||||
"""
|
"""
|
||||||
A provider iterator that fails mid-stream by raising (Bedrock surfaces
|
A provider iterator that fails mid-stream by raising (Bedrock surfaces
|
||||||
|
|
@ -454,7 +454,7 @@ def _anthropic_stream_fallback_error_for_raised(
|
||||||
return MidStreamFallbackError(
|
return MidStreamFallbackError(
|
||||||
message=str(error),
|
message=str(error),
|
||||||
model=model,
|
model=model,
|
||||||
llm_provider=llm_provider,
|
llm_provider="anthropic",
|
||||||
original_exception=error,
|
original_exception=error,
|
||||||
is_pre_first_chunk=True,
|
is_pre_first_chunk=True,
|
||||||
)
|
)
|
||||||
|
|
@ -5059,8 +5059,6 @@ class Router:
|
||||||
has_generated_content = False # rebind-ok: set once real content is seen, or the buffer cap is hit
|
has_generated_content = False # rebind-ok: set once real content is seen, or the buffer cap is hit
|
||||||
buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline
|
buffered_lifecycle_chunks: tuple[bytes, ...] = () # rebind-ok: flushed once committed or on decline
|
||||||
model: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group
|
model: Final = cast(str, initial_kwargs.get("model")) # cast-ok: kwargs always carries the model group
|
||||||
custom_llm_provider: Final = initial_kwargs.get("custom_llm_provider")
|
|
||||||
llm_provider: Final = custom_llm_provider if isinstance(custom_llm_provider, str) else "anthropic"
|
|
||||||
try:
|
try:
|
||||||
async for chunk in source_iterator:
|
async for chunk in source_iterator:
|
||||||
if _anthropic_stream_forwards_ping_live(
|
if _anthropic_stream_forwards_ping_live(
|
||||||
|
|
@ -5109,7 +5107,6 @@ class Router:
|
||||||
has_generated_content,
|
has_generated_content,
|
||||||
buffered_lifecycle_chunks,
|
buffered_lifecycle_chunks,
|
||||||
model,
|
model,
|
||||||
llm_provider,
|
|
||||||
initial_kwargs,
|
initial_kwargs,
|
||||||
wrapper,
|
wrapper,
|
||||||
):
|
):
|
||||||
|
|
@ -5130,7 +5127,6 @@ class Router:
|
||||||
has_generated_content: bool,
|
has_generated_content: bool,
|
||||||
buffered_lifecycle_chunks: tuple[bytes, ...],
|
buffered_lifecycle_chunks: tuple[bytes, ...],
|
||||||
model: str,
|
model: str,
|
||||||
llm_provider: str,
|
|
||||||
initial_kwargs: dict[str, Any], # mutable-ok: handed to _aanthropic_messages_fallback_attempt, which mutates it
|
initial_kwargs: dict[str, Any], # mutable-ok: handed to _aanthropic_messages_fallback_attempt, which mutates it
|
||||||
wrapper: "FallbackAwareAnthropicMessagesStream",
|
wrapper: "FallbackAwareAnthropicMessagesStream",
|
||||||
) -> AsyncGenerator[bytes, None]:
|
) -> AsyncGenerator[bytes, None]:
|
||||||
|
|
@ -5158,7 +5154,7 @@ class Router:
|
||||||
fallback_error: Final = (
|
fallback_error: Final = (
|
||||||
stream_error
|
stream_error
|
||||||
if isinstance(stream_error, MidStreamFallbackError)
|
if isinstance(stream_error, MidStreamFallbackError)
|
||||||
else _anthropic_stream_fallback_error_for_raised(stream_error, model, llm_provider, has_generated_content)
|
else _anthropic_stream_fallback_error_for_raised(stream_error, model, has_generated_content)
|
||||||
)
|
)
|
||||||
if fallback_error is None:
|
if fallback_error is None:
|
||||||
raise stream_error
|
raise stream_error
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,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
|
||||||
|
|
@ -10264,14 +10265,28 @@ async def test_anthropic_messages_raised_provider_error_before_content_triggers_
|
||||||
assert source.closed 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.asyncio
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"raised_error",
|
"raised_error",
|
||||||
[
|
[
|
||||||
BedrockError(status_code=400, message='validationException {"message": "Malformed input"}'),
|
BedrockError(status_code=400, message='validationException {"message": "Malformed input"}'),
|
||||||
BedrockError(status_code=424, message='modelStreamErrorException {"message": "Model stream error"}'),
|
BedrockError(status_code=424, message='modelStreamErrorException {"message": "Model stream error"}'),
|
||||||
|
_AnthropicMessagesStringStatusError(),
|
||||||
|
_AnthropicMessagesResponseOnlyStatusError(),
|
||||||
],
|
],
|
||||||
ids=["400", "424"],
|
ids=["400", "424", "str-400", "response-only-400"],
|
||||||
)
|
)
|
||||||
async def test_anthropic_messages_raised_non_retriable_provider_error_propagates_unchanged(raised_error):
|
async def test_anthropic_messages_raised_non_retriable_provider_error_propagates_unchanged(raised_error):
|
||||||
"""A raised 4xx (other than 429) is a client error no other deployment can
|
"""A raised 4xx (other than 429) is a client error no other deployment can
|
||||||
|
|
@ -10295,7 +10310,7 @@ async def test_anthropic_messages_raised_non_retriable_provider_error_propagates
|
||||||
async for chunk in wrapped:
|
async for chunk in wrapped:
|
||||||
collected.append(chunk)
|
collected.append(chunk)
|
||||||
|
|
||||||
with pytest.raises(BedrockError) as exc_info:
|
with pytest.raises(type(raised_error)) as exc_info:
|
||||||
await _consume()
|
await _consume()
|
||||||
|
|
||||||
assert collected == []
|
assert collected == []
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue