mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
fix(proxy): retry rate-limit fallbacks from a pristine request snapshot
Backport of #40596 to rc/1.102.0.
Cherry-picked from merge commit c5325b1492 (main), originally by app/devin-ai-integration.
This commit is contained in:
parent
e6f29fbe9f
commit
8f8b47d6f0
3 changed files with 237 additions and 28 deletions
|
|
@ -32,7 +32,11 @@ from litellm.constants import (
|
|||
UNSAFE_PROXY_RESPONSE_HEADERS,
|
||||
)
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket, is_expected_client_error
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
get_or_create_metadata_bucket,
|
||||
independent_snapshot,
|
||||
is_expected_client_error,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import NullTracer, tracer
|
||||
from litellm.litellm_core_utils.get_supported_openai_params import (
|
||||
get_supported_openai_params,
|
||||
|
|
@ -2034,6 +2038,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
) -> tuple[dict, LiteLLMLoggingObj]:
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
|
||||
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 configured_fallbacks else None
|
||||
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
request=request,
|
||||
|
|
@ -2052,14 +2063,19 @@ class ProxyBaseLLMRequestProcessing:
|
|||
llm_router=llm_router,
|
||||
)
|
||||
except ProxyRateLimitError as original_exc:
|
||||
original_model: Final = self.data.get("model")
|
||||
if not original_model or not llm_router or self.data.get("disable_fallbacks"):
|
||||
rate_limited_data: Final = self.data
|
||||
original_model: Final = rate_limited_data.get("model")
|
||||
if (
|
||||
pristine is None
|
||||
or not configured_fallbacks
|
||||
or rate_limited_data.get("disable_fallbacks")
|
||||
or not isinstance(original_model, str)
|
||||
):
|
||||
raise
|
||||
|
||||
fallback_models: Final = self._resolve_fallback_models(
|
||||
model=original_model,
|
||||
llm_router=llm_router,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
fallbacks=configured_fallbacks,
|
||||
)
|
||||
if not fallback_models:
|
||||
raise
|
||||
|
|
@ -2074,6 +2090,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
for fallback_model in fallback_models:
|
||||
if fallback_model == original_model:
|
||||
continue
|
||||
self.data = independent_snapshot(pristine)
|
||||
self.data["model"] = fallback_model
|
||||
try:
|
||||
return await self.common_processing_pre_call_logic(
|
||||
|
|
@ -2095,39 +2112,30 @@ class ProxyBaseLLMRequestProcessing:
|
|||
except ProxyRateLimitError:
|
||||
continue
|
||||
except BaseException:
|
||||
self.data["model"] = original_model
|
||||
self.data = rate_limited_data
|
||||
raise
|
||||
|
||||
self.data["model"] = original_model
|
||||
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:
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ class TestSkipPreCallLogic:
|
|||
await processor.base_process_llm_request(
|
||||
request=MagicMock(spec=Request),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
route_type="aresponses",
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
llm_router=MagicMock(),
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.proxy.common_request_processing import (
|
|||
create_response,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy._types import UserAPIKeyAuth as ProxyUserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
|
@ -6235,6 +6236,206 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
|
|||
call_type="acompletion",
|
||||
)
|
||||
|
||||
@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
|
||||
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
limiter_models: list[str] = []
|
||||
|
||||
async def run_limiter(
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str
|
||||
) -> dict[str, object]:
|
||||
limiter_models.append(str(data["model"]))
|
||||
await limiter.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type=call_type,
|
||||
)
|
||||
return data
|
||||
|
||||
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
|
||||
|
||||
@staticmethod
|
||||
def _otel_key(
|
||||
rpm_limit: int | None = None,
|
||||
model_rpm_limit: dict[str, int] | None = None,
|
||||
disable_fallbacks: bool = False,
|
||||
) -> ProxyUserAPIKeyAuth:
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
||||
span = TracerProvider().get_tracer("test").start_span("proxy-request")
|
||||
return ProxyUserAPIKeyAuth(
|
||||
api_key="hashed-key",
|
||||
parent_otel_span=span,
|
||||
rpm_limit=rpm_limit,
|
||||
metadata={
|
||||
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
|
||||
**({"disable_fallbacks": True} if disable_fallbacks else {}),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _chat_request() -> Request:
|
||||
return Request({"type": "http", "method": "POST", "path": "/v1/chat/completions", "headers": []})
|
||||
|
||||
async def _pre_call(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: ProxyUserAPIKeyAuth,
|
||||
rig: tuple[ProxyLogging, litellm.Router, ProxyConfig, list[str]],
|
||||
) -> tuple[ProxyBaseLLMRequestProcessing, tuple[dict[str, object], LiteLLMLoggingObj]]:
|
||||
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
|
||||
|
||||
@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(model_rpm_limit={primary_model: 1})
|
||||
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
|
||||
|
||||
def client_request() -> dict[str, object]:
|
||||
return {
|
||||
"model": primary_model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"metadata": {"tags": ["client-tag"]},
|
||||
}
|
||||
|
||||
_, (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={},
|
||||
proxy_logging_obj=rig[0],
|
||||
user_api_key_dict=key,
|
||||
version=None,
|
||||
proxy_config=rig[2],
|
||||
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=rig[1],
|
||||
)
|
||||
|
||||
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_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(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"}]}
|
||||
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
_, (data, _) = await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert data["model"] == fallback_model
|
||||
assert rig[3] == [primary_model, primary_model, fallback_model]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_metadata_disable_fallbacks_returns_429_instead_of_retrying(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""``disable_fallbacks`` set in key metadata only lands on ``data`` during the first
|
||||
pre-call pass (``add_key_level_controls``), so it must be honored after that pass."""
|
||||
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(model_rpm_limit={primary_model: 1}, disable_fallbacks=True)
|
||||
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)
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
await self._pre_call(dict(request), key, rig)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert rig[3] == [primary_model, primary_model]
|
||||
|
||||
|
||||
class _RecordingSuccessLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue