From b768b6206779a6dd5e2060be62fabfaad7e964ed Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 1 Jul 2026 06:42:51 +0000 Subject: [PATCH] fix(proxy): restore model state on non-rate-limit exceptions in fallback loop Addresses Greptile review feedback: wrap the fallback loop in try/except BaseException to always restore self.data['model'] to the original value when a non-ProxyRateLimitError exception escapes a fallback attempt. Add regression test for this edge case --- litellm/proxy/common_request_processing.py | 50 ++++++++++--------- .../proxy/test_common_request_processing.py | 46 +++++++++++++++++ 2 files changed, 73 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index ddae83e50ae..df2d1de48b6 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1229,29 +1229,33 @@ class ProxyBaseLLMRequestProcessing: fallback_models, ) - for fallback_model in fallback_models: - if fallback_model == original_model: - continue - self.data["model"] = fallback_model - try: - return await self.common_processing_pre_call_logic( - request=request, - general_settings=general_settings, - proxy_logging_obj=proxy_logging_obj, - user_api_key_dict=user_api_key_dict, - version=version, - proxy_config=proxy_config, - user_model=user_model, - user_temperature=user_temperature, - user_request_timeout=user_request_timeout, - user_max_tokens=user_max_tokens, - user_api_base=user_api_base, - model=fallback_model, - route_type=route_type, - llm_router=llm_router, - ) - except ProxyRateLimitError: - continue + try: + for fallback_model in fallback_models: + if fallback_model == original_model: + continue + self.data["model"] = fallback_model + try: + return await self.common_processing_pre_call_logic( + request=request, + general_settings=general_settings, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + version=version, + proxy_config=proxy_config, + user_model=user_model, + user_temperature=user_temperature, + user_request_timeout=user_request_timeout, + user_max_tokens=user_max_tokens, + user_api_base=user_api_base, + model=fallback_model, + route_type=route_type, + llm_router=llm_router, + ) + except ProxyRateLimitError: + continue + except BaseException: + self.data["model"] = original_model + raise self.data["model"] = original_model raise original_exc diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index a8a74200c25..ef26d9e3e87 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -4575,3 +4575,49 @@ class TestPreCallWithFallbacksOnLocalRateLimit: route_type="acompletion", llm_router=mock_router, ) + + @pytest.mark.asyncio + async def test_model_restored_on_non_rate_limit_exception(self): + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + primary_model = "gpt-4" + + processor = ProxyBaseLLMRequestProcessing(data={"model": primary_model}) + + async def mock_pre_call_logic(**kwargs): + model_in_data = processor.data.get("model") + if model_in_data == primary_model: + raise ProxyRateLimitError( + detail="TPM limit exceeded for gpt-4", + headers={"retry-after": "30"}, + ) + raise ValueError("unexpected auth failure on fallback") + + mock_router = MagicMock() + mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo"]}] + + with patch.object( + processor, + "common_processing_pre_call_logic", + side_effect=mock_pre_call_logic, + ): + with pytest.raises(ValueError, match="unexpected auth failure"): + await processor._pre_call_with_fallbacks( + request=MagicMock(), + general_settings={}, + proxy_logging_obj=MagicMock(), + user_api_key_dict=MagicMock(router_settings=None), + version=None, + proxy_config=MagicMock(), + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + model="gpt-4", + route_type="acompletion", + llm_router=mock_router, + ) + + assert processor.data["model"] == primary_model