diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index abaf82fb661..02bb66388ca 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c61fdec370c..aa1911f80bc 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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", + )