diff --git a/litellm/proxy/hooks/tag_rate_limiter.py b/litellm/proxy/hooks/tag_rate_limiter.py index b19eaff96dd..5c7c77f0777 100644 --- a/litellm/proxy/hooks/tag_rate_limiter.py +++ b/litellm/proxy/hooks/tag_rate_limiter.py @@ -1,24 +1,11 @@ -""" -Tag-scoped token, request, dollar, and concurrency rate limits. - -Each limit entry is keyed by an arbitrary caller-supplied tag value (not a -DB-provisioned entity, not composed with the calling API key) and enforced on -every routing attempt for a chain/model-group -- the primary hop and every -fallback hop, each checked against its own configuration. - -Opt-in via `litellm_settings.callbacks: ["tag_rate_limiter"]` (not part of -`PROXY_HOOKS`), following the `dynamic_rate_limiter_v3` precedent: this hook -reuses `_PROXY_MaxParallelRequestsHandler_v3`'s Redis/TTL-preserving increment -machinery rather than duplicating it, and is never joined onto the default -limiter every proxy already runs. -""" +"""Tag-scoped token, request, dollar, and concurrency rate limits.""" import asyncio import contextvars from collections.abc import Callable from dataclasses import dataclass from datetime import datetime -from typing import TYPE_CHECKING, Any, Literal, Optional +from typing import TYPE_CHECKING, Any, Literal from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache @@ -123,10 +110,10 @@ class _ConfiguredLimit: # None => chain-wide (every deployment in the model_group shares one # bucket). Otherwise the sorted deployment ids that declared this exact # value -- the bucket is shared among only those deployments. - deployment_scope: Optional[tuple[str, ...]] + deployment_scope: tuple[str, ...] | None -def _extract_identity(tags: list[str], tag_id: str) -> Optional[str]: +def _extract_identity(tags: list[str], tag_id: str) -> str | None: """ First tag matching `f"{tag_id}:"`, value after the colon. Tags starting with `!` are tag-routing negation markers, not identity tags, and are @@ -141,11 +128,11 @@ def _extract_identity(tags: list[str], tag_id: str) -> Optional[str]: return None -def _deployment_id(deployment: dict) -> Optional[str]: +def _deployment_id(deployment: dict) -> str | None: return (deployment.get("model_info") or {}).get("id") -def _extract_team_id(request_kwargs: dict) -> Optional[str]: +def _extract_team_id(request_kwargs: dict) -> str | None: """Same two-channel lookup Router itself uses to resolve a caller's own team-scoped deployment (see `Router._common_checks_available_deployment`, which reads `user_api_key_team_id` from `metadata` falling back to @@ -156,7 +143,7 @@ def _extract_team_id(request_kwargs: dict) -> Optional[str]: return team_id if isinstance(team_id, str) else None -def _extract_key_hash(request_kwargs: dict) -> Optional[str]: +def _extract_key_hash(request_kwargs: dict) -> str | None: """Same two-channel lookup as `_extract_team_id`, but for the calling virtual key's hash: `LiteLLMProxyRequestSetup` sets `metadata["user_api_key"]` to `user_api_key_dict.api_key`, which despite the plain name is already @@ -262,7 +249,7 @@ class _LimitsIndex: by_model_name: dict[str, list[_ConfiguredLimit]] by_team_alias: dict[tuple[str, str], list[_ConfiguredLimit]] - def resolve(self, model: str, team_id: Optional[str]) -> list[_ConfiguredLimit]: + def resolve(self, model: str, team_id: str | None) -> list[_ConfiguredLimit]: if team_id is not None: scoped = self.by_team_alias.get((team_id, model)) if scoped is not None: @@ -339,45 +326,43 @@ _INDEX_TTL_SECONDS = 5.0 # slowly. period_seconds can still raise the TTL further, never lower it. _CONCURRENCY_MIN_SAFETY_TTL_SECONDS = 3600 + # Concurrency reservation keys accumulated for the current logical request, -# not yet released. A `ContextVar` rather than a plain module-level -# collection or a dict keyed by anything from `kwargs`, because every -# candidate for "correlate this hop with its logical request" that litellm -# itself exposes turns out to be either caller-controlled (`litellm_call_id` -# is `request.headers["x-litellm-call-id"]`, falling back to a fresh uuid -# only when absent -- two unrelated concurrent requests reusing the same -# caller-chosen value would merge their reservations under one key) or -# task-discontinuous (the success path runs `async_log_success_event` from -# inside a process-global `LoggingWorker` task, never the admission-time -# task, so `id(asyncio.current_task())` differs even for one hop's own -# success). `ContextVar` is the one mechanism immune to both problems: its -# value is pure Python-runtime state, never caller-visible or -# caller-settable, and litellm's own logging pipeline is already built to -# propagate it correctly across every task boundary a hop crosses -- -# `asyncio.create_task()` copies the calling context by default (used for -# `wrapper_async`'s success dispatch in `litellm/utils.py` and for this -# hook's own rejections propagating through `Router.async_callback_filter_ -# deployments`), and `LoggingWorker.enqueue()` (`litellm/litellm_core_utils/ -# logging_worker.py`) explicitly calls `contextvars.copy_context()` at -# enqueue time and later runs the queued coroutine via -# `task["context"].run(asyncio.create_task, ...)`, so a value set during -# admission is still visible when the eventual release callback executes, -# however many hops or worker hops later that turns out to be. Each -# concurrent request gets its own isolated context (forked at whatever -# `create_task` call started it), so two unrelated requests never share a -# value regardless of what identifiers they happen to reuse. -_pending_concurrency_keys: contextvars.ContextVar[tuple[str, ...]] = contextvars.ContextVar( - "tag_rate_limiter_pending_concurrency_keys", default=() +# not yet released. Held via a ContextVar bound to a mutable holder object +# (not an immutable tuple rebound with `.set()`) because `asyncio.create_task` +# only copies which *object* a ContextVar is bound to, not a snapshot of that +# object's contents: a `.set()` performed inside a task forked off this +# context mutates only that task's own binding, invisible to the parent task +# that continues on to a fallback hop. Mutating a shared holder in place is +# visible from every task forked after the holder was first created, +# regardless of which task performs the mutation. +class _PendingConcurrencyKeys: + __slots__ = ("keys",) + + def __init__(self) -> None: + self.keys: list[str] = [] + + +_pending_concurrency_keys: contextvars.ContextVar[_PendingConcurrencyKeys | None] = contextvars.ContextVar( + "tag_rate_limiter_pending_concurrency_keys", default=None ) +def _pending_concurrency_holder() -> _PendingConcurrencyKeys: + holder = _pending_concurrency_keys.get() + if holder is None: + holder = _PendingConcurrencyKeys() + _pending_concurrency_keys.set(holder) + return holder + + class _TagRateLimitIndex: """Rebuilds the limits index when `llm_router.model_list` changes, or at least every `_INDEX_TTL_SECONDS`, whichever comes first.""" def __init__(self, time_provider: Callable[[], datetime]) -> None: self._time_provider = time_provider - self._cache_key: Optional[tuple[int, int]] = None + self._cache_key: tuple[int, int] | None = None self._built_at: float = 0.0 self._index: _LimitsIndex = _LimitsIndex(by_model_name={}, by_team_alias={}) @@ -392,7 +377,7 @@ class _TagRateLimitIndex: return self._index -def _scope_suffix(deployment_scope: Optional[tuple[str, ...]]) -> str: +def _scope_suffix(deployment_scope: tuple[str, ...] | None) -> str: return "chain" if deployment_scope is None else "dep:" + "+".join(deployment_scope) @@ -401,7 +386,7 @@ def _bucket_key( configured: _ConfiguredLimit, tag_value: str, bucket_id: int, - key_hash: Optional[str] = None, + key_hash: str | None = None, ) -> str: scope = _scope_suffix(configured.deployment_scope) key_suffix = f":key:{key_hash}" if key_hash is not None else "" @@ -413,7 +398,7 @@ def _inflight_key( model_group: str, configured: _ConfiguredLimit, tag_value: str, - key_hash: Optional[str] = None, + key_hash: str | None = None, ) -> str: """Concurrency counter key: not epoch-bucketed, since "how many are in flight right now" has no window to reset on -- it's released explicitly @@ -428,14 +413,14 @@ class _PROXY_TagRateLimiter(CustomLogger): def __init__( self, internal_usage_cache: DualCache, - time_provider: Optional[Callable[[], datetime]] = None, + time_provider: Callable[[], datetime] | None = None, ): self.internal_usage_cache = InternalUsageCache(dual_cache=internal_usage_cache) self._v3 = _PROXY_MaxParallelRequestsHandler_v3(self.internal_usage_cache, time_provider=time_provider) self._time_provider = time_provider or datetime.now self._index = _TagRateLimitIndex(time_provider=self._time_provider) self._lock = asyncio.Lock() - self.llm_router: Optional[Router] = None + self.llm_router: Router | None = None redis_cache = self.internal_usage_cache.dual_cache.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 @@ -479,7 +464,7 @@ class _PROXY_TagRateLimiter(CustomLogger): async def _atomic_check_and_increment( self, checks: list[tuple[str, float, float, int]], - ) -> tuple[Optional[int], list[float]]: + ) -> tuple[int | None, list[float]]: """ All-or-nothing across every (key, limit, increment, ttl) in `checks`: if any would exceed its limit, none are incremented -- a single hop's @@ -524,9 +509,9 @@ class _PROXY_TagRateLimiter(CustomLogger): self, model: str, healthy_deployments: list[dict], - messages: Optional[list[AllMessageValues]], - request_kwargs: Optional[dict] = None, - parent_otel_span: Optional[Span] = None, + messages: list[AllMessageValues] | None, + request_kwargs: dict | None = None, + parent_otel_span: Span | None = None, ) -> list[dict]: if not healthy_deployments or not isinstance(healthy_deployments, list) or self.llm_router is None: return healthy_deployments @@ -579,11 +564,11 @@ class _PROXY_TagRateLimiter(CustomLogger): configured_limit, tag_value, _key = atomic_checks[failing_index] self._raise_over_limit(configured_limit, tag_value, model, current=values[0]) - concurrency_keys = tuple( + concurrency_keys = [ key for configured_limit, _tag_value, key in atomic_checks if configured_limit.unit == "concurrency" - ) + ] if concurrency_keys: - _pending_concurrency_keys.set(_pending_concurrency_keys.get() + concurrency_keys) + _pending_concurrency_holder().keys.extend(concurrency_keys) return healthy_deployments @@ -602,8 +587,8 @@ class _PROXY_TagRateLimiter(CustomLogger): async def _read_only_values( self, read_only_checks: list[tuple[_ConfiguredLimit, str, str]], - parent_otel_span: Optional[Span], - ) -> list[Optional[float]]: + parent_otel_span: Span | None, + ) -> list[float | None]: if not read_only_checks: return [] keys = [key for _cfg, _tag_value, key in read_only_checks] @@ -617,7 +602,7 @@ class _PROXY_TagRateLimiter(CustomLogger): def _raise_if_over_limit( self, read_only_checks: list[tuple[_ConfiguredLimit, str, str]], - current_values: list[Optional[float]], + current_values: list[float | None], model: str, ) -> None: for (configured_limit, tag_value, _key), current_value in zip(read_only_checks, current_values): @@ -676,61 +661,43 @@ class _PROXY_TagRateLimiter(CustomLogger): except Exception as e: # noqa: BLE001 - releasing a slot must never raise into the caller's request path verbose_proxy_logger.warning("tag_rate_limiter: failed to release concurrency slot %s: %s", key, e) + @staticmethod + def _pop_pending_concurrency_keys() -> list[str]: + # Snapshot then remove only those exact keys, never a blanket clear: + # a sibling hop can still be live and appending to the same shared + # holder concurrently (see the holder's own comment above), so + # wiping the whole list here would silently strand that hop's + # reservation instead of releasing it later. + holder = _pending_concurrency_keys.get() + if holder is None or not holder.keys: + return [] + keys = list(holder.keys) + for key in keys: + try: + holder.keys.remove(key) + except ValueError: + pass + return keys + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: - """ - Release every concurrency slot accumulated onto `_pending_concurrency_keys` - for the current logical request. Never recomputes a key from - `standard_logging_object`: only releases exactly what admission - itself accumulated, so a rejection this hook raises for being over - its own limit -- which `_atomic_check_and_increment` already - refunded synchronously, inside that same call, before ever adding - anything here -- naturally has nothing new to release, by - construction, rather than needing a special case for it. - - The explicit `ProxyRateLimitError` check below is belt-and-suspenders - on top of that: hops of one logical request run strictly - sequentially today (a fallback is only ever attempted after the - previous hop has fully concluded, including firing its own - completion event), so a rejected hop's own `_pending_concurrency_keys` - is provably empty by the time this fires. If a future routing - strategy ever dispatches hops concurrently instead, that invariant - would break silently; this check means a rejection never releases - anything even if it does. `ProxyRateLimitError.detail` carries - `{"error": "tag_rate_limit_exceeded", ...}`, a string unique to this - module, so it's distinguishable from a genuine provider failure. - - litellm dedupes this event to fire once per logical request (the - first failed hop only, via `Logging.has_run_logging`'s - `has_logged_async_failure` guard). That no longer matters for - correctness here: whichever event fires next for this request -- - this one, `async_log_success_event`, or another failed hop's -- pops - and releases whatever has accumulated in `_pending_concurrency_keys` - since the last release, covering every hop this event's dedup would - otherwise skip. See that variable's module-level docstring for why a - `ContextVar` is what makes this safe: it survives every task - boundary a hop crosses (litellm's own logging pipeline is built to - propagate it), without ever depending on anything a caller supplies. - """ if isinstance(kwargs.get("exception"), ProxyRateLimitError): detail = kwargs["exception"].detail if isinstance(kwargs["exception"].detail, dict) else {} if detail.get("error") == "tag_rate_limit_exceeded": return - release_keys = _pending_concurrency_keys.get() + release_keys = self._pop_pending_concurrency_keys() if release_keys: - _pending_concurrency_keys.set(()) - await self._release_keys(list(release_keys)) + await self._release_keys(release_keys) async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: - release_keys = _pending_concurrency_keys.get() + release_keys = self._pop_pending_concurrency_keys() if release_keys: - _pending_concurrency_keys.set(()) - asyncio.create_task(self._release_keys(list(release_keys))) + asyncio.create_task(self._release_keys(release_keys)) if self.llm_router is None: return - standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get("standard_logging_object") + standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object") if standard_logging_object is None: return diff --git a/litellm/types/router.py b/litellm/types/router.py index 924b0762bb5..5c721229cdf 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -138,55 +138,26 @@ def _as_utc(value: datetime.datetime | None) -> datetime.datetime | None: class TagRateLimitEntry(BaseModel): - """ - One tag-scoped limit: a caller-supplied tag value (identified by `tag_id`, - e.g. `end_user_id` in a request tag like `end_user_id:user-123`) is capped - at `limit` units per rolling `period_seconds`-second window. Bucketing is - `epoch_second // period_seconds`, so `period_seconds=86400` resets at UTC - midnight and `period_seconds=60` resets on real clock-minute boundaries. - - For a `concurrency_limits` entry specifically, `period_seconds` is not a - window: it is a floor under the safety TTL a reserved in-flight slot - self-heals after, in case a worker crashes before releasing it (the - counter, not a window). The effective TTL is at least one hour regardless - of this value, so a slow but genuinely still-running request never has - its reservation expire out from under it; set this higher only if an - even longer self-heal window is wanted. `concurrency_limits` also only - supports chain-wide entries (declared identically by every deployment - sharing a `model_name`) -- a divergent per-deployment value is dropped - with a warning, not silently scoped to a subset of deployments. - """ - name: str tag_id: str = "end_user_id" limit: float period_seconds: int scope_by_key_hash: bool = False - """ - When `True`, the bucket is additionally scoped by the calling virtual - key's hash, on top of the existing `tag_id`/tag-value match. Without - this, two different keys (e.g. two separate services) that both happen - to send the same tag value (e.g. the same `end_user_id`) share one - bucket and one counter; opting in gives each calling key its own - independent counter for the same tag value. Defaults to `False`, which - is today's existing behavior: the bucket is scoped by tag value alone, - shared across every key that sends it. - """ model_config = ConfigDict(protected_namespaces=()) + @model_validator(mode="after") + def _validate_period_seconds(self) -> "TagRateLimitEntry": + if self.period_seconds <= 0: + raise ValueError("period_seconds must be a positive integer") + return self + class TagRateLimitGroup(BaseModel): limits: list[TagRateLimitEntry] = Field(default_factory=list) class TagRateLimits(BaseModel): - """ - Per-chain/model-group tag rate limits, set under a deployment's - `model_info.tag_rate_limits`. Each entry carries its own `tag_id`, so two - entries of the same unit on the same chain can key by different tags. - """ - token_limits: TagRateLimitGroup | None = None request_limits: TagRateLimitGroup | None = None dollar_limits: TagRateLimitGroup | None = None diff --git a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py index 50b1278f77d..43055d715a7 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py @@ -12,11 +12,12 @@ import litellm from litellm.caching.dual_cache import DualCache from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.tag_rate_limiter import ( + _CONCURRENCY_MIN_SAFETY_TTL_SECONDS, _build_group_limits, _build_limits_index, _ConfiguredLimit, - _CONCURRENCY_MIN_SAFETY_TTL_SECONDS, _extract_identity, + _pending_concurrency_holder, _PROXY_TagRateLimiter, ) from litellm.types.router import TagRateLimitEntry @@ -71,6 +72,26 @@ def test_extract_identity_skips_negation_tags(): assert _extract_identity(["!end_user_id:u1"], "end_user_id") is None +# --------------------------------------------------------------------------- +# TagRateLimitEntry -- period_seconds validation +# --------------------------------------------------------------------------- + + +def test_tag_rate_limit_entry_rejects_zero_period_seconds(): + with pytest.raises(Exception): + TagRateLimitEntry(name="n", limit=1, period_seconds=0) + + +def test_tag_rate_limit_entry_rejects_negative_period_seconds(): + with pytest.raises(Exception): + TagRateLimitEntry(name="n", limit=1, period_seconds=-1) + + +def test_tag_rate_limit_entry_accepts_positive_period_seconds(): + entry = TagRateLimitEntry(name="n", limit=1, period_seconds=60) + assert entry.period_seconds == 60 + + # --------------------------------------------------------------------------- # _build_group_limits -- chain-wide vs per-deployment scoping # --------------------------------------------------------------------------- @@ -78,8 +99,12 @@ def test_extract_identity_skips_negation_tags(): def test_build_group_limits_chain_wide_when_all_deployments_agree(): deployments = [ - _deployment("grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}), - _deployment("grp", "dep-2", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}), + _deployment( + "grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}} + ), + _deployment( + "grp", "dep-2", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}} + ), ] configured = _build_group_limits(deployments, "tokens") assert len(configured) == 1 @@ -94,8 +119,12 @@ def test_build_group_limits_per_deployment_when_values_diverge(): per-deployment-scoped entries instead. """ deployments = [ - _deployment("grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}), - _deployment("grp", "dep-2", {"token_limits": {"limits": [{"name": "daily", "limit": 999, "period_seconds": 86400}]}}), + _deployment( + "grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}} + ), + _deployment( + "grp", "dep-2", {"token_limits": {"limits": [{"name": "daily", "limit": 999, "period_seconds": 86400}]}} + ), ] configured = _build_group_limits(deployments, "tokens") assert len(configured) == 2 @@ -108,7 +137,9 @@ def test_build_group_limits_per_deployment_when_values_diverge(): def test_build_group_limits_per_deployment_when_only_some_declare_it(): deployments = [ - _deployment("grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}), + _deployment( + "grp", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}} + ), _deployment("grp", "dep-2", {}), ] configured = _build_group_limits(deployments, "tokens") @@ -158,7 +189,11 @@ async def test_filter_deployments_allows_under_limit_and_rejects_at_limit(time_c _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 2, "period_seconds": 60}]}}, + { + "request_limits": { + "limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 2, "period_seconds": 60}] + } + }, ) ] ) @@ -261,7 +296,10 @@ async def test_different_tag_ids_with_same_name_do_not_share_a_counter(time_cont # end_user_id "u1" makes its one allowed request. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) # team_id "u1" -- identical value, different tag_id, its own untouched @@ -274,11 +312,17 @@ async def test_different_tag_ids_with_same_name_do_not_share_a_counter(time_cont # Both identities are now genuinely at their own limit of 1. with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["team_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["team_id:u1"]}}, ) @@ -302,12 +346,20 @@ async def test_load_balanced_group_per_deployment_breach_rejects_whole_hop(time_ _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]}}, + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, ), _deployment( "grp", "dep-2", - {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 999, "period_seconds": 86400}]}}, + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 999, "period_seconds": 86400}] + } + }, ), ] ) @@ -340,9 +392,17 @@ async def test_log_success_event_increments_configured_units(time_controller): "grp", "dep-1", { - "token_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 500000, "period_seconds": 86400}]}, - "request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 100, "period_seconds": 86400}]}, - "dollar_limits": {"limits": [{"name": "monthly", "tag_id": "end_user_id", "limit": 50.0, "period_seconds": 2592000}]}, + "token_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 500000, "period_seconds": 86400}] + }, + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 100, "period_seconds": 86400}] + }, + "dollar_limits": { + "limits": [ + {"name": "monthly", "tag_id": "end_user_id", "limit": 50.0, "period_seconds": 2592000} + ] + }, }, ) ] @@ -366,8 +426,12 @@ async def test_log_success_event_increments_configured_units(time_controller): token_key = f"{{tag_rl:grp:tokens:daily:end_user_id:chain:u1}}:{int(now) // 86400}" dollar_key = f"{{tag_rl:grp:dollars:monthly:end_user_id:chain:u1}}:{int(now) // 2592000}" - assert float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 - assert float(await limiter.internal_usage_cache.async_get_cache(key=dollar_key, litellm_parent_otel_span=None)) == 0.01 + assert ( + float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 + ) + assert ( + float(await limiter.internal_usage_cache.async_get_cache(key=dollar_key, litellm_parent_otel_span=None)) == 0.01 + ) # "requests" is accounted atomically at admission (async_filter_deployments), # not here -- async_log_success_event must not touch its bucket at all. @@ -429,7 +493,10 @@ async def test_cross_unit_rejection_does_not_leave_a_phantom_increment(time_cont # Occupy the one concurrency slot. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) # A second attempt: requests-unit alone would admit (well under 10), but @@ -437,7 +504,10 @@ async def test_cross_unit_rejection_does_not_leave_a_phantom_increment(time_cont # requests counter must remain untouched by this rejected attempt. with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert exc_info.value.detail["type"] == "concurrency" @@ -456,12 +526,19 @@ async def test_concurrency_limit_rejects_third_concurrent_reservation(time_contr kwargs_1 = {"metadata": {"tags": ["end_user_id:u1"]}} kwargs_2 = {"metadata": {"tags": ["end_user_id:u1"]}} - await limiter.async_filter_deployments(model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs_1) - await limiter.async_filter_deployments(model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs_2) + await limiter.async_filter_deployments( + model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs_1 + ) + await limiter.async_filter_deployments( + model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs_2 + ) with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert exc_info.value.detail["type"] == "concurrency" @@ -482,7 +559,11 @@ async def test_requests_admission_is_race_free_under_genuine_concurrency(time_co _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}]}}, + { + "request_limits": { + "limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}] + } + }, ) ] ) @@ -519,7 +600,11 @@ async def test_index_refreshes_after_ttl_for_length_preserving_update(time_contr _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]}}, + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, ) ] ) @@ -527,22 +612,33 @@ async def test_index_refreshes_after_ttl_for_length_preserving_update(time_contr healthy = router.model_list await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) # Same length, deployment mutated in place -- raise the limit to 100. router.model_list[0]["model_info"]["tag_rate_limits"] = { - "request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 100, "period_seconds": 86400}]} + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 100, "period_seconds": 86400}] + } } time_controller.advance(6) # past _INDEX_TTL_SECONDS result = await limiter.async_filter_deployments( - model="grp", healthy_deployments=router.model_list, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=router.model_list, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert result == router.model_list @@ -555,21 +651,34 @@ async def test_concurrency_slot_released_on_success_frees_capacity(time_controll healthy = router.model_list kwargs = {"metadata": {"tags": ["end_user_id:u1"]}} - await limiter.async_filter_deployments(model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs) + await limiter.async_filter_deployments( + model="grp", healthy_deployments=healthy, messages=None, request_kwargs=kwargs + ) # At capacity: a second concurrent request is rejected. with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) # The first request completes -- its slot is released -- freeing capacity again. - kwargs["standard_logging_object"] = {"model_group": "grp", "model_id": "dep-1", "total_tokens": 0, "response_cost": 0} + kwargs["standard_logging_object"] = { + "model_group": "grp", + "model_id": "dep-1", + "total_tokens": 0, + "response_cost": 0, + } await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) await asyncio.sleep(0) result = await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert result == healthy @@ -582,7 +691,10 @@ async def test_concurrency_slot_released_on_failure_frees_capacity(time_controll healthy = router.model_list await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) await limiter.async_log_failure_event( @@ -593,7 +705,10 @@ async def test_concurrency_slot_released_on_failure_frees_capacity(time_controll ) result = await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert result == healthy @@ -617,17 +732,26 @@ async def test_concurrency_slot_released_on_fallback_recovered_hop_failure(time_ healthy = router.model_list await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) await limiter.async_log_failure_event( - kwargs={"standard_logging_object": {"model_group": "grp", "model_id": "dep-1"}, "metadata": {"tags": ["end_user_id:u1"]}}, + kwargs={ + "standard_logging_object": {"model_group": "grp", "model_id": "dep-1"}, + "metadata": {"tags": ["end_user_id:u1"]}, + }, response_obj=None, start_time=0, end_time=0, ) result = await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert result == healthy @@ -667,7 +791,12 @@ async def test_pending_concurrency_context_does_not_leak_across_concurrent_tasks async def _release(tag_value): await limiter.async_log_success_event( kwargs={ - "standard_logging_object": {"model_group": "grp", "model_id": "dep-1", "total_tokens": 0, "response_cost": 0}, + "standard_logging_object": { + "model_group": "grp", + "model_id": "dep-1", + "total_tokens": 0, + "response_cost": 0, + }, "metadata": {"tags": [f"end_user_id:{tag_value}"]}, }, response_obj=None, @@ -688,14 +817,20 @@ async def test_pending_concurrency_context_does_not_leak_across_concurrent_tasks # Exactly one slot was freed: a fresh request is admitted (back to 2 in flight)... await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:a"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:a"]}}, ) # ...but a second one does not, since B's reservation is genuinely still # held. If task isolation were broken, task A's release would have # drained B's reservation too, and this would wrongly admit. with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:a"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:a"]}}, ) @@ -730,7 +865,10 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti # (dedup allows exactly the first failure through), releasing its # own key immediately. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) await limiter.async_log_failure_event( kwargs={"standard_logging_object": {"model_group": "grp"}, "metadata": {"tags": ["end_user_id:u1"]}}, @@ -742,14 +880,20 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti # Hop 2 (a retry or fallback) admits and also fails, but -- per # litellm's dedup -- no async_log_failure_event call follows it. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) # Hop 3 admits and succeeds. Its success event, dispatched as a # child task (mirroring the real worker hop), must release both # hop 2's still-pending reservation and its own. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) async def _hop_3_success_event(): @@ -776,10 +920,16 @@ async def test_concurrency_released_for_every_hop_across_a_real_task_boundary(ti # hop 3's own were released. If the earlier hop's leaked reservation # hadn't been released too, only one of these two admissions would succeed. await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) result = await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert result == healthy @@ -817,7 +967,10 @@ async def test_own_rejection_does_not_release_a_live_reservation(time_controller # One legitimate request, in its own task, holds the only slot. async def _admit(): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) await asyncio.create_task(_admit()) @@ -828,7 +981,10 @@ async def test_own_rejection_does_not_release_a_live_reservation(time_controller async def _reject_and_fire_failure_event(): with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) await limiter.async_log_failure_event( kwargs={ @@ -848,7 +1004,10 @@ async def test_own_rejection_does_not_release_a_live_reservation(time_controller # it, this would wrongly admit instead. with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) @@ -865,7 +1024,11 @@ async def test_token_limit_rejects_once_bucket_is_seeded_at_limit(time_controlle _deployment( "grp", "dep-1", - {"token_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1000, "period_seconds": 86400}]}}, + { + "token_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1000, "period_seconds": 86400}] + } + }, ) ] ) @@ -878,7 +1041,10 @@ async def test_token_limit_rejects_once_bucket_is_seeded_at_limit(time_controlle with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:u1"]}}, ) assert exc_info.value.detail["type"] == "tokens" assert exc_info.value.detail["limit_name"] == "daily" @@ -892,7 +1058,11 @@ async def test_dollar_limit_rejects_once_bucket_is_seeded_at_limit(time_controll _deployment( "grp", "dep-1", - {"dollar_limits": {"limits": [{"name": "monthly", "tag_id": "team_id", "limit": 50.0, "period_seconds": 2592000}]}}, + { + "dollar_limits": { + "limits": [{"name": "monthly", "tag_id": "team_id", "limit": 50.0, "period_seconds": 2592000}] + } + }, ) ] ) @@ -905,7 +1075,10 @@ async def test_dollar_limit_rejects_once_bucket_is_seeded_at_limit(time_controll with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["team_id:t1"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["team_id:t1"]}}, ) assert exc_info.value.detail["type"] == "dollars" assert exc_info.value.detail["tag_value"] == "t1" @@ -942,14 +1115,18 @@ async def test_redis_backed_requests_admission_is_race_free_under_genuine_concur try: await redis_cache.ping() except Exception as e: - pytest.skip(f"Redis connection failed: {str(e)}") + pytest.skip(f"Redis connection failed: {e!s}") router = litellm.Router( model_list=[ _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}]}}, + { + "request_limits": { + "limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 5, "period_seconds": 60}] + } + }, ) ] ) @@ -960,7 +1137,10 @@ async def test_redis_backed_requests_admission_is_race_free_under_genuine_concur async def attempt(): try: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}}, ) return True except ProxyRateLimitError: @@ -977,7 +1157,7 @@ async def test_redis_backed_cross_unit_rejection_does_not_leave_a_phantom_increm try: await redis_cache.ping() except Exception as e: - pytest.skip(f"Redis connection failed: {str(e)}") + pytest.skip(f"Redis connection failed: {e!s}") router = litellm.Router( model_list=[ @@ -1000,11 +1180,17 @@ async def test_redis_backed_cross_unit_rejection_does_not_leave_a_phantom_increm tag = f"redis-phantom-check-{uuid.uuid4().hex}" await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}}, ) with pytest.raises(ProxyRateLimitError) as exc_info: await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": [f"end_user_id:{tag}"]}}, ) assert exc_info.value.detail["type"] == "concurrency" @@ -1031,7 +1217,11 @@ def test_build_limits_index_is_also_keyed_by_team_public_model_name(): configured limits, or a team-aliased chain's limits are silently never checked. """ - deployment = _deployment("real-model-name", "dep-1", {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}) + deployment = _deployment( + "real-model-name", + "dep-1", + {"token_limits": {"limits": [{"name": "daily", "limit": 500, "period_seconds": 86400}]}}, + ) deployment["model_info"]["team_id"] = "team-1" deployment["model_info"]["team_public_model_name"] = "team-alias-name" index = _build_limits_index([deployment]) @@ -1047,11 +1237,15 @@ def test_build_limits_index_keeps_different_teams_same_alias_separate(): `(team_id, name)` rather than by name alone. Keying the limits index by name alone would let one team's config silently overwrite another's. """ - team_a = _deployment("model-a", "dep-a", {"token_limits": {"limits": [{"name": "daily", "limit": 100, "period_seconds": 86400}]}}) + team_a = _deployment( + "model-a", "dep-a", {"token_limits": {"limits": [{"name": "daily", "limit": 100, "period_seconds": 86400}]}} + ) team_a["model_info"]["team_id"] = "team-a" team_a["model_info"]["team_public_model_name"] = "shared-alias" - team_b = _deployment("model-b", "dep-b", {"token_limits": {"limits": [{"name": "daily", "limit": 999, "period_seconds": 86400}]}}) + team_b = _deployment( + "model-b", "dep-b", {"token_limits": {"limits": [{"name": "daily", "limit": 999, "period_seconds": 86400}]}} + ) team_b["model_info"]["team_id"] = "team-b" team_b["model_info"]["team_public_model_name"] = "shared-alias" @@ -1075,12 +1269,18 @@ def test_build_limits_index_merges_alias_limits_across_different_model_names(): the same alias, only the entry declared by whichever model_name group is processed last would survive. """ - dep_a = _deployment("model_name_team1_aaa", "dep-a", {"token_limits": {"limits": [{"name": "daily", "limit": 100, "period_seconds": 86400}]}}) + dep_a = _deployment( + "model_name_team1_aaa", + "dep-a", + {"token_limits": {"limits": [{"name": "daily", "limit": 100, "period_seconds": 86400}]}}, + ) dep_a["model_info"]["team_id"] = "team-1" dep_a["model_info"]["team_public_model_name"] = "shared-alias" dep_b = _deployment( - "model_name_team1_bbb", "dep-b", {"dollar_limits": {"limits": [{"name": "monthly", "limit": 50.0, "period_seconds": 2592000}]}} + "model_name_team1_bbb", + "dep-b", + {"dollar_limits": {"limits": [{"name": "monthly", "limit": 50.0, "period_seconds": 2592000}]}}, ) dep_b["model_info"]["team_id"] = "team-1" dep_b["model_info"]["team_public_model_name"] = "shared-alias" @@ -1094,7 +1294,15 @@ def test_build_limits_index_merges_alias_limits_across_different_model_names(): @pytest.mark.asyncio async def test_filter_deployments_enforces_limit_when_called_with_team_alias(time_controller): limiter = _make_limiter(time_controller) - deployment = _deployment("real-model-name", "dep-1", {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]}}) + deployment = _deployment( + "real-model-name", + "dep-1", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, + ) deployment["model_info"]["team_id"] = "team-1" deployment["model_info"]["team_public_model_name"] = "team-alias-name" router = litellm.Router(model_list=[deployment]) @@ -1104,9 +1312,13 @@ async def test_filter_deployments_enforces_limit_when_called_with_team_alias(tim # Router passes the alias as `model`, not "real-model-name", and threads # the caller's team_id through request metadata. request_kwargs = {"metadata": {"tags": ["end_user_id:u1"], "user_api_key_team_id": "team-1"}} - await limiter.async_filter_deployments(model="team-alias-name", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs) + await limiter.async_filter_deployments( + model="team-alias-name", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs + ) with pytest.raises(ProxyRateLimitError): - await limiter.async_filter_deployments(model="team-alias-name", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs) + await limiter.async_filter_deployments( + model="team-alias-name", healthy_deployments=healthy, messages=None, request_kwargs=request_kwargs + ) @pytest.mark.asyncio @@ -1117,11 +1329,27 @@ async def test_filter_deployments_does_not_cross_team_alias_boundary(time_contro counted) by team-a's configured limit and usage. """ limiter = _make_limiter(time_controller) - team_a = _deployment("model-a", "dep-a", {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]}}) + team_a = _deployment( + "model-a", + "dep-a", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}] + } + }, + ) team_a["model_info"]["team_id"] = "team-a" team_a["model_info"]["team_public_model_name"] = "shared-alias" - team_b = _deployment("model-b", "dep-b", {"request_limits": {"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 5, "period_seconds": 86400}]}}) + team_b = _deployment( + "model-b", + "dep-b", + { + "request_limits": { + "limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 5, "period_seconds": 86400}] + } + }, + ) team_b["model_info"]["team_id"] = "team-b" team_b["model_info"]["team_public_model_name"] = "shared-alias" @@ -1167,8 +1395,12 @@ def test_concurrency_divergent_config_is_dropped_not_scoped_per_deployment(): with no concurrency entry at all. """ deployments = [ - _deployment("grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}}), - _deployment("grp", "dep-2", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 5, "period_seconds": 60}]}}), + _deployment( + "grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}} + ), + _deployment( + "grp", "dep-2", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 5, "period_seconds": 60}]}} + ), ] configured = _build_group_limits(deployments, "concurrency") assert configured == [] @@ -1176,7 +1408,9 @@ def test_concurrency_divergent_config_is_dropped_not_scoped_per_deployment(): def test_concurrency_partial_declaration_is_dropped_not_scoped_per_deployment(): deployments = [ - _deployment("grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}}), + _deployment( + "grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}} + ), _deployment("grp", "dep-2", {}), ] configured = _build_group_limits(deployments, "concurrency") @@ -1185,8 +1419,12 @@ def test_concurrency_partial_declaration_is_dropped_not_scoped_per_deployment(): def test_concurrency_identical_across_all_deployments_is_still_chain_wide(): deployments = [ - _deployment("grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}}), - _deployment("grp", "dep-2", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}}), + _deployment( + "grp", "dep-1", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}} + ), + _deployment( + "grp", "dep-2", {"concurrency_limits": {"limits": [{"name": "inflight", "limit": 2, "period_seconds": 60}]}} + ), ] configured = _build_group_limits(deployments, "concurrency") assert len(configured) == 1 @@ -1206,11 +1444,65 @@ def test_concurrency_ttl_floor_overrides_a_too_short_period_seconds(): def test_concurrency_ttl_floor_does_not_shorten_a_longer_period_seconds(): - entry = TagRateLimitEntry(name="inflight", tag_id="end_user_id", limit=1, period_seconds=_CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100) + entry = TagRateLimitEntry( + name="inflight", tag_id="end_user_id", limit=1, period_seconds=_CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 + ) configured_limit = _ConfiguredLimit(unit="concurrency", entry=entry, deployment_scope=None) assert _PROXY_TagRateLimiter._ttl_for(configured_limit) == _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100 +# --------------------------------------------------------------------------- +# pending-concurrency-key holder must survive a detached asyncio.create_task +# fork (e.g. litellm's own failure-logging dispatch) without a rebind in that +# forked task hiding the release from the parent, and a release must never +# sweep up a key a still-live sibling hop appended in the meantime +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_release_in_a_forked_task_is_visible_to_the_parent_context(): + _pending_concurrency_holder().keys.clear() + _pending_concurrency_holder().keys.append("key1") + + async def detached_release(): + return _PROXY_TagRateLimiter._pop_pending_concurrency_keys() + + released = await asyncio.create_task(detached_release()) + assert released == ["key1"] + + # The parent's own binding must see the same, now-empty holder -- + # not a stale copy still holding "key1". + assert _pending_concurrency_holder().keys == [] + + +@pytest.mark.asyncio +async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot(): + _pending_concurrency_holder().keys.clear() + _pending_concurrency_holder().keys.append("key1") + + async def detached_release_then_sibling_admits(): + released = _PROXY_TagRateLimiter._pop_pending_concurrency_keys() + # A sibling hop's admission, appending to the same shared holder, + # interleaved right after this release's snapshot was taken. + _pending_concurrency_holder().keys.append("key2") + return released + + released = await asyncio.create_task(detached_release_then_sibling_admits()) + assert released == ["key1"] + # key2 must still be pending for its own hop's eventual release. + assert _pending_concurrency_holder().keys == ["key2"] + + +@pytest.mark.asyncio +async def test_release_is_not_repeated_for_the_same_snapshot(): + _pending_concurrency_holder().keys.clear() + _pending_concurrency_holder().keys.append("key1") + first = _PROXY_TagRateLimiter._pop_pending_concurrency_keys() + second = _PROXY_TagRateLimiter._pop_pending_concurrency_keys() + assert first == ["key1"] + assert second == [] + + # --------------------------------------------------------------------------- # refund-on-rollback across differently-hash-tagged keys (Redis Cluster safety) # --------------------------------------------------------------------------- @@ -1233,8 +1525,12 @@ async def test_cross_unit_refund_leaves_no_phantom_increment_in_memory(time_cont "grp", "dep-1", { - "request_limits": {"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 10, "period_seconds": 60}]}, - "concurrency_limits": {"limits": [{"name": "inflight", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}]}, + "request_limits": { + "limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 10, "period_seconds": 60}] + }, + "concurrency_limits": { + "limits": [{"name": "inflight", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}] + }, }, ) ] @@ -1243,11 +1539,17 @@ async def test_cross_unit_refund_leaves_no_phantom_increment_in_memory(time_cont healthy = router.model_list await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:refund-check"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:refund-check"]}}, ) with pytest.raises(ProxyRateLimitError): await limiter.async_filter_deployments( - model="grp", healthy_deployments=healthy, messages=None, request_kwargs={"metadata": {"tags": ["end_user_id:refund-check"]}} + model="grp", + healthy_deployments=healthy, + messages=None, + request_kwargs={"metadata": {"tags": ["end_user_id:refund-check"]}}, ) now = time_controller.now().timestamp() @@ -1422,7 +1724,11 @@ async def test_request_limit_without_scope_by_key_hash_still_shares_one_counter( _deployment( "grp", "dep-1", - {"request_limits": {"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 2, "period_seconds": 60}]}}, + { + "request_limits": { + "limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 2, "period_seconds": 60}] + } + }, ) ] )