diff --git a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index 8cfef5b31bd..05e696bd33c 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -660,6 +660,29 @@ def _fixed_length_identity(tag_value: str) -> str: return hashlib.sha256(tag_value.encode()).hexdigest() +def _policy_fingerprint(entry: TagRateLimitEntry) -> str: + """ + Two entries can share a `name` and `tag_id` while genuinely disagreeing + on `limit`, `period_seconds`, or any of the four scoping fields -- + `_DedupSignature`/`resolve_any` already treat that as two distinct + policies (see `distinct_signature_count_by_name` in `_build_group_limits`), + so the Redis/in-memory bucket key must too, or two differently-configured + entries that happen to share a name check and charge the identical + counter. Hashed to a fixed-length digest for the same reason + `_fixed_length_identity` hashes `tag_value`: an operator's own + `included_values`/`excluded_values` list has no length bound. + """ + fingerprint_source: Final = ( + entry.limit, + entry.period_seconds, + entry.included_values, + entry.excluded_values, + _scope_signature(entry.enabled_for), + _scope_signature(entry.disabled_for), + ) + return hashlib.sha256(repr(fingerprint_source).encode()).hexdigest()[:16] + + 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 @@ -675,9 +698,13 @@ def _hash_tag(model_group: str, configured: _ConfiguredLimit, tag_value: str, ke # 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 "" + # Two entries can share `name`/`tag_id` while disagreeing on limit, + # period_seconds, or scoping (see _policy_fingerprint) -- included so + # they never collide onto the same counter despite the shared name. + policy_suffix: Final = f":policy:{_policy_fingerprint(configured.entry)}" return ( f"tag_rl:{effective_model_group}:{configured.unit}:{configured.entry.name}:{configured.entry.tag_id}:" - f"{scope}{team_suffix}:{_fixed_length_identity(tag_value)}{key_suffix}" + f"{scope}{team_suffix}:{_fixed_length_identity(tag_value)}{key_suffix}{policy_suffix}" ) diff --git a/litellm/types/router.py b/litellm/types/router.py index ab42a7966ec..5ad1f09adf7 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -227,6 +227,11 @@ class TagRateLimitEntry(BaseModel): # defeats the entry; reject it at config load time instead. if math.isnan(self.limit): raise ValueError("limit must not be NaN") + if math.isinf(self.limit): + raise ValueError( + "limit must be finite -- positive infinity makes admission never reject (current + increment " + "> limit is always false), negative infinity makes it always reject every tagged request" + ) return self @model_validator(mode="after") diff --git a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index f2e23fc9b00..7053ab6f9b8 100644 --- a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -104,16 +104,38 @@ def _expected_bucket_key( team_scope: str | None = None, resolved_group: str | None = None, key_hash: str | None = None, + limit: float = 1, + included_values: tuple | None = None, + excluded_values: tuple | None = None, + enabled_for: dict | None = None, + disabled_for: dict | 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. + + `limit` and the four scoping fields default to values that produce a + stable fingerprint for tests that don't care about it, but must be + passed matching the real entry's own configuration whenever a test's + router declares a `limit` other than 1 (or any scoping) for the entry + whose key this reproduces -- see _policy_fingerprint, which folds them + into the key precisely so two differently-configured entries sharing a + name never collide onto the same counter. """ configured = _ConfiguredLimit( unit=unit, - entry=TagRateLimitEntry(name=name, tag_id=tag_id, limit=1, period_seconds=period_seconds), + entry=TagRateLimitEntry( + name=name, + tag_id=tag_id, + limit=limit, + period_seconds=period_seconds, + included_values=included_values, + excluded_values=excluded_values, + enabled_for=enabled_for, + disabled_for=disabled_for, + ), deployment_scope=deployment_scope, team_scope=team_scope, resolved_group=resolved_group, @@ -252,6 +274,20 @@ def test_tag_rate_limit_entry_rejects_nan_limit(): TagRateLimitEntry(name="n", limit=float("nan"), period_seconds=60) +def test_tag_rate_limit_entry_rejects_infinite_limit(): + """ + Positive infinity makes the atomic requests/concurrency + current + increment > limit check always false, so admission never + rejects; negative infinity makes it always true, rejecting every tagged + request. Same silent-misconfiguration class as NaN, just via a different + non-finite float rather than a non-ordering one. + """ + with pytest.raises(ValidationError, match="limit must be finite"): + TagRateLimitEntry(name="n", limit=float("inf"), period_seconds=60) + with pytest.raises(ValidationError, match="limit must be finite"): + TagRateLimitEntry(name="n", limit=float("-inf"), period_seconds=60) + + # --------------------------------------------------------------------------- # TagRateLimitEntry -- period_seconds validation # --------------------------------------------------------------------------- @@ -483,6 +519,48 @@ def test_tag_rate_limit_scope_normalizes_values_order_and_duplicates(): assert scope.values == ("1001", "1032") +# --------------------------------------------------------------------------- +# _hash_tag / _bucket_key -- policy identity folds into the Redis key itself +# --------------------------------------------------------------------------- + + +def test_bucket_key_differs_for_same_named_entries_with_different_limits(): + """ + A plain, unscoped entry and a stricter, scoped override can legitimately + share a `name` (the worked example in the docs uses distinct names, but + nothing in validation requires that) -- resolve_any/_build_group_limits + already treat differing limit/scoping as genuinely distinct policies for + dedup purposes, so the actual counter key must too, or two + differently-configured entries that happen to share a name check and + charge the identical Redis/in-memory bucket. + """ + now = 0.0 + default_key = _expected_bucket_key("grp", "requests", "daily", "end_user_id", "u1", 86400, now, limit=2500) + override_key = _expected_bucket_key( + "grp", + "requests", + "daily", + "end_user_id", + "u1", + 86400, + now, + limit=1, + enabled_for={"tag_id": "company_id", "values": ["1032"]}, + ) + assert default_key != override_key + + +def test_bucket_key_differs_for_same_named_entries_with_different_scoping_only(): + now = 0.0 + excluding_u1 = _expected_bucket_key( + "grp", "requests", "daily", "end_user_id", "u2", 86400, now, limit=100, excluded_values=("u1",) + ) + excluding_u2 = _expected_bucket_key( + "grp", "requests", "daily", "end_user_id", "u2", 86400, now, limit=100, excluded_values=("u2",) + ) + assert excluding_u1 != excluding_u2 + + # --------------------------------------------------------------------------- # _build_group_limits -- scoping fields fold into the dedup signature # --------------------------------------------------------------------------- @@ -1248,8 +1326,8 @@ async def test_log_success_event_increments_configured_units(time_controller): await asyncio.sleep(0) now = time_controller.now().timestamp() - 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) + token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=500000) + dollar_key = _expected_bucket_key("grp", "dollars", "monthly", "end_user_id", "u1", 2592000, now, limit=50.0) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 @@ -1260,7 +1338,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 = _expected_bucket_key("grp", "requests", "daily", "end_user_id", "u1", 86400, now) + request_key = _expected_bucket_key("grp", "requests", "daily", "end_user_id", "u1", 86400, now, limit=100) assert await limiter.internal_usage_cache.async_get_cache(key=request_key, litellm_parent_otel_span=None) is None @@ -1305,7 +1383,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 = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now) + token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=500000) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 ) @@ -1351,7 +1429,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 = _expected_bucket_key("backend-a", "tokens", "daily", "end_user_id", "u1", 86400, now) + token_key = _expected_bucket_key("backend-a", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=500000) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 ) @@ -1418,7 +1496,7 @@ async def test_log_success_event_accounts_against_the_same_bucket_admission_chec now = time_controller.now().timestamp() token_key = _expected_bucket_key( - "my-group", "tokens", "daily", "end_user_id", "u1", 86400, now, resolved_group=admission_bucket_group + "my-group", "tokens", "daily", "end_user_id", "u1", 86400, now, resolved_group=admission_bucket_group, limit=500000 ) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 @@ -1461,7 +1539,7 @@ async def test_admission_dedups_against_the_full_group_not_just_currently_health # since backend-a is the only one excluded below) stays empty. now = time_controller.now().timestamp() over_limit_key = _expected_bucket_key( - "my-group", "tokens", "daily", "end_user_id", "u1", 86400, now, resolved_group="backend-a" + "my-group", "tokens", "daily", "end_user_id", "u1", 86400, now, resolved_group="backend-a", limit=10 ) await limiter.internal_usage_cache.async_set_cache(key=over_limit_key, value=20.0, litellm_parent_otel_span=None) @@ -1521,8 +1599,12 @@ async def test_log_success_event_accounts_against_the_key_hash_admission_checked await asyncio.sleep(0) now = time_controller.now().timestamp() - keyed_bucket = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash="keyA") - unkeyed_bucket = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash=None) + keyed_bucket = _expected_bucket_key( + "grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash="keyA", limit=500000 + ) + unkeyed_bucket = _expected_bucket_key( + "grp", "tokens", "daily", "end_user_id", "u1", 86400, now, key_hash=None, limit=500000 + ) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=keyed_bucket, litellm_parent_otel_span=None)) == 42.0 @@ -1569,9 +1651,11 @@ async def test_log_success_event_charges_the_window_admission_checked_not_a_late await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0) await asyncio.sleep(0) - admitted_window_bucket = _expected_bucket_key("grp", "tokens", "per_minute", "end_user_id", "u1", 60, admission_time) + admitted_window_bucket = _expected_bucket_key( + "grp", "tokens", "per_minute", "end_user_id", "u1", 60, admission_time, limit=500 + ) later_window_bucket = _expected_bucket_key( - "grp", "tokens", "per_minute", "end_user_id", "u1", 60, time_controller.now().timestamp() + "grp", "tokens", "per_minute", "end_user_id", "u1", 60, time_controller.now().timestamp(), limit=500 ) assert ( float( @@ -1636,9 +1720,11 @@ async def test_log_success_event_accounts_against_the_team_id_admission_checked( now = time_controller.now().timestamp() correct_bucket = _expected_bucket_key( - "team-alias-name", "tokens", "daily", "end_user_id", "u1", 86400, now, team_scope="team-1" + "team-alias-name", "tokens", "daily", "end_user_id", "u1", 86400, now, team_scope="team-1", limit=500 + ) + wrong_bucket = _expected_bucket_key( + "team-alias-name", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=500 ) - wrong_bucket = _expected_bucket_key("team-alias-name", "tokens", "daily", "end_user_id", "u1", 86400, now) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=correct_bucket, litellm_parent_otel_span=None)) == 42.0 @@ -1719,7 +1805,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 = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "u1", 60, now) + request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "u1", 60, now, limit=10) 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 @@ -2028,7 +2114,7 @@ async def test_success_event_token_accounting_is_wired_through_the_background_re assert len(_BACKGROUND_TASKS) == 0 now = time_controller.now().timestamp() - token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now) + token_key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=500000) assert ( float(await limiter.internal_usage_cache.async_get_cache(key=token_key, litellm_parent_otel_span=None)) == 42.0 ) @@ -2460,7 +2546,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 = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now) + key = _expected_bucket_key("grp", "tokens", "daily", "end_user_id", "u1", 86400, now, limit=1000) 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: @@ -2494,7 +2580,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 = _expected_bucket_key("grp", "dollars", "monthly", "team_id", "t1", 2592000, now) + key = _expected_bucket_key("grp", "dollars", "monthly", "team_id", "t1", 2592000, now, limit=50.0) 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: @@ -2619,7 +2705,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 = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", tag, 60, now) + request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", tag, 60, now, limit=10) 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 @@ -3159,7 +3245,7 @@ async def test_cross_unit_refund_leaves_no_phantom_increment_in_memory(time_cont ) now = time_controller.now().timestamp() - request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "refund-check", 60, now) + request_key = _expected_bucket_key("grp", "requests", "per_minute", "end_user_id", "refund-check", 60, now, limit=10) 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