mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
402c73cb62
commit
6663924b2f
7 changed files with 1320 additions and 26 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
613
litellm/proxy/hooks/global_tag_rate_limits_hook.py
Normal file
613
litellm/proxy/hooks/global_tag_rate_limits_hook.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
@ -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"}},
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue