From 170aa0604bfa859be9313ebb29bb679b8b1ab760 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Tue, 11 Aug 2026 13:16:45 -0400 Subject: [PATCH] fix(rate-limiting): scope alias buckets by team, stop caller-forged metadata bypassing key-scoped limits - _ConfiguredLimit now carries team_scope: two teams can publish the identical team_public_model_name alias, and the limits index already scopes lookup by (team_id, alias) correctly -- but the Redis bucket key itself never included team_id, so identically-named, identically-configured limits from two different teams collided on the same counter. Fold team_scope into the hash tag for both admission and concurrency keys. - _extract_key_hash and _extract_team_id no longer OR across metadata and litellm_metadata. litellm_pre_call_utils.py writes the real, server-authenticated value into only whichever one field is authoritative for a given route, leaving the other exactly as the caller sent it -- so an OR-fallback let a caller-forged metadata.user_api_key (on a route where litellm_metadata is authoritative) win over the real hash and bypass every scope_by_key_hash=True limit by sending a fresh forged value per request. Both extractors now read only the field get_metadata_variable_name_from_kwargs names as authoritative, matching the pattern this file already uses correctly for tags. --- litellm/proxy/hooks/tag_rate_limiter.py | 98 ++++++++++++------- .../proxy/hooks/test_tag_rate_limiter.py | 80 ++++++++++++++- 2 files changed, 140 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/hooks/tag_rate_limiter.py b/litellm/proxy/hooks/tag_rate_limiter.py index d9171ce01b3..993893d452a 100644 --- a/litellm/proxy/hooks/tag_rate_limiter.py +++ b/litellm/proxy/hooks/tag_rate_limiter.py @@ -3,7 +3,7 @@ import asyncio import contextvars from collections.abc import Callable, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, replace from datetime import datetime from itertools import groupby from types import MappingProxyType @@ -122,6 +122,14 @@ class _ConfiguredLimit: # bucket). Otherwise the sorted deployment ids that declared this exact # value -- the bucket is shared among only those deployments. deployment_scope: tuple[str, ...] | None + # The team_id this limit was resolved under via `by_team_alias`, or None + # when resolved via `by_model_name`. team_public_model_name is only + # unique per team, so two teams can publish the identical alias string; + # without the team_id folded into the bucket key too, both teams' + # identically-named, identically-configured limits would collide on the + # same Redis counter despite the index itself correctly scoping the + # lookup by (team_id, alias). + team_scope: str | None = None def _extract_identity(tags: Sequence[str], tag_id: str) -> str | None: @@ -143,25 +151,29 @@ 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: 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: 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") +def _extract_team_id(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None: + """Reads `user_api_key_team_id` from only the one field + `get_metadata_variable_name_from_kwargs` names as authoritative for this + request -- never falling back to the other field, since + `litellm_pre_call_utils.py` writes the real, server-authenticated value + into that one field alone and leaves the other exactly as the caller + sent it. An OR-fallback across both would let a caller's own + `metadata.user_api_key_team_id` (still present, unvalidated, on a route + where `litellm_metadata` is the authoritative field) win over the real + value.""" + active: Final = request_kwargs.get(metadata_variable_name) or _EMPTY_MAPPING + team_id: Final = active.get("user_api_key_team_id") return team_id if isinstance(team_id, str) else 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: 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") +def _extract_key_hash(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None: + """Same single-authoritative-field lookup as `_extract_team_id`, but for + the calling virtual key's 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`). + """ + active: Final = request_kwargs.get(metadata_variable_name) or _EMPTY_MAPPING + key_hash: Final = active.get("user_api_key") return key_hash if isinstance(key_hash, str) else None @@ -367,7 +379,9 @@ def _build_limits_index(model_list: Sequence[Mapping[str, object]]) -> _LimitsIn 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) + replace(limit, team_scope=alias_key[0]) + for unit in _LIMIT_UNITS + for limit in _build_group_limits(aliased_group, unit) ) ) } @@ -447,6 +461,20 @@ def _scope_suffix(deployment_scope: tuple[str, ...] | None) -> str: return "chain" if deployment_scope is None else "dep:" + "+".join(deployment_scope) +def _hash_tag(model_group: str, configured: _ConfiguredLimit, tag_value: str, key_hash: str | None) -> str: + scope: Final = _scope_suffix(configured.deployment_scope) + # team_scope disambiguates two teams that publish the identical + # team_public_model_name alias with identically-configured limits -- + # without it their buckets would collide despite the index correctly + # scoping the lookup by (team_id, alias). See _ConfiguredLimit.team_scope. + team_suffix: Final = f":team:{configured.team_scope}" if configured.team_scope is not None else "" + key_suffix: Final = f":key:{key_hash}" if key_hash is not None else "" + return ( + f"tag_rl:{model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:" + f"{scope}{team_suffix}:{tag_value}{key_suffix}" + ) + + def _bucket_key( model_group: str, configured: _ConfiguredLimit, @@ -454,13 +482,7 @@ def _bucket_key( bucket_id: int, key_hash: str | None = None, ) -> str: - 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}" + return f"{{{_hash_tag(model_group, configured, tag_value, key_hash)}}}:{bucket_id}" def _inflight_key( @@ -472,13 +494,7 @@ 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: 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" + return f"{{{_hash_tag(model_group, configured, tag_value, key_hash)}}}:inflight" class _ClassifiedCheck(NamedTuple): @@ -494,6 +510,7 @@ def _classify_check( tags: Sequence[str], present_deployment_ids: frozenset[str], request_kwargs: Mapping[str, object], + metadata_variable_name: str, now: float, ) -> _ClassifiedCheck | None: if configured_limit.deployment_scope is not None and not ( @@ -503,7 +520,9 @@ def _classify_check( 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 + key_hash: Final = ( + _extract_key_hash(request_kwargs, metadata_variable_name) 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) @@ -667,11 +686,12 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer return healthy_deployments 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)) + metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(resolved_request_kwargs) + team_id: Final = _extract_team_id(resolved_request_kwargs, metadata_variable_name) + configured: Final = self._index.get(self.llm_router).resolve(model, team_id) if not configured: return healthy_deployments - 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 ) @@ -686,7 +706,13 @@ class _PROXY_TagRateLimiter( # pyright: ignore[reportUnusedClass] # only refer for configured_limit in configured if ( check := _classify_check( - configured_limit, model, tags, present_deployment_ids, resolved_request_kwargs, now + configured_limit, + model, + tags, + present_deployment_ids, + resolved_request_kwargs, + metadata_variable_name, + now, ) ) is not 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 65ed8c5bb92..9344da1bf72 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_tag_rate_limiter.py @@ -13,10 +13,14 @@ 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, + _bucket_key, _build_group_limits, _build_limits_index, _ConfiguredLimit, _extract_identity, + _extract_key_hash, + _extract_team_id, + _inflight_key, _pending_concurrency_holder, _PROXY_TagRateLimiter, ) @@ -72,6 +76,38 @@ def test_extract_identity_skips_negation_tags(): assert _extract_identity(["!end_user_id:u1"], "end_user_id") is None +# --------------------------------------------------------------------------- +# _extract_key_hash / _extract_team_id -- must read only the one field the +# server actually authenticates into, never fall back to the other +# --------------------------------------------------------------------------- + + +def test_extract_key_hash_ignores_a_forged_value_in_the_non_authoritative_field(): + """ + On a route where litellm_metadata is authoritative, the server writes + the real hash there and never touches metadata -- so a caller-supplied + metadata.user_api_key must not be read at all, let alone win. + """ + request_kwargs = { + "metadata": {"user_api_key": "forged-by-caller"}, + "litellm_metadata": {"user_api_key": "real-authenticated-hash"}, + } + assert _extract_key_hash(request_kwargs, "litellm_metadata") == "real-authenticated-hash" + + +def test_extract_key_hash_reads_metadata_when_it_is_the_authoritative_field(): + request_kwargs = {"metadata": {"user_api_key": "real-hash"}} + assert _extract_key_hash(request_kwargs, "metadata") == "real-hash" + + +def test_extract_team_id_ignores_a_forged_value_in_the_non_authoritative_field(): + request_kwargs = { + "metadata": {"user_api_key_team_id": "forged-team"}, + "litellm_metadata": {"user_api_key_team_id": "real-team"}, + } + assert _extract_team_id(request_kwargs, "litellm_metadata") == "real-team" + + # --------------------------------------------------------------------------- # TagRateLimitEntry -- period_seconds validation # --------------------------------------------------------------------------- @@ -1225,8 +1261,15 @@ def test_build_limits_index_is_also_keyed_by_team_public_model_name(): deployment["model_info"]["team_id"] = "team-1" deployment["model_info"]["team_public_model_name"] = "team-alias-name" index = _build_limits_index([deployment]) - assert index.resolve("real-model-name", team_id=None) == index.resolve("team-alias-name", team_id="team-1") - assert index.resolve("real-model-name", team_id=None) != [] + by_name = index.resolve("real-model-name", team_id=None) + by_alias = index.resolve("team-alias-name", team_id="team-1") + assert by_name != () + assert [c.entry for c in by_name] == [c.entry for c in by_alias] + # The alias resolution must carry the team_id into the bucket scope -- + # see test_build_limits_index_keeps_different_teams_same_alias_separate + # for why (two teams can publish the identical alias string). + assert by_name[0].team_scope is None + assert by_alias[0].team_scope == "team-1" def test_build_limits_index_keeps_different_teams_same_alias_separate(): @@ -1256,6 +1299,39 @@ def test_build_limits_index_keeps_different_teams_same_alias_separate(): assert resolved_b[0].entry.limit == 999 +def test_bucket_key_differs_across_teams_sharing_an_alias_and_identical_limit_config(): + """ + Two teams that happen to publish the identical team_public_model_name + AND configure an identically-named, identically-valued limit must not + land on the same Redis bucket -- team_public_model_name is only unique + per team, so this is a realistic collision, not a contrived one. + """ + team_a = _deployment( + "model-a", "dep-a", {"request_limits": {"limits": [{"name": "per_minute", "limit": 5, "period_seconds": 60}]}} + ) + 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": "per_minute", "limit": 5, "period_seconds": 60}]}} + ) + team_b["model_info"]["team_id"] = "team-b" + team_b["model_info"]["team_public_model_name"] = "shared-alias" + + index = _build_limits_index([team_a, team_b]) + limit_a = index.resolve("shared-alias", team_id="team-a")[0] + limit_b = index.resolve("shared-alias", team_id="team-b")[0] + assert limit_a.entry == limit_b.entry # identical configuration, by construction + + key_a = _bucket_key("shared-alias", limit_a, tag_value="same-caller-tag", bucket_id=0) + key_b = _bucket_key("shared-alias", limit_b, tag_value="same-caller-tag", bucket_id=0) + assert key_a != key_b + + inflight_a = _inflight_key("shared-alias", limit_a, tag_value="same-caller-tag") + inflight_b = _inflight_key("shared-alias", limit_b, tag_value="same-caller-tag") + assert inflight_a != inflight_b + + def test_build_limits_index_merges_alias_limits_across_different_model_names(): """ litellm auto-generates each team-added deployment's own internal