mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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.
This commit is contained in:
parent
38feedba65
commit
ba4bde4fcd
5 changed files with 536 additions and 41 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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, ...] = ()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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"]}},
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue