From 7e835a99dd30dcfa5e907fb34ba4df3ed7441c35 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 15 Jul 2026 00:50:00 +0000 Subject: [PATCH] fix(proxy): secure Anthropic fallback normalization --- litellm/llms/anthropic/common_utils.py | 7 +- .../proxy/anthropic_endpoints/endpoints.py | 28 +++-- litellm/proxy/auth/user_api_key_auth.py | 15 +-- .../anthropic_endpoints/test_endpoints.py | 101 ++++++++++++------ .../test_router_override_fallback_auth.py | 41 +------ 5 files changed, 95 insertions(+), 97 deletions(-) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index d5861226856..5b00b5a5a9c 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -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"], } diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 767117a65e0..91864ab9cd1 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -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( diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 49714251da7..744d8182715 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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))) diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index f115182ed81..664817ab523 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -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): diff --git a/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py index 0121fb32f03..fc0e9aec501 100644 --- a/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py +++ b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py @@ -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