fix(proxy): secure Anthropic fallback normalization

This commit is contained in:
Devin AI 2026-07-15 00:50:00 +00:00
parent 431140dc40
commit 7e835a99dd
5 changed files with 95 additions and 97 deletions

View file

@ -53,13 +53,16 @@ def normalize_anthropic_server_side_fallbacks(
request_data: Mapping[str, object],
headers: Mapping[str, str],
) -> dict[str, object]:
sanitized_request_data = {
key: value for key, value in request_data.items() if key != ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM
}
if not is_anthropic_server_side_fallback_request(
request_data=request_data,
headers=headers,
):
return dict(request_data)
return sanitized_request_data
return {
**{key: value for key, value in request_data.items() if key != "fallbacks"},
**{key: value for key, value in sanitized_request_data.items() if key != "fallbacks"},
ANTHROPIC_SERVER_SIDE_FALLBACKS_PARAM: request_data["fallbacks"],
}

View file

@ -24,7 +24,10 @@ from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
create_response,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_set_request_parsed_body,
)
from litellm.types.utils import TokenCountResponse
router = APIRouter()
@ -64,10 +67,26 @@ def _strip_total_tokens_from_anthropic_response(response: Any) -> None:
usage.pop("total_tokens", None)
async def _normalize_anthropic_server_side_fallback_request(
request: Request,
) -> None:
normalized_request_data = normalize_anthropic_server_side_fallbacks(
request_data=await _read_request_body(request=request),
headers=dict(request.headers),
)
_safe_set_request_parsed_body(
request=request,
parsed_body=normalized_request_data,
)
@router.post(
"/v1/messages",
tags=["[beta] Anthropic `/v1/messages`"],
dependencies=[Depends(user_api_key_auth)],
dependencies=[
Depends(_normalize_anthropic_server_side_fallback_request),
Depends(user_api_key_auth),
],
)
async def anthropic_response(
fastapi_response: Response,
@ -92,10 +111,7 @@ async def anthropic_response(
version,
)
data = normalize_anthropic_server_side_fallbacks(
request_data=await _read_request_body(request=request),
headers=dict(request.headers),
)
data = await _read_request_body(request=request)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
result = await base_llm_response_processor.base_process_llm_request(

View file

@ -28,9 +28,6 @@ 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,
@ -2812,18 +2809,8 @@ 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:
if _fb_key != "fallbacks" or not has_anthropic_server_side_fallbacks:
fallback_names.extend(iter_router_fallback_model_names(request_data.get(_fb_key)))
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

@ -7,6 +7,7 @@ import unittest
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI, Request
from fastapi.testclient import TestClient
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -144,48 +145,78 @@ class TestEventLoggingBatchEndpoint:
assert response.json() == {"status": "ok"}
@pytest.mark.asyncio
async def test_server_side_fallbacks_are_normalized_before_routing():
@pytest.mark.parametrize(
"request_body,headers,expected_data",
[
(
{
"model": "claude-fable-5",
"fallbacks": [{"model": "claude-opus-4-8"}],
"anthropic_server_fallbacks": [{"model": "restricted-model"}],
},
{
"Anthropic-Beta": "other-beta, server-side-fallback-2026-06-01"
},
{
"model": "claude-fable-5",
"anthropic_server_fallbacks": [{"model": "claude-opus-4-8"}],
},
),
(
{
"model": "claude-fable-5",
"anthropic_server_fallbacks": [{"model": "restricted-model"}],
},
{},
{"model": "claude-fable-5"},
),
(
{
"model": "claude-fable-5",
"fallbacks": [{"model": "litellm-fallback"}],
},
{},
{
"model": "claude-fable-5",
"fallbacks": [{"model": "litellm-fallback"}],
},
),
],
)
def test_server_side_fallbacks_are_normalized_before_auth_and_routing(
request_body,
headers,
expected_data,
):
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"
}
auth_request_data = {}
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(),
async def fake_auth(request: Request):
auth_request_data.update(await ep._read_request_body(request=request))
return MagicMock()
app = FastAPI()
app.include_router(ep.router)
app.dependency_overrides[ep.user_api_key_auth] = fake_auth
with patch.object(
ep,
"ProxyBaseLLMRequestProcessing",
return_value=processor,
) as processor_factory:
response = TestClient(app).post(
"/v1/messages",
json=request_body,
headers=headers,
)
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
assert response.status_code == 200
assert response.json() == {"id": "msg_test"}
assert auth_request_data == expected_data
assert processor_factory.call_args.kwargs["data"] == expected_data
class TestStripTotalTokens(unittest.TestCase):

View file

@ -6,7 +6,7 @@ execute requests against models their API key cannot call.
"""
from typing import List
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, patch
import pytest
@ -200,45 +200,6 @@ 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