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
This commit is contained in:
devin-ai-integration[bot] 2026-07-01 06:42:51 +00:00 committed by GitHub
parent 9ea149b49e
commit b768b62067
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 73 additions and 23 deletions

View file

@ -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

View file

@ -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