From 6663924b2f832b8e8a2d13ee80f6e9fdc52f4140 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Mon, 24 Aug 2026 14:25:06 -0400 Subject: [PATCH] feat(rate-limiting): global-scope, model-independent tag rate limits Adds apply_to_key_alias to TagRateLimitEntry (unset applies to every request; set, restricts an entry to specific virtual keys) and a new global_tag_rate_limits_hook enforcing tag rate limits via a single litellm_settings.global_tag_rate_limits config block, once per request in async_pre_call_hook, before routing. Composes with the existing scope_by_key_hash field to optionally split a bucket per calling key. Generalizes the already-shipped, per-deployment model_based_tag_rate_limits_hook to a config surface that is neither model-scoped nor key-scoped by default. --- litellm/__init__.py | 3 + litellm/litellm_core_utils/litellm_logging.py | 26 + .../hooks/global_tag_rate_limits_hook.py | 613 ++++++++++++++++++ .../hooks/model_based_tag_rate_limits_hook.py | 65 +- litellm/types/router.py | 10 + .../hooks/test_global_tag_rate_limits_hook.py | 505 +++++++++++++++ .../test_model_based_tag_rate_limits_hook.py | 124 +++- 7 files changed, 1320 insertions(+), 26 deletions(-) create mode 100644 litellm/proxy/hooks/global_tag_rate_limits_hook.py create mode 100644 tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py diff --git a/litellm/__init__.py b/litellm/__init__.py index f9348f68f1b..b82c918230d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -120,6 +120,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "dynamic_rate_limiter", "dynamic_rate_limiter_v3", "model_based_tag_rate_limits_hook", + "global_tag_rate_limits_hook", "langsmith", "prometheus", "otel", @@ -393,6 +394,8 @@ default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None model_based_tag_rate_limits_max_in_memory_cache_size: Optional[int] = None +global_tag_rate_limits: Optional["TagRateLimits"] = None +global_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 4f7a7510f96..4003b5f4d9a 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -4496,6 +4496,23 @@ def _init_custom_logger_compatible_class( 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 == "global_tag_rate_limits_hook": + from litellm.proxy.hooks.global_tag_rate_limits_hook import ( + _PROXY_GlobalTagRateLimitsHook, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here + ) + + for callback in _in_memory_loggers: + if isinstance(callback, _PROXY_GlobalTagRateLimitsHook): + return callback + + if internal_usage_cache is None: + raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}") + + global_tag_rate_limits_hook_obj: Final = _PROXY_GlobalTagRateLimitsHook( + internal_usage_cache=internal_usage_cache + ) + _in_memory_loggers.append(global_tag_rate_limits_hook_obj) + return global_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") @@ -4945,6 +4962,15 @@ def get_custom_logger_compatible_class( if isinstance(callback, _PROXY_ModelBasedTagRateLimitsHook): return callback + elif logging_integration == "global_tag_rate_limits_hook": + from litellm.proxy.hooks.global_tag_rate_limits_hook import ( + _PROXY_GlobalTagRateLimitsHook, # pyright: ignore[reportPrivateUsage] # resolved by name like every other opt-in callback here + ) + + for callback in _in_memory_loggers: + if isinstance(callback, _PROXY_GlobalTagRateLimitsHook): + return callback + elif logging_integration == "langtrace": from litellm.integrations.opentelemetry import OpenTelemetry diff --git a/litellm/proxy/hooks/global_tag_rate_limits_hook.py b/litellm/proxy/hooks/global_tag_rate_limits_hook.py new file mode 100644 index 00000000000..be90e2ff2a4 --- /dev/null +++ b/litellm/proxy/hooks/global_tag_rate_limits_hook.py @@ -0,0 +1,613 @@ +""" +Tag-scoped token, request, dollar, and concurrency rate limits declared once, +globally, in `litellm_settings.global_tag_rate_limits` -- enforced once per +request in `async_pre_call_hook`, before Router does any routing, so a limit +applies regardless of which model or fallback chain the request ends up +hitting. + +This is the model-independent sibling of `model_based_tag_rate_limits_hook`, +which enforces the same `TagRateLimitEntry` shape but nested per-deployment +under `model_info.tag_rate_limits`, once per routing hop +(`async_filter_deployments`). A global entry has no deployment/routing-group +to reconcile -- there is exactly one config value, read once -- so this hook +reuses that sibling's free, already-hardened helper functions +(`_entry_applies`, the Lua atomic check-and-increment scripts, cache +partitioning, bucket-key hashing primitives) directly rather than duplicating +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 +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. +- `scope_by_key_hash` (already exists on `TagRateLimitEntry`): whether the + keys an entry applies to share one bucket, or each gets its own. + +`async_pre_call_hook` runs before Router constructs `Logging`/`litellm_logging_obj` +for this request (see `common_request_processing.py`: `pre_call_hook` fires +well before `base_process_llm_request` builds the logging object), so unlike +`model_based_tag_rate_limits_hook` this hook cannot stash pending concurrency +reservations on `data["litellm_logging_obj"].model_call_details` -- that +object doesn't exist yet. Per-request state is instead kept on a +`ContextVar`-based stash, the same established pattern +`parallel_request_limiter_v3.py`'s v3 handler already uses for exactly this +problem: the ContextVar is inherited by every asyncio Task forked from this +request's own task (the SDK call, streaming generators, the logging worker), +so concurrent requests never see each other's stash regardless of a +caller-supplied `litellm_call_id` colliding, and `owner_litellm_call_id` only +exists to tell a nested LiteLLM call (e.g. a guardrail's own LLM judge call) +apart from the owning request. +""" + +import asyncio +from collections.abc import Callable, Mapping, Sequence +from contextvars import ContextVar +from dataclasses import dataclass, field +from datetime import datetime +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, TypeAlias + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.model_based_tag_rate_limits_hook import ( + _ATOMIC_UNITS, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _BACKGROUND_TASKS, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _CONCURRENCY_MIN_SAFETY_TTL_SECONDS, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _EMPTY_MAPPING, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _LIMIT_UNITS, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _UNIT_TO_GROUP_FIELD, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _UNIT_TO_RATE_LIMIT_TYPE, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + TAG_RL_CHECK_AND_INCR_SCRIPT, + TAG_RL_DECR_FLOOR_ZERO_SCRIPT, + _bucket_ttl_seconds, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _entry_applies, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _extract_identity, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _extract_key_alias, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _extract_key_hash, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _fixed_length_identity, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _LimitUnit, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _partition_key, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _PartitionKey, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _PartitionOperations, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring + _policy_fingerprint, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring +) +from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, matching model_based_tag_rate_limits_hook's identical import +) +from litellm.proxy.utils import InternalUsageCache +from litellm.router_strategy.tag_based_routing import ( + _get_tags_from_request_kwargs, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, matching model_based_tag_rate_limits_hook's identical import +) +from litellm.types.caching import RedisPipelineIncrementOperation +from litellm.types.router import TagRateLimitEntry, TagRateLimits +from litellm.types.utils import StandardLoggingPayload + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + + Span: TypeAlias = _Span +else: + Span: TypeAlias = object + + +def _hash_tag(entry: TagRateLimitEntry, unit: _LimitUnit, tag_value: str, key_hash: str | None) -> str: + """ + Global-hook equivalent of `model_based_tag_rate_limits_hook._hash_tag`, + without a `model_group`/deployment-scope/team-scope dimension -- a global + entry has none of those. Namespaced under `tag_rl:global:` so it can never + collide with that sibling hook's own `tag_rl:{model_group}:...` keys even + if an operator names a deployment "global": every key also differs by + `unit`/`name`/`tag_id`/`_policy_fingerprint`, and the two hooks' entries + are never meant to share a bucket in the first place. + """ + key_suffix: Final = f":key:{key_hash}" if key_hash is not None else "" + policy_suffix: Final = f":policy:{_policy_fingerprint(entry)}" + return f"tag_rl:global:{unit}:{entry.name}:{entry.tag_id}:{_fixed_length_identity(tag_value)}{key_suffix}{policy_suffix}" + + +def _bucket_key( + entry: TagRateLimitEntry, unit: _LimitUnit, tag_value: str, bucket_id: int, key_hash: str | None +) -> str: + return f"{{{_hash_tag(entry, unit, tag_value, key_hash)}}}:{bucket_id}" + + +def _inflight_key(entry: TagRateLimitEntry, unit: _LimitUnit, tag_value: str, key_hash: str | None) -> str: + return f"{{{_hash_tag(entry, unit, tag_value, key_hash)}}}:inflight" + + +@dataclass(frozen=True, slots=True) +class _ClassifiedGlobalCheck: + unit: _LimitUnit + entry: TagRateLimitEntry + tag_value: str + key: str + is_atomic: bool + + +@dataclass(frozen=True, slots=True) +class _CachePartition: + internal_usage_cache: InternalUsageCache + v3: _PROXY_MaxParallelRequestsHandler_v3 + + +@dataclass(slots=True) +class _GlobalTagRateLimitStash: + """Per-request bookkeeping `async_pre_call_hook` hands to the success/ + failure/disconnect callbacks -- see module docstring for why this is a + `ContextVar`, not `model_call_details`.""" + + owner_litellm_call_id: str | None = None + admission_time: float | None = None + pending_concurrency_keys: list[tuple[str, _PartitionKey]] = field(default_factory=list) # mutable-ok: queue + + +_request_stash: Final[ContextVar[_GlobalTagRateLimitStash | None]] = ContextVar( + "global_tag_rate_limits_request_stash", default=None +) + + +def _claim_stash_for_data(data: Mapping[str, object]) -> _GlobalTagRateLimitStash: + stash = _request_stash.get() + if stash is None: + stash = _GlobalTagRateLimitStash() + _request_stash.set(stash) + owner_call_id: Final = data.get("litellm_call_id") + if isinstance(owner_call_id, str): + stash.owner_litellm_call_id = owner_call_id + return stash + + +def _stash_for_call(litellm_call_id: str | None) -> _GlobalTagRateLimitStash | None: + stash: Final = _request_stash.get() + if stash is None: + return None + if stash.owner_litellm_call_id is None or litellm_call_id is None: + return stash + return stash if litellm_call_id == stash.owner_litellm_call_id else None + + +def _call_id_from_kwargs(kwargs: object) -> str | None: + if not isinstance(kwargs, dict): + return None + call_id: Final = kwargs.get("litellm_call_id") + return call_id if isinstance(call_id, str) else None + + +def _resolve_max_in_memory_cache_size() -> int | None: + """Same shape as `model_based_tag_rate_limits_hook`'s own function, reading + this hook's own `litellm_settings` knob instead.""" + configured: Final = litellm.global_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( + "global_tag_rate_limits_hook: global_tag_rate_limits_max_in_memory_cache_size=%r is not a positive " + "integer; falling back to the default in-memory cache size.", + configured, + ) + return None + + +class _PROXY_GlobalTagRateLimitsHook( # 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__( + self, + internal_usage_cache: DualCache, + time_provider: Callable[[], datetime] | None = None, + ) -> None: + self._redis_cache: Final = internal_usage_cache.redis_cache + self._time_provider = time_provider or datetime.now + self._partitions: dict[_PartitionKey, _CachePartition] = {} # mutable-ok: lazily memoized; see _partition_for + self._partitions_lock = asyncio.Lock() + default_partition: Final = self._build_partition(_resolve_max_in_memory_cache_size()) + self._partitions[None] = default_partition + self.internal_usage_cache = default_partition.internal_usage_cache + self._lock = asyncio.Lock() + redis_cache: Final = self._redis_cache + self._check_and_incr_script = ( + redis_cache.async_register_script(TAG_RL_CHECK_AND_INCR_SCRIPT) if redis_cache is not None else None + ) + self._decr_floor_zero_script = ( + redis_cache.async_register_script(TAG_RL_DECR_FLOOR_ZERO_SCRIPT) if redis_cache is not None else None + ) + self._config_cache_key: object | None = None + self._config: TagRateLimits | None = None + + def _refresh_config(self) -> TagRateLimits | None: + """Re-validates `litellm.global_tag_rate_limits` whenever the object + identity changes (a config reload replaces it wholesale via + `setattr(litellm, key, value)`), so a hot-reloaded config takes effect + on the very next request with no staleness window and no TTL to tune.""" + raw: Final = getattr(litellm, "global_tag_rate_limits", None) + if raw is not self._config_cache_key: + self._config = TagRateLimits.model_validate(raw) if raw else None + self._config_cache_key = raw + return self._config + + def _build_partition(self, cache_size_override: int | None) -> _CachePartition: + dual_cache: Final = DualCache( + in_memory_cache=InMemoryCache(max_size_in_memory=cache_size_override), + redis_cache=self._redis_cache, + ) + cache: Final = InternalUsageCache(dual_cache=dual_cache) + return _CachePartition( + internal_usage_cache=cache, + v3=_PROXY_MaxParallelRequestsHandler_v3(cache, time_provider=self._time_provider), + ) + + async def _partition_for(self, partition_key: _PartitionKey) -> _CachePartition: + existing: Final = self._partitions.get(partition_key) + if existing is not None: + return existing + async with self._partitions_lock: + existing_after_lock: Final = self._partitions.get(partition_key) + if existing_after_lock is not None: + return existing_after_lock + cache_size_override: Final = partition_key[-1] if partition_key is not None else None + built: Final = self._build_partition(cache_size_override) + self._partitions[partition_key] = ( + built # mutable-ok: lazily memoized per distinct partition key, guarded by _partitions_lock above + ) + return built + + async def _check_and_increment_one( + self, cache: InternalUsageCache, key: str, limit: float, increment: float, ttl: int + ) -> tuple[bool, float]: + if self._check_and_incr_script is not None: + raw: Final = await self._check_and_incr_script(keys=(key,), args=(limit, increment, ttl)) + return bool(raw[0]), float(raw[1]) + async with self._lock: + current_value: Final = await cache.async_get_cache(key=key, litellm_parent_otel_span=None) + current: Final = float(current_value) if current_value is not None else 0.0 + if current + increment > limit: + return False, current + new_value: Final = current + increment + await cache.async_set_cache(key=key, value=new_value, ttl=ttl, litellm_parent_otel_span=None) + return True, new_value + + async def _decrement_floor_zero(self, cache: InternalUsageCache, key: str, delta: float) -> None: + if self._decr_floor_zero_script is not None: + await self._decr_floor_zero_script(keys=(key,), args=(delta,)) + return + async with self._lock: + current_value: Final = await cache.async_get_cache(key=key, litellm_parent_otel_span=None) + current: Final = float(current_value) if current_value is not None else 0.0 + await cache.async_set_cache(key=key, value=max(0.0, current + delta), litellm_parent_otel_span=None) + + async def _atomic_check_and_increment( + self, + checks: Sequence[tuple[InternalUsageCache, str, float, float, int]], + ) -> tuple[int | None, tuple[float, ...]]: + """All-or-nothing atomic admission across `checks` -- see + `model_based_tag_rate_limits_hook._PROXY_ModelBasedTagRateLimitsHook._atomic_check_and_increment`'s + own docstring for the full rationale (refund-on-rollback, why a + raising key's own outcome is never refunded); identical logic, + duplicated rather than shared since it lives as instance methods + rather than free functions.""" + if not checks: + return None, () + admitted_values: Final = [] # mutable-ok: sequential async accumulator, discardable on early rejection + for index, (cache, key, limit, increment, ttl) in enumerate(checks): + admitted = False + try: + admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl) + finally: + if not admitted: + await self._refund_admitted(checks, up_to_index=index) + if admitted: + admitted_values.append(value) # mutable-ok: see accumulator comment above + continue + return index, (value,) + return None, tuple(admitted_values) + + async def _refund_admitted( + self, checks: Sequence[tuple[InternalUsageCache, str, float, float, int]], up_to_index: int + ) -> None: + for refund_index in range(up_to_index): + refund_cache, refund_key, _limit, refund_increment, _ttl = checks[refund_index] + 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( + "global_tag_rate_limits_hook: failed to refund %s on rollback: %s", refund_key, e + ) + + async def _release_keys(self, reservations: Sequence[tuple[str, _PartitionKey]]) -> None: + for key, partition_key in reservations: + try: + 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( + "global_tag_rate_limits_hook: failed to release concurrency slot %s: %s", key, e + ) + + @staticmethod + def _ttl_for(unit: _LimitUnit, entry: TagRateLimitEntry) -> int: + if unit == "concurrency": + requested_ttl: Final = entry.key_ttl_seconds if entry.key_ttl_seconds is not None else entry.period_seconds + return max(requested_ttl, _CONCURRENCY_MIN_SAFETY_TTL_SECONDS) + return _bucket_ttl_seconds(entry) + + def _classify( + self, config: TagRateLimits, tags: Sequence[str], key_alias: str | None, key_hash: str | None, now: float + ) -> tuple[_ClassifiedGlobalCheck, ...]: + classified: Final = [] # mutable-ok: sequential accumulator, immediately frozen into a tuple below + for unit in _LIMIT_UNITS: + group = getattr(config, _UNIT_TO_GROUP_FIELD[unit]) + if group is None: + continue + for entry in group.limits: + tag_value = _extract_identity(tags, entry.tag_id) + if tag_value is None: + continue + if not _entry_applies(entry, tag_value, tags, key_alias): + continue + effective_key_hash = key_hash if entry.scope_by_key_hash else None + if unit == "concurrency": + key = _inflight_key(entry, unit, tag_value, key_hash=effective_key_hash) + classified.append( + _ClassifiedGlobalCheck(unit, entry, tag_value, key, is_atomic=True) + ) # mutable-ok: see comment above + continue + bucket_id = int(now) // entry.period_seconds + key = _bucket_key(entry, unit, tag_value, bucket_id, key_hash=effective_key_hash) + classified.append( # mutable-ok: see comment above + _ClassifiedGlobalCheck(unit, entry, tag_value, key, is_atomic=unit in _ATOMIC_UNITS) + ) + return tuple(classified) + + async def _read_only_values( + self, read_only_checks: Sequence[_ClassifiedGlobalCheck], parent_otel_span: Span | None + ) -> tuple[float | None, ...]: + if not read_only_checks: + return () + indices_by_partition: Final[dict[_PartitionKey, list[int]]] = {} # mutable-ok: grouped, reassembled below + for index, check in enumerate(read_only_checks): + partition_key = _partition_key(check.entry) + indices = indices_by_partition.setdefault(partition_key, []) # mutable-ok: see above + indices.append(index) # mutable-ok: see comment above + values_by_index: Final[dict[int, float | None]] = {} # mutable-ok: see comment above + for partition_key, indices in indices_by_partition.items(): + partition = await self._partition_for(partition_key) # not Final: rebound each loop iteration + keys = [read_only_checks[i].key for i in indices] # mutable-ok: async_batch_get_cache needs a real list + redis_cache = partition.internal_usage_cache.dual_cache.redis_cache + if redis_cache is not None: + redis_values: Mapping[str, object] = await redis_cache.async_batch_get_cache( + key_list=keys, parent_otel_span=parent_otel_span + ) + resolved = [redis_values.get(key) for key in keys] # mutable-ok: needs a real list + else: + current_values = await partition.internal_usage_cache.async_batch_get_cache( + keys=keys, parent_otel_span=parent_otel_span, local_only=True + ) + missing = [None] * len(keys) # mutable-ok: async_batch_get_cache requires a real list; see above + resolved = current_values if current_values is not None else missing + for i, value in zip(indices, resolved): + values_by_index[i] = value # mutable-ok: see comment above + return tuple(values_by_index[i] for i in range(len(read_only_checks))) + + def _raise_if_over_limit( + self, + read_only_checks: Sequence[_ClassifiedGlobalCheck], + current_values: Sequence[float | None], + model: str | None, + ) -> None: + for check, current_value in zip(read_only_checks, current_values): + current = float(current_value) if current_value is not None else 0.0 + if current < check.entry.limit: + continue + self._raise_over_limit(check.unit, check.entry, check.tag_value, model, current=current) + + def _raise_over_limit( + self, unit: _LimitUnit, entry: TagRateLimitEntry, tag_value: str, model: str | None, current: float + ) -> None: + verbose_proxy_logger.debug( + "global_tag_rate_limits_hook: OVER_LIMIT unit=%s name=%s tag_id=%s tag_value=%s current=%s limit=%s", + unit, + entry.name, + entry.tag_id, + tag_value, + current, + entry.limit, + ) + raise ProxyRateLimitError( + detail={ # mutable-ok: async_log_failure_event and generic proxy exception rendering branch on isinstance(exc.detail, dict) + "error": "tag_rate_limit_exceeded", + "type": unit, + "tag_id": entry.tag_id, + "tag_value": tag_value, + "limit_name": entry.name, + "limit": entry.limit, + "period_seconds": entry.period_seconds, + }, + headers={"retry-after": str(entry.period_seconds)}, # mutable-ok: same as detail + rate_limit_type=_UNIT_TO_RATE_LIMIT_TYPE[unit], + model=model, + llm_provider="litellm_proxy", + ) + + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: DualCache, + data: dict, # mutable-ok: must match CustomLogger.async_pre_call_hook's own base signature exactly + call_type: str, + ) -> dict: # mutable-ok: must match CustomLogger.async_pre_call_hook's own base signature exactly + config: Final = self._refresh_config() + if config is None: + return data + + # Unlike model_based_tag_rate_limits_hook's async_filter_deployments + # (called once per routing hop, so a still-queued reservation can + # legitimately belong to an earlier, already-failed hop of the same + # request), async_pre_call_hook fires exactly once per request -- + # there is no "prior hop" case here, so no stale-reservation release + # is needed at the top of admission. + stash: Final = _claim_stash_for_data(data) + + metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(data) + 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 + + now: Final = self._time_provider().timestamp() + stash.admission_time = now + classified: Final = self._classify(config, tags, key_alias, key_hash, now) + 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) + + if atomic_checks: + atomic_partitions_list: Final = [] # mutable-ok: sequential async lookups, one per atomic_checks entry + for check in atomic_checks: + atomic_partitions_list.append( + await self._partition_for(_partition_key(check.entry)) + ) # mutable-ok: see comment above + atomic_partitions: Final = tuple(atomic_partitions_list) + failing_index, values = await self._atomic_check_and_increment( + tuple( + ( + partition.internal_usage_cache, + check.key, + check.entry.limit, + 1.0, + self._ttl_for(check.unit, check.entry), + ) + for partition, check in zip(atomic_partitions, atomic_checks) + ) + ) + if failing_index is not None: + failing_check: Final = atomic_checks[failing_index] + self._raise_over_limit( + failing_check.unit, failing_check.entry, failing_check.tag_value, model, current=values[0] + ) + + concurrency_reservations: Final = tuple( + (check.key, _partition_key(check.entry)) for check in atomic_checks if check.unit == "concurrency" + ) + if concurrency_reservations: + stash.pending_concurrency_keys.extend(concurrency_reservations) # mutable-ok: see field's own docstring + + return data + + async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None: + stash: Final = _stash_for_call(_call_id_from_kwargs(dict(request_data))) + if stash is None or not stash.pending_concurrency_keys: + return + release_keys: Final = tuple(stash.pending_concurrency_keys) + stash.pending_concurrency_keys.clear() + await self._release_keys(release_keys) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: + if isinstance(kwargs.get("exception"), ProxyRateLimitError): + detail: Final = ( + kwargs["exception"].detail if isinstance(kwargs["exception"].detail, dict) else _EMPTY_MAPPING + ) + if detail.get("error") == "tag_rate_limit_exceeded": + return + + stash: Final = _stash_for_call(_call_id_from_kwargs(kwargs)) + if stash is None or not stash.pending_concurrency_keys: + return + release_keys: Final = tuple(stash.pending_concurrency_keys) + stash.pending_concurrency_keys.clear() + await self._release_keys(release_keys) + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + stash: Final = _stash_for_call(_call_id_from_kwargs(kwargs)) + if stash is not None and stash.pending_concurrency_keys: + release_keys: Final = tuple(stash.pending_concurrency_keys) + stash.pending_concurrency_keys.clear() + release_task: Final = asyncio.create_task(self._release_keys(release_keys)) + _BACKGROUND_TASKS.add(release_task) # mutable-ok: see _BACKGROUND_TASKS's own docstring + release_task.add_done_callback(_BACKGROUND_TASKS.discard) + + config: Final = self._refresh_config() + if config is None: + return + + standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") + if standard_logging_object is None: + return + + litellm_params_for_metadata: Final = kwargs.get("litellm_params") or kwargs + metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(litellm_params_for_metadata) + key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name) + key_alias: Final = _extract_key_alias(litellm_params_for_metadata, metadata_variable_name) + + tags: Final = _get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name) + if not tags: + return + + now: Final = ( + stash.admission_time + if stash is not None and stash.admission_time is not None + else self._time_provider().timestamp() + ) + increment_by_unit: Final[Mapping[_LimitUnit, float]] = MappingProxyType( + { + "tokens": float(standard_logging_object.get("total_tokens") or 0), + "dollars": float(standard_logging_object.get("response_cost") or 0), + } + ) + + operation_by_entry: Final = [] # mutable-ok: sequential accumulator over config groups, immediately used below + for unit in ("tokens", "dollars"): + group = getattr(config, _UNIT_TO_GROUP_FIELD[unit]) + if group is None: + continue + for entry in group.limits: + tag_value = _extract_identity(tags, entry.tag_id) + if tag_value is None: + continue + if not _entry_applies(entry, tag_value, tags, key_alias): + continue + increment_value = increment_by_unit[unit] + if increment_value == 0: + continue + bucket_id = int(now) // entry.period_seconds + key_hash_for_entry = key_hash if entry.scope_by_key_hash else None + key = _bucket_key(entry, unit, tag_value, bucket_id, key_hash=key_hash_for_entry) + operation_by_entry.append( # mutable-ok: see comment above + ( + entry, + RedisPipelineIncrementOperation( + key=key, increment_value=increment_value, ttl=_bucket_ttl_seconds(entry) + ), + ) + ) + + if not operation_by_entry: + return + + operations_by_partition: Final[_PartitionOperations] = {} # mutable-ok: grouped by cache partition below + for entry, operation in operation_by_entry: + partition_key = _partition_key(entry) + operations = operations_by_partition.setdefault(partition_key, []) # mutable-ok: see above + operations.append(operation) # mutable-ok: see comment above + + for partition_key, group_operations in operations_by_partition.items(): + partition = await self._partition_for(partition_key) # not Final: rebound each loop iteration + accounting_task = asyncio.create_task( # not Final: rebound each loop iteration + partition.v3.async_increment_tokens_with_ttl_preservation( + pipeline_operations=tuple(group_operations), parent_otel_span=None + ) + ) + _BACKGROUND_TASKS.add(accounting_task) # mutable-ok: see _BACKGROUND_TASKS's own docstring + accounting_task.add_done_callback(_BACKGROUND_TASKS.discard) 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 05e696bd33c..cb7c14da99e 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -47,15 +47,24 @@ _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, 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`. +# excluded_values, 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`. _DedupSignature: TypeAlias = tuple[ - str, str, float, int, bool, tuple[str, ...] | None, tuple[str, ...] | None, _ScopeSignature, _ScopeSignature + str, + str, + float, + int, + bool, + tuple[str, ...] | None, + tuple[str, ...] | None, + _ScopeSignature, + _ScopeSignature, + 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). @@ -204,11 +213,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, tag_value: str, tags: Sequence[str]) -> bool: +def _entry_applies(entry: TagRateLimitEntry, tag_value: str, tags: Sequence[str], key_alias: str | None) -> 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: + `enabled_for`/`disabled_for`/`apply_to_key_alias`), 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. @@ -220,8 +229,12 @@ def _entry_applies(entry: TagRateLimitEntry, tag_value: str, tags: Sequence[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". + 5. `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. - An entry with none of the four fields set always applies -- this is the + An entry with none of these 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: @@ -236,7 +249,9 @@ def _entry_applies(entry: TagRateLimitEntry, tag_value: str, tags: Sequence[str] 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 + if entry.apply_to_key_alias is None: + return True + return key_alias in entry.apply_to_key_alias def _deployment_id(deployment: Mapping[str, object]) -> str | None: @@ -269,6 +284,16 @@ def _extract_key_hash(request_kwargs: Mapping[str, object], metadata_variable_na return key_hash if isinstance(key_hash, str) else None +def _extract_key_alias(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None: + """Same single-authoritative-field lookup as `_extract_team_id`, but for + the calling virtual key's own `key_alias`: `LiteLLMProxyRequestSetup` sets + `metadata["user_api_key_alias"]` to `user_api_key_dict.key_alias` + (see `litellm_pre_call_utils.py`).""" + active: Final = request_kwargs.get(metadata_variable_name) or _EMPTY_MAPPING + key_alias: Final = active.get("user_api_key_alias") + return key_alias if isinstance(key_alias, str) else None + + def _entries_for_unit(deployment: Mapping[str, object], unit: _LimitUnit) -> tuple[TagRateLimitEntry, ...]: raw_tag_rate_limits: Final = (deployment.get("model_info") or _EMPTY_MAPPING).get("tag_rate_limits") if not raw_tag_rate_limits: @@ -357,6 +382,7 @@ def _build_group_limits(deployments: Sequence[Mapping[str, object]], unit: _Limi entry.excluded_values, _scope_signature(entry.enabled_for), _scope_signature(entry.disabled_for), + entry.apply_to_key_alias, ) ids_for_signature = declaring_ids_by_signature.setdefault(signature, []) # mutable-ok: see comment above # One deployment declaring the identical entry twice (a config @@ -472,6 +498,7 @@ class _LimitsIndex: limit.entry.excluded_values, _scope_signature(limit.entry.enabled_for), _scope_signature(limit.entry.disabled_for), + limit.entry.apply_to_key_alias, limit.deployment_scope, limit.team_scope, ) @@ -679,6 +706,7 @@ def _policy_fingerprint(entry: TagRateLimitEntry) -> str: entry.excluded_values, _scope_signature(entry.enabled_for), _scope_signature(entry.disabled_for), + entry.apply_to_key_alias, ) return hashlib.sha256(repr(fingerprint_source).encode()).hexdigest()[:16] @@ -745,6 +773,7 @@ def _classify_check( request_kwargs: Mapping[str, object], metadata_variable_name: str, now: float, + key_alias: str | None, ) -> _ClassifiedCheck | None: if configured_limit.deployment_scope is not None and not ( present_deployment_ids & frozenset(configured_limit.deployment_scope) @@ -753,7 +782,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, tag_value, tags): + if not _entry_applies(configured_limit.entry, tag_value, tags, key_alias): return None key_hash: Final = ( _extract_key_hash(request_kwargs, metadata_variable_name) if configured_limit.entry.scope_by_key_hash else None @@ -781,6 +810,7 @@ def _increment_operation_for_limit( tags: Sequence[str], deployment_id: str | None, key_hash: str | None, + key_alias: str | None, increment_by_unit: Mapping[_LimitUnit, float], now: float, ) -> RedisPipelineIncrementOperation | None: @@ -791,7 +821,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, tag_value, tags): + if not _entry_applies(configured_limit.entry, tag_value, tags, key_alias): return None if configured_limit.unit not in increment_by_unit: return None # "requests" is accounted atomically at admission, not here @@ -1139,6 +1169,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] dep_id for d in healthy_deployments if (dep_id := _deployment_id(d)) is not None ) + key_alias: Final = _extract_key_alias(resolved_request_kwargs, metadata_variable_name) now: Final = self._time_provider().timestamp() _record_admission_time(resolved_request_kwargs, now) classified: Final = tuple( @@ -1153,6 +1184,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] resolved_request_kwargs, metadata_variable_name, now, + key_alias, ) ) is not None @@ -1460,6 +1492,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # here instead keeps this bucket identical to the one admission # already scoped the check against. key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name) + key_alias: Final = _extract_key_alias(litellm_params_for_metadata, metadata_variable_name) # model_group is the caller-visible name, which Router deliberately # keeps distinct from the serving deployment's own model_name for a # routing-group call (see resolve_any's docstring). Passing only the @@ -1516,7 +1549,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass] for configured_limit in configured if ( operation := _increment_operation_for_limit( - configured_limit, model_group, tags, deployment_id, key_hash, increment_by_unit, now + configured_limit, model_group, tags, deployment_id, key_hash, key_alias, increment_by_unit, now ) ) is not None diff --git a/litellm/types/router.py b/litellm/types/router.py index 5ad1f09adf7..941abf79344 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -213,6 +213,12 @@ class TagRateLimitEntry(BaseModel): # either (nothing to match against a denylist). enabled_for: TagRateLimitScope | None = None disabled_for: TagRateLimitScope | None = None + # Restrict this entry to requests authenticated with one of these virtual + # keys' own `key_alias`. Unset (the default) means the entry applies to + # every request regardless of which key made it. A key with no alias set + # never satisfies this allowlist, same "absent gate never matches an + # allowlist" precedent as `enabled_for`. + apply_to_key_alias: tuple[str, ...] | None = None model_config = ConfigDict(protected_namespaces=()) @@ -263,6 +269,8 @@ class TagRateLimitEntry(BaseModel): 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") + if self.apply_to_key_alias is not None and not self.apply_to_key_alias: + raise ValueError("apply_to_key_alias must be a non-empty list of strings when set") return self @model_validator(mode="after") @@ -276,6 +284,8 @@ class TagRateLimitEntry(BaseModel): self.included_values = tuple(sorted(set(self.included_values))) if self.excluded_values is not None: self.excluded_values = tuple(sorted(set(self.excluded_values))) + if self.apply_to_key_alias is not None: + self.apply_to_key_alias = tuple(sorted(set(self.apply_to_key_alias))) return self 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 new file mode 100644 index 00000000000..934f177a352 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_global_tag_rate_limits_hook.py @@ -0,0 +1,505 @@ +""" +Unit tests for the global-scope, model-independent tag rate limiter. +""" + +import asyncio +from datetime import datetime, timedelta + +import pytest +from pydantic import ValidationError + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.global_tag_rate_limits_hook import ( + _PROXY_GlobalTagRateLimitsHook, +) + + +class TimeController: + def __init__(self): + self._current = datetime(2026, 1, 1, 0, 0, 0) + + def now(self) -> datetime: + return self._current + + def advance(self, seconds: float) -> None: + self._current += timedelta(seconds=seconds) + + +@pytest.fixture +def time_controller(): + return TimeController() + + +def _make_hook(time_controller: TimeController) -> _PROXY_GlobalTagRateLimitsHook: + return _PROXY_GlobalTagRateLimitsHook( + internal_usage_cache=DualCache(), + time_provider=time_controller.now, + ) + + +def _key(alias: str | None = None, api_key: str = "hash") -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key=api_key, key_alias=alias) + + +def _data(tags: list[str], call_id: str = "call-1") -> dict: + return {"metadata": {"tags": tags}, "litellm_call_id": call_id} + + +# --------------------------------------------------------------------------- +# No-op when unconfigured +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_no_op_when_no_config_set(time_controller, monkeypatch): + monkeypatch.setattr(litellm, "global_tag_rate_limits", None) + hook = _make_hook(time_controller) + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion" + ) + assert result == _data(["end_user_id:u1"]) + + +@pytest.mark.asyncio +async def test_malformed_config_raises_at_first_use(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, "global_tag_rate_limits", {"dollar_limits": {"limits": [{"name": "bad", "limit": "not-a-number"}]}} + ) + hook = _make_hook(time_controller) + with pytest.raises(ValidationError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion" + ) + + +# --------------------------------------------------------------------------- +# Global scope: applies to every key by default +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_request_limit_shared_across_keys_by_default(time_controller, monkeypatch): + """No apply_to_key_alias -> the entry is one shared bucket regardless of + which key made the request.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, + ) + hook = _make_hook(time_controller) + + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="key-a"), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion" + ) + # A different key, identical tag value: must be rejected too -- proves + # the bucket is genuinely shared, not per-key by default. + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="key-b"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_request_limit_is_independent_of_model(time_controller, monkeypatch): + """The hook never reads `data["model"]` for identity -- two different + "models" (irrelevant to this hook) must still share the same bucket.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, + ) + hook = _make_hook(time_controller) + + data_model_a = {**_data(["end_user_id:u1"]), "model": "gpt-4o"} + data_model_b = {**_data(["end_user_id:u1"]), "model": "claude-3"} + await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=data_model_a, call_type="completion" + ) + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=data_model_b, call_type="completion" + ) + + +# --------------------------------------------------------------------------- +# apply_to_key_alias -- narrows which keys an entry applies to +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_apply_to_key_alias_ignores_non_matching_keys(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "daily", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_key_alias": ["premium-key"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + for _ in range(3): + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="other-key"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + assert result is not None + + +@pytest.mark.asyncio +async def test_apply_to_key_alias_enforces_for_the_listed_key(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "daily", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_key_alias": ["premium-key"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_apply_to_key_alias_composes_with_scope_by_key_hash(time_controller, monkeypatch): + """Both listed keys are subject to the entry, but scope_by_key_hash + splits their buckets: exhausting one must not affect the other.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [ + { + "name": "daily", + "tag_id": "end_user_id", + "limit": 1, + "period_seconds": 86400, + "apply_to_key_alias": ["key-a", "key-b"], + "scope_by_key_hash": True, + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="key-a", api_key="hashA"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + with pytest.raises(ProxyRateLimitError): + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="key-a", api_key="hashA"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + # key-b is unaffected by key-a's exhausted bucket. + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="key-b", api_key="hashB"), + cache=DualCache(), + data=_data(["end_user_id:u1"]), + call_type="completion", + ) + assert result is not None + + +# --------------------------------------------------------------------------- +# Concurrency: reservation at admission, release on success/failure/disconnect +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_concurrency_limit_rejects_second_admission_until_release(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "concurrency_limits": { + "limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}] + } + }, + ) + hook = _make_hook(time_controller) + + await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data=_data(["end_user_id:u1"], call_id="call-1"), + call_type="completion", + ) + 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"), + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_concurrency_reservation_released_on_success(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "concurrency_limits": { + "limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}] + } + }, + ) + hook = _make_hook(time_controller) + + async def one_request(call_id: str) -> None: + data = _data(["end_user_id:u1"], call_id=call_id) + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion") + kwargs = {"litellm_call_id": call_id, "metadata": {"tags": ["end_user_id:u1"]}} + await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) + + await one_request("call-1") + await asyncio.sleep(0) # let the fire-and-forget release task run + + # The slot was released, so a fresh request must be admitted again. + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data=_data(["end_user_id:u1"], call_id="call-2"), + call_type="completion", + ) + assert result is not None + + +@pytest.mark.asyncio +async def test_concurrency_reservation_released_on_disconnect(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "concurrency_limits": { + "limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}] + } + }, + ) + hook = _make_hook(time_controller) + + data = _data(["end_user_id:u1"], call_id="call-1") + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion") + await hook.async_release_disconnect_state_hook({"litellm_call_id": "call-1"}) + + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data=_data(["end_user_id:u1"], call_id="call-2"), + call_type="completion", + ) + assert result is not None + + +@pytest.mark.asyncio +async def test_concurrent_requests_do_not_share_each_others_reservation_state(time_controller, monkeypatch): + """Two logically distinct requests running as separate asyncio Tasks must + not see each other's pending-concurrency stash, even though both share + this hook instance -- the whole point of the ContextVar-based stash.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "concurrency_limits": { + "limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}] + } + }, + ) + hook = _make_hook(time_controller) + + async def one_request(call_id: str) -> int: + data = _data(["end_user_id:u1"], call_id=call_id) + await hook.async_pre_call_hook(user_api_key_dict=_key(), cache=DualCache(), data=data, call_type="completion") + kwargs = {"litellm_call_id": call_id, "metadata": {"tags": ["end_user_id:u1"]}} + await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) + return 1 + + results = await asyncio.gather(one_request("call-a"), one_request("call-b")) + await asyncio.sleep(0) + assert results == [1, 1] + + +# --------------------------------------------------------------------------- +# Accounting: tokens/dollars via async_log_success_event +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_dollar_limit_accounts_usage_and_rejects_once_over(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "dollar_limits": { + "limits": [{"name": "daily_spend", "tag_id": "end_user_id", "limit": 10.0, "period_seconds": 86400}] + } + }, + ) + hook = _make_hook(time_controller) + + data = _data(["end_user_id:u1"], call_id="call-1") + 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": 12.0}, + } + await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) + await asyncio.sleep(0) + + 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"), + call_type="completion", + ) + + +@pytest.mark.asyncio +async def test_dollar_limit_respects_apply_to_key_alias_at_accounting_time(time_controller, monkeypatch): + """The entry only applies to `premium-key`; a non-listed key's spend must + not be charged against this bucket at all.""" + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "dollar_limits": { + "limits": [ + { + "name": "daily_spend", + "tag_id": "end_user_id", + "limit": 10.0, + "period_seconds": 86400, + "apply_to_key_alias": ["premium-key"], + } + ] + } + }, + ) + hook = _make_hook(time_controller) + + data = _data(["end_user_id:u1"], call_id="call-1") + await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="other-key"), cache=DualCache(), data=data, call_type="completion" + ) + kwargs = { + "litellm_call_id": "call-1", + "metadata": {"tags": ["end_user_id:u1"], "user_api_key_alias": "other-key"}, + "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) + + # premium-key was never charged -- still fully under its own limit. + result = await hook.async_pre_call_hook( + user_api_key_dict=_key(alias="premium-key"), + cache=DualCache(), + data=_data(["end_user_id:u1"], call_id="call-2"), + call_type="completion", + ) + assert result is not None + + +# --------------------------------------------------------------------------- +# Config hot-reload: identity-based re-validation, no restart needed +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_config_reload_takes_effect_on_next_request(time_controller, monkeypatch): + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 100, "period_seconds": 86400}] + } + }, + ) + hook = _make_hook(time_controller) + await hook.async_pre_call_hook( + user_api_key_dict=_key(), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion" + ) + + # Reload to a stricter config -- a fresh dict object, matching how a + # proxy config reload replaces litellm_settings.global_tag_rate_limits + # wholesale via setattr(litellm, key, value). A changed `limit` folds + # into the bucket's own policy fingerprint, so this is a fresh counter; + # the new, stricter limit=1 is still reachable in exactly one more call. + monkeypatch.setattr( + litellm, + "global_tag_rate_limits", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, + ) + await hook.async_pre_call_hook( + user_api_key_dict=_key(), + cache=DualCache(), + data=_data(["end_user_id:u1"], call_id="call-2"), + call_type="completion", + ) + 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-3"), + call_type="completion", + ) 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 7053ab6f9b8..662fb366718 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 @@ -375,35 +375,35 @@ def test_build_group_limits_empty_when_no_deployment_configures_unit(): 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 + assert _entry_applies(entry, "u1", ["end_user_id:u1"], None) 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 + assert _entry_applies(entry, "u1", ["end_user_id:u1"], None) 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 + assert _entry_applies(entry, "u2", ["end_user_id:u2"], None) 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 + assert _entry_applies(entry, "u1", ["end_user_id:u1"], None) 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 + assert _entry_applies(entry, "u2", ["end_user_id:u2", "company_id:1032"], None) is True def test_entry_applies_matches_an_enabled_for_gate(): @@ -414,7 +414,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, "u1", ["end_user_id:u1", "company_id:1032"]) is True + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"], None) is True def test_entry_applies_skips_when_enabled_for_gate_tag_is_absent(): @@ -429,7 +429,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, "u1", ["end_user_id:u1"]) is False + assert _entry_applies(entry, "u1", ["end_user_id:u1"], None) is False def test_entry_applies_skips_when_disabled_for_gate_matches(): @@ -440,7 +440,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, "u1", ["end_user_id:u1", "company_id:1032"]) is False + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"], None) is False def test_entry_applies_when_disabled_for_gate_tag_is_absent(): @@ -453,7 +453,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, "u1", ["end_user_id:u1"]) is True + assert _entry_applies(entry, "u1", ["end_user_id:u1"], None) is True def test_entry_applies_excluded_values_overrides_a_matching_enabled_for_gate(): @@ -467,7 +467,36 @@ def test_entry_applies_excluded_values_overrides_a_matching_enabled_for_gate(): 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 + assert _entry_applies(entry, "u1", ["end_user_id:u1", "company_id:1032"], 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, "u1", ["end_user_id:u1"], "any-key-alias") is True + assert _entry_applies(entry, "u1", ["end_user_id:u1"], 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, "u1", ["end_user_id:u1"], "team-a-key") 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, "u1", ["end_user_id:u1"], "team-b-key") is False + + +def test_entry_applies_rejects_when_key_has_no_alias_but_allowlist_is_set(): + """apply_to_key_alias is an allowlist gate: a key with no alias at all + never satisfies it, same as enabled_for's absent-gate-tag semantics.""" + 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, "u1", ["end_user_id:u1"], None) is False # --------------------------------------------------------------------------- @@ -519,6 +548,18 @@ def test_tag_rate_limit_scope_normalizes_values_order_and_duplicates(): assert scope.values == ("1001", "1032") +def test_tag_rate_limit_entry_rejects_empty_apply_to_key_alias(): + with pytest.raises(ValidationError, match="apply_to_key_alias must be a non-empty list"): + TagRateLimitEntry(name="daily", limit=1, period_seconds=60, apply_to_key_alias=()) + + +def test_tag_rate_limit_entry_normalizes_apply_to_key_alias_order_and_duplicates(): + entry = TagRateLimitEntry( + name="daily", limit=1, period_seconds=60, apply_to_key_alias=("team-b-key", "team-a-key", "team-a-key") + ) + assert entry.apply_to_key_alias == ("team-a-key", "team-b-key") + + # --------------------------------------------------------------------------- # _hash_tag / _bucket_key -- policy identity folds into the Redis key itself # --------------------------------------------------------------------------- @@ -4214,3 +4255,66 @@ async def test_token_accounting_with_a_cache_size_override_lands_on_that_entrys_ messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) + + +# --------------------------------------------------------------------------- +# apply_to_key_alias -- shared TagRateLimitEntry field, also usable on a +# per-model entry (the global_tag_rate_limits_hook is its primary motivation, +# but the field composes with async_filter_deployments unmodified) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_apply_to_key_alias_restricts_a_per_model_entry_to_the_listed_key(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_key_alias": ["premium-key"], + } + ] + } + }, + ) + ] + ) + limiter.update_variables(llm_router=router) + healthy = router.model_list + + # A key with no matching alias is entirely unaffected -- the entry never + # applies to it, so it can call repeatedly with no rejection. + 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"], "user_api_key_alias": "other-key"}}, + ) + assert result == healthy + + # The listed key alias is admitted once, then rejected on its 2nd call. + result = await limiter.async_filter_deployments( + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"], "user_api_key_alias": "premium-key"}}, + ) + 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"], "user_api_key_alias": "premium-key"}}, + )