mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(rate-limiting): hash the caller-controlled tag identity before it enters a cache key
tag_value has no length or content bound before this hook embeds it directly into an in-memory dict key (bypassing max_in_memory_cache_size, which caps item count, not key bytes) and an uncapped Redis key. Hashing to a fixed-length digest bounds this hook's own contribution to key size regardless of the caller's input, while preserving distinctness. Existing tests that hand-wrote the raw tag value into an expected key string now build it through the real key-construction helper instead of hardcoding the (now-hashed) internal format.
This commit is contained in:
parent
c6f940312c
commit
d9eee566fa
2 changed files with 135 additions and 16 deletions
|
|
@ -2,6 +2,7 @@
|
|||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import hashlib
|
||||
from collections.abc import Callable, Iterable, Mapping, Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
from datetime import datetime
|
||||
|
|
@ -522,6 +523,20 @@ def _scope_suffix(deployment_scope: tuple[str, ...] | None) -> str:
|
|||
return "chain" if deployment_scope is None else "dep:" + "+".join(deployment_scope)
|
||||
|
||||
|
||||
def _fixed_length_identity(tag_value: str) -> str:
|
||||
"""
|
||||
`tag_value` is caller-controlled (whatever follows the tag_id prefix in
|
||||
a caller-supplied tag) with no length or content bound. Embedding it
|
||||
directly would let a caller inflate this hook's own in-memory dict keys
|
||||
past what `max_in_memory_cache_size` bounds (that caps item *count*, not
|
||||
key bytes) and grow unbounded Redis keys with no cap at all. Hashing to
|
||||
a fixed-length digest bounds this hook's own contribution to key size
|
||||
regardless of the caller's input, while still preserving distinctness
|
||||
(two different tag values still resolve to two different buckets).
|
||||
"""
|
||||
return hashlib.sha256(tag_value.encode()).hexdigest()
|
||||
|
||||
|
||||
def _hash_tag(model_group: str, configured: _ConfiguredLimit, tag_value: str, key_hash: str | None) -> str:
|
||||
# resolved_group overrides the caller-visible model_group when this
|
||||
# limit was found via resolve_any()'s per-deployment fallback (routing
|
||||
|
|
@ -539,7 +554,7 @@ def _hash_tag(model_group: str, configured: _ConfiguredLimit, tag_value: str, ke
|
|||
key_suffix: Final = f":key:{key_hash}" if key_hash is not None else ""
|
||||
return (
|
||||
f"tag_rl:{effective_model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:"
|
||||
f"{scope}{team_suffix}:{tag_value}{key_suffix}"
|
||||
f"{scope}{team_suffix}:{_fixed_length_identity(tag_value)}{key_suffix}"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ from litellm.proxy.hooks.tag_rate_limiter import (
|
|||
_extract_identity,
|
||||
_extract_key_hash,
|
||||
_extract_team_id,
|
||||
_fixed_length_identity,
|
||||
_inflight_key,
|
||||
_partition_key,
|
||||
_pending_concurrency_holder,
|
||||
|
|
@ -61,6 +62,36 @@ def _deployment(model_name: str, deployment_id: str, tag_rate_limits: dict) -> d
|
|||
}
|
||||
|
||||
|
||||
def _expected_bucket_key(
|
||||
model_group: str,
|
||||
unit: str,
|
||||
name: str,
|
||||
tag_id: str,
|
||||
tag_value: str,
|
||||
period_seconds: int,
|
||||
now: float,
|
||||
deployment_scope: tuple | None = None,
|
||||
team_scope: str | None = None,
|
||||
resolved_group: str | None = None,
|
||||
key_hash: str | None = None,
|
||||
) -> str:
|
||||
"""
|
||||
Builds the exact key the real code would compute (via _hash_tag's
|
||||
fixed-length hashing of tag_value), instead of hand-writing the raw
|
||||
tag value into a literal string -- the internal key format (hashed or
|
||||
not) is an implementation detail these tests shouldn't hardcode.
|
||||
"""
|
||||
configured = _ConfiguredLimit(
|
||||
unit=unit,
|
||||
entry=TagRateLimitEntry(name=name, tag_id=tag_id, limit=1, period_seconds=period_seconds),
|
||||
deployment_scope=deployment_scope,
|
||||
team_scope=team_scope,
|
||||
resolved_group=resolved_group,
|
||||
)
|
||||
bucket_id = int(now) // period_seconds
|
||||
return _bucket_key(model_group, configured, tag_value, bucket_id, key_hash=key_hash)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _extract_identity
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
@ -79,6 +110,69 @@ def test_extract_identity_skips_negation_tags():
|
|||
assert _extract_identity(["!end_user_id:u1"], "end_user_id") is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _fixed_length_identity -- tag_value is caller-controlled with no length
|
||||
# bound; this hook's own contribution to a cache key must not grow with it
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_fixed_length_identity_bounds_key_contribution_regardless_of_input_size():
|
||||
"""
|
||||
A caller can submit an arbitrarily long tag value (no length or content
|
||||
bound is enforced upstream of this hook). Without hashing, that value
|
||||
would go straight into an in-memory dict key (bypassing
|
||||
max_in_memory_cache_size, which caps item *count* not key bytes) and an
|
||||
unbounded-length Redis key (Redis has no key-count or key-size cap at
|
||||
all here). A fixed-length digest bounds this hook's own contribution to
|
||||
the key regardless of input size.
|
||||
"""
|
||||
huge_value = "x" * 5_000_000
|
||||
digest = _fixed_length_identity(huge_value)
|
||||
assert len(digest) == 64 # sha256 hex digest length, independent of input size
|
||||
|
||||
|
||||
def test_fixed_length_identity_preserves_distinctness():
|
||||
"""Hashing must not collapse two different tag values onto one bucket."""
|
||||
assert _fixed_length_identity("user-a") != _fixed_length_identity("user-b")
|
||||
assert _fixed_length_identity("user-a") == _fixed_length_identity("user-a")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_an_oversized_tag_value_does_not_inflate_the_bucket_key(time_controller):
|
||||
"""
|
||||
End-to-end: a request tagged with a multi-megabyte end_user_id value
|
||||
must still resolve to a short, fixed-length bucket key, not one whose
|
||||
size scales with the caller's input.
|
||||
"""
|
||||
limiter = _make_limiter(time_controller)
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
_deployment(
|
||||
"grp",
|
||||
"dep-1",
|
||||
{
|
||||
"request_limits": {
|
||||
"limits": [{"name": "per_minute", "tag_id": "end_user_id", "limit": 1000, "period_seconds": 60}]
|
||||
}
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
limiter.update_variables(llm_router=router)
|
||||
healthy = router.model_list
|
||||
huge_tag_value = "y" * 2_000_000
|
||||
|
||||
await limiter.async_filter_deployments(
|
||||
model="grp",
|
||||
healthy_deployments=healthy,
|
||||
messages=None,
|
||||
request_kwargs={"metadata": {"tags": [f"end_user_id:{huge_tag_value}"]}},
|
||||
)
|
||||
|
||||
(only_key,) = limiter.internal_usage_cache.dual_cache.in_memory_cache.cache_dict.keys()
|
||||
assert len(only_key) < 200
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _extract_key_hash / _extract_team_id -- must read only the one field the
|
||||
# server actually authenticates into, never fall back to the other
|
||||
|
|
@ -514,8 +608,9 @@ async def test_load_balanced_group_per_deployment_breach_rejects_whole_hop(time_
|
|||
request_kwargs = {"metadata": {"tags": ["end_user_id:u1"]}}
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
bucket_id = int(now) // 86400
|
||||
dep1_key = f"{{tag_rl:grp:requests:daily:end_user_id:dep:dep-1:u1}}:{bucket_id}"
|
||||
dep1_key = _expected_bucket_key(
|
||||
"grp", "requests", "daily", "end_user_id", "u1", 86400, now, deployment_scope=("dep-1",)
|
||||
)
|
||||
await limiter.internal_usage_cache.async_set_cache(key=dep1_key, value=1, ttl=86400, litellm_parent_otel_span=None)
|
||||
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
|
|
@ -569,8 +664,8 @@ async def test_log_success_event_increments_configured_units(time_controller):
|
|||
await asyncio.sleep(0)
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
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}"
|
||||
token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now)
|
||||
dollar_key = _expected_bucket_key("grp", "dollars", "monthly", "end_user_id", "u1", 2592000, now)
|
||||
|
||||
assert (
|
||||
float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0
|
||||
|
|
@ -581,7 +676,7 @@ async def test_log_success_event_increments_configured_units(time_controller):
|
|||
|
||||
# "requests" is accounted atomically at admission (async_filter_deployments),
|
||||
# not here -- async_log_success_event must not touch its bucket at all.
|
||||
request_key = f"{{tag_rl:grp:requests:daily:end_user_id:chain:u1}}:{int(now) // 86400}"
|
||||
request_key = _expected_bucket_key("grp", "requests", "daily", "end_user_id", "u1", 86400, now)
|
||||
assert await limiter.internal_usage_cache.async_get_cache(key=request_key, litellm_parent_otel_span=None) is None
|
||||
|
||||
|
||||
|
|
@ -626,7 +721,7 @@ async def test_log_success_event_reads_nested_litellm_metadata_when_that_is_auth
|
|||
await asyncio.sleep(0)
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
token_key = f"{{tag_rl:grp:tokens:daily:end_user_id:chain:u1}}:{int(now) // 86400}"
|
||||
token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now)
|
||||
assert (
|
||||
float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0
|
||||
)
|
||||
|
|
@ -672,7 +767,7 @@ async def test_log_success_event_falls_back_to_serving_deployment_model_name_for
|
|||
await asyncio.sleep(0)
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
token_key = f"{{tag_rl:backend-a:tokens:daily:end_user_id:chain:u1}}:{int(now) // 86400}"
|
||||
token_key = _expected_bucket_key("backend-a", "tokens", "daily", "end_user_id", "u1", 86400, now)
|
||||
assert (
|
||||
float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0
|
||||
)
|
||||
|
|
@ -751,7 +846,7 @@ async def test_cross_unit_rejection_does_not_leave_a_phantom_increment(time_cont
|
|||
assert exc_info.value.detail["type"] == "concurrency"
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
request_key = f"{{tag_rl:grp:requests:per_minute:end_user_id:chain:u1}}:{int(now) // 60}"
|
||||
request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "u1", 60, now)
|
||||
requests_value = await limiter.internal_usage_cache.async_get_cache(key=request_key, litellm_parent_otel_span=None)
|
||||
assert (float(requests_value) if requests_value is not None else 0.0) == 1.0
|
||||
|
||||
|
|
@ -1275,7 +1370,7 @@ async def test_token_limit_rejects_once_bucket_is_seeded_at_limit(time_controlle
|
|||
healthy = router.model_list
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
key = f"{{tag_rl:grp:tokens:daily:end_user_id:chain:u1}}:{int(now) // 86400}"
|
||||
key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now)
|
||||
await limiter.internal_usage_cache.async_set_cache(key=key, value=1000, ttl=86400, litellm_parent_otel_span=None)
|
||||
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
|
|
@ -1309,7 +1404,7 @@ async def test_dollar_limit_rejects_once_bucket_is_seeded_at_limit(time_controll
|
|||
healthy = router.model_list
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
key = f"{{tag_rl:grp:dollars:monthly:team_id:chain:t1}}:{int(now) // 2592000}"
|
||||
key = _expected_bucket_key("grp", "dollars", "monthly", "team_id", "t1", 2592000, now)
|
||||
await limiter.internal_usage_cache.async_set_cache(key=key, value=50.0, ttl=2592000, litellm_parent_otel_span=None)
|
||||
|
||||
with pytest.raises(ProxyRateLimitError) as exc_info:
|
||||
|
|
@ -1434,7 +1529,7 @@ async def test_redis_backed_cross_unit_rejection_does_not_leave_a_phantom_increm
|
|||
assert exc_info.value.detail["type"] == "concurrency"
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
request_key = f"{{tag_rl:grp:requests:per_minute:end_user_id:chain:{tag}}}:{int(now) // 60}"
|
||||
request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", tag, 60, now)
|
||||
requests_value = await limiter.internal_usage_cache.async_get_cache(key=request_key, litellm_parent_otel_span=None)
|
||||
assert (float(requests_value) if requests_value is not None else 0.0) == 1.0
|
||||
|
||||
|
|
@ -1895,7 +1990,7 @@ async def test_cross_unit_refund_leaves_no_phantom_increment_in_memory(time_cont
|
|||
)
|
||||
|
||||
now = time_controller.now().timestamp()
|
||||
request_key = f"{{tag_rl:grp:requests:per_minute:end_user_id:chain:refund-check}}:{int(now) // 60}"
|
||||
request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "refund-check", 60, now)
|
||||
value = await limiter.internal_usage_cache.async_get_cache(key=request_key, litellm_parent_otel_span=None)
|
||||
assert (float(value) if value is not None else 0.0) == 1.0
|
||||
|
||||
|
|
@ -2485,7 +2580,11 @@ def _request_limit(period_seconds: int, key_ttl_seconds: int | None = None) -> _
|
|||
return _ConfiguredLimit(
|
||||
unit="requests",
|
||||
entry=TagRateLimitEntry(
|
||||
name="per_minute", tag_id="end_user_id", limit=1, period_seconds=period_seconds, key_ttl_seconds=key_ttl_seconds
|
||||
name="per_minute",
|
||||
tag_id="end_user_id",
|
||||
limit=1,
|
||||
period_seconds=period_seconds,
|
||||
key_ttl_seconds=key_ttl_seconds,
|
||||
),
|
||||
deployment_scope=None,
|
||||
)
|
||||
|
|
@ -2511,7 +2610,10 @@ def test_bucket_ttl_seconds_honors_key_ttl_seconds_override():
|
|||
|
||||
def test_ttl_for_concurrency_honors_key_ttl_seconds_above_the_safety_floor():
|
||||
above_floor: Final = _CONCURRENCY_MIN_SAFETY_TTL_SECONDS + 100
|
||||
assert _PROXY_TagRateLimiter._ttl_for(_concurrency_limit(period_seconds=60, key_ttl_seconds=above_floor)) == above_floor
|
||||
assert (
|
||||
_PROXY_TagRateLimiter._ttl_for(_concurrency_limit(period_seconds=60, key_ttl_seconds=above_floor))
|
||||
== above_floor
|
||||
)
|
||||
|
||||
|
||||
def test_ttl_for_concurrency_never_drops_below_the_safety_floor_even_with_a_lower_override():
|
||||
|
|
@ -2721,7 +2823,9 @@ async def test_concurrency_slot_with_a_cache_size_override_is_released_against_t
|
|||
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 reservation for the same tag is rejected.
|
||||
with pytest.raises(ProxyRateLimitError):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue