mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
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:
parent
9ea149b49e
commit
b768b62067
2 changed files with 73 additions and 23 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue