Merge pull request #41783 from BerriAI/litellm_rate_limit_fallback_guardrails

fix(proxy): keep requested model guardrails and key disable_fallbacks on rate-limit fallback
This commit is contained in:
yucheng-berri 2026-09-18 14:59:16 -07:00 committed by GitHub
commit 84df4c0d1b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 103 additions and 9 deletions

View file

@ -1934,6 +1934,7 @@ class ProxyBaseLLMRequestProcessing:
user_api_base: str | None = None,
model: str | None = None,
llm_router: Router | None = None,
rate_limited_model: str | None = None,
) -> tuple[dict, LiteLLMLoggingObj]:
start_time: Final = datetime.now() # start before calling guardrail hooks
@ -2097,8 +2098,15 @@ class ProxyBaseLLMRequestProcessing:
# model_info when allow_client_pricing_override is set, so a caller
# could otherwise spoof an unguarded model_info.id while requesting
# a guarded alias and bypass guardrails (veria-ai HIGH on #29654).
merged_for_requested: Final = (
self.data
if rate_limited_model is None
else _check_and_merge_model_level_guardrails(
data=self.data, llm_router=llm_router, trust_client_model_info=False, model_alias=rate_limited_model
)
)
self.data = _check_and_merge_model_level_guardrails(
data=self.data,
data=merged_for_requested,
llm_router=llm_router,
trust_client_model_info=False,
)
@ -2163,7 +2171,7 @@ class ProxyBaseLLMRequestProcessing:
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")
if llm_router is not None
else None
)
pristine: Final = independent_snapshot(self.data) if configured_fallbacks else None
@ -2208,7 +2216,6 @@ class ProxyBaseLLMRequestProcessing:
original_model,
fallback_models,
)
try:
for fallback_model in fallback_models:
if fallback_model == original_model:
@ -2231,6 +2238,7 @@ class ProxyBaseLLMRequestProcessing:
model=fallback_model,
route_type=route_type,
llm_router=llm_router,
rate_limited_model=original_model,
)
except ProxyRateLimitError:
continue

View file

