mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(proxy): trigger gateway fallbacks on local rate limit errors
When pre-call hooks (parallel_request_limiter, dynamic_rate_limiter_v3) reject a request with ProxyRateLimitError, the router's fallback logic was never reached because the exception was raised before route_request was called. Add _pre_call_with_fallbacks that catches ProxyRateLimitError, resolves configured fallbacks (key-level router_settings -> router-level), and retries with each fallback model in order. If all fallbacks are also rate-limited, the original error is re-raised.
This commit is contained in:
parent
cca71a07c2
commit
1d89e65731
2 changed files with 362 additions and 1 deletions
|
|
@ -1174,6 +1174,118 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
return self.data, logging_obj
|
||||
|
||||
async def _pre_call_with_fallbacks(
|
||||
self,
|
||||
request: Request,
|
||||
general_settings: dict,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
version: Optional[str],
|
||||
proxy_config: ProxyConfig,
|
||||
user_model: Optional[str],
|
||||
user_temperature: Optional[float],
|
||||
user_request_timeout: Optional[float],
|
||||
user_max_tokens: Optional[int],
|
||||
user_api_base: Optional[str],
|
||||
model: Optional[str],
|
||||
route_type: str,
|
||||
llm_router: Optional[Router],
|
||||
) -> Tuple[dict, LiteLLMLoggingObj]:
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
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=model,
|
||||
route_type=route_type,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyRateLimitError as original_exc:
|
||||
original_model = self.data.get("model")
|
||||
if not original_model or not llm_router or self.data.get("disable_fallbacks"):
|
||||
raise
|
||||
|
||||
fallback_models = self._resolve_fallback_models(
|
||||
model=original_model,
|
||||
llm_router=llm_router,
|
||||
proxy_config=proxy_config,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
if not fallback_models:
|
||||
raise
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"Local rate limit hit for model=%s, attempting fallbacks: %s",
|
||||
original_model,
|
||||
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
|
||||
|
||||
self.data["model"] = original_model
|
||||
raise original_exc
|
||||
|
||||
def _resolve_fallback_models(
|
||||
self,
|
||||
model: str,
|
||||
llm_router: Router,
|
||||
proxy_config: ProxyConfig,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Optional[list]:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallbacks = None
|
||||
|
||||
key_router_settings = getattr(user_api_key_dict, "router_settings", None)
|
||||
if isinstance(key_router_settings, dict) and "fallbacks" in key_router_settings:
|
||||
fallbacks = key_router_settings["fallbacks"]
|
||||
|
||||
if fallbacks is None:
|
||||
fallbacks = llm_router.fallbacks
|
||||
|
||||
if not fallbacks:
|
||||
return None
|
||||
|
||||
fallback_model_group, generic_fallback_idx = get_fallback_model_group(
|
||||
fallbacks=fallbacks,
|
||||
model_group=model,
|
||||
)
|
||||
if fallback_model_group is None and generic_fallback_idx is not None:
|
||||
fallback_model_group = fallbacks[generic_fallback_idx]["*"]
|
||||
return fallback_model_group
|
||||
|
||||
@staticmethod
|
||||
def _get_model_id_from_response(hidden_params: dict, data: dict) -> str:
|
||||
"""Extract model_id from hidden_params with fallback to litellm_metadata."""
|
||||
|
|
@ -1349,7 +1461,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"Ensure common_processing_pre_call_logic was called before using this parameter."
|
||||
)
|
||||
else:
|
||||
self.data, logging_obj = await self.common_processing_pre_call_logic(
|
||||
self.data, logging_obj = await self._pre_call_with_fallbacks(
|
||||
request=request,
|
||||
general_settings=general_settings,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
|
|||
|
|
@ -4352,3 +4352,252 @@ class TestResponseCostHeaderForTypedDictResponses:
|
|||
|
||||
assert "x-litellm-response-cost" not in fastapi_response.headers
|
||||
recompute.assert_not_called()
|
||||
|
||||
|
||||
class TestPreCallWithFallbacksOnLocalRateLimit:
|
||||
"""
|
||||
Regression tests for LIT-3890: proxy fallbacks must trigger when local rate
|
||||
limits (key-level TPM/RPM or dynamic_rate_limiter_v3) reject a request.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_triggered_on_local_rate_limit(self):
|
||||
"""
|
||||
When the primary model is locally rate-limited, the request should
|
||||
proceed with a configured fallback model.
|
||||
"""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
primary_model = "gpt-4"
|
||||
fallback_model = "gpt-3.5-turbo"
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": primary_model})
|
||||
|
||||
call_count = 0
|
||||
|
||||
async def mock_pre_call_logic(**kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
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"},
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
return processor.data, logging_obj
|
||||
|
||||
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,
|
||||
):
|
||||
data, logging_obj = 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=primary_model,
|
||||
route_type="acompletion",
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
assert processor.data["model"] == fallback_model
|
||||
assert call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_no_fallbacks_configured(self):
|
||||
"""
|
||||
When no fallbacks are configured, the original rate limit error
|
||||
should propagate unchanged.
|
||||
"""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4"})
|
||||
|
||||
async def mock_pre_call_logic(**kwargs):
|
||||
raise ProxyRateLimitError(
|
||||
detail="TPM limit exceeded",
|
||||
headers={"retry-after": "30"},
|
||||
)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.fallbacks = None
|
||||
|
||||
with patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
side_effect=mock_pre_call_logic,
|
||||
):
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
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,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_all_fallbacks_also_rate_limited(self):
|
||||
"""
|
||||
When all fallback models are also locally rate-limited, the original
|
||||
error for the primary model should be re-raised.
|
||||
"""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4"})
|
||||
|
||||
async def mock_pre_call_logic(**kwargs):
|
||||
raise ProxyRateLimitError(
|
||||
detail=f"TPM limit exceeded for {processor.data.get('model')}",
|
||||
headers={"retry-after": "30"},
|
||||
)
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo", "claude-3-haiku"]}]
|
||||
|
||||
with patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
side_effect=mock_pre_call_logic,
|
||||
):
|
||||
with pytest.raises(ProxyRateLimitError, match="gpt-4"):
|
||||
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,
|
||||
)
|
||||
|
||||
# Model should be restored to original
|
||||
assert processor.data["model"] == "gpt-4"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_uses_key_level_router_settings(self):
|
||||
"""
|
||||
Key-level router_settings fallbacks should take precedence over
|
||||
router-level fallbacks.
|
||||
"""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4"})
|
||||
|
||||
async def mock_pre_call_logic(**kwargs):
|
||||
if processor.data.get("model") == "gpt-4":
|
||||
raise ProxyRateLimitError(
|
||||
detail="TPM limit exceeded",
|
||||
headers={"retry-after": "30"},
|
||||
)
|
||||
return processor.data, MagicMock()
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo"]}]
|
||||
|
||||
user_api_key_dict = MagicMock()
|
||||
user_api_key_dict.router_settings = {
|
||||
"fallbacks": [{"gpt-4": ["claude-3-haiku"]}]
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
side_effect=mock_pre_call_logic,
|
||||
):
|
||||
data, _ = await processor._pre_call_with_fallbacks(
|
||||
request=MagicMock(),
|
||||
general_settings={},
|
||||
proxy_logging_obj=MagicMock(),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
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,
|
||||
)
|
||||
|
||||
# Should use key-level fallback, not router-level
|
||||
assert processor.data["model"] == "claude-3-haiku"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_fallbacks_flag_respected(self):
|
||||
"""
|
||||
When disable_fallbacks is set in request data, local rate limit
|
||||
errors should not trigger fallback logic.
|
||||
"""
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(
|
||||
data={"model": "gpt-4", "disable_fallbacks": True}
|
||||
)
|
||||
|
||||
async def mock_pre_call_logic(**kwargs):
|
||||
raise ProxyRateLimitError(
|
||||
detail="TPM limit exceeded",
|
||||
headers={"retry-after": "30"},
|
||||
)
|
||||
|
||||
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(ProxyRateLimitError):
|
||||
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,
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue