diff --git a/litellm/__init__.py b/litellm/__init__.py index 3f68dd0cbd1..f9348f68f1b 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -119,7 +119,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "litellm_agent", "dynamic_rate_limiter", "dynamic_rate_limiter_v3", - "tag_rate_limiter", + "model_based_tag_rate_limits_hook", "langsmith", "prometheus", "otel", @@ -392,7 +392,7 @@ cache: Optional["Cache"] = None # cache object <- use this - https://docs.litel default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None -tag_rate_limiter_max_in_memory_cache_size: Optional[int] = None +model_based_tag_rate_limits_max_in_memory_cache_size: Optional[int] = None model_alias_map: Dict[str, str] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 140449d2e54..4f7a7510f96 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4476,24 +4476,26 @@ def _init_custom_logger_compatible_class( dynamic_rate_limiter_obj_v3.update_variables(llm_router=llm_router) _in_memory_loggers.append(dynamic_rate_limiter_obj_v3) return dynamic_rate_limiter_obj_v3 - elif logging_integration == "tag_rate_limiter": - from litellm.proxy.hooks.tag_rate_limiter import ( - _PROXY_TagRateLimiter, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here + elif logging_integration == "model_based_tag_rate_limits_hook": + from litellm.proxy.hooks.model_based_tag_rate_limits_hook import ( + _PROXY_ModelBasedTagRateLimitsHook, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_TagRateLimiter): + if isinstance(callback, _PROXY_ModelBasedTagRateLimitsHook): return callback if internal_usage_cache is None: raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") - tag_rate_limiter_obj: Final = _PROXY_TagRateLimiter(internal_usage_cache=internal_usage_cache) + model_based_tag_rate_limits_hook_obj: Final = _PROXY_ModelBasedTagRateLimitsHook( + internal_usage_cache=internal_usage_cache + ) if llm_router is not None and isinstance(llm_router, litellm.Router): - tag_rate_limiter_obj.update_variables(llm_router=llm_router) - _in_memory_loggers.append(tag_rate_limiter_obj) - return tag_rate_limiter_obj + model_based_tag_rate_limits_hook_obj.update_variables(llm_router=llm_router) + _in_memory_loggers.append(model_based_tag_rate_limits_hook_obj) + return model_based_tag_rate_limits_hook_obj elif logging_integration == "langtrace": if "LANGTRACE_API_KEY" not in os.environ: raise ValueError("LANGTRACE_API_KEY not found in environment variables") @@ -4934,13 +4936,13 @@ def get_custom_logger_compatible_class( if isinstance(callback, _PROXY_DynamicRateLimitHandlerV3): return callback - elif logging_integration == "tag_rate_limiter": - from litellm.proxy.hooks.tag_rate_limiter import ( - _PROXY_TagRateLimiter, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here + elif logging_integration == "model_based_tag_rate_limits_hook": + from litellm.proxy.hooks.model_based_tag_rate_limits_hook import ( + _PROXY_ModelBasedTagRateLimitsHook, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here ) for callback in _in_memory_loggers: - if isinstance(callback, _PROXY_TagRateLimiter): + if isinstance(callback, _PROXY_ModelBasedTagRateLimitsHook): return callback elif logging_integration == "langtrace": diff --git a/litellm/proxy/hooks/tag_rate_limiter.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py similarity index 93% rename from litellm/proxy/hooks/tag_rate_limiter.py rename to litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index e8b727b6736..25dbe9aa420 100644 --- a/litellm/proxy/hooks/tag_rate_limiter.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -30,7 +30,7 @@ from litellm.router_strategy.tag_based_routing import ( ) from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import AllMessageValues -from litellm.types.router import TagRateLimitEntry, TagRateLimits +from litellm.types.router import TagRateLimitEntry, TagRateLimits, TagRateLimitScope from litellm.types.utils import StandardLoggingPayload if TYPE_CHECKING: @@ -42,10 +42,21 @@ else: _LimitUnit: TypeAlias = Literal["tokens", "requests", "dollars", "concurrency"] _LIMIT_UNITS: Final[tuple[_LimitUnit, ...]] = ("tokens", "requests", "dollars", "concurrency") -# (tag_id, name, limit, period_seconds, scope_by_key_hash) -- the fields that -# decide whether two deployments' entries are the same rate limit for dedup -# purposes; see _build_group_limits. -_DedupSignature: TypeAlias = tuple[str, str, float, int, bool] +# A (tag_id, values) pair mirroring TagRateLimitScope's own fields, used only +# to fold `enabled_for`/`disabled_for` into `_DedupSignature` below without +# depending on TagRateLimitScope's own hashability. +_ScopeSignature: TypeAlias = tuple[str, tuple[str, ...]] | None +# (tag_id, name, limit, period_seconds, scope_by_key_hash, included_values, +# excluded_values, enabled_for, disabled_for) -- 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, float, int, bool, tuple[str, ...] | None, tuple[str, ...] | None, _ScopeSignature, _ScopeSignature +] # Units whose admission must be atomic (check-and-increment in one Redis # round trip) because the increment amount is known upfront (always 1). # tokens/dollars can't be: real usage is only known after the response, so @@ -185,6 +196,49 @@ def _extract_identity(tags: Sequence[str], tag_id: str) -> str | None: return None +def _scope_signature(scope: TagRateLimitScope | None) -> _ScopeSignature: + """Normalizes a `TagRateLimitScope` into a plain, hashable tuple for use + in `_DedupSignature` -- see that alias's own comment for why two + deployments disagreeing on `enabled_for`/`disabled_for` must be treated + as genuinely different policies rather than merged into one bucket.""" + return None if scope is None else (scope.tag_id, scope.values) + + +def _entry_applies(entry: TagRateLimitEntry, tag_value: str, tags: Sequence[str]) -> bool: + """ + Applies `entry`'s own scoping fields (`included_values`/`excluded_values`/ + `enabled_for`/`disabled_for`), evaluated in this order -- deny overrides + allow, checked before either allowlist: + + 1. `excluded_values`: `tag_value` is in it -> doesn't apply. + 2. `included_values`: `tag_value` is NOT in it -> doesn't apply. + 3. `disabled_for`: the gate tag (a tag OTHER than `entry.tag_id`, + resolved via `disabled_for.tag_id`) is present and its value is in + `disabled_for.values` -> doesn't apply. Absent gate tag never + triggers this -- nothing to match against a denylist. + 4. `enabled_for`: the gate tag is absent, or present but its value is + 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". + + An entry with none of the four fields set always applies -- this is the + unscoped behavior every existing entry has today, unchanged. + """ + if entry.excluded_values is not None and tag_value in entry.excluded_values: + return False + if entry.included_values is not None and tag_value not in entry.included_values: + return False + if entry.disabled_for is not None: + gate_value = _extract_identity(tags, entry.disabled_for.tag_id) + if gate_value is not None and gate_value in entry.disabled_for.values: + return False + if entry.enabled_for is not None: + gate_value = _extract_identity(tags, entry.enabled_for.tag_id) + if gate_value is None or gate_value not in entry.enabled_for.values: + return False + return True + + def _deployment_id(deployment: Mapping[str, object]) -> str | None: return (deployment.get("model_info") or _EMPTY_MAPPING).get("id") @@ -232,7 +286,7 @@ def _configured_limit_for_signature( ) -> _ConfiguredLimit | None: if unit == "concurrency" and not is_chain_wide: verbose_proxy_logger.warning( - "tag_rate_limiter: concurrency_limits entry %r (tag_id=%s) is not declared identically by every " + "model_based_tag_rate_limits_hook: concurrency_limits entry %r (tag_id=%s) is not declared identically by every " "deployment sharing this model_name; per-deployment-scoped concurrency limits are not supported " "and this entry is being skipped entirely.", entry.name, @@ -293,7 +347,17 @@ def _build_group_limits(deployments: Sequence[Mapping[str, object]], unit: _Limi if dep_id is None: continue for entry in _entries_for_unit(deployment, unit): - signature = (entry.tag_id, entry.name, entry.limit, entry.period_seconds, entry.scope_by_key_hash) + signature = ( + entry.tag_id, + entry.name, + entry.limit, + entry.period_seconds, + entry.scope_by_key_hash, + entry.included_values, + entry.excluded_values, + _scope_signature(entry.enabled_for), + _scope_signature(entry.disabled_for), + ) ids_for_signature = declaring_ids_by_signature.setdefault(signature, []) # mutable-ok: see comment above # One deployment declaring the identical entry twice (a config # duplicate) must count once, or len(declaring_ids) inflates past @@ -537,7 +601,7 @@ _CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600 # server-side per logical request (and shared across that request's own # fallback hops, matching the original chain-wide release semantics), so it # can't be forged or guessed. -_PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_tag_rate_limiter_pending_concurrency_keys" +_PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_model_based_tag_rate_limits_pending_concurrency_keys" # The admission-time timestamp a hop's token/dollar checks classified their # bucket against, stashed on the same model_call_details object so success @@ -550,7 +614,7 @@ _PENDING_CONCURRENCY_KEYS_FIELD: Final[str] = "_tag_rate_limiter_pending_concurr # rollover. Overwritten by each hop's own admission (last-write-wins), which # is correct: success only ever fires for whichever hop actually served the # request, so its own most recent admission timestamp is the right one. -_ADMISSION_TIME_FIELD: Final[str] = "_tag_rate_limiter_admission_time" +_ADMISSION_TIME_FIELD: Final[str] = "_model_based_tag_rate_limits_admission_time" class _TagRateLimitIndex: @@ -658,6 +722,8 @@ 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, tag_value, tags): + return None key_hash: Final = ( _extract_key_hash(request_kwargs, metadata_variable_name) if configured_limit.entry.scope_by_key_hash else None ) @@ -694,6 +760,8 @@ 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, tag_value, tags): + return None if configured_limit.unit not in increment_by_unit: return None # "requests" is accounted atomically at admission, not here increment_value: Final = increment_by_unit[configured_limit.unit] @@ -711,7 +779,7 @@ def _increment_operation_for_limit( def _resolve_max_in_memory_cache_size() -> int | None: """ - `litellm_settings` values reach `litellm.tag_rate_limiter_max_in_memory_cache_size` + `litellm_settings` values reach `litellm.model_based_tag_rate_limits_max_in_memory_cache_size` via a plain, unvalidated `setattr`, so a config typo (a negative number, or a string like "500" from an unresolved os.environ/ substitution) can reach here. InMemoryCache raises when comparing its size against a non-positive-int @@ -719,12 +787,12 @@ def _resolve_max_in_memory_cache_size() -> int | None: an invalid value would otherwise silently disable every counter write for this hook rather than fail loudly -- rejected here in favor of the safe default instead. """ - configured: Final = litellm.tag_rate_limiter_max_in_memory_cache_size + configured: Final = litellm.model_based_tag_rate_limits_max_in_memory_cache_size if isinstance(configured, int) and not isinstance(configured, bool) and configured > 0: return configured if configured is not None: verbose_proxy_logger.warning( - "tag_rate_limiter: tag_rate_limiter_max_in_memory_cache_size=%r is not a positive integer; " + "model_based_tag_rate_limits_hook: model_based_tag_rate_limits_max_in_memory_cache_size=%r is not a positive integer; " "falling back to the default in-memory cache size.", configured, ) @@ -737,7 +805,7 @@ def _resolve_max_in_memory_cache_size() -> int | None: # entries that happen to choose the identical max_in_memory_cache_size don't # get merged into one shared partition; the same entry (same config content) # always resolves to the same signature across index rebuilds, which is what -# keeps _PROXY_TagRateLimiter._partitions from leaking a fresh partition +# keeps _PROXY_ModelBasedTagRateLimitsHook._partitions from leaking a fresh partition # every time _TagRateLimitIndex rebuilds and reconstructs `_ConfiguredLimit`s. _PartitionKey: TypeAlias = tuple[str, str, float, int, bool, int] | None # Grouping type for async_log_success_event's per-partition tokens/dollars @@ -801,7 +869,7 @@ class _CachePartition: v3: _PROXY_MaxParallelRequestsHandler_v3 -class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only referenced via the deferred import in litellm_logging.py's callback resolver; basedpyright doesn't trace that usage +class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # only referenced via the deferred import in litellm_logging.py's callback resolver; basedpyright doesn't trace that usage CustomLogger ): def __init__( @@ -825,7 +893,7 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer # here -- see _partition_for. None (the key every entry uses unless # it sets its own max_in_memory_cache_size) is this hook's single # default partition, sized by - # litellm.tag_rate_limiter_max_in_memory_cache_size (200 if that's + # litellm.model_based_tag_rate_limits_max_in_memory_cache_size (200 if that's # also unset), matching today's behavior for every entry that doesn't # opt into its own partition. self._partitions: dict[_PartitionKey, _CachePartition] = {} # mutable-ok: lazily memoized; see _partition_for @@ -988,7 +1056,9 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer try: await self._decrement_floor_zero(refund_cache, refund_key, -refund_increment) except Exception as e: # noqa: BLE001 - one failed refund must not block refunding the rest - verbose_proxy_logger.warning("tag_rate_limiter: failed to refund %s on rollback: %s", refund_key, e) + verbose_proxy_logger.warning( + "model_based_tag_rate_limits_hook: failed to refund %s on rollback: %s", refund_key, e + ) async def async_filter_deployments( self, @@ -1192,7 +1262,7 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer current: float, ) -> None: verbose_proxy_logger.debug( - "tag_rate_limiter: OVER_LIMIT model=%s unit=%s name=%s tag_id=%s tag_value=%s current=%s limit=%s", + "model_based_tag_rate_limits_hook: OVER_LIMIT model=%s unit=%s name=%s tag_id=%s tag_value=%s current=%s limit=%s", model, configured_limit.unit, configured_limit.entry.name, @@ -1237,7 +1307,9 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer partition = await self._partition_for(partition_key) # not Final: rebound each loop iteration await self._decrement_floor_zero(partition.internal_usage_cache, key, -1.0) except Exception as e: # noqa: BLE001 - releasing a slot must never raise into the caller's request path - verbose_proxy_logger.warning("tag_rate_limiter: failed to release concurrency slot %s: %s", key, e) + verbose_proxy_logger.warning( + "model_based_tag_rate_limits_hook: failed to release concurrency slot %s: %s", key, e + ) async def _release_stale_hop_reservations(self, request_kwargs: Mapping[str, object]) -> None: """ diff --git a/litellm/types/router.py b/litellm/types/router.py index 5c693354629..8cf676f1a04 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -137,6 +137,27 @@ def _as_utc(value: datetime.datetime | None) -> datetime.datetime | None: return value.astimezone(datetime.timezone.utc) +class TagRateLimitScope(BaseModel): + """ + A gate on a tag OTHER than the entry's own `tag_id` -- e.g. scoping an + entry to `tag_id: company_id, values: ["1032"]` so it only applies to + requests tagged as belonging to company 1032, independent of whichever + tag the entry itself keys its bucket by. See `TagRateLimitEntry.enabled_for`/ + `disabled_for`, which are the only two fields that construct this. + """ + + tag_id: str + values: tuple[str, ...] + + model_config = ConfigDict(frozen=True) + + @model_validator(mode="after") + def _validate_values(self) -> "TagRateLimitScope": + if not self.values: + raise ValueError("values must be a non-empty list of strings") + return self + + class TagRateLimitEntry(BaseModel): name: str tag_id: str = "end_user_id" @@ -145,7 +166,7 @@ class TagRateLimitEntry(BaseModel): scope_by_key_hash: bool = False # Overrides this entry's bucket/reservation key TTL (Redis, and the # in-memory fallback when Redis isn't configured). Defaults to - # period_seconds + 3600 when unset -- see _PROXY_TagRateLimiter._ttl_for. + # period_seconds + 3600 when unset -- see _PROXY_ModelBasedTagRateLimitsHook._ttl_for. # A high-cardinality tag_id can keep many keys alive at once; lowering # this lets an operator shed them sooner without shortening # period_seconds itself. @@ -154,11 +175,27 @@ class TagRateLimitEntry(BaseModel): # entry's own keys live in, when Redis isn't configured (or as a local # fast-path cache when it is). Unset means this entry shares the hook's # single default partition, sized by - # litellm.tag_rate_limiter_max_in_memory_cache_size (200 if that's also + # litellm.model_based_tag_rate_limits_max_in_memory_cache_size (200 if that's also # unset). A high-cardinality tag_id can churn past that shared cap and # evict another entry's active counters; setting this gives the entry # its own dedicated partition instead. max_in_memory_cache_size: int | None = None + # Scope this entry to a subset of its own resolved `tag_id` value -- + # e.g. hand-picking a handful of identities without needing a second + # tag at all. `excluded_values` is checked before `included_values` + # (deny overrides allow) when both happen to be set on the same entry. + included_values: tuple[str, ...] | None = None + excluded_values: tuple[str, ...] | None = None + # Gate this entry on a SECOND, independent tag rather than its own + # `tag_id` -- e.g. `enabled_for: {tag_id: company_id, values: ["1032"]}` + # to scope an override to one company's traffic without enumerating + # every one of that company's end_user_id values by hand. + # `disabled_for` is checked first (deny overrides allow) when both are + # set. An absent gate tag never satisfies `enabled_for` (an allowlist + # gate requires an explicit match) but never triggers `disabled_for` + # either (nothing to match against a denylist). + enabled_for: TagRateLimitScope | None = None + disabled_for: TagRateLimitScope | None = None model_config = ConfigDict(protected_namespaces=()) @@ -185,6 +222,14 @@ class TagRateLimitEntry(BaseModel): raise ValueError("max_in_memory_cache_size must be a positive integer when set") return self + @model_validator(mode="after") + def _validate_included_and_excluded_values(self) -> "TagRateLimitEntry": + if self.included_values is not None and not self.included_values: + raise ValueError("included_values must be a non-empty list of strings when set") + if self.excluded_values is not None and not self.excluded_values: + raise ValueError("excluded_values must be a non-empty list of strings when set") + return self + class TagRateLimitGroup(BaseModel): limits: tuple[TagRateLimitEntry, ...] = () diff --git a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py similarity index 91% rename from tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py rename to tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index d6b7b744657..996b56f706e 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -17,13 +17,14 @@ from pydantic import ValidationError import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError -from litellm.proxy.hooks.tag_rate_limiter import ( +from litellm.proxy.hooks.model_based_tag_rate_limits_hook import ( _CONCURRENCY_MIN_SAFETY_TTL_SECONDS, _bucket_key, _bucket_ttl_seconds, _build_group_limits, _build_limits_index, _ConfiguredLimit, + _entry_applies, _extract_identity, _extract_key_hash, _extract_team_id, @@ -32,10 +33,10 @@ from litellm.proxy.hooks.tag_rate_limiter import ( _inflight_key, _partition_key, _PENDING_CONCURRENCY_KEYS_FIELD, - _PROXY_TagRateLimiter, + _PROXY_ModelBasedTagRateLimitsHook, _queue_pending_concurrency_reservations, ) -from litellm.types.router import RoutingGroup, TagRateLimitEntry +from litellm.types.router import RoutingGroup, TagRateLimitEntry, TagRateLimitScope class TimeController: @@ -54,8 +55,8 @@ def time_controller(): return TimeController() -def _make_limiter(time_controller: TimeController) -> _PROXY_TagRateLimiter: - return _PROXY_TagRateLimiter( +def _make_limiter(time_controller: TimeController) -> _PROXY_ModelBasedTagRateLimitsHook: + return _PROXY_ModelBasedTagRateLimitsHook( internal_usage_cache=DualCache(), time_provider=time_controller.now, ) @@ -314,6 +315,206 @@ def test_build_group_limits_empty_when_no_deployment_configures_unit(): assert _build_group_limits(deployments, "tokens") == () +# --------------------------------------------------------------------------- +# _entry_applies -- included_values / excluded_values / enabled_for / disabled_for +# --------------------------------------------------------------------------- + + +def test_entry_applies_with_none_of_the_four_fields_set(): + entry = TagRateLimitEntry(name="daily", tag_id="end_user_id", limit=500, period_seconds=86400) + assert _entry_applies(entry, "u1", ["end_user_id:u1"]) is True + + +def test_entry_applies_excludes_a_listed_value(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, excluded_values=("u1",) + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1"]) is False + + +def test_entry_applies_admits_a_value_not_on_the_exclusion_list(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, excluded_values=("u1",) + ) + assert _entry_applies(entry, "u2", ["end_user_id:u2"]) is True + + +def test_entry_applies_rejects_a_value_missing_from_the_inclusion_list(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, included_values=("u2", "u3") + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1"]) is False + + +def test_entry_applies_admits_a_value_on_the_inclusion_list(): + entry = TagRateLimitEntry( + name="daily", tag_id="end_user_id", limit=500, period_seconds=86400, included_values=("u2", "u3") + ) + assert _entry_applies(entry, "u2", ["end_user_id:u2", "company_id:1032"]) is True + + +def test_entry_applies_matches_an_enabled_for_gate(): + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"]) is True + + +def test_entry_applies_skips_when_enabled_for_gate_tag_is_absent(): + """ + enabled_for is an allowlist gate: absence of the gate tag must not + satisfy it, unlike disabled_for below. + """ + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1"]) is False + + +def test_entry_applies_skips_when_disabled_for_gate_matches(): + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + disabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"]) is False + + +def test_entry_applies_when_disabled_for_gate_tag_is_absent(): + """disabled_for is a denylist gate: absence of the gate tag has nothing + to match against, so the entry still applies.""" + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + disabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1"]) is True + + +def test_entry_applies_excluded_values_overrides_a_matching_enabled_for_gate(): + """Deny (identity-level excluded_values) takes effect independently of + whether the enabled_for gate itself matched.""" + entry = TagRateLimitEntry( + name="daily", + tag_id="end_user_id", + limit=500, + period_seconds=86400, + enabled_for=TagRateLimitScope(tag_id="company_id", values=("1032",)), + excluded_values=("u1",), + ) + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"]) is False + + +# --------------------------------------------------------------------------- +# TagRateLimitEntry / TagRateLimitScope -- scoping field validation +# --------------------------------------------------------------------------- + + +def test_tag_rate_limit_entry_rejects_empty_included_values(): + with pytest.raises(ValidationError, match="included_values must be a non-empty list"): + TagRateLimitEntry(name="daily", limit=1, period_seconds=60, included_values=()) + + +def test_tag_rate_limit_entry_rejects_empty_excluded_values(): + with pytest.raises(ValidationError, match="excluded_values must be a non-empty list"): + TagRateLimitEntry(name="daily", limit=1, period_seconds=60, excluded_values=()) + + +def test_tag_rate_limit_scope_rejects_empty_values(): + with pytest.raises(ValidationError, match="values must be a non-empty list"): + TagRateLimitScope(tag_id="company_id", values=()) + + +def test_tag_rate_limit_entry_rejects_enabled_for_missing_values(): + with pytest.raises(ValidationError): + TagRateLimitEntry(name="daily", limit=1, period_seconds=60, enabled_for={"tag_id": "company_id"}) + + +# --------------------------------------------------------------------------- +# _build_group_limits -- scoping fields fold into the dedup signature +# --------------------------------------------------------------------------- + + +def test_build_group_limits_per_deployment_when_excluded_values_diverge(): + """ + Regression test: two deployments agreeing on tag_id/limit/period_seconds + but declaring different excluded_values are genuinely different + policies and must not be silently merged into one shared bucket -- the + same class of bug test_build_group_limits_per_deployment_when_values_diverge + already guards against for a plain divergent limit value. + """ + deployments = [ + _deployment( + "grp", + "dep-1", + { + "token_limits": { + "limits": [ + {"name": "daily", "limit": 500, "period_seconds": 86400, "excluded_values": ["u1"]} + ] + } + }, + ), + _deployment( + "grp", + "dep-2", + { + "token_limits": { + "limits": [ + {"name": "daily", "limit": 500, "period_seconds": 86400, "excluded_values": ["u2"]} + ] + } + }, + ), + ] + configured = _build_group_limits(deployments, "tokens") + assert len(configured) == 2 + scopes = {c.deployment_scope for c in configured} + assert scopes == {("dep-1",), ("dep-2",)} + + +def test_build_group_limits_chain_wide_when_excluded_values_agree(): + deployments = [ + _deployment( + "grp", + "dep-1", + { + "token_limits": { + "limits": [ + {"name": "daily", "limit": 500, "period_seconds": 86400, "excluded_values": ["u1"]} + ] + } + }, + ), + _deployment( + "grp", + "dep-2", + { + "token_limits": { + "limits": [ + {"name": "daily", "limit": 500, "period_seconds": 86400, "excluded_values": ["u1"]} + ] + } + }, + ), + ] + configured = _build_group_limits(deployments, "tokens") + assert len(configured) == 1 + assert configured[0].deployment_scope is None + + # --------------------------------------------------------------------------- # async_filter_deployments -- enforcement # --------------------------------------------------------------------------- @@ -578,7 +779,7 @@ def test_resolve_any_picks_the_same_resolved_group_regardless_of_hash_seed(): see the bug report this regression-tests for the exact reproduction. """ script = ( - "from litellm.proxy.hooks.tag_rate_limiter import _build_limits_index\n" + "from litellm.proxy.hooks.model_based_tag_rate_limits_hook import _build_limits_index\n" "def _deployment(model_name, deployment_id, tag_rate_limits):\n" " return {'model_name': model_name, 'litellm_params': {'model': 'gpt-4o'}," " 'model_info': {'id': deployment_id, 'tag_rate_limits': tag_rate_limits}}\n" @@ -648,6 +849,129 @@ async def test_filter_deployments_per_entry_fail_open_when_tag_absent(time_contr assert await limiter.internal_usage_cache.async_get_cache(key=team_key, litellm_parent_otel_span=None) is None +def _company_tiered_cap_router(default_limit: int, override_limit: int) -> "litellm.Router": + return litellm.Router( + model_list=[ + _deployment( + "grp", + "dep-1", + { + "request_limits": { + "limits": [ + { + "name": "default_daily", + "tag_id": "end_user_id", + "limit": default_limit, + "period_seconds": 86400, + }, + { + "name": "company_1032_daily", + "tag_id": "end_user_id", + "limit": override_limit, + "period_seconds": 86400, + "enabled_for": {"tag_id": "company_id", "values": ["1032"]}, + "excluded_values": ["u1"], + }, + ] + } + }, + ) + ] + ) + + +@pytest.mark.asyncio +async def test_filter_deployments_scoped_override_skips_for_an_excluded_identity(time_controller): + """ + Company-tiered-cap example from the plan: a stricter override entry + gated to one company via enabled_for, with a handful of named users + excluded from it via excluded_values. An excluded user must fall + through to the unscoped default entry entirely -- the override never + enforces or accounts for them. + """ + limiter = _make_limiter(time_controller) + router = _company_tiered_cap_router(default_limit=3, override_limit=1) + 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", "company_id:1032"]}}, + ) + assert result == healthy + + with pytest.raises(ProxyRateLimitError) as exc_info: + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1", "company_id:1032"]}}, + ) + assert exc_info.value.detail["limit_name"] == "default_daily" + + +@pytest.mark.asyncio +async def test_filter_deployments_scoped_override_enforces_for_a_non_excluded_identity_in_scope(time_controller): + """ + The same override applies, and enforces its own stricter limit, for a + company-1032 user who is not on excluded_values, proving the two + entries are independently enforced rather than one silently replacing + the other. + """ + limiter = _make_limiter(time_controller) + router = _company_tiered_cap_router(default_limit=3, override_limit=1) + 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:u2", "company_id:1032"]}}, + ) + assert result == healthy + + with pytest.raises(ProxyRateLimitError) as exc_info: + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u2", "company_id:1032"]}}, + ) + assert exc_info.value.detail["limit_name"] == "company_1032_daily" + + +@pytest.mark.asyncio +async def test_filter_deployments_scoped_override_does_not_apply_outside_its_enabled_for_gate(time_controller): + """A user not tagged with the gate company at all only ever hits the + unscoped default entry, even though the override's own limit is looser + and would otherwise still have room.""" + limiter = _make_limiter(time_controller) + router = _company_tiered_cap_router(default_limit=1, override_limit=5) + 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:u3"]}}, + ) + assert result == healthy + + with pytest.raises(ProxyRateLimitError) as exc_info: + await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u3"]}}, + ) + assert exc_info.value.detail["limit_name"] == "default_daily" + + @pytest.mark.asyncio async def test_different_tag_ids_with_same_name_do_not_share_a_counter(time_controller): """ @@ -2086,7 +2410,7 @@ def _redis_limiter(time_controller: TimeController): pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set") redis_cache = RedisCache(host=redis_host, port=int(redis_port), password=os.getenv("REDIS_PASSWORD")) dual_cache = DualCache(redis_cache=redis_cache) - return _PROXY_TagRateLimiter(internal_usage_cache=dual_cache, time_provider=time_controller.now), redis_cache + return _PROXY_ModelBasedTagRateLimitsHook(internal_usage_cache=dual_cache, time_provider=time_controller.now), redis_cache @pytest.mark.asyncio @@ -2608,7 +2932,7 @@ def test_concurrency_identical_across_all_deployments_is_still_chain_wide(): def test_concurrency_ttl_floor_overrides_a_too_short_period_seconds(): entry = TagRateLimitEntry(name="inflight", tag_id="end_user_id", limit=1, period_seconds=5) configured_limit = _ConfiguredLimit(unit="concurrency", entry=entry, deployment_scope=None) - assert _PROXY_TagRateLimiter._ttl_for(configured_limit) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + assert _PROXY_ModelBasedTagRateLimitsHook._ttl_for(configured_limit) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS def test_concurrency_ttl_floor_does_not_shorten_a_longer_period_seconds(): @@ -2616,7 +2940,7 @@ def test_concurrency_ttl_floor_does_not_shorten_a_longer_period_seconds(): name="inflight", tag_id="end_user_id", limit=1, period_seconds=_CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 ) configured_limit = _ConfiguredLimit(unit="concurrency", entry=entry, deployment_scope=None) - assert _PROXY_TagRateLimiter._ttl_for(configured_limit) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 + assert _PROXY_ModelBasedTagRateLimitsHook._ttl_for(configured_limit) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 # --------------------------------------------------------------------------- @@ -2636,7 +2960,7 @@ async def test_release_in_a_forked_task_is_visible_to_the_parent_context(): model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]} async def detached_release(): - return _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details) + return _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details) released = await asyncio.create_task(detached_release()) assert released == ("key1",) @@ -2650,7 +2974,7 @@ async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot(): model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]} async def detached_release_then_sibling_admits(): - released = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details) + released = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details) # A sibling hop's admission, appending to the same shared dict, # interleaved right after this release's snapshot was taken. model_call_details[_PENDING_CONCURRENCY_KEYS_FIELD].append("key2") @@ -2665,8 +2989,8 @@ async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot(): @pytest.mark.asyncio async def test_release_is_not_repeated_for_the_same_snapshot(): model_call_details: dict = {_PENDING_CONCURRENCY_KEYS_FIELD: ["key1"]} - first = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details) - second = _PROXY_TagRateLimiter._pop_pending_concurrency_keys(model_call_details) + first = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details) + second = _PROXY_ModelBasedTagRateLimitsHook._pop_pending_concurrency_keys(model_call_details) assert first == ("key1",) assert second == () @@ -2758,7 +3082,7 @@ async def test_refund_failure_on_one_key_does_not_block_others_or_raise(time_con other_key = "{tag_rl:test:refund-fail:b}:requests" rejecting_key = "{tag_rl:test:refund-fail:c}:requests" - class _FlakyLimiter(_PROXY_TagRateLimiter): + class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook): async def _decrement_floor_zero(self, cache, key: str, delta: float) -> None: if key == failing_key: raise RuntimeError("simulated transient redis failure") @@ -2794,7 +3118,7 @@ async def test_exception_mid_batch_refunds_every_earlier_admission_before_propag admitted_key = "{tag_rl:test:exception-refund:a}:requests" raising_key = "{tag_rl:test:exception-refund:b}:requests" - class _FlakyLimiter(_PROXY_TagRateLimiter): + class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook): async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int): if key == raising_key: raise RuntimeError("simulated transient redis failure") @@ -2832,7 +3156,7 @@ async def test_a_raising_keys_own_ambiguous_outcome_is_never_refunded(time_contr admitted_key = "{tag_rl:test:ambiguous-no-refund:a}:requests" raising_key = "{tag_rl:test:ambiguous-no-refund:b}:requests" - class _FlakyLimiter(_PROXY_TagRateLimiter): + class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook): async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int): if key == raising_key: # Simulate Redis committing the increment before the @@ -3219,7 +3543,7 @@ async def test_flooding_tag_buckets_does_not_evict_the_shared_cache_authenticati shared_cache = DualCache() await shared_cache.async_set_cache(key="authentication_bound_counter", value="do-not-evict") - limiter = _PROXY_TagRateLimiter(internal_usage_cache=shared_cache, time_provider=time_controller.now) + limiter = _PROXY_ModelBasedTagRateLimitsHook(internal_usage_cache=shared_cache, time_provider=time_controller.now) router = litellm.Router( model_list=[ _deployment( @@ -3271,14 +3595,14 @@ async def test_max_in_memory_cache_size_setting_lets_high_cardinality_tags_avoid This hook's own isolated cache still defaults to 200 items, shared across every distinct tag value it sees. A deployment rate-limiting on a high-cardinality tag_id (e.g. per end user) without Redis can raise - `litellm_settings.tag_rate_limiter_max_in_memory_cache_size` so an + `litellm_settings.model_based_tag_rate_limits_max_in_memory_cache_size` so an earlier bucket survives churn from later, unrelated tag values: with limit=1, a still-live bucket rejects a second request instead of having been evicted back to a fresh count of 0. """ - monkeypatch.setattr(litellm, "tag_rate_limiter_max_in_memory_cache_size", 500) + monkeypatch.setattr(litellm, "model_based_tag_rate_limits_max_in_memory_cache_size", 500) - limiter = _PROXY_TagRateLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now) + limiter = _PROXY_ModelBasedTagRateLimitsHook(internal_usage_cache=DualCache(), time_provider=time_controller.now) router = _single_request_per_minute_router() limiter.update_variables(llm_router=router) healthy = router.model_list @@ -3327,9 +3651,9 @@ async def test_invalid_max_in_memory_cache_size_falls_back_to_the_safe_default( of failing loudly. Each of these must be rejected in favor of the safe default: a limit=1 bucket must still reject a second, immediate request. """ - monkeypatch.setattr(litellm, "tag_rate_limiter_max_in_memory_cache_size", invalid_configured_size) + monkeypatch.setattr(litellm, "model_based_tag_rate_limits_max_in_memory_cache_size", invalid_configured_size) - limiter = _PROXY_TagRateLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now) + limiter = _PROXY_ModelBasedTagRateLimitsHook(internal_usage_cache=DualCache(), time_provider=time_controller.now) router = _single_request_per_minute_router() limiter.update_variables(llm_router=router) healthy = router.model_list @@ -3390,7 +3714,7 @@ def test_bucket_ttl_seconds_honors_key_ttl_seconds_override(): def test_ttl_for_concurrency_honors_key_ttl_seconds_above_the_safety_floor(): above_floor: Final = _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 assert ( - _PROXY_TagRateLimiter._ttl_for(_concurrency_limit(period_seconds=60, key_ttl_seconds=above_floor)) + _PROXY_ModelBasedTagRateLimitsHook._ttl_for(_concurrency_limit(period_seconds=60, key_ttl_seconds=above_floor)) == above_floor ) @@ -3404,7 +3728,7 @@ def test_ttl_for_concurrency_never_drops_below_the_safety_floor_even_with_a_lowe """ below_floor: Final = 10 assert ( - _PROXY_TagRateLimiter._ttl_for(_concurrency_limit(period_seconds=5, key_ttl_seconds=below_floor)) + _PROXY_ModelBasedTagRateLimitsHook._ttl_for(_concurrency_limit(period_seconds=5, key_ttl_seconds=below_floor)) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS ) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 4dc51f026e4..f854d8f94e5 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -4071,7 +4071,7 @@ class TestCancelOnDisconnect: never reaches litellm.utils.wrapper_async's own except block -- the cancelled call's async_log_failure_event never fires, and the 499 this raises is later handled by post_call_failure_hook, a different - hook a CustomLogger like tag_rate_limiter doesn't implement. Without + hook a CustomLogger like model_based_tag_rate_limits_hook doesn't implement. Without an explicit release here, a callback that reserved per-request state at admission (a concurrency slot) leaks it until that state's own safety TTL. This mirrors the streaming disconnect case diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1f5bffbf00b..018b98ef95b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -35112,6 +35112,12 @@ export interface components { }; /** TagRateLimitEntry */ TagRateLimitEntry: { + disabled_for?: components["schemas"]["TagRateLimitScope"] | null; + enabled_for?: components["schemas"]["TagRateLimitScope"] | null; + /** Excluded Values */ + excluded_values?: string[] | null; + /** Included Values */ + included_values?: string[] | null; /** Key Ttl Seconds */ key_ttl_seconds?: number | null; /** Limit */ @@ -35141,6 +35147,20 @@ export interface components { */ limits: components["schemas"]["TagRateLimitEntry"][]; }; + /** + * TagRateLimitScope + * @description A gate on a tag OTHER than the entry's own `tag_id` -- e.g. scoping an + * entry to `tag_id: company_id, values: ["1032"]` so it only applies to + * requests tagged as belonging to company 1032, independent of whichever + * tag the entry itself keys its bucket by. See `TagRateLimitEntry.enabled_for`/ + * `disabled_for`, which are the only two fields that construct this. + */ + TagRateLimitScope: { + /** Tag Id */ + tag_id: string; + /** Values */ + values: string[]; + }; /** TagRateLimits */ TagRateLimits: { concurrency_limits?: components["schemas"]["TagRateLimitGroup"] | null;