fix(rate-limiting): stop apply_to_models fallback bypass on tag rejection

veria-ai finding on this PR: a global entry scoped with apply_to_models is
meant to cap an entire fallback chain as one shared unit, but the proxy's
generic local-rate-limit fallback retry (_pre_call_with_fallbacks) caught
that rejection and quietly served the request via any fallback model not
also listed in apply_to_models, defeating the whole point of the field.

global_tag_rate_limits_hook now marks a rejection raised by an
apply_to_models-scoped entry with detail["cross_model_scope"], and
_pre_call_with_fallbacks re-raises immediately on that marker instead of
retrying fallbacks. Every other ProxyRateLimitError caller (parallel
request limiter, budget limiters, model-local tag limits, etc.) is
untouched, since none of them ever set this marker.

Live-verified against a real proxy: an opus-chain to sonnet-chain fallback
configured with apply_to_models: [opus-chain] now gets a 429 on both legs
once the chain-wide cap is hit, instead of silently succeeding via
sonnet-chain.
This commit is contained in:
Deepanshu 2026-08-25 16:45:19 -04:00
parent 6e6226906c
commit 57cb129be9
3 changed files with 68 additions and 1 deletions

View file

@ -2003,7 +2003,10 @@ class ProxyBaseLLMRequestProcessing:
)
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"):
cross_model_scope: Final = (
isinstance(original_exc.detail, Mapping) and original_exc.detail.get("cross_model_scope") is True
)
if not original_model or not llm_router or self.data.get("disable_fallbacks") or cross_model_scope:
raise
fallback_models: Final = self._resolve_fallback_models(

View file

@ -487,6 +487,11 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
"limit_name": entry.name,
"limit": entry.limit,
"period_seconds": entry.period_seconds,
# ProxyBaseLLMRequestProcessing._pre_call_with_fallbacks reads this:
# an apply_to_models entry caps an entire named chain as one unit, so
# retrying against a fallback model outside that list would silently
# defeat the very policy that just rejected this request.
**({"cross_model_scope": True} if entry.apply_to_models is not None else {}),
},
headers={"retry-after": str(entry.period_seconds)}, # mutable-ok: same as detail
rate_limit_type=_UNIT_TO_RATE_LIMIT_TYPE[unit],

View file

@ -5476,6 +5476,65 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
assert processor.data["model"] == primary_model
@pytest.mark.asyncio
async def test_cross_model_scoped_rejection_is_not_retried_via_fallback(self):
"""
veria-ai finding on PR #36541: an entry using ``apply_to_models`` to cap
an entire fallback chain as one unit is defeated by this exact mechanism
if a fallback model isn't also listed in ``apply_to_models`` -- the
rejection here is a deliberate "this whole chain is capped" decision,
not a "this one model is unhealthy" signal, so retrying against an
unlisted fallback silently serves a request the operator's policy meant
to block. ``detail["cross_model_scope"]`` is the marker
global_tag_rate_limits_hook sets for exactly this case; the fallback
handler must re-raise immediately instead of trying any fallback model.
"""
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
primary_model = "opus-chain"
processor = ProxyBaseLLMRequestProcessing(data={"model": primary_model})
call_count = 0
async def mock_pre_call_logic(**kwargs):
nonlocal call_count
call_count += 1
raise ProxyRateLimitError(
detail={"error": "tag_rate_limit_exceeded", "cross_model_scope": True},
headers={"retry-after": "30"},
)
mock_router = MagicMock()
mock_router.fallbacks = [{"opus-chain": ["sonnet-chain"]}]
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=primary_model,
route_type="acompletion",
llm_router=mock_router,
)
assert call_count == 1
assert processor.data["model"] == primary_model
@pytest.mark.asyncio
async def test_real_parallel_request_limiter_model_tpm_limit_triggers_fallback(self):
"""