From 2b6184d76867fd38a990d2df124b7d1cd808ca6c Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 16 Sep 2026 07:08:24 +0000 Subject: [PATCH] fix(proxy): resolve rate-limit fallbacks after model normalization and retry from a client-request snapshot The fallback retry in _pre_call_with_fallbacks re-entered common_processing_pre_call_logic with data already enriched by the first pass, so add_litellm_data_to_request deep-copied a metadata dict holding the live OTel span and the request failed with a 500 (cannot pickle '_thread.RLock') instead of the intended 429 or fallback. Capture the configured fallbacks and a snapshot of the client request before the first pass, look up the fallback chain by the normalized model group after the limiter raises, and run each fallback attempt on a fresh copy of that snapshot. Replaces the mock-heavy tests with a rig that runs the real v3 limiter and a live OTel span through the proxy_logging_obj seam Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 62 ++-- .../test_response_polling_pre_call_checks.py | 8 +- .../proxy/test_common_request_processing.py | 347 +++++++----------- 3 files changed, 166 insertions(+), 251 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 3c8bb9b3c92..7c9a296965c 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2066,20 +2066,12 @@ class ProxyBaseLLMRequestProcessing: ) -> tuple[dict, LiteLLMLoggingObj]: from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError - original_model: Final = self.data.get("model") - fallback_models: Final = ( - self._resolve_fallback_models( - model=original_model, - llm_router=llm_router, - user_api_key_dict=user_api_key_dict, - ) - if original_model - and isinstance(original_model, str) - and llm_router - and not self.data.get("disable_fallbacks") + configured_fallbacks: Final = ( + self._configured_fallbacks(llm_router=llm_router, user_api_key_dict=user_api_key_dict) + if llm_router is not None and not self.data.get("disable_fallbacks") else None ) - pristine: Final = independent_snapshot(self.data) if fallback_models else None + pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None try: return await self.common_processing_pre_call_logic( @@ -2099,7 +2091,16 @@ class ProxyBaseLLMRequestProcessing: llm_router=llm_router, ) except ProxyRateLimitError as original_exc: - if not fallback_models or pristine is None: + rate_limited_data: Final = self.data + original_model: Final = rate_limited_data.get("model") + if pristine is None or not configured_fallbacks or not isinstance(original_model, str): + raise + + fallback_models: Final = self._resolve_fallback_models( + model=original_model, + fallbacks=configured_fallbacks, + ) + if not fallback_models: raise verbose_proxy_logger.info( @@ -2133,39 +2134,30 @@ class ProxyBaseLLMRequestProcessing: except ProxyRateLimitError: continue except BaseException: - self.data = pristine + self.data = rate_limited_data raise - self.data = pristine + self.data = rate_limited_data raise original_exc - def _resolve_fallback_models( - self, - model: str, - llm_router: Router, - user_api_key_dict: UserAPIKeyAuth, - ) -> list | None: - from litellm.router_utils.fallback_event_handlers import get_fallback_model_group - - fallbacks = None - + @staticmethod + def _configured_fallbacks(llm_router: Router, user_api_key_dict: UserAPIKeyAuth) -> list | None: key_router_settings: Final = user_api_key_dict.router_settings - if isinstance(key_router_settings, dict) and "fallbacks" in key_router_settings: - fallbacks = key_router_settings["fallbacks"] + key_fallbacks: Final = key_router_settings.get("fallbacks") if isinstance(key_router_settings, dict) else None + fallbacks: Final = key_fallbacks if key_fallbacks is not None else llm_router.fallbacks + return fallbacks if isinstance(fallbacks, list) and fallbacks else None - if fallbacks is None: - fallbacks = llm_router.fallbacks - - if not fallbacks: - return None + @staticmethod + def _resolve_fallback_models(model: str, fallbacks: list) -> list | None: + from litellm.router_utils.fallback_event_handlers import get_fallback_model_group 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 + if fallback_model_group is not None: + return fallback_model_group + return fallbacks[generic_fallback_idx]["*"] if generic_fallback_idx is not None else None @staticmethod def _get_model_id_from_response(hidden_params: Mapping[str, object], data: Mapping[str, object]) -> str: diff --git a/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py index 38f087f51ca..459834d0fd2 100644 --- a/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py +++ b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py @@ -48,10 +48,10 @@ class TestSkipPreCallLogic: await processor.base_process_llm_request( request=MagicMock(spec=Request), fastapi_response=MagicMock(spec=Response), - user_api_key_dict=MagicMock(spec=UserAPIKeyAuth, router_settings=None), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), route_type="aresponses", proxy_logging_obj=mock_proxy_logging, - llm_router=MagicMock(fallbacks=None), + llm_router=MagicMock(), general_settings={}, proxy_config=MagicMock(), skip_pre_call_logic=True, @@ -87,10 +87,10 @@ class TestSkipPreCallLogic: await processor.base_process_llm_request( request=MagicMock(spec=Request), fastapi_response=MagicMock(spec=Response), - user_api_key_dict=MagicMock(spec=UserAPIKeyAuth, router_settings=None), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), route_type="aresponses", proxy_logging_obj=mock_proxy_logging, - llm_router=MagicMock(fallbacks=None), + llm_router=MagicMock(), general_settings={}, proxy_config=MagicMock(), ) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 92af65b2637..f689dd62df6 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -6365,247 +6365,170 @@ class TestPreCallWithFallbacksOnLocalRateLimit: call_type="acompletion", ) - @pytest.mark.asyncio - async def test_fallback_retries_from_pristine_request_data(self): - import threading + @staticmethod + def _v3_limiter_rig( + monkeypatch: pytest.MonkeyPatch, + user_api_key_dict: ProxyUserAPIKeyAuth, + fallbacks: list[dict[str, list[str]]], + ) -> tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]]: + """Real v3 limiter (the default ``parallel_request_limiter``) wired in through the + ``proxy_logging_obj`` seam, so ``common_processing_pre_call_logic`` runs for real: + ``add_litellm_data_to_request`` with a live OTel span, ``function_setup``, then the limiter.""" + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy.hooks.parallel_request_limiter_v3 import _PROXY_MaxParallelRequestsHandler_v3 + from litellm.proxy.utils import InternalUsageCache - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + monkeypatch.setattr(proxy_server, "prisma_client", None) + limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache())) + limiter_models: list[str] = [] - primary_model = "gpt-4" - fallback_model = "gpt-3.5-turbo" + async def run_limiter(**kwargs): + limiter_models.append(kwargs["data"]["model"]) + await limiter.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=kwargs["data"], + call_type=kwargs["call_type"], + ) + return kwargs["data"] - processor = ProxyBaseLLMRequestProcessing( - data={ - "model": primary_model, - "messages": [{"role": "user", "content": "hi"}], - "metadata": {"tags": ["a"]}, - } + proxy_logging_obj = MagicMock(spec=ProxyLogging) + proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter) + router = litellm.Router( + model_list=[ + {"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}} + for chain in fallbacks + for group in (*chain.keys(), *(m for models in chain.values() for m in models)) + ], + fallbacks=fallbacks, ) + return proxy_logging_obj, router, proxy_server.ProxyConfig(), limiter_models - metadata_at_entry = [] - - async def mock_pre_call_logic(**kwargs): - copy.deepcopy(processor.data["metadata"]) - metadata_at_entry.append(dict(processor.data["metadata"])) - processor.data["metadata"]["litellm_parent_otel_span"] = threading.RLock() - processor.data["litellm_logging_obj"] = object() - if processor.data.get("model") == primary_model: - raise ProxyRateLimitError( - detail="TPM limit exceeded for gpt-4", - headers={"retry-after": "30"}, - ) - return processor.data, MagicMock() - - mock_router = MagicMock() - mock_router.fallbacks = [{primary_model: [fallback_model]}] - - 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 metadata_at_entry[1] == {"tags": ["a"]} - - @pytest.mark.asyncio - async def test_exhausted_fallbacks_restore_pristine_request_data(self): - import threading - - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError - - primary_model = "gpt-4" - original_data = { - "model": primary_model, - "messages": [{"role": "user", "content": "hi"}], - "metadata": {"tags": ["a"]}, - } - processor = ProxyBaseLLMRequestProcessing(data=copy.deepcopy(original_data)) - - async def mock_pre_call_logic(**kwargs): - processor.data["metadata"]["litellm_parent_otel_span"] = threading.RLock() - processor.data["litellm_logging_obj"] = object() - raise ProxyRateLimitError( - detail=f"TPM limit exceeded for {processor.data.get('model')}", - headers={"retry-after": "30"}, - ) - - mock_router = MagicMock() - mock_router.fallbacks = [{primary_model: ["gpt-3.5-turbo"]}] - - 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=primary_model, - route_type="acompletion", - llm_router=mock_router, - ) - - assert processor.data == original_data - - @pytest.mark.asyncio - async def test_real_add_litellm_data_to_request_rerun_with_otel_span_falls_back(self): - from opentelemetry import trace + @staticmethod + def _otel_key(**limits) -> ProxyUserAPIKeyAuth: from opentelemetry.sdk.trace import TracerProvider - 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.litellm_pre_call_utils import add_litellm_data_to_request - from litellm.proxy.proxy_server import ProxyConfig + span = TracerProvider().get_tracer("test").start_span("proxy-request") + return ProxyUserAPIKeyAuth(api_key="hashed-key", parent_otel_span=span, **limits) - trace.set_tracer_provider(TracerProvider()) + @staticmethod + def _chat_request() -> Request: + return Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []}) - primary_model = "gpt-4" - fallback_model = "gpt-3.5-turbo" - - request_mock = MagicMock(spec=Request) - request_mock.url = MagicMock() - request_mock.url.path = "/v1/chat/completions" - request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" - request_mock.method = "POST" - request_mock.query_params = {} - request_mock.headers = {"Content-Type": "application/json"} - request_mock.client = MagicMock() - request_mock.client.host = "127.0.0.1" - - user_api_key_dict = UserAPIKeyAuth( - parent_otel_span=trace.get_tracer("x").start_span("s"), - api_key="hashed-key", - user_id="u1", - team_id="t1", - metadata={}, - team_metadata={}, - team_member_tpm_limit=1000, + async def _pre_call( + self, + data: dict, + user_api_key_dict: ProxyUserAPIKeyAuth, + rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]], + ) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict, object]]: + proxy_logging_obj, router, proxy_config, _ = rig + processor = ProxyBaseLLMRequestProcessing(data=data) + result = await processor._pre_call_with_fallbacks( + request=self._chat_request(), + general_settings={}, + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=user_api_key_dict, + version=None, + proxy_config=proxy_config, + user_model=None, + user_temperature=None, + user_request_timeout=None, + user_max_tokens=None, + user_api_base=None, + model=None, + route_type="acompletion", + llm_router=router, ) + return processor, result - processor = ProxyBaseLLMRequestProcessing( - data={ + @pytest.mark.asyncio + async def test_v3_limiter_with_otel_span_falls_back_from_client_request(self, monkeypatch: pytest.MonkeyPatch): + """Customer path: OTel on, per-key model RPM cap on the primary, a router fallback configured. + The first pass enriches ``data["metadata"]`` with the live span, then the limiter raises. The + fallback pass must start from the client's request again, so ``add_litellm_data_to_request`` + never deep-copies the span (the ``cannot pickle '_thread.RLock'`` 500).""" + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + key = self._otel_key(metadata={"model_rpm_limit": {primary_model: 1}}) + rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}]) + + def client_request() -> dict: + return { "model": primary_model, "messages": [{"role": "user", "content": "hi"}], - "metadata": {"tags": ["a"]}, + "metadata": {"tags": ["client-tag"]}, } - ) - async def real_add_litellm_data_pre_call(**kwargs): - await add_litellm_data_to_request( - data=processor.data, - request=request_mock, - user_api_key_dict=user_api_key_dict, - proxy_config=ProxyConfig(), + _, (first_data, _) = await self._pre_call(client_request(), key, rig) + processor, (data, logging_obj) = await self._pre_call(client_request(), key, rig) + + assert first_data["model"] == primary_model + assert data["model"] == fallback_model + assert processor.data is data + assert data["litellm_logging_obj"] is logging_obj + assert logging_obj.model == fallback_model + requester_metadata = data["metadata"]["requester_metadata"] + assert requester_metadata["tags"] == ["client-tag"] + assert "litellm_parent_otel_span" not in requester_metadata + assert "user_api_key_auth" not in requester_metadata + assert data["metadata"]["litellm_parent_otel_span"] is key.parent_otel_span + assert rig[3] == [primary_model, primary_model, fallback_model] + + @pytest.mark.asyncio + async def test_v3_limiter_with_otel_span_returns_429_when_fallbacks_exhausted( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + key = self._otel_key(rpm_limit=1) + rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}]) + request = {"model": primary_model, "messages": [{"role": "user", "content": "hi"}]} + + await self._pre_call(dict(request), key, rig) + processor = ProxyBaseLLMRequestProcessing(data=dict(request)) + with pytest.raises(ProxyRateLimitError) as exc_info: + await processor._pre_call_with_fallbacks( + request=self._chat_request(), general_settings={}, - version="test", - ) - if processor.data.get("model") == primary_model: - raise ProxyRateLimitError( - detail="TPM limit exceeded for gpt-4", - headers={"retry-after": "30"}, - ) - return processor.data, MagicMock() - - mock_router = MagicMock() - mock_router.fallbacks = [{primary_model: [fallback_model]}] - - with patch.object( - processor, - "common_processing_pre_call_logic", - side_effect=real_add_litellm_data_pre_call, - ): - data, logging_obj = await processor._pre_call_with_fallbacks( - request=request_mock, - general_settings={}, - proxy_logging_obj=MagicMock(), - user_api_key_dict=user_api_key_dict, + proxy_logging_obj=rig[0], + user_api_key_dict=key, version=None, - proxy_config=MagicMock(), + proxy_config=rig[2], user_model=None, user_temperature=None, user_request_timeout=None, user_max_tokens=None, user_api_base=None, - model=primary_model, + model=None, route_type="acompletion", - llm_router=mock_router, + llm_router=rig[1], ) - assert processor.data["model"] == fallback_model + assert rig[3] == [primary_model, primary_model, fallback_model] + assert exc_info.value.status_code == 429 + assert "Rate limit exceeded" in str(exc_info.value.detail) + assert exc_info.value.headers["retry-after"] + assert processor.data["model"] == primary_model + assert processor.data["litellm_logging_obj"].model == primary_model + assert processor.data["litellm_call_id"] @pytest.mark.asyncio - async def test_no_fallbacks_skips_snapshot(self): - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError + async def test_fallback_lookup_uses_alias_resolved_model_group(self, monkeypatch: pytest.MonkeyPatch): + primary_model = "gpt-4.1" + fallback_model = "gpt-4.1-mini" + monkeypatch.setattr(litellm, "model_alias_map", {"my-alias": primary_model}) + key = self._otel_key(metadata={"model_rpm_limit": {primary_model: 1}}) + rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}]) + request = {"model": "my-alias", "messages": [{"role": "user", "content": "hi"}]} - processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4"}) + await self._pre_call(dict(request), key, rig) + _, (data, _) = await self._pre_call(dict(request), key, rig) - 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( # test-quality-ok: spying the snapshot seam is the only observable check that the no-fallback path skips it - "litellm.proxy.common_request_processing.independent_snapshot" - ) as snapshot_mock: - 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, - ) - - snapshot_mock.assert_not_called() + assert data["model"] == fallback_model + assert rig[3] == [primary_model, primary_model, fallback_model] class _RecordingSuccessLogger(CustomLogger):