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.
This commit is contained in:
Deepanshu 2026-08-24 14:25:06 -04:00
parent 402c73cb62
commit 6663924b2f
7 changed files with 1320 additions and 26 deletions

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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",
)

View file

@ -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"}},
)