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:
yuneng-jiang 2026-07-06 21:30:47 -07:00 committed by GitHub
commit f8606b8b03
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 520 additions and 1 deletions

View file

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

View file

@ -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",
)