mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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.
This commit is contained in:
parent
48b71a8a32
commit
170aa0604b
2 changed files with 140 additions and 38 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue