fix(rate-limiting): satisfy type-discipline, basedpyright, and eager-logging gates

- Rewrite tag_rate_limiter.py's dict/list usage to immutable equivalents
  (Mapping/Sequence params, tuple/frozenset/MappingProxyType returns and
  locals, Final everywhere) to clear the ruff-strict type-discipline
  budget (LIT001/LIT002/LIT010), keeping the pending-concurrency-keys
  holder mutable by design with a documented # mutable-ok.
- Rebuild _build_limits_index's grouping via a stable sort + groupby
  instead of a setdefault accumulator; caught and fixed a real bug in
  that rewrite where a per-unit loop re-consumed groupby's already-
  exhausted sub-iterator, silently returning empty limits for every
  group after the first unit.
- Fix a basedpyright reportIncompatibleMethodOverride: match
  async_filter_deployments's signature to CustomLogger's base exactly
  (list/dict, not Mapping/Sequence) since it's an override.
- Fix a second reportGeneralTypeIssues: two Final-annotated locals
  named `key` in sibling branches of the same function tripped
  "previously declared as Final" despite being on mutually exclusive
  paths; renamed them apart.
- Re-fix an eager-built f-string log message (%-style args instead)
  that a prior rewrite pass had inadvertently reintroduced.
This commit is contained in:
Deepanshu 2026-08-11 12:13:48 -04:00
parent 0004fe1b92
commit 48b71a8a32
3 changed files with 366 additions and 232 deletions

View file

@ -4486,7 +4486,7 @@ def _init_custom_logger_compatible_class(
if internal_usage_cache is None:
raise Exception(f"Internal Error: Cache cannot be empty - internal_usage_cache={internal_usage_cache}")
tag_rate_limiter_obj = _PROXY_TagRateLimiter(internal_usage_cache=internal_usage_cache)
tag_rate_limiter_obj: Final = _PROXY_TagRateLimiter(internal_usage_cache=internal_usage_cache)
if llm_router is not None and isinstance(llm_router, litellm.Router):
tag_rate_limiter_obj.update_variables(llm_router=llm_router)

View file

@ -2,10 +2,12 @@
import asyncio
import contextvars
from collections.abc import Callable
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from typing import TYPE_CHECKING, Literal
from itertools import groupby
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, NamedTuple, TypeAlias
from litellm._logging import verbose_proxy_logger
from litellm.caching.dual_cache import DualCache
@ -30,31 +32,40 @@ from litellm.types.utils import StandardLoggingPayload
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
Span = _Span
Span: TypeAlias = _Span
else:
Span = object
Span: TypeAlias = object
_LimitUnit = Literal["tokens", "requests", "dollars", "concurrency"]
_LIMIT_UNITS: tuple[_LimitUnit, ...] = ("tokens", "requests", "dollars", "concurrency")
_LimitUnit: TypeAlias = Literal["tokens", "requests", "dollars", "concurrency"]
_LIMIT_UNITS: Final[tuple[_LimitUnit, ...]] = ("tokens", "requests", "dollars", "concurrency")
# Units whose admission must be atomic (check-and-increment in one Redis
# round trip) because the increment amount is known upfront (always 1).
# tokens/dollars can't be: real usage is only known after the response, so
# they stay a read-then-account-on-success check with a documented,
# unavoidable admit-vs-account race.
_ATOMIC_UNITS: frozenset[_LimitUnit] = frozenset({"requests", "concurrency"})
_ATOMIC_UNITS: Final[frozenset[_LimitUnit]] = frozenset({"requests", "concurrency"})
_UNIT_TO_GROUP_FIELD: dict[_LimitUnit, str] = {
"tokens": "token_limits",
"requests": "request_limits",
"dollars": "dollar_limits",
"concurrency": "concurrency_limits",
}
_UNIT_TO_RATE_LIMIT_TYPE: dict[_LimitUnit, RateLimitType] = {
"tokens": RateLimitType.TOKENS,
"requests": RateLimitType.REQUESTS,
"dollars": RateLimitType.BUDGET,
"concurrency": RateLimitType.CONCURRENT_REQUESTS,
}
_UNIT_TO_GROUP_FIELD: Final[Mapping[_LimitUnit, str]] = MappingProxyType(
{
"tokens": "token_limits",
"requests": "request_limits",
"dollars": "dollar_limits",
"concurrency": "concurrency_limits",
}
)
_UNIT_TO_RATE_LIMIT_TYPE: Final[Mapping[_LimitUnit, RateLimitType]] = MappingProxyType(
{
"tokens": RateLimitType.TOKENS,
"requests": RateLimitType.REQUESTS,
"dollars": RateLimitType.BUDGET,
"concurrency": RateLimitType.CONCURRENT_REQUESTS,
}
)
# Shared read-only fallback for an absent/None mapping (request_kwargs,
# metadata, model_info, ...): avoids constructing a fresh mutable `{}` at
# every one of these call sites just to immediately call `.get()` on it.
_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
# Single-key atomic check-and-increment. Deliberately one key per script call
# (never a batch of differently-hash-tagged keys in one call): every tag_rl
@ -68,7 +79,7 @@ _UNIT_TO_RATE_LIMIT_TYPE: dict[_LimitUnit, RateLimitType] = {
# `atomic_check_and_increment_by_n` in parallel_request_limiter_v3.py, applied
# per-key instead of per-descriptor since each key already is one hash-tag
# group by construction.
TAG_RL_CHECK_AND_INCR_SCRIPT = """
TAG_RL_CHECK_AND_INCR_SCRIPT: Final = """
local key = KEYS[1]
local limit = tonumber(ARGV[1])
local increment = tonumber(ARGV[2])
@ -91,7 +102,7 @@ return { 1, new_value }
# be attributed to the exact reservation that caused it (see
# `_release_keys`'s docstring) degrades to under-counting rather than a
# negative counter that would admit unlimited requests.
TAG_RL_DECR_FLOOR_ZERO_SCRIPT = """
TAG_RL_DECR_FLOOR_ZERO_SCRIPT: Final = """
local key = KEYS[1]
local delta = tonumber(ARGV[1])
local new_value = redis.call('INCRBY', key, delta)
@ -113,13 +124,13 @@ class _ConfiguredLimit:
deployment_scope: tuple[str, ...] | None
def _extract_identity(tags: list[str], tag_id: str) -> str | None:
def _extract_identity(tags: Sequence[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
skipped so they can never be misread as an identity value.
"""
prefix = f"{tag_id}:"
prefix: Final = f"{tag_id}:"
for tag in tags:
if tag.startswith("!"):
continue
@ -128,42 +139,71 @@ def _extract_identity(tags: list[str], tag_id: str) -> str | None:
return None
def _deployment_id(deployment: dict) -> str | None:
return (deployment.get("model_info") or {}).get("id")
def _deployment_id(deployment: Mapping[str, object]) -> str | None:
return (deployment.get("model_info") or _EMPTY_MAPPING).get("id")
def _extract_team_id(request_kwargs: dict) -> str | None:
def _extract_team_id(request_kwargs: Mapping[str, object]) -> 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
`litellm_metadata`)."""
metadata = request_kwargs.get("metadata") or {}
litellm_metadata = request_kwargs.get("litellm_metadata") or {}
team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
metadata: Final = request_kwargs.get("metadata") or _EMPTY_MAPPING
litellm_metadata: Final = request_kwargs.get("litellm_metadata") or _EMPTY_MAPPING
team_id: Final = metadata.get("user_api_key_team_id") or litellm_metadata.get("user_api_key_team_id")
return team_id if isinstance(team_id, str) else None
def _extract_key_hash(request_kwargs: dict) -> str | None:
def _extract_key_hash(request_kwargs: Mapping[str, object]) -> 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
the hashed token (see `litellm_pre_call_utils.py`)."""
metadata = request_kwargs.get("metadata") or {}
litellm_metadata = request_kwargs.get("litellm_metadata") or {}
key_hash = metadata.get("user_api_key") or litellm_metadata.get("user_api_key")
metadata: Final = request_kwargs.get("metadata") or _EMPTY_MAPPING
litellm_metadata: Final = request_kwargs.get("litellm_metadata") or _EMPTY_MAPPING
key_hash: Final = metadata.get("user_api_key") or litellm_metadata.get("user_api_key")
return key_hash if isinstance(key_hash, str) else None
def _entries_for_unit(deployment: dict, unit: _LimitUnit) -> list[TagRateLimitEntry]:
raw_tag_rate_limits = (deployment.get("model_info") or {}).get("tag_rate_limits")
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:
return []
tag_rate_limits = TagRateLimits.model_validate(raw_tag_rate_limits)
group = getattr(tag_rate_limits, _UNIT_TO_GROUP_FIELD[unit])
return group.limits if group is not None else []
return ()
tag_rate_limits: Final = TagRateLimits.model_validate(raw_tag_rate_limits)
group: Final = getattr(tag_rate_limits, _UNIT_TO_GROUP_FIELD[unit])
return tuple(group.limits) if group is not None else ()
def _build_group_limits(deployments: list[dict], unit: _LimitUnit) -> list[_ConfiguredLimit]:
def _configured_limit_for_signature(
unit: _LimitUnit,
signature: tuple[str, str, float, int, bool],
declaring_ids: Sequence[str],
is_chain_wide: bool,
) -> _ConfiguredLimit | None:
tag_id, name, limit, period_seconds, scope_by_key_hash = signature
if unit == "concurrency" and not is_chain_wide:
verbose_proxy_logger.warning(
"tag_rate_limiter: concurrency_limits entry %r (tag_id=%s) is not declared identically by every "
"deployment sharing this model_name; per-deployment-scoped concurrency limits are not supported "
"and this entry is being skipped entirely.",
name,
tag_id,
)
return None
return _ConfiguredLimit(
unit=unit,
entry=TagRateLimitEntry(
name=name,
tag_id=tag_id,
limit=limit,
period_seconds=period_seconds,
scope_by_key_hash=scope_by_key_hash,
),
deployment_scope=None if is_chain_wide else tuple(sorted(declaring_ids)),
)
def _build_group_limits(deployments: Sequence[Mapping[str, object]], unit: _LimitUnit) -> tuple[_ConfiguredLimit, ...]:
"""
One `_ConfiguredLimit` per distinct (tag_id, name, limit, period_seconds)
declared for `unit` across `deployments` (all sharing one `model_name`).
@ -190,47 +230,50 @@ def _build_group_limits(deployments: list[dict], unit: _LimitUnit) -> list[_Conf
creating a bucket that can leak; only chain-wide concurrency entries
(identical across every deployment in the group) are supported.
"""
declaring_ids_by_signature: dict[tuple[str, str, float, int, bool], list[str]] = {}
# Insertion order here is load-bearing: it decides which limit's
# ProxyRateLimitError surfaces first when several are breached by the
# same hop (see async_filter_deployments). A presort-based
# itertools.groupby would need to sort by signature to group it, which
# would scramble that first-seen order, so this stays a plain
# accumulator instead.
declaring_ids_by_signature: Final = {} # mutable-ok: first-seen order here decides which limit's error raises first (see comment above); sorting to use groupby would scramble it
for deployment in deployments:
dep_id = _deployment_id(deployment)
if dep_id is None:
continue
for entry in _entries_for_unit(deployment, unit):
signature = (entry.tag_id, entry.name, entry.limit, entry.period_seconds, entry.scope_by_key_hash)
declaring_ids_by_signature.setdefault(signature, []).append(dep_id)
ids_for_signature = declaring_ids_by_signature.setdefault(signature, []) # mutable-ok: see comment above
ids_for_signature.append(dep_id)
distinct_signatures_by_name: dict[tuple[str, str], int] = {}
for tag_id, name, _limit, _period, _scope_by_key_hash in declaring_ids_by_signature:
key = (tag_id, name)
distinct_signatures_by_name[key] = distinct_signatures_by_name.get(key, 0) + 1
total_deployments = len(deployments)
configured: list[_ConfiguredLimit] = []
for signature, declaring_ids in declaring_ids_by_signature.items():
tag_id, name, limit, period_seconds, scope_by_key_hash = signature
is_chain_wide = distinct_signatures_by_name[(tag_id, name)] == 1 and len(declaring_ids) == total_deployments
if unit == "concurrency" and not is_chain_wide:
verbose_proxy_logger.warning(
"tag_rate_limiter: concurrency_limits entry %r (tag_id=%s) is not declared identically by every "
"deployment sharing this model_name; per-deployment-scoped concurrency limits are not supported "
"and this entry is being skipped entirely.",
name,
tag_id,
distinct_signature_count_by_name: Final[Mapping[tuple[str, str], int]] = MappingProxyType(
{
(tag_id, name): sum(
1
for other_tag_id, other_name, *_rest in declaring_ids_by_signature
if (other_tag_id, other_name) == (tag_id, name)
)
continue
configured.append(
_ConfiguredLimit(
unit=unit,
entry=TagRateLimitEntry(
name=name,
tag_id=tag_id,
limit=limit,
period_seconds=period_seconds,
scope_by_key_hash=scope_by_key_hash,
for tag_id, name, *_rest in declaring_ids_by_signature
}
)
total_deployments: Final = len(deployments)
configured: Final = tuple(
configured_limit
for signature, declaring_ids in declaring_ids_by_signature.items()
if (
configured_limit := _configured_limit_for_signature(
unit,
signature,
declaring_ids,
is_chain_wide=(
distinct_signature_count_by_name[(signature[0], signature[1])] == 1
and len(declaring_ids) == total_deployments
),
deployment_scope=None if is_chain_wide else tuple(sorted(declaring_ids)),
)
)
is not None
)
return configured
@ -246,18 +289,27 @@ class _LimitsIndex:
overwrite another's.
"""
by_model_name: dict[str, list[_ConfiguredLimit]]
by_team_alias: dict[tuple[str, str], list[_ConfiguredLimit]]
by_model_name: Mapping[str, tuple[_ConfiguredLimit, ...]]
by_team_alias: Mapping[tuple[str, str], tuple[_ConfiguredLimit, ...]]
def resolve(self, model: str, team_id: str | None) -> list[_ConfiguredLimit]:
def resolve(self, model: str, team_id: str | None) -> tuple[_ConfiguredLimit, ...]:
if team_id is not None:
scoped = self.by_team_alias.get((team_id, model))
scoped: Final = self.by_team_alias.get((team_id, model))
if scoped is not None:
return scoped
return self.by_model_name.get(model, [])
return self.by_model_name.get(model, ())
def _build_limits_index(model_list: list[dict]) -> _LimitsIndex:
def _team_alias_key(deployment: Mapping[str, object]) -> tuple[str, str] | None:
model_info: Final = deployment.get("model_info") or _EMPTY_MAPPING
team_id: Final = model_info.get("team_id")
team_public_model_name: Final = model_info.get("team_public_model_name")
if team_id and team_public_model_name:
return (team_id, team_public_model_name)
return None
def _build_limits_index(model_list: Sequence[Mapping[str, object]]) -> _LimitsIndex:
"""
`by_model_name` is keyed by every deployment's own `model_name`, grouping
deployments that share one.
@ -285,28 +337,41 @@ def _build_limits_index(model_list: list[dict]) -> _LimitsIndex:
`_build_group_limits` has no `model_name`-specific logic (it only reads
each deployment's own id and its own entries), so it's safe to reuse
unchanged for a deployment set spanning multiple `model_name` values.
Deployments are grouped via a stable sort + itertools.groupby rather than
a setdefault-in-a-loop accumulator: `sorted` is stable, so deployments
sharing a key keep the exact same relative order `_build_group_limits`
would have seen them in without the sort, which is what keeps this safe
(that relative order decides first-seen signature order downstream).
"""
groups: dict[str, list[dict]] = {}
alias_groups: dict[tuple[str, str], list[dict]] = {}
for deployment in model_list:
groups.setdefault(deployment["model_name"], []).append(deployment)
model_info = deployment.get("model_info") or {}
team_id = model_info.get("team_id")
team_public_model_name = model_info.get("team_public_model_name")
if team_id and team_public_model_name:
alias_groups.setdefault((team_id, team_public_model_name), []).append(deployment)
sorted_by_model_name: Final = sorted(model_list, key=lambda deployment: deployment["model_name"])
by_model_name: Final[Mapping[str, tuple[_ConfiguredLimit, ...]]] = MappingProxyType(
{
model_name: configured
for model_name, deployment_group in groupby(
sorted_by_model_name, key=lambda deployment: deployment["model_name"]
)
for group in (tuple(deployment_group),)
if (configured := tuple(limit for unit in _LIMIT_UNITS for limit in _build_group_limits(group, unit)))
}
)
by_model_name: dict[str, list[_ConfiguredLimit]] = {}
for model_name, deployments in groups.items():
configured = [limit for unit in _LIMIT_UNITS for limit in _build_group_limits(deployments, unit)]
if configured:
by_model_name[model_name] = configured
by_team_alias: dict[tuple[str, str], list[_ConfiguredLimit]] = {}
for alias_key, deployments in alias_groups.items():
configured = [limit for unit in _LIMIT_UNITS for limit in _build_group_limits(deployments, unit)]
if configured:
by_team_alias[alias_key] = configured
aliased: Final = tuple(
(alias_key, deployment) for deployment in model_list if (alias_key := _team_alias_key(deployment)) is not None
)
sorted_by_alias: Final = sorted(aliased, key=lambda pair: pair[0])
by_team_alias: Final[Mapping[tuple[str, str], tuple[_ConfiguredLimit, ...]]] = MappingProxyType(
{
alias_key: alias_configured
for alias_key, alias_group in groupby(sorted_by_alias, key=lambda pair: pair[0])
for aliased_group in (tuple(dep for _key, dep in alias_group),)
if (
alias_configured := tuple(
limit for unit in _LIMIT_UNITS for limit in _build_group_limits(aliased_group, unit)
)
)
}
)
return _LimitsIndex(by_model_name=by_model_name, by_team_alias=by_team_alias)
@ -316,7 +381,7 @@ def _build_limits_index(model_list: list[dict]) -> _LimitsIndex:
# in place via the admin API, which never changes len(model_list)). Router
# exposes no generic "config changed" version counter to key off instead, so
# this bounds staleness by simply re-checking periodically.
_INDEX_TTL_SECONDS = 5.0
_INDEX_TTL_SECONDS: Final = 5.0
# Floor for a concurrency reservation's self-heal TTL, regardless of the
# configured period_seconds. A reservation that expires while its request is
@ -324,7 +389,7 @@ _INDEX_TTL_SECONDS = 5.0
# generous floor keeps that window far larger than any realistic request
# duration, at the cost of a leaked (crashed-worker) slot self-healing more
# slowly. period_seconds can still raise the TTL further, never lower it.
_CONCURRENCY_MIN_SAFETY_TTL_SECONDS = 3600
_CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600
# Concurrency reservation keys accumulated for the current logical request,
@ -340,19 +405,20 @@ class _PendingConcurrencyKeys:
__slots__ = ("keys",)
def __init__(self) -> None:
self.keys: list[str] = []
self.keys: list[str] = [] # mutable-ok: shared across asyncio.create_task forks by design; see class docstring
_pending_concurrency_keys: contextvars.ContextVar[_PendingConcurrencyKeys | None] = contextvars.ContextVar(
_pending_concurrency_keys: Final[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)
existing: Final = _pending_concurrency_keys.get()
if existing is not None:
return existing
holder: Final = _PendingConcurrencyKeys()
_pending_concurrency_keys.set(holder)
return holder
@ -364,12 +430,12 @@ class _TagRateLimitIndex:
self._time_provider = time_provider
self._cache_key: tuple[int, int] | None = None
self._built_at: float = 0.0
self._index: _LimitsIndex = _LimitsIndex(by_model_name={}, by_team_alias={})
self._index: _LimitsIndex = _LimitsIndex(by_model_name=MappingProxyType({}), by_team_alias=MappingProxyType({}))
def get(self, llm_router: Router) -> _LimitsIndex:
model_list = llm_router.model_list or []
cache_key = (id(llm_router), len(model_list))
now = self._time_provider().timestamp()
model_list: Final = llm_router.model_list or ()
cache_key: Final = (id(llm_router), len(model_list))
now: Final = self._time_provider().timestamp()
if cache_key != self._cache_key or (now - self._built_at) >= _INDEX_TTL_SECONDS:
self._index = _build_limits_index(model_list)
self._cache_key = cache_key
@ -388,9 +454,12 @@ def _bucket_key(
bucket_id: int,
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 ""
hash_tag = f"tag_rl:{model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:{scope}:{tag_value}{key_suffix}"
scope: Final = _scope_suffix(configured.deployment_scope)
key_suffix: Final = f":key:{key_hash}" if key_hash is not None else ""
hash_tag: Final = (
f"tag_rl:{model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:"
f"{scope}:{tag_value}{key_suffix}"
)
return f"{{{hash_tag}}}:{bucket_id}"
@ -403,13 +472,82 @@ def _inflight_key(
"""Concurrency counter key: not epoch-bucketed, since "how many are in
flight right now" has no window to reset on -- it's released explicitly
on completion, with a TTL fallback only for a leaked (crashed) reservation."""
scope = _scope_suffix(configured.deployment_scope)
key_suffix = f":key:{key_hash}" if key_hash is not None else ""
hash_tag = f"tag_rl:{model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:{scope}:{tag_value}{key_suffix}"
scope: Final = _scope_suffix(configured.deployment_scope)
key_suffix: Final = f":key:{key_hash}" if key_hash is not None else ""
hash_tag: Final = (
f"tag_rl:{model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:"
f"{scope}:{tag_value}{key_suffix}"
)
return f"{{{hash_tag}}}:inflight"
class _PROXY_TagRateLimiter(CustomLogger):
class _ClassifiedCheck(NamedTuple):
configured_limit: _ConfiguredLimit
tag_value: str
key: str
is_atomic: bool
def _classify_check(
configured_limit: _ConfiguredLimit,
model: str,
tags: Sequence[str],
present_deployment_ids: frozenset[str],
request_kwargs: Mapping[str, object],
now: float,
) -> _ClassifiedCheck | None:
if configured_limit.deployment_scope is not None and not (
present_deployment_ids & frozenset(configured_limit.deployment_scope)
):
return None
tag_value: Final = _extract_identity(tags, configured_limit.entry.tag_id)
if tag_value is None:
return None
key_hash: Final = _extract_key_hash(request_kwargs) if configured_limit.entry.scope_by_key_hash else None
if configured_limit.unit == "concurrency":
inflight_key: Final = _inflight_key(model, configured_limit, tag_value, key_hash=key_hash)
return _ClassifiedCheck(configured_limit, tag_value, inflight_key, is_atomic=True)
bucket_id: Final = int(now) // configured_limit.entry.period_seconds
bucket_key_value: Final = _bucket_key(model, configured_limit, tag_value, bucket_id, key_hash=key_hash)
return _ClassifiedCheck(
configured_limit, tag_value, bucket_key_value, is_atomic=configured_limit.unit in _ATOMIC_UNITS
)
def _increment_operation_for_limit(
configured_limit: _ConfiguredLimit,
model_group: str,
tags: Sequence[str],
deployment_id: str | None,
key_hash: str | None,
increment_by_unit: Mapping[_LimitUnit, float],
now: float,
) -> RedisPipelineIncrementOperation | None:
if configured_limit.unit == "concurrency":
return None # released above, from _pending_concurrency_keys
if configured_limit.deployment_scope is not None and deployment_id not in configured_limit.deployment_scope:
return None
tag_value: Final = _extract_identity(tags, configured_limit.entry.tag_id)
if tag_value is None:
return None
if configured_limit.unit not in increment_by_unit:
return None # "requests" is accounted atomically at admission, not here
increment_value: Final = increment_by_unit[configured_limit.unit]
if increment_value == 0:
return None
bucket_id: Final = int(now) // configured_limit.entry.period_seconds
key_hash_for_limit: Final = key_hash if configured_limit.entry.scope_by_key_hash else None
key: Final = _bucket_key(model_group, configured_limit, tag_value, bucket_id, key_hash=key_hash_for_limit)
return RedisPipelineIncrementOperation(
key=key,
increment_value=increment_value,
ttl=configured_limit.entry.period_seconds + 3600,
)
class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only referenced via the deferred import in litellm_logging.py's callback resolver; basedpyright doesn't trace that usage
CustomLogger
):
def __init__(
self,
internal_usage_cache: DualCache,
@ -421,7 +559,7 @@ class _PROXY_TagRateLimiter(CustomLogger):
self._index = _TagRateLimitIndex(time_provider=self._time_provider)
self._lock = asyncio.Lock()
self.llm_router: Router | None = None
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
redis_cache: Final = 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
)
@ -436,15 +574,17 @@ class _PROXY_TagRateLimiter(CustomLogger):
"""Single-key atomic check-and-increment. Always one key per Lua
call -- see TAG_RL_CHECK_AND_INCR_SCRIPT's module docstring for why."""
if self._check_and_incr_script is not None:
raw = await self._check_and_incr_script(keys=[key], args=[limit, increment, ttl])
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 = await self.internal_usage_cache.async_get_cache(key=key, litellm_parent_otel_span=None)
current = float(current_value) if current_value is not None else 0.0
current_value: Final = await self.internal_usage_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 = current + increment
new_value: Final = current + increment
await self.internal_usage_cache.async_set_cache(
key=key, value=new_value, ttl=ttl, litellm_parent_otel_span=None
)
@ -452,19 +592,21 @@ class _PROXY_TagRateLimiter(CustomLogger):
async def _decrement_floor_zero(self, key: str, delta: float) -> None:
if self._decr_floor_zero_script is not None:
await self._decr_floor_zero_script(keys=[key], args=[delta])
await self._decr_floor_zero_script(keys=(key,), args=(delta,))
return
async with self._lock:
current_value = await self.internal_usage_cache.async_get_cache(key=key, litellm_parent_otel_span=None)
current = float(current_value) if current_value is not None else 0.0
current_value: Final = await self.internal_usage_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 self.internal_usage_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: list[tuple[str, float, float, int]],
) -> tuple[int | None, list[float]]:
checks: Sequence[tuple[str, float, float, int]],
) -> tuple[int | None, tuple[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
@ -487,86 +629,88 @@ class _PROXY_TagRateLimiter(CustomLogger):
one key's current (unmodified) value.
"""
if not checks:
return None, []
return None, ()
admitted_values: list[float] = []
# Sequential async admission: each element needs its own awaited
# Redis round trip, and a rejection mid-loop discards everything
# accumulated so far in favor of refunding and returning early, so
# this can't be expressed as a one-shot comprehension.
admitted_values: Final = [] # mutable-ok: sequential async accumulator, discardable on early rejection; see comment above
for index, (key, limit, increment, ttl) in enumerate(checks):
admitted, value = await self._check_and_increment_one(key, limit, increment, ttl)
if admitted:
admitted_values.append(value)
admitted_values.append(value) # mutable-ok: see accumulator comment above
continue
for refund_index in range(index):
refund_key, _limit, refund_increment, _ttl = checks[refund_index]
try:
await self._decrement_floor_zero(refund_key, -refund_increment)
except Exception as e: # noqa: BLE001 - one failed refund must not block refunding the rest
verbose_proxy_logger.warning(f"tag_rate_limiter: failed to refund {refund_key} on rollback: {e}")
return index, [value]
verbose_proxy_logger.warning("tag_rate_limiter: failed to refund %s on rollback: %s", refund_key, e)
return index, (value,)
return None, admitted_values
return None, tuple(admitted_values)
async def async_filter_deployments(
self,
model: str,
healthy_deployments: list[dict],
messages: list[AllMessageValues] | None,
request_kwargs: dict | None = None,
healthy_deployments: list, # mutable-ok: must match CustomLogger's base signature exactly, or basedpyright flags reportIncompatibleMethodOverride
messages: list[AllMessageValues] | None, # mutable-ok: see reason above
request_kwargs: dict | None = None, # mutable-ok: see reason above
parent_otel_span: Span | None = None,
) -> list[dict]:
if not healthy_deployments or not isinstance(healthy_deployments, list) or self.llm_router is None:
) -> list[dict]: # mutable-ok: see reason above
if (
not healthy_deployments
or not isinstance(healthy_deployments, list) # pyright: ignore[reportUnnecessaryIsInstance] # defensive at runtime despite the static list annotation Router's own callers aren't guaranteed to honor
or self.llm_router is None
):
return healthy_deployments
request_kwargs = request_kwargs or {}
configured = self._index.get(self.llm_router).resolve(model, _extract_team_id(request_kwargs))
resolved_request_kwargs: Final = request_kwargs or _EMPTY_MAPPING
configured: Final = self._index.get(self.llm_router).resolve(model, _extract_team_id(resolved_request_kwargs))
if not configured:
return healthy_deployments
metadata_variable_name = get_metadata_variable_name_from_kwargs(request_kwargs)
tags = _get_tags_from_request_kwargs(request_kwargs, metadata_variable_name=metadata_variable_name)
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(resolved_request_kwargs)
tags: Final = _get_tags_from_request_kwargs(
resolved_request_kwargs, metadata_variable_name=metadata_variable_name
)
present_deployment_ids = {dep_id for d in healthy_deployments if (dep_id := _deployment_id(d)) is not None}
present_deployment_ids: Final[frozenset[str]] = frozenset(
dep_id for d in healthy_deployments if (dep_id := _deployment_id(d)) is not None
)
now = self._time_provider().timestamp()
read_only_checks: list[tuple[_ConfiguredLimit, str, str]] = []
atomic_checks: list[tuple[_ConfiguredLimit, str, str]] = []
for configured_limit in configured:
if configured_limit.deployment_scope is not None and not (
present_deployment_ids & set(configured_limit.deployment_scope)
):
continue
tag_value = _extract_identity(tags, configured_limit.entry.tag_id)
if tag_value is None:
continue
key_hash = _extract_key_hash(request_kwargs) if configured_limit.entry.scope_by_key_hash else None
if configured_limit.unit == "concurrency":
key = _inflight_key(model, configured_limit, tag_value, key_hash=key_hash)
atomic_checks.append((configured_limit, tag_value, key))
elif configured_limit.unit in _ATOMIC_UNITS:
bucket_id = int(now) // configured_limit.entry.period_seconds
key = _bucket_key(model, configured_limit, tag_value, bucket_id, key_hash=key_hash)
atomic_checks.append((configured_limit, tag_value, key))
else:
bucket_id = int(now) // configured_limit.entry.period_seconds
key = _bucket_key(model, configured_limit, tag_value, bucket_id, key_hash=key_hash)
read_only_checks.append((configured_limit, tag_value, key))
now: Final = self._time_provider().timestamp()
classified: Final = tuple(
check
for configured_limit in configured
if (
check := _classify_check(
configured_limit, model, tags, present_deployment_ids, resolved_request_kwargs, now
)
)
is not None
)
read_only_checks: Final = tuple((c.configured_limit, c.tag_value, c.key) for c in classified if not c.is_atomic)
atomic_checks: Final = tuple((c.configured_limit, c.tag_value, c.key) for c in classified if c.is_atomic)
current_values = await self._read_only_values(read_only_checks, parent_otel_span)
current_values: Final = await self._read_only_values(read_only_checks, parent_otel_span)
self._raise_if_over_limit(read_only_checks, current_values, model)
if atomic_checks:
failing_index, values = await self._atomic_check_and_increment(
[
tuple(
(key, configured_limit.entry.limit, 1.0, self._ttl_for(configured_limit))
for configured_limit, _tag_value, key in atomic_checks
]
)
)
if failing_index is not None:
configured_limit, tag_value, _key = atomic_checks[failing_index]
self._raise_over_limit(configured_limit, tag_value, model, current=values[0])
concurrency_keys = [
concurrency_keys: Final = tuple(
key for configured_limit, _tag_value, key in atomic_checks if configured_limit.unit == "concurrency"
]
)
if concurrency_keys:
_pending_concurrency_holder().keys.extend(concurrency_keys)
@ -586,23 +730,23 @@ class _PROXY_TagRateLimiter(CustomLogger):
async def _read_only_values(
self,
read_only_checks: list[tuple[_ConfiguredLimit, str, str]],
read_only_checks: Sequence[tuple[_ConfiguredLimit, str, str]],
parent_otel_span: Span | None,
) -> list[float | None]:
) -> tuple[float | None, ...]:
if not read_only_checks:
return []
keys = [key for _cfg, _tag_value, key in read_only_checks]
current_values = await self.internal_usage_cache.async_batch_get_cache(
keys=keys,
return ()
keys: Final = tuple(key for _cfg, _tag_value, key in read_only_checks)
current_values: Final = await self.internal_usage_cache.async_batch_get_cache(
keys=list(keys), # mutable-ok: async_batch_get_cache requires a real list; converted only at this boundary
parent_otel_span=parent_otel_span,
local_only=False,
)
return current_values if current_values is not None else [None] * len(keys)
return tuple(current_values) if current_values is not None else tuple(None for _ in keys)
def _raise_if_over_limit(
self,
read_only_checks: list[tuple[_ConfiguredLimit, str, str]],
current_values: list[float | None],
read_only_checks: Sequence[tuple[_ConfiguredLimit, str, str]],
current_values: Sequence[float | None],
model: str,
) -> None:
for (configured_limit, tag_value, _key), current_value in zip(read_only_checks, current_values):
@ -629,7 +773,7 @@ class _PROXY_TagRateLimiter(CustomLogger):
configured_limit.entry.limit,
)
raise ProxyRateLimitError(
detail={
detail={ # mutable-ok: must stay a real dict -- async_log_failure_event below (and generic proxy exception rendering, e.g. proxy/utils.py, guardrail hooks) branch on isinstance(exc.detail, dict); a MappingProxyType silently falls through those checks
"error": "tag_rate_limit_exceeded",
"type": configured_limit.unit,
"tag_id": configured_limit.entry.tag_id,
@ -638,13 +782,13 @@ class _PROXY_TagRateLimiter(CustomLogger):
"limit": configured_limit.entry.limit,
"period_seconds": configured_limit.entry.period_seconds,
},
headers={"retry-after": str(configured_limit.entry.period_seconds)},
headers={"retry-after": str(configured_limit.entry.period_seconds)}, # mutable-ok: same as detail
rate_limit_type=_UNIT_TO_RATE_LIMIT_TYPE[configured_limit.unit],
model=model,
llm_provider="litellm_proxy",
)
async def _release_keys(self, keys: list[str]) -> None:
async def _release_keys(self, keys: Sequence[str]) -> None:
"""
Release each key by one slot. This does not verify the completing
request still owns a live reservation (no per-request slot identity
@ -662,16 +806,16 @@ class _PROXY_TagRateLimiter(CustomLogger):
verbose_proxy_logger.warning("tag_rate_limiter: failed to release concurrency slot %s: %s", key, e)
@staticmethod
def _pop_pending_concurrency_keys() -> list[str]:
def _pop_pending_concurrency_keys() -> tuple[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()
holder: Final = _pending_concurrency_keys.get()
if holder is None or not holder.keys:
return []
keys = list(holder.keys)
return ()
keys: Final = tuple(holder.keys)
for key in keys:
try:
holder.keys.remove(key)
@ -681,73 +825,63 @@ class _PROXY_TagRateLimiter(CustomLogger):
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
if isinstance(kwargs.get("exception"), ProxyRateLimitError):
detail = kwargs["exception"].detail if isinstance(kwargs["exception"].detail, dict) else {}
detail: Final = (
kwargs["exception"].detail if isinstance(kwargs["exception"].detail, dict) else _EMPTY_MAPPING
)
if detail.get("error") == "tag_rate_limit_exceeded":
return
release_keys = self._pop_pending_concurrency_keys()
release_keys: Final = self._pop_pending_concurrency_keys()
if 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 = self._pop_pending_concurrency_keys()
release_keys: Final = self._pop_pending_concurrency_keys()
if release_keys:
asyncio.create_task(self._release_keys(release_keys))
if self.llm_router is None:
return
standard_logging_object: StandardLoggingPayload | None = kwargs.get("standard_logging_object")
standard_logging_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object")
if standard_logging_object is None:
return
model_group = standard_logging_object.get("model_group")
model_group: Final = standard_logging_object.get("model_group")
if not model_group:
return
standard_logging_metadata = standard_logging_object.get("metadata") or {}
team_id = standard_logging_metadata.get("user_api_key_team_id")
key_hash = standard_logging_metadata.get("user_api_key_hash")
configured = self._index.get(self.llm_router).resolve(model_group, team_id)
standard_logging_metadata: Final = standard_logging_object.get("metadata") or _EMPTY_MAPPING
team_id: Final = standard_logging_metadata.get("user_api_key_team_id")
key_hash: Final = standard_logging_metadata.get("user_api_key_hash")
configured: Final = self._index.get(self.llm_router).resolve(model_group, team_id)
if not configured:
return
metadata_variable_name = get_metadata_variable_name_from_kwargs(kwargs)
tags = _get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(kwargs)
tags: Final = _get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)
if not tags:
return
deployment_id = standard_logging_object.get("model_id")
now = self._time_provider().timestamp()
increment_by_unit: dict[_LimitUnit, float] = {
"tokens": float(standard_logging_object.get("total_tokens") or 0),
"dollars": float(standard_logging_object.get("response_cost") or 0),
}
deployment_id: Final = standard_logging_object.get("model_id")
now: Final = 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),
}
)
operations: list[RedisPipelineIncrementOperation] = []
for configured_limit in configured:
if configured_limit.unit == "concurrency":
continue # released above, from _pending_concurrency_keys
if configured_limit.deployment_scope is not None and deployment_id not in configured_limit.deployment_scope:
continue
tag_value = _extract_identity(tags, configured_limit.entry.tag_id)
if tag_value is None:
continue
if configured_limit.unit not in increment_by_unit:
continue # "requests" is accounted atomically at admission, not here
increment_value = increment_by_unit[configured_limit.unit]
if increment_value == 0:
continue
bucket_id = int(now) // configured_limit.entry.period_seconds
key_hash_for_limit = key_hash if configured_limit.entry.scope_by_key_hash else None
key = _bucket_key(model_group, configured_limit, tag_value, bucket_id, key_hash=key_hash_for_limit)
operations.append(
RedisPipelineIncrementOperation(
key=key,
increment_value=increment_value,
ttl=configured_limit.entry.period_seconds + 3600,
operations: Final = tuple(
operation
for configured_limit in configured
if (
operation := _increment_operation_for_limit(
configured_limit, model_group, tags, deployment_id, key_hash, increment_by_unit, now
)
)
is not None
)
if not operations:
return

View file

@ -149,7 +149,7 @@ def test_build_group_limits_per_deployment_when_only_some_declare_it():
def test_build_group_limits_empty_when_no_deployment_configures_unit():
deployments = [_deployment("grp", "dep-1", {}), _deployment("grp", "dep-2", {})]
assert _build_group_limits(deployments, "tokens") == []
assert _build_group_limits(deployments, "tokens") == ()
# ---------------------------------------------------------------------------
@ -1403,7 +1403,7 @@ def test_concurrency_divergent_config_is_dropped_not_scoped_per_deployment():
),
]
configured = _build_group_limits(deployments, "concurrency")
assert configured == []
assert configured == ()
def test_concurrency_partial_declaration_is_dropped_not_scoped_per_deployment():
@ -1414,7 +1414,7 @@ def test_concurrency_partial_declaration_is_dropped_not_scoped_per_deployment():
_deployment("grp", "dep-2", {}),
]
configured = _build_group_limits(deployments, "concurrency")
assert configured == []
assert configured == ()
def test_concurrency_identical_across_all_deployments_is_still_chain_wide():
@ -1468,7 +1468,7 @@ async def test_release_in_a_forked_task_is_visible_to_the_parent_context():
return _PROXY_TagRateLimiter._pop_pending_concurrency_keys()
released = await asyncio.create_task(detached_release())
assert released == ["key1"]
assert released == ("key1",)
# The parent's own binding must see the same, now-empty holder --
# not a stale copy still holding "key1".
@ -1488,7 +1488,7 @@ async def test_release_does_not_sweep_up_a_key_appended_after_its_snapshot():
return released
released = await asyncio.create_task(detached_release_then_sibling_admits())
assert released == ["key1"]
assert released == ("key1",)
# key2 must still be pending for its own hop's eventual release.
assert _pending_concurrency_holder().keys == ["key2"]
@ -1499,8 +1499,8 @@ async def test_release_is_not_repeated_for_the_same_snapshot():
_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 == []
assert first == ("key1",)
assert second == ()
# ---------------------------------------------------------------------------