From ba4bde4fcd2328db04a8f8be304da4dd4674c579 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Tue, 25 Aug 2026 15:02:40 -0400 Subject: [PATCH] feat(rate-limiting): add apply_to_models scoping to tag rate limit entries Lets a global_tag_rate_limits (or per-deployment model_info.tag_rate_limits) entry restrict itself to a named list of caller-facing model strings, so one entry can rate-limit an entire fallback chain as a single shared bucket instead of requiring the per-hop model_based_tag_rate_limits_hook mechanism. apply_to_models composes with the existing apply_to_key_alias/enabled_for/ disabled_for gates, all AND together. It is evaluated once, at admission, against the caller-requested model, and is not re-evaluated if Router later falls back to an unlisted model for the same request; this is a documented limitation, not a bug, and is covered by a dedicated regression test. --- .../hooks/global_tag_rate_limits_hook.py | 35 ++- .../hooks/model_based_tag_rate_limits_hook.py | 39 +-- litellm/types/router.py | 18 ++ .../hooks/test_global_tag_rate_limits_hook.py | 255 ++++++++++++++++++ .../test_model_based_tag_rate_limits_hook.py | 230 ++++++++++++++-- 5 files changed, 536 insertions(+), 41 deletions(-) diff --git a/litellm/proxy/hooks/global_tag_rate_limits_hook.py b/litellm/proxy/hooks/global_tag_rate_limits_hook.py index de5d181adf4..152fe257b9c 100644 --- a/litellm/proxy/hooks/global_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/global_tag_rate_limits_hook.py @@ -17,11 +17,21 @@ them, but implements its own, much smaller admission/accounting engine: no `_LimitsIndex`, no routing-group or team-alias resolution, no per-deployment dedup signatures. -Two independent entry-level knobs decide who a global entry applies to and +Three independent entry-level knobs decide who a global entry applies to and how its bucket is shared: - `apply_to_key_alias`: unset means every request, any key, any model. Set to a list of virtual-key aliases, only those keys' requests count. +- `apply_to_models`: unset means every model. Set to a list of model names, + only requests whose caller-facing `model` field is in that list count -- + letting one entry rate-limit a whole fallback chain as a single unit by + naming every model in the chain. This is evaluated once, against the + caller-requested `model`, before Router does any routing: if the request's + own model fails and Router falls back to a model not in `apply_to_models`, + that fallback hop is not re-evaluated -- the original admission already + stands. An operator who needs the limit to track whichever model actually + ends up serving a request, including after a fallback, needs + `model_info.tag_rate_limits` instead. - `scope_by_key_hash` (already exists on `TagRateLimitEntry`): whether the keys an entry applies to share one bucket, or each gets its own. @@ -159,6 +169,11 @@ class _GlobalTagRateLimitStash: """ admission_time: float | None = None + # The caller-facing `model` admission read from `data.get("model")`, so + # async_log_success_event's tokens/dollars accounting gates + # apply_to_models against the same, originally-requested model admission + # decided on -- not whatever model a later fallback actually served. + model: str | None = None pending_concurrency_keys: list[tuple[str, _PartitionKey]] = field(default_factory=list) # mutable-ok: queue @@ -361,7 +376,13 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o return _bucket_ttl_seconds(entry) def _classify( - self, config: TagRateLimits, tags: Sequence[str], key_alias: str | None, key_hash: str | None, now: float + self, + config: TagRateLimits, + tags: Sequence[str], + key_alias: str | None, + key_hash: str | None, + now: float, + model: str | None, ) -> tuple[_ClassifiedGlobalCheck, ...]: classified: Final = [] # mutable-ok: sequential accumulator, immediately frozen into a tuple below for unit in _LIMIT_UNITS: @@ -372,7 +393,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o tag_value = _extract_identity(tags, entry.tag_id) if tag_value is None: continue - if not _entry_applies(entry, tags, key_alias): + if not _entry_applies(entry, tags, key_alias, model): continue effective_key_hash = key_hash if entry.scope_by_key_hash else None if unit == "concurrency": @@ -481,17 +502,18 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o tags: Final = _get_tags_from_request_kwargs(data, metadata_variable_name=metadata_variable_name) key_alias: Final = user_api_key_dict.key_alias key_hash: Final = user_api_key_dict.api_key + model: Final = data.get("model") if isinstance(data.get("model"), str) else None now: Final = self._time_provider().timestamp() stash.admission_time = now - classified: Final = self._classify(config, tags, key_alias, key_hash, now) + stash.model = model + classified: Final = self._classify(config, tags, key_alias, key_hash, now, model) if not classified: return data read_only_checks: Final = tuple(c for c in classified if not c.is_atomic) atomic_checks: Final = tuple(c for c in classified if c.is_atomic) - model: Final = data.get("model") if isinstance(data.get("model"), str) else None current_values: Final = await self._read_only_values(read_only_checks, parent_otel_span=None) self._raise_if_over_limit(read_only_checks, current_values, model) @@ -588,6 +610,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o if stash is not None and stash.admission_time is not None else self._time_provider().timestamp() ) + model: Final = stash.model if stash is not None else None increment_by_unit: Final[Mapping[_LimitUnit, float]] = MappingProxyType( { "tokens": float(standard_logging_object.get("total_tokens") or 0), @@ -604,7 +627,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o tag_value = _extract_identity(tags, entry.tag_id) if tag_value is None: continue - if not _entry_applies(entry, tags, key_alias): + if not _entry_applies(entry, tags, key_alias, model): continue increment_value = increment_by_unit[unit] if increment_value == 0: diff --git a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index f328fb24231..3c1c6ef9d6e 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -49,13 +49,13 @@ _LIMIT_UNITS: Final[tuple[_LimitUnit, ...]] = ("tokens", "requests", "dollars", # depending on TagRateLimitScope's own hashability. _ScopeSignature: TypeAlias = tuple[str, tuple[str, ...]] | None # (tag_id, name, limit, period_seconds, scope_by_key_hash, enabled_for, -# disabled_for, apply_to_key_alias) -- the fields that decide whether two -# deployments' entries are the same rate limit for dedup purposes; see -# _build_group_limits. Two deployments that agree on the first five but -# disagree on any scoping field are declaring genuinely different policies -# (e.g. one excludes a user the other doesn't) and must not be merged into -# one shared bucket -- the same class of bug this signature already guards -# against for a plain divergent `limit`. +# disabled_for, apply_to_key_alias, apply_to_models) -- the fields that +# decide whether two deployments' entries are the same rate limit for dedup +# purposes; see _build_group_limits. Two deployments that agree on the first +# five but disagree on any scoping field are declaring genuinely different +# policies (e.g. one excludes a user the other doesn't) and must not be +# merged into one shared bucket -- the same class of bug this signature +# already guards against for a plain divergent `limit`. _DedupSignature: TypeAlias = tuple[ str, str, @@ -65,6 +65,7 @@ _DedupSignature: TypeAlias = tuple[ _ScopeSignature, _ScopeSignature, tuple[str, ...] | None, + tuple[str, ...] | None, ] # Units whose admission must be atomic (check-and-increment in one Redis # round trip) because the increment amount is known upfront (always 1). @@ -213,11 +214,11 @@ def _scope_signature(scope: TagRateLimitScope | None) -> _ScopeSignature: return None if scope is None else (scope.tag_id, scope.values) -def _entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str | None) -> bool: +def _entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str | None, model: str | None) -> bool: """ Applies `entry`'s own scoping fields (`enabled_for`/`disabled_for`/ - `apply_to_key_alias`), evaluated in this order -- deny overrides allow, - checked before either allowlist: + `apply_to_key_alias`/`apply_to_models`), evaluated in this order -- deny + overrides allow, checked before any allowlist: 1. `disabled_for`: the gate tag (often a SECOND, independent tag, but `disabled_for.tag_id` can equally be set to this entry's own @@ -229,7 +230,10 @@ def _entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str NOT in `enabled_for.values` -> doesn't apply. Unlike `disabled_for`, absence DOES fail this check -- an allowlist gate requires an explicit match, so "not tagged at all" means "not in scope". - 3. `apply_to_key_alias`: the calling key's own alias is absent, or + 3. `apply_to_models`: `model` is absent, or present but not in the + list -> doesn't apply. Same allowlist semantics as `enabled_for` -- + a request with no `model` never satisfies this gate. + 4. `apply_to_key_alias`: the calling key's own alias is absent, or present but not in the list -> doesn't apply. Same allowlist semantics as `enabled_for` -- a key with no alias set never satisfies this gate. @@ -245,6 +249,8 @@ def _entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str enabled_gate_value: Final = _extract_identity(tags, entry.enabled_for.tag_id) if enabled_gate_value is None or enabled_gate_value not in entry.enabled_for.values: return False + if entry.apply_to_models is not None and model not in entry.apply_to_models: + return False if entry.apply_to_key_alias is None: return True return key_alias in entry.apply_to_key_alias @@ -397,6 +403,7 @@ def _build_group_limits(deployments: Sequence[Mapping[str, object]], unit: _Limi _scope_signature(entry.enabled_for), _scope_signature(entry.disabled_for), entry.apply_to_key_alias, + entry.apply_to_models, ) ids_for_signature = declaring_ids_by_signature.setdefault(signature, []) # mutable-ok: see comment above # One deployment declaring the identical entry twice (a config @@ -511,6 +518,7 @@ class _LimitsIndex: _scope_signature(limit.entry.enabled_for), _scope_signature(limit.entry.disabled_for), limit.entry.apply_to_key_alias, + limit.entry.apply_to_models, limit.deployment_scope, limit.team_scope, ) @@ -797,7 +805,7 @@ def _policy_fingerprint(entry: TagRateLimitEntry) -> str: unscoped entry with a key-hash-scoped one that agrees on every other field. Hashed to a fixed-length digest for the same reason `_fixed_length_identity` hashes `tag_value`: an operator's own - `enabled_for`/`disabled_for`/`apply_to_key_alias` list has no length bound. + `enabled_for`/`disabled_for`/`apply_to_key_alias`/`apply_to_models` list has no length bound. """ fingerprint_source: Final = ( entry.limit, @@ -806,6 +814,7 @@ def _policy_fingerprint(entry: TagRateLimitEntry) -> str: _scope_signature(entry.enabled_for), _scope_signature(entry.disabled_for), entry.apply_to_key_alias, + entry.apply_to_models, ) return hashlib.sha256(repr(fingerprint_source).encode()).hexdigest()[:16] @@ -881,7 +890,7 @@ def _classify_check( tag_value: Final = _extract_identity(tags, configured_limit.entry.tag_id) if tag_value is None: return None - if not _entry_applies(configured_limit.entry, tags, key_alias): + if not _entry_applies(configured_limit.entry, tags, key_alias, model): return None key_hash: Final = ( _extract_key_hash(request_kwargs, metadata_variable_name) if configured_limit.entry.scope_by_key_hash else None @@ -920,7 +929,7 @@ def _increment_operation_for_limit( tag_value: Final = _extract_identity(tags, configured_limit.entry.tag_id) if tag_value is None: return None - if not _entry_applies(configured_limit.entry, tags, key_alias): + if not _entry_applies(configured_limit.entry, tags, key_alias, model_group): return None if configured_limit.unit not in increment_by_unit: return None # "requests" is accounted atomically at admission, not here @@ -971,7 +980,7 @@ def _resolve_max_in_memory_cache_size() -> int | None: # entries can share tag_id/name (see # test_bucket_key_differs_for_same_named_entries_with_different_scoping_only) # while disagreeing on limit/period_seconds/scope_by_key_hash/enabled_for/ -# disabled_for/apply_to_key_alias -- _DedupSignature already treats that as +# disabled_for/apply_to_key_alias/apply_to_models -- _DedupSignature already treats that as # two distinct policies, so a shared max_in_memory_cache_size must not route # them onto the same partition either, or one entry's high-cardinality # traffic can evict the other's active counters from a cache neither entry diff --git a/litellm/types/router.py b/litellm/types/router.py index 28f30e7e0ee..7ab53431755 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -214,6 +214,12 @@ class TagRateLimitEntry(BaseModel): # never satisfies this allowlist, same "absent gate never matches an # allowlist" precedent as `enabled_for`. apply_to_key_alias: tuple[str, ...] | None = None + # Restrict this entry to requests whose caller-facing `model` matches one + # of these names. Unset (the default) means the entry applies to every + # model. A request with no `model` field never satisfies this allowlist, + # same "absent gate never matches an allowlist" precedent as + # `apply_to_key_alias`. + apply_to_models: tuple[str, ...] | None = None model_config = ConfigDict(protected_namespaces=()) @@ -282,6 +288,18 @@ class TagRateLimitEntry(BaseModel): self.apply_to_key_alias = tuple(sorted(set(self.apply_to_key_alias))) # mutable-ok: frozen before escaping return self + @model_validator(mode="after") + def _validate_apply_to_models(self) -> "TagRateLimitEntry": + if self.apply_to_models is not None and not self.apply_to_models: + raise ValueError("apply_to_models must be a non-empty list of strings when set") + return self + + @model_validator(mode="after") + def _normalize_apply_to_models(self) -> "TagRateLimitEntry": + if self.apply_to_models is not None: + self.apply_to_models = tuple(sorted(set(self.apply_to_models))) # mutable-ok: frozen before escaping + return self + class TagRateLimitGroup(BaseModel): limits: tuple[TagRateLimitEntry, ...] = () diff --git a/tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py index 2dbe0a9bd48..c8be1d2e778 100644 --- a/tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py @@ -254,6 +254,261 @@ async def test_apply_to_key_alias_composes_with_scope_by_key_hash(time_controlle assert result is not None +# --------------------------------------------------------------------------- +# apply_to_models -- narrows which requested model an entry applies to, +# letting one entry rate-limit a whole fallback chain as a single unit +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_apply_to_models_ignores_non_matching_model(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "chain_cap", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_models": ["opus-chain"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + for i in range(3): + data = {**_data(["end_user_id:u1"], call_id=f"call-{i}"), "model": "sonnet-chain"} + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion" + ) + assert result is not None + + +@pytest.mark.asyncio +async def test_apply_to_models_enforces_for_the_listed_model(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "chain_cap", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_models": ["opus-chain"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + data1 = {**_data(["end_user_id:u1"], call_id="call-1"), "model": "opus-chain"} + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data1, call_type="completion") + data2 = {**_data(["end_user_id:u1"], call_id="call-2"), "model": "opus-chain"} + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data2, call_type="completion") + + +@pytest.mark.asyncio +async def test_apply_to_models_shares_one_bucket_across_every_listed_model(time_controller, monkeypatch): + """The core "rate limit the whole chain" use case: a single limit shared + across every model named in apply_to_models, not one bucket per model.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "chain_cap", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_models": ["opus-chain", "sonnet-chain"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + data1 = {**_data(["end_user_id:u1"], call_id="call-1"), "model": "opus-chain"} + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data1, call_type="completion") + + data2 = {**_data(["end_user_id:u1"], call_id="call-2"), "model": "sonnet-chain"} + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data2, call_type="completion") + + +@pytest.mark.asyncio +async def test_apply_to_models_composes_with_apply_to_key_alias(time_controller, monkeypatch): + """Both gates must pass -- the listed key requesting a non-listed model + is unaffected, and only the listed key requesting the listed model is + actually enforced.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "chain_cap", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_models": ["opus-chain"], + "apply_to_key_alias": ["premium-key"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + # premium-key requesting a non-listed model: apply_to_models alone must + # still exclude it, even though apply_to_key_alias matches. + for i in range(3): + data = {**_data(["end_user_id:u1"], call_id=f"wrong-model-{i}"), "model": "sonnet-chain"} + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), cache=DualCache(), data=data, call_type="completion" + ) + assert result is not None + + # A non-listed key requesting the listed model: apply_to_key_alias alone + # must still exclude it, even though apply_to_models matches. + for i in range(3): + data = {**_data(["end_user_id:u1"], call_id=f"wrong-key-{i}"), "model": "opus-chain"} + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="other-key"), cache=DualCache(), data=data, call_type="completion" + ) + assert result is not None + + # Both gates match: enforced. + data1 = {**_data(["end_user_id:u1"], call_id="call-1"), "model": "opus-chain"} + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), cache=DualCache(), data=data1, call_type="completion" + ) + data2 = {**_data(["end_user_id:u1"], call_id="call-2"), "model": "opus-chain"} + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), cache=DualCache(), data=data2, call_type="completion" + ) + + +@pytest.mark.asyncio +async def test_dollar_limit_respects_apply_to_models_at_accounting_time(time_controller, monkeypatch): + """The entry only applies to opus-chain; a non-listed model's spend must + not be charged against this bucket at all -- proves apply_to_models + gates async_log_success_event's tokens/dollars accounting, not just + admission.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "dollar_limits": { + "limits": [ + { + "name": "chain_spend", + "tag_id": "end_user_id", + "limit": 10.0, + "period_seconds": 86400, + "apply_to_models": ["opus-chain"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + data = {**_data(["end_user_id:u1"], call_id="call-1"), "model": "sonnet-chain"} + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion") + kwargs = { + "litellm_call_id": "call-1", + "metadata": {"tags": ["end_user_id:u1"]}, + "standard_logging_object": {"total_tokens": 0, "response_cost": 999.0}, + } + await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) + await asyncio.sleep(0) + + # opus-chain was never charged -- still fully under its own limit. + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data={**_data(["end_user_id:u1"], call_id="call-2"), "model": "opus-chain"}, + call_type="completion", + ) + assert result is not None + + +@pytest.mark.asyncio +async def test_apply_to_models_fallback_does_not_re_narrow_accounting_to_the_serving_model(time_controller, monkeypatch): + """ + Documented limitation, not a bug: apply_to_models is evaluated exactly + once, at admission, against the caller-requested model -- it is never + re-evaluated against whichever model a later fallback actually serves. + This request names "opus-chain" at admission (the entry applies), but its + response accounting reports "sonnet-chain" as the model that actually + served it, simulating Router falling back after opus-chain failed. The + spend must still land in the opus-chain-scoped bucket: the check already + ran and decided at admission, and is not re-run for the fallback target. + """ + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "dollar_limits": { + "limits": [ + { + "name": "chain_spend", + "tag_id": "end_user_id", + "limit": 10.0, + "period_seconds": 86400, + "apply_to_models": ["opus-chain"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + data = {**_data(["end_user_id:u1"], call_id="call-1"), "model": "opus-chain"} + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion") + + # The response accounting reports the fallback target as the model that + # actually served the request -- not the "opus-chain" admission decided. + kwargs = { + "litellm_call_id": "call-1", + "metadata": {"tags": ["end_user_id:u1"]}, + "model": "sonnet-chain", + "standard_logging_object": { + "total_tokens": 0, + "response_cost": 12.0, + "model": "sonnet-chain", + "model_group": "sonnet-chain", + }, + } + await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) + await asyncio.sleep(0) + + # The spend landed in the opus-chain-scoped bucket regardless -- a fresh + # opus-chain request is now over the limit. + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data={**_data(["end_user_id:u1"], call_id="call-2"), "model": "opus-chain"}, + call_type="completion", + ) + + # --------------------------------------------------------------------------- # Concurrency: reservation at admission, release on success/failure/disconnect # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index 66d74dfc12b..9a429c26467 100644 --- a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -388,7 +388,7 @@ def test_build_group_limits_empty_when_no_deployment_configures_unit(): def test_entry_applies_with_none_of_the_scoping_fields_set(): entry = TagRateLimitEntry(name="daily", tag_id="end_user_id", limit=500, period_seconds=86400) - assert _entry_applies(entry, ["end_user_id:u1"], None) is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is True def test_entry_applies_disabled_for_on_its_own_tag_id_excludes_a_listed_value(): @@ -401,8 +401,8 @@ def test_entry_applies_disabled_for_on_its_own_tag_id_excludes_a_listed_value(): period_seconds=86400, disabled_for=TagRateLimitScope(tag_id="end_user_id", values=("u1",)), ) - assert _entry_applies(entry, ["end_user_id:u1"], None) is False - assert _entry_applies(entry, ["end_user_id:u2"], None) is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is False + assert _entry_applies(entry, ["end_user_id:u2"], None, None) is True def test_entry_applies_enabled_for_on_its_own_tag_id_restricts_to_a_listed_value(): @@ -415,8 +415,8 @@ def test_entry_applies_enabled_for_on_its_own_tag_id_restricts_to_a_listed_value period_seconds=86400, enabled_for=TagRateLimitScope(tag_id="end_user_id", values=("u2", "u3")), ) - assert _entry_applies(entry, ["end_user_id:u1"], None) is False - assert _entry_applies(entry, ["end_user_id:u2"], None) is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is False + assert _entry_applies(entry, ["end_user_id:u2"], None, None) is True def test_entry_applies_matches_an_enabled_for_gate(): @@ -427,7 +427,7 @@ def test_entry_applies_matches_an_enabled_for_gate(): period_seconds=86400, enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), ) - assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None) is True + assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None, None) is True def test_entry_applies_skips_when_enabled_for_gate_tag_is_absent(): @@ -442,7 +442,7 @@ def test_entry_applies_skips_when_enabled_for_gate_tag_is_absent(): period_seconds=86400, enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), ) - assert _entry_applies(entry, ["end_user_id:u1"], None) is False + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is False def test_entry_applies_skips_when_disabled_for_gate_matches(): @@ -453,7 +453,7 @@ def test_entry_applies_skips_when_disabled_for_gate_matches(): period_seconds=86400, disabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), ) - assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None) is False + assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None, None) is False def test_entry_applies_when_disabled_for_gate_tag_is_absent(): @@ -466,7 +466,7 @@ def test_entry_applies_when_disabled_for_gate_tag_is_absent(): period_seconds=86400, disabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), ) - assert _entry_applies(entry, ["end_user_id:u1"], None) is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is True def test_entry_applies_disabled_for_overrides_a_matching_enabled_for_gate(): @@ -480,27 +480,27 @@ def test_entry_applies_disabled_for_overrides_a_matching_enabled_for_gate(): enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), disabled_for=TagRateLimitScope(tag_id="end_user_id", values=("u1",)), ) - assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None) is False + assert _entry_applies(entry, ["end_user_id:u1", "company_id:1032"], None, None) is False def test_entry_applies_with_apply_to_key_alias_unset_applies_to_every_key(): entry = TagRateLimitEntry(name="daily", tag_id="end_user_id", limit=500, period_seconds=86400) - assert _entry_applies(entry, ["end_user_id:u1"], "any-key-alias") is True - assert _entry_applies(entry, ["end_user_id:u1"], None) is True + assert _entry_applies(entry, ["end_user_id:u1"], "any-key-alias", None) is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is True def test_entry_applies_admits_a_key_alias_on_the_allowlist(): entry = TagRateLimitEntry( name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_key_alias=("team-a-key",) ) - assert _entry_applies(entry, ["end_user_id:u1"], "team-a-key") is True + assert _entry_applies(entry, ["end_user_id:u1"], "team-a-key", None) is True def test_entry_applies_rejects_a_key_alias_missing_from_the_allowlist(): entry = TagRateLimitEntry( name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_key_alias=("team-a-key",) ) - assert _entry_applies(entry, ["end_user_id:u1"], "team-b-key") is False + assert _entry_applies(entry, ["end_user_id:u1"], "team-b-key", None) is False def test_entry_applies_rejects_when_key_has_no_alias_but_allowlist_is_set(): @@ -509,7 +509,53 @@ def test_entry_applies_rejects_when_key_has_no_alias_but_allowlist_is_set(): entry = TagRateLimitEntry( name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_key_alias=("team-a-key",) ) - assert _entry_applies(entry, ["end_user_id:u1"], None) is False + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is False + + +def test_entry_applies_with_apply_to_models_unset_applies_to_every_model(): + entry = TagRateLimitEntry(name="daily", tag_id="end_user_id", limit=500, period_seconds=86400) + assert _entry_applies(entry, ["end_user_id:u1"], None, "opus-chain") is True + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is True + + +def test_entry_applies_admits_a_model_on_the_apply_to_models_allowlist(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_models=("opus-chain",) + ) + assert _entry_applies(entry, ["end_user_id:u1"], None, "opus-chain") is True + + +def test_entry_applies_rejects_a_model_missing_from_the_apply_to_models_allowlist(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_models=("opus-chain",) + ) + assert _entry_applies(entry, ["end_user_id:u1"], None, "sonnet-chain") is False + + +def test_entry_applies_rejects_when_model_is_absent_but_apply_to_models_is_set(): + """apply_to_models is an allowlist gate: a request with no model at all + never satisfies it, same as apply_to_key_alias's absent-key semantics.""" + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, apply_to_models=("opus-chain",) + ) + assert _entry_applies(entry, ["end_user_id:u1"], None, None) is False + + +def test_entry_applies_apply_to_models_composes_with_apply_to_key_alias(): + """Both gates must pass: a request against the listed model but a + non-listed key alias must not apply, even though apply_to_models alone + would have admitted it.""" + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + apply_to_models=("opus-chain",), + apply_to_key_alias=("premium-key",), + ) + assert _entry_applies(entry, ["end_user_id:u1"], "premium-key", "opus-chain") is True + assert _entry_applies(entry, ["end_user_id:u1"], "other-key", "opus-chain") is False + assert _entry_applies(entry, ["end_user_id:u1"], "premium-key", "sonnet-chain") is False # --------------------------------------------------------------------------- @@ -544,6 +590,18 @@ def test_tag_rate_limit_entry_normalizes_apply_to_key_alias_order_and_duplicates assert entry.apply_to_key_alias == ("team-a-key", "team-b-key") +def test_tag_rate_limit_entry_rejects_empty_apply_to_models(): + with pytest.raises(ValidationError, match="apply_to_models must be a non-empty list"): + TagRateLimitEntry(name="daily", limit=1, period_seconds=60, apply_to_models=()) + + +def test_tag_rate_limit_entry_normalizes_apply_to_models_order_and_duplicates(): + entry = TagRateLimitEntry( + name="daily", limit=1, period_seconds=60, apply_to_models=("sonnet-chain", "opus-chain", "opus-chain") + ) + assert entry.apply_to_models == ("opus-chain", "sonnet-chain") + + # --------------------------------------------------------------------------- # _hash_tag / _bucket_key -- policy identity folds into the Redis key itself # --------------------------------------------------------------------------- @@ -1061,6 +1119,47 @@ def test_resolve_any_keeps_divergent_disabled_for_across_member_model_names_sepa assert {c.entry.disabled_for.values for c in resolved} == {("u1",), ("u2",)} +def test_resolve_any_keeps_divergent_apply_to_models_across_member_model_names_separate(): + """Same class of bug as the disabled_for test above, for apply_to_models: + two routing-group members agreeing on tag_id/limit/period_seconds but + scoped to different apply_to_models lists must not collapse to one.""" + concurrency_limits_for_opus = { + "concurrency_limits": { + "limits": [ + { + "name": "inflight", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 300, + "apply_to_models": ["opus-chain"], + } + ] + } + } + concurrency_limits_for_sonnet = { + "concurrency_limits": { + "limits": [ + { + "name": "inflight", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 300, + "apply_to_models": ["sonnet-chain"], + } + ] + } + } + index = _build_limits_index( + [ + _deployment("backend-a", "dep-a", concurrency_limits_for_opus), + _deployment("backend-b", "dep-b", concurrency_limits_for_sonnet), + ] + ) + resolved = index.resolve_any("my-group", team_id=None, candidate_model_names=("backend-a", "backend-b")) + assert len(resolved) == 2 + assert {c.entry.apply_to_models for c in resolved} == {("opus-chain",), ("sonnet-chain",)} + + def test_resolve_any_picks_the_same_resolved_group_regardless_of_hash_seed(): """ Two members with an identical signature dedup to whichever one @@ -3830,10 +3929,10 @@ def test_partition_key_distinguishes_entries_that_differ_only_by_scoping_fields( A plain, unscoped entry and a scoped override can legitimately share name/tag_id/limit/period_seconds/scope_by_key_hash (see test_bucket_key_differs_for_same_named_entries_with_different_scoping_only) - while disagreeing on enabled_for/disabled_for/apply_to_key_alias -- - _DedupSignature and _policy_fingerprint already treat that as two - distinct policies, so a shared max_in_memory_cache_size must not route - them onto the same in-memory partition either, or one entry's + while disagreeing on enabled_for/disabled_for/apply_to_key_alias/ + apply_to_models -- _DedupSignature and _policy_fingerprint already treat + that as two distinct policies, so a shared max_in_memory_cache_size must + not route them onto the same in-memory partition either, or one entry's high-cardinality traffic can evict the other's active counters from a cache neither entry asked to share. """ @@ -3846,14 +3945,16 @@ def test_partition_key_distinguishes_entries_that_differ_only_by_scoping_fields( **base_kwargs, disabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)) ) alias_scoped = TagRateLimitEntry(**base_kwargs, apply_to_key_alias=("premium-key",)) + models_scoped = TagRateLimitEntry(**base_kwargs, apply_to_models=("opus-chain",)) keys = { _partition_key(unscoped), _partition_key(enabled_for_scoped), _partition_key(disabled_for_scoped), _partition_key(alias_scoped), + _partition_key(models_scoped), } - assert len(keys) == 4 + assert len(keys) == 5 @pytest.mark.asyncio @@ -4626,3 +4727,92 @@ async def test_apply_to_key_alias_restricts_a_per_model_entry_to_the_listed_key( messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"], "user_api_key_alias": "premium-key"}}, ) + + +# --------------------------------------------------------------------------- +# apply_to_models -- shared TagRateLimitEntry field, also usable on a +# per-model entry (expected to be rarely useful there, since a per-deployment +# entry is already implicitly scoped to whichever model_name declares it, but +# it must compose identically to every other shared scoping field) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_apply_to_models_ignores_a_non_matching_model_group(time_controller): + limiter = _make_limiter(time_controller) + router = litellm.Router( + model_list=[ + _deployment( + "grp", + "dep-1", + { + "request_limits": { + "limits": [ + { + "name": "per_minute", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 60, + "apply_to_models": ["other-group"], + } + ] + } + }, + ) + ] + ) + limiter.update_variables(llm_router=router) + healthy = router.model_list + + for _ in range(3): + result = await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, + ) + assert result == healthy + + +@pytest.mark.asyncio +async def test_apply_to_models_restricts_a_per_model_entry_to_the_listed_model_group(time_controller): + limiter = _make_limiter(time_controller) + router = litellm.Router( + model_list=[ + _deployment( + "grp", + "dep-1", + { + "request_limits": { + "limits": [ + { + "name": "per_minute", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 60, + "apply_to_models": ["grp"], + } + ] + } + }, + ) + ] + ) + limiter.update_variables(llm_router=router) + healthy = router.model_list + + result = await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, + ) + assert result == healthy + + with pytest.raises(ProxyRateLimitError): + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, + )