fix(proxy): pass through Anthropic server fallbacks

This commit is contained in:
Devin AI 2026-07-15 00:17:41 +00:00
parent 8c776605d8
commit 431140dc40
8 changed files with 180 additions and 6 deletions

View file

@ -4,7 +4,7 @@ This file contains common utils for anthropic calls.
import copy
import re
from typing import Any, Dict, List, Optional, Union
from typing import Any, Dict, List, Mapping, Optional, Union
import httpx
@ -18,6 +18,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.anthropic import (
ANTHROPIC_BETA_HEADER_VALUES,
ANTHROPIC_HOSTED_TOOLS,
ANTHROPIC_OAUTH_BETA_HEADER,
ANTHROPIC_OAUTH_TOKEN_PREFIX,
@ -30,6 +31,37 @@ _BEDROCK_VERSION_SUFFIX_RE = re.compile(r"-v\d+(?::\d+)?$")
_INFERENCE_PROFILE_MINOR_RE = re.compile(r":\d+$")
_DATED_RELEASE_SUFFIX_RE = re.compile(r"-\d{8}$")
_DOTTED_VERSION_RE = re.compile(r"(\d)\.(\d)")
ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM = "anthropic_server_fallbacks"
def is_anthropic_server_side_fallback_request(
request_data: Mapping[str, object],
headers: Mapping[str, str],
) -> bool:
beta_header = next(
(value for key, value in headers.items() if key.lower() == "anthropic-beta"),
None,
)
if beta_header is None or not isinstance(request_data.get("fallbacks"), list):
return False
return ANTHROPIC_BETA_HEADER_VALUES.SERVER_SIDE_FALLBACK_2026_06_01.value in {
value.strip() for value in beta_header.split(",")
}
def normalize_anthropic_server_side_fallbacks(
request_data: Mapping[str, object],
headers: Mapping[str, str],
) -> dict[str, object]:
if not is_anthropic_server_side_fallback_request(
request_data=request_data,
headers=headers,
):
return dict(request_data)
return {
**{key: value for key, value in request_data.items() if key != "fallbacks"},
ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM: request_data["fallbacks"],
}
def _strip_bedrock_id_suffixes(model: str) -> str:

View file

@ -23,6 +23,7 @@ from typing import (
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.common_utils import (
ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM,
sanitize_tool_use_ids_in_anthropic_messages,
strip_empty_text_blocks_from_anthropic_messages,
)
@ -429,6 +430,7 @@ def anthropic_messages_handler(
metadata = validate_anthropic_api_metadata(metadata)
anthropic_server_fallbacks = kwargs.pop(ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM, None)
local_vars = locals()
is_async = kwargs.pop("is_async", False)
# Use provided client or create a new one
@ -532,7 +534,7 @@ def anthropic_messages_handler(
)
local_vars.update(kwargs)
anthropic_messages_optional_request_params = (
anthropic_messages_optional_request_params = dict(
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
params=local_vars,
model=model,
@ -540,6 +542,8 @@ def anthropic_messages_handler(
custom_llm_provider=custom_llm_provider,
)
)
if custom_llm_provider == LlmProviders.ANTHROPIC.value and isinstance(anthropic_server_fallbacks, list):
anthropic_messages_optional_request_params["fallbacks"] = anthropic_server_fallbacks
if is_reasoning_auto_summary_enabled():
thinking_param = anthropic_messages_optional_request_params.get("thinking")
if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled":
@ -552,7 +556,7 @@ def anthropic_messages_handler(
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params),
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
_is_async=is_async,
client=client,
custom_llm_provider=custom_llm_provider,

View file

@ -9,6 +9,9 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.anthropic.common_utils import (
normalize_anthropic_server_side_fallbacks,
)
from litellm.llms.anthropic.experimental_pass_through.context_management import (
AnthropicContextManagementError,
)
@ -89,7 +92,10 @@ async def anthropic_response(
version,
)
data = await _read_request_body(request=request)
data = normalize_anthropic_server_side_fallbacks(
request_data=await _read_request_body(request=request),
headers=dict(request.headers),
)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
result = await base_llm_response_processor.base_process_llm_request(

View file

@ -28,6 +28,9 @@ from litellm.integrations.otel.model.config import is_otel_v2_enabled
from litellm.integrations.otel.runtime import phase_span, seed_request_identity
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.llms.anthropic.common_utils import (
is_anthropic_server_side_fallback_request,
)
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
@ -2809,8 +2812,18 @@ async def _enforce_key_and_fallback_model_access(
# allowlist or a caller can smuggle a restricted model. VERIA-44.
fallback_names: List[str] = []
override_settings = request_data.get("router_settings_override")
normalized_route = normalize_route_for_root_path(route)
has_anthropic_server_side_fallbacks = (
normalized_route == "/v1/messages"
and request is not None
and is_anthropic_server_side_fallback_request(
request_data=request_data,
headers=_safe_get_request_headers(request),
)
)
for _fb_key in ROUTER_FALLBACK_FIELDS:
fallback_names.extend(iter_router_fallback_model_names(request_data.get(_fb_key)))
if _fb_key != "fallbacks" or not has_anthropic_server_side_fallbacks:
fallback_names.extend(iter_router_fallback_model_names(request_data.get(_fb_key)))
if isinstance(override_settings, dict):
fallback_names.extend(iter_router_fallback_model_names(override_settings.get(_fb_key)))

View file

@ -705,6 +705,7 @@ class ANTHROPIC_BETA_HEADER_VALUES(str, Enum):
ADVANCED_TOOL_USE_2025_11_20 = "advanced-tool-use-2025-11-20"
FAST_MODE_2026_02_01 = "fast-mode-2026-02-01"
ADVISOR_TOOL_2026_03_01 = "advisor-tool-2026-03-01"
SERVER_SIDE_FALLBACK_2026_06_01 = "server-side-fallback-2026-06-01"
# Tool search beta header constant (for Anthropic direct API and Microsoft Foundry)

View file

@ -64,6 +64,41 @@ def test_anthropic_experimental_pass_through_messages_handler_dynamic_api_key_an
assert mock_completion.call_args.kwargs["custom_key"] == "custom_value"
def test_anthropic_server_side_fallbacks_forwarded_to_anthropic_request():
from litellm.llms.anthropic.experimental_pass_through.messages import handler
fallbacks = [{"model": "claude-opus-4-8"}]
with patch.object(
handler.base_llm_http_handler,
"anthropic_messages_handler",
return_value={"id": "msg_test"},
) as mock_anthropic_messages:
handler.anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "Hello"}],
model="anthropic/claude-fable-5",
custom_llm_provider="anthropic",
api_key="test-api-key",
anthropic_server_fallbacks=fallbacks,
)
call_kwargs = mock_anthropic_messages.call_args.kwargs
optional_params = call_kwargs["anthropic_messages_optional_request_params"]
request_body = call_kwargs[
"anthropic_messages_provider_config"
].transform_anthropic_messages_request(
model=call_kwargs["model"],
messages=call_kwargs["messages"],
anthropic_messages_optional_request_params=dict(optional_params),
litellm_params=call_kwargs["litellm_params"],
headers={},
)
assert optional_params["fallbacks"] == fallbacks
assert request_body["fallbacks"] == fallbacks
assert "anthropic_server_fallbacks" not in call_kwargs["kwargs"]
@pytest.mark.asyncio
async def test_anthropic_messages_sanitizes_empty_text_blocks_before_dispatch():
"""Regression test for #22930. The unified /v1/messages path must

View file

@ -144,6 +144,50 @@ class TestEventLoggingBatchEndpoint:
assert response.json() == {"status": "ok"}
@pytest.mark.asyncio
async def test_server_side_fallbacks_are_normalized_before_routing():
import litellm.proxy.anthropic_endpoints.endpoints as ep
from litellm.llms.anthropic.common_utils import (
ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM,
)
processor = MagicMock()
processor.base_process_llm_request = AsyncMock(return_value={"id": "msg_test"})
fallbacks = [{"model": "claude-opus-4-8"}]
request = MagicMock()
request.headers = {
"Anthropic-Beta": "other-beta, server-side-fallback-2026-06-01"
}
with (
patch.object(
ep,
"_read_request_body",
new=AsyncMock(
return_value={
"model": "claude-fable-5",
"fallbacks": fallbacks,
}
),
),
patch.object(
ep,
"ProxyBaseLLMRequestProcessing",
return_value=processor,
) as processor_factory,
):
result = await ep.anthropic_response(
fastapi_response=MagicMock(),
request=request,
user_api_key_dict=MagicMock(),
)
routed_data = processor_factory.call_args.kwargs["data"]
assert result == {"id": "msg_test"}
assert "fallbacks" not in routed_data
assert routed_data[ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM] == fallbacks
class TestStripTotalTokens(unittest.TestCase):
"""Cover ``_strip_total_tokens_from_anthropic_response``.

View file

@ -6,7 +6,7 @@ execute requests against models their API key cannot call.
"""
from typing import List
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -200,6 +200,45 @@ async def test_top_level_fallback_fields_validated(fallback_field):
assert "top-level-smuggled" in seen
@pytest.mark.asyncio
async def test_anthropic_server_side_fallbacks_are_not_routed_by_litellm():
valid_token = _key_with_models(["claude-fable-5"])
request_data = {
"model": "claude-fable-5",
"fallbacks": [{"model": "claude-opus-4-8"}],
}
request = MagicMock()
request.headers = {
"anthropic-beta": "other-beta,server-side-fallback-2026-06-01"
}
seen: List[str] = []
async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router):
seen.append(model)
with (
patch(
"litellm.proxy.auth.user_api_key_auth.can_key_call_model",
side_effect=fake_can_key_call_model,
),
patch(
"litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model",
new=AsyncMock(),
) as mock_is_valid_fallback,
):
await _enforce_key_and_fallback_model_access(
valid_token=valid_token,
request_data=request_data,
route="/v1/messages",
request=request,
llm_model_list=None,
llm_router=None,
)
assert seen == ["claude-fable-5"]
mock_is_valid_fallback.assert_not_awaited()
@pytest.mark.asyncio
async def test_router_override_without_fallbacks_does_not_break_auth():
"""``router_settings_override`` set without any fallback fields is a