@ -7497,6 +7497,7 @@ def _check_and_merge_model_level_guardrails(
data: dict,
llm_router: Router | None,
trust_client_model_info: bool = True,
model_alias: str | None = None,
) -> dict:
"""
Check if the model has guardrails defined and merge them with existing guardrails in the request data.
@ -7504,6 +7505,7 @@ def _check_and_merge_model_level_guardrails(
Args:
data: The request data dict
llm_router: The LLM router instance to get deployment info from
model_alias: Resolve guardrails for this model group instead of data["model"]
trust_client_model_info: If False, ignore metadata.model_info.id and
resolve guardrails by alias-union only. Set to False on the
pre_call path because add_litellm_data_to_request preserves
@ -7548,13 +7550,13 @@ def _check_and_merge_model_level_guardrails(
# set on ANY eligible deployment still fires (#29652; addresses
# veria-ai HIGH on the single-deployment fallback that would skip
# non-first deployments).
model_alias: Final = data.get("model")
if not isinstance(model_alias, str) or not model_alias:
alias: Final = model_alias if model_alias is not None else data.get("model")
if not isinstance(alias, str) or not alias:
return data
# Pass team_id so team-scoped public model names resolve the same way
# route_request resolves them; otherwise team-scoped deployments are
# invisible to this lookup and their guardrails are silently dropped.
deployments: Final = llm_router.get_model_list(model_name=model_alias, team_id=team_id) or []
deployments: Final = llm_router.get_model_list(model_name=alias, team_id=team_id) or []
seen: Final[set] = set()
union: Final[list] = []
for dep in deployments:

View file

@ -6711,6 +6711,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
monkeypatch: pytest.MonkeyPatch,
user_api_key_dict: ProxyUserAPIKeyAuth,
fallbacks: list[dict[str, list[str]]],
model_guardrails: dict[str, list[str]] | None = None,
) -> 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:
@ -6738,9 +6739,17 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=run_limiter)
guardrails_by_group = model_guardrails or {}
router = litellm.Router(
model_list=[
{"model_name": group, "litellm_params": {"model": "openai/gpt-4.1-nano", "api_key": "fake"}}
{
"model_name": group,
"litellm_params": {
"model": "openai/gpt-4.1-nano",
"api_key": "fake",
**({"guardrails": guardrails_by_group[group]} if group in guardrails_by_group else {}),
},
}
for chain in fallbacks
for group in (*chain.keys(), *(m for models in chain.values() for m in models))
],
@ -6752,7 +6761,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
def _otel_key(
rpm_limit: int | None = None,
model_rpm_limit: dict[str, int] | None = None,
disable_fallbacks: bool = False,
disable_fallbacks: bool | None = None,
) -> ProxyUserAPIKeyAuth:
from opentelemetry.sdk.trace import TracerProvider
@ -6763,7 +6772,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
rpm_limit=rpm_limit,
metadata={
**({"model_rpm_limit": model_rpm_limit} if model_rpm_limit else {}),
**({"disable_fallbacks": True} if disable_fallbacks else {}),
**({"disable_fallbacks": disable_fallbacks} if disable_fallbacks is not None else {}),
},
)
@ -6906,6 +6915,81 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
assert exc_info.value.status_code == 429
assert rig[3] == [primary_model, primary_model]
@pytest.mark.asyncio
async def test_key_metadata_disable_fallbacks_false_overrides_request_body(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
key = self._otel_key(model_rpm_limit={primary_model: 1}, disable_fallbacks=False)
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"disable_fallbacks": True,
}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert data["disable_fallbacks"] is False
assert rig[3] == [primary_model, primary_model, fallback_model]
@pytest.mark.asyncio
async def test_fallback_keeps_requested_model_guardrails(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
guardrail = "pii-guard-for-primary"
key = self._otel_key(model_rpm_limit={primary_model: 1})
rig = self._v3_limiter_rig(
monkeypatch, key, [{primary_model: [fallback_model]}], model_guardrails={primary_model: [guardrail]}
)
run_limiter = rig[0].pre_call_hook
async def limiter_then_guardrail(
user_api_key_dict: ProxyUserAPIKeyAuth, data: dict[str, object], call_type: str
) -> dict[str, object]:
limited = await run_limiter(user_api_key_dict=user_api_key_dict, data=data, call_type=call_type)
if guardrail not in (limited["metadata"].get("guardrails") or []):
return limited
return {
**limited,
"messages": [
{**m, "content": str(m["content"]).replace("123-45-6789", "[REDACTED-SSN]")}
for m in limited["messages"]
],
}
rig[0].pre_call_hook = AsyncMock(side_effect=limiter_then_guardrail)
request = {"model": primary_model, "messages": [{"role": "user", "content": "my ssn is 123-45-6789"}]}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert guardrail in data["metadata"]["guardrails"]
assert data["messages"] == [{"role": "user", "content": "my ssn is [REDACTED-SSN]"}]
assert rig[3] == [primary_model, primary_model, fallback_model]
@pytest.mark.asyncio
async def test_fallback_keeps_structured_request_guardrails(self, monkeypatch: pytest.MonkeyPatch):
primary_model = "gpt-4.1"
fallback_model = "gpt-4.1-mini"
structured_guardrail = {"pii-guard": {"extra_body": {"threshold": 0.5}}}
key = self._otel_key(model_rpm_limit={primary_model: 1})
rig = self._v3_limiter_rig(monkeypatch, key, [{primary_model: [fallback_model]}])
request = {
"model": primary_model,
"messages": [{"role": "user", "content": "hi"}],
"guardrails": [structured_guardrail],
}
await self._pre_call(dict(request), key, rig)
_, (data, _) = await self._pre_call(dict(request), key, rig)
assert data["model"] == fallback_model
assert data["metadata"]["guardrails"] == [structured_guardrail]
assert rig[3] == [primary_model, primary_model, fallback_model]
class _RecordingSuccessLogger(CustomLogger):
def __init__(self):