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>
This commit is contained in:
yucheng 2026-09-16 07:08:24 +00:00
parent 2e699914e1
commit 2b6184d768
3 changed files with 166 additions and 251 deletions

View file

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

View file

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

View file

@ -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):