mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
fix(proxy): keep requested model guardrails and key disable_fallbacks on rate-limit fallback
The local rate-limit fallback path introduced in #40596 restores the request from a snapshot taken before the pre-call pass, so the guardrails resolved for the requested model were dropped when another deployment was selected, and the raw request-body disable_fallbacks field gated the fallback before the key-level disable_fallbacks override had been applied Carry the requested model's merged guardrail list onto every fallback pass and read the effective disable_fallbacks value after the pre-call pass Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
c256c3c1a6
commit
a187a5bfc6
2 changed files with 87 additions and 6 deletions
|
|
@ -94,7 +94,11 @@ from litellm.proxy.common_utils.sse_keepalive import (
|
|||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
|
||||
from litellm.proxy.utils import (
|
||||
ProxyLogging,
|
||||
_check_and_merge_model_level_guardrails,
|
||||
_merge_guardrails_with_existing,
|
||||
)
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict
|
||||
from litellm.router_utils.common_utils import resolve_model_group_alias
|
||||
|
|
@ -1934,6 +1938,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_base: str | None = None,
|
||||
model: str | None = None,
|
||||
llm_router: Router | None = None,
|
||||
requested_model_guardrails: list | None = None,
|
||||
) -> tuple[dict, LiteLLMLoggingObj]:
|
||||
start_time: Final = datetime.now() # start before calling guardrail hooks
|
||||
|
||||
|
|
@ -2098,7 +2103,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# could otherwise spoof an unguarded model_info.id while requesting
|
||||
# a guarded alias and bypass guardrails (veria-ai HIGH on #29654).
|
||||
self.data = _check_and_merge_model_level_guardrails(
|
||||
data=self.data,
|
||||
data=(
|
||||
self.data
|
||||
if requested_model_guardrails is None
|
||||
else _merge_guardrails_with_existing(data=self.data, model_level_guardrails=requested_model_guardrails)
|
||||
),
|
||||
llm_router=llm_router,
|
||||
trust_client_model_info=False,
|
||||
)
|
||||
|
|
@ -2163,7 +2172,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,6 +2217,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
original_model,
|
||||
fallback_models,
|
||||
)
|
||||
requested_model_guardrails: Final = self._request_guardrails(rate_limited_data)
|
||||
|
||||
try:
|
||||
for fallback_model in fallback_models:
|
||||
|
|
@ -2231,6 +2241,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
model=fallback_model,
|
||||
route_type=route_type,
|
||||
llm_router=llm_router,
|
||||
requested_model_guardrails=requested_model_guardrails,
|
||||
)
|
||||
except ProxyRateLimitError:
|
||||
continue
|
||||
|
|
@ -2241,6 +2252,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
self.data = rate_limited_data
|
||||
raise original_exc
|
||||
|
||||
@staticmethod
|
||||
def _request_guardrails(data: dict) -> list | None:
|
||||
metadata: Final = data.get("metadata")
|
||||
guardrails: Final = metadata.get("guardrails") if isinstance(metadata, dict) else None
|
||||
return guardrails if isinstance(guardrails, list) else 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
|
||||
|
|
|
|||
|
|
@ -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,61 @@ 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]
|
||||
|
||||
|
||||
class _RecordingSuccessLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue