mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
Merge pull request #31788 from BerriAI/litellm_local-rate-limit-fallbacks
fix(proxy): trigger gateway fallbacks on local rate limit errors
This commit is contained in:
commit
f8606b8b03
2 changed files with 520 additions and 1 deletions
|
|
@ -1198,6 +1198,120 @@ 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,
|
||||
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,
|
||||
)
|
||||
|
||||
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
|
||||
|
||||
def _resolve_fallback_models(
|
||||
self,
|
||||
model: str,
|
||||
llm_router: Router,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Optional[list]:
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
|
||||
fallbacks = None
|
||||
|
||||
key_router_settings = user_api_key_dict.router_settings
|
||||
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."""
|
||||
|
|
@ -1373,7 +1487,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,
|
||||
|
|
|
|||
|
|
@ -4466,3 +4466,408 @@ class TestResponseCostHeaderForTypedDictResponses:
|
|||
assert "_hidden_params" not in result
|
||||
assert fastapi_response.headers["x-ratelimit-limit-input-tokens"] == "25"
|
||||
assert fastapi_response.headers["x-litellm-response-cost"] == "0.00123"
|
||||
|
||||
|
||||
class TestPreCallWithFallbacksOnLocalRateLimit:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_triggered_on_local_rate_limit(self):
|
||||
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):
|
||||
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):
|
||||
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,
|
||||
)
|
||||
|
||||
assert processor.data["model"] == "gpt-4"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fallback_uses_key_level_router_settings(self):
|
||||
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,
|
||||
)
|
||||
|
||||
assert processor.data["model"] == "claude-3-haiku"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_fallbacks_flag_respected(self):
|
||||
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,
|
||||
)
|
||||
|
||||
@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
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_real_parallel_request_limiter_model_tpm_limit_triggers_fallback(self):
|
||||
"""
|
||||
Customer-reported scenario from LIT-3890 / GH #8822.
|
||||
|
||||
The prior tests in this class hand-build a ``ProxyRateLimitError``. The
|
||||
customer's production setup is different: they set a *per-key per-model*
|
||||
TPM cap on the key itself::
|
||||
|
||||
Model TPM Limits: {"gpt-4.1-20250414-test": 100}
|
||||
|
||||
and configure a proxy-side fallback (gpt-4.1-...-test -> gpt-4.1-...).
|
||||
When the per-model TPM cap trips, the real
|
||||
``parallel_request_limiter`` raises ``ProxyRateLimitError`` from inside
|
||||
``proxy_logging_obj.pre_call_hook`` — the seam ``_pre_call_with_fallbacks``
|
||||
wraps. This test drives that *real* limiter (not a mock error) end-to-end
|
||||
to prove the customer's exact knob triggers the gateway fallback instead
|
||||
of returning a 429 to the client.
|
||||
"""
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_request_processing import (
|
||||
ProxyBaseLLMRequestProcessing,
|
||||
)
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import (
|
||||
ProxyRateLimitError,
|
||||
)
|
||||
from litellm.proxy.hooks.parallel_request_limiter import (
|
||||
_PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
from litellm.proxy.utils import InternalUsageCache
|
||||
|
||||
primary_model = "gpt-4"
|
||||
fallback_model = "gpt-3.5-turbo"
|
||||
|
||||
# Freeze the limiter's clock so the per-minute counter key is stable and
|
||||
# the pre-seeded counter is guaranteed to be the one it reads.
|
||||
class _FrozenClock(datetime.datetime):
|
||||
@classmethod
|
||||
def now(cls, tz=None):
|
||||
return cls(2026, 1, 1, 12, 30, 0)
|
||||
|
||||
precise_minute = "2026-01-01-12-30"
|
||||
|
||||
# Real per-key per-model TPM limiter + a key carrying the customer's
|
||||
# `model_tpm_limit` metadata (only the primary is capped).
|
||||
limiter = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(DualCache())
|
||||
)
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="sk-lit3890",
|
||||
metadata={"model_tpm_limit": {primary_model: 100}},
|
||||
)
|
||||
|
||||
# Pre-seed the primary's per-model token counter at the cap so the very
|
||||
# next request trips it. The counter key uses the *hashed* api_key.
|
||||
counter_key = (
|
||||
f"{user_api_key_dict.api_key}::{primary_model}"
|
||||
f"::{precise_minute}::request_count"
|
||||
)
|
||||
await limiter.internal_usage_cache.async_set_cache(
|
||||
key=counter_key,
|
||||
value={"current_requests": 0, "current_tpm": 100, "current_rpm": 0},
|
||||
litellm_parent_otel_span=None,
|
||||
local_only=True,
|
||||
)
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data={"model": primary_model})
|
||||
|
||||
# Stand in for common_processing_pre_call_logic's pre_call_hook step by
|
||||
# invoking the real limiter for whatever model is currently selected.
|
||||
limiter_calls = []
|
||||
|
||||
async def real_limiter_pre_call(**kwargs):
|
||||
current_model = processor.data["model"]
|
||||
limiter_calls.append(current_model)
|
||||
await limiter.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={
|
||||
"model": current_model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
call_type="acompletion",
|
||||
)
|
||||
return processor.data, MagicMock()
|
||||
|
||||
mock_router = MagicMock()
|
||||
mock_router.fallbacks = [{primary_model: [fallback_model]}]
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
|
||||
):
|
||||
with patch.object(
|
||||
processor,
|
||||
"common_processing_pre_call_logic",
|
||||
side_effect=real_limiter_pre_call,
|
||||
):
|
||||
data, logging_obj = 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=primary_model,
|
||||
route_type="acompletion",
|
||||
llm_router=mock_router,
|
||||
)
|
||||
|
||||
# The capped primary tripped the real limiter, and the fallback (which
|
||||
# has no per-model cap) served the request — no 429 to the client.
|
||||
assert processor.data["model"] == fallback_model
|
||||
assert limiter_calls == [primary_model, fallback_model]
|
||||
|
||||
# Sanity-check the premise: the limiter genuinely raises a
|
||||
# ProxyRateLimitError for the capped primary under the frozen clock.
|
||||
with patch(
|
||||
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
|
||||
):
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
await limiter.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={
|
||||
"model": primary_model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
call_type="acompletion",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue