mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): fix key-hash/alias extraction at log time, trim comments, add Lua script tests
Greptile/Bugbot findings from the first review round: - extract_key_hash/extract_key_alias only checked the top-level metadata bucket, unlike order_tags_for_identity_resolution's _active_metadata_bucket helper. By async_log_success_event/async_log_failure_event time, metadata is only nested under litellm_params, so both silently returned None there, diverging from the value read at admission time. Fixed by routing both through _active_metadata_bucket, with regression tests for the nested case. - Condensed the module's docstrings/comments to one concise line each, keeping the non-obvious behavior they document, matching the trimming already applied to litellm/types/router.py on the stacked PR #39900. - Added real Redis execution coverage for both Lua scripts (admission, rejection, ttl refresh semantics, decrement floor) against a throwaway local redis-server, since fakeredis has no EVALSHA support without the optional lupa dependency.
This commit is contained in:
parent
74c2dd3e22
commit
38ec960064
2 changed files with 209 additions and 229 deletions
|
|
@ -1,20 +1,7 @@
|
|||
"""
|
||||
Primitives shared by both tag-scoped rate-limit hooks:
|
||||
`model_based_tag_rate_limits_hook.py` (per-deployment limits nested under
|
||||
`model_info.tag_rate_limits`, admitted once per routing hop) and
|
||||
`global_tag_rate_limits_hook.py` (model-independent limits declared once
|
||||
under `litellm_settings.global_tag_rate_limits`, admitted once per request).
|
||||
|
||||
Both hooks enforce the identical `TagRateLimitEntry` shape and need the same
|
||||
identity/scope extraction, policy fingerprinting, bucket-key hashing, and
|
||||
cache-partitioning primitives, so those live here rather than in either
|
||||
hook's own module -- keeping neither hook reaching into the other's private
|
||||
internals to reuse them. Every name here is this module's own public
|
||||
interface (no leading underscore): each hook imports what it needs aliased
|
||||
back to its own historical, underscore-prefixed local name (e.g.
|
||||
`entry_applies as _entry_applies`), so this is a genuine export, not a
|
||||
private symbol either hook reaches across a module boundary to grab.
|
||||
"""
|
||||
"""Primitives shared by both tag-scoped rate-limit hooks (model_based_tag_rate_limits_hook.py,
|
||||
global_tag_rate_limits_hook.py): identity/scope extraction, policy fingerprinting, bucket-key
|
||||
hashing, and cache-partitioning. Every name here is a genuine public export (no leading
|
||||
underscore); each hook imports what it needs aliased back to its own historical private name."""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
|
|
@ -29,11 +16,8 @@ from litellm.types.router import TagRateLimitEntry, TagRateLimitScope
|
|||
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.
|
||||
# requests/concurrency admit via atomic check-and-increment; tokens/dollars are only known
|
||||
# after the response, so they stay a read-then-account-on-success check with an unavoidable race.
|
||||
ATOMIC_UNITS: Final[frozenset[LimitUnit]] = frozenset({"requests", "concurrency"})
|
||||
|
||||
UNIT_TO_GROUP_FIELD: Final[Mapping[LimitUnit, str]] = MappingProxyType(
|
||||
|
|
@ -53,56 +37,22 @@ UNIT_TO_RATE_LIMIT_TYPE: Final[Mapping[LimitUnit, RateLimitType]] = MappingProxy
|
|||
}
|
||||
)
|
||||
|
||||
# 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.
|
||||
# shared read-only fallback for an absent/None mapping, so call sites don't build a fresh {}
|
||||
EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
# `asyncio.create_task`'s own docs: "Save a reference to the result of this
|
||||
# function, to avoid a task disappearing mid-execution. The event loop only
|
||||
# keeps weak references to tasks. A task that isn't referenced elsewhere may
|
||||
# get garbage collected at any time, even before it's done." The success path
|
||||
# deliberately fires-and-forgets its concurrency release and its token/dollar
|
||||
# accounting increment (unlike the failure/disconnect paths, which await
|
||||
# concurrency release directly) to keep the hot success-response path from
|
||||
# waiting on a Redis round trip; by the time either background task would
|
||||
# run, the state it needs (popped pending keys, or the request's own usage
|
||||
# figures) is only available in that task's own closure, so a collected
|
||||
# task's work is unrecoverable, not just delayed. Holding a strong reference
|
||||
# here until each task's own completion callback discards it is the standard
|
||||
# fix, shared by every fire-and-forget task either hook creates.
|
||||
# holds a strong reference to every fire-and-forget background task (concurrency release,
|
||||
# token/dollar accounting) so the event loop's weak-ref-only tracking can't garbage-collect
|
||||
# one mid-execution; each task's own completion callback discards its entry
|
||||
BACKGROUND_TASKS: Final[set["asyncio.Task[None]"]] = set() # mutable-ok: see comment above
|
||||
|
||||
# Floor for a concurrency reservation's self-heal TTL, regardless of the
|
||||
# configured period_seconds. A reservation that expires while its request is
|
||||
# still genuinely in flight silently admits requests past the limit; this
|
||||
# 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.
|
||||
# floor for a concurrency reservation's self-heal ttl, regardless of period_seconds, so an
|
||||
# expiring-while-in-flight reservation can't silently admit past the limit
|
||||
CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600
|
||||
|
||||
# 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
|
||||
# key carries its own self-contained {..} hash tag so unrelated buckets never
|
||||
# forcibly co-locate on the same Redis Cluster shard, which means a single Lua
|
||||
# invocation can never span more than one key's slot without risking a
|
||||
# cross-slot error. All-or-nothing across a hop's multiple atomic checks
|
||||
# (e.g. requests + concurrency checked together) is achieved in Python by
|
||||
# calling this once per key and refunding every earlier admission in the same
|
||||
# batch if a later one is rejected -- the same refund-on-rollback shape as
|
||||
# `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.
|
||||
#
|
||||
# refresh_ttl (ARGV[4]) distinguishes the two callers of this script:
|
||||
# "requests" is an epoch-bucketed fixed window, whose TTL must be set once
|
||||
# (at first write) and never extended, or the bucket outlives the epoch it's
|
||||
# meant to reset at. "concurrency" is not windowed at all -- its TTL exists
|
||||
# purely as a crash-safety net for a reservation whose explicit release never
|
||||
# runs -- so a still-active bucket must keep pushing that TTL out on every
|
||||
# admission, or a long-lived burst of continuous traffic expires the whole
|
||||
# counter mid-flight (silently admitting past the cap, and letting a release
|
||||
# for a since-reset counter decrement an unrelated, newer cohort).
|
||||
# single-key atomic check-and-increment (one key per call: every tag_rl key carries its own
|
||||
# {..} hash tag, so a multi-key call could span shards and cross-slot error). refresh_ttl
|
||||
# (ARGV[4]) distinguishes callers: "requests" is a fixed window whose ttl is set once and never
|
||||
# extended; "concurrency" isn't windowed, so its crash-safety ttl must refresh on every admission.
|
||||
TAG_RL_CHECK_AND_INCR_SCRIPT: Final = """
|
||||
local key = KEYS[1]
|
||||
local limit = tonumber(ARGV[1])
|
||||
|
|
@ -127,16 +77,8 @@ end
|
|||
return { 1, new_value }
|
||||
"""
|
||||
|
||||
# Atomic decrement that never leaves a counter negative. Used both to refund
|
||||
# an earlier admission when a later key in the same batch is rejected, and to
|
||||
# release a concurrency reservation -- floors at 0 so a decrement that can't
|
||||
# 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. Floors via DEL, not
|
||||
# `SET key 0`: releasing a reservation whose key already expired makes
|
||||
# INCRBY recreate it with no TTL, and a plain SET would leave that recreated
|
||||
# key permanently in Redis (SET clears any TTL); DEL removes it outright,
|
||||
# which reads back identically to 0 everywhere this key is read (`GET key or 0`).
|
||||
# atomic decrement floored at 0 (refund or concurrency release); floors via DEL rather than
|
||||
# `SET key 0`, since a release on an already-expired key would otherwise recreate it with no ttl
|
||||
TAG_RL_DECR_FLOOR_ZERO_SCRIPT: Final = """
|
||||
local key = KEYS[1]
|
||||
local delta = tonumber(ARGV[1])
|
||||
|
|
@ -148,28 +90,19 @@ end
|
|||
return new_value
|
||||
"""
|
||||
|
||||
# A (tag_id, values) pair mirroring TagRateLimitScope's own fields, used to
|
||||
# fold `enabled_for`/`disabled_for` into a hashable form (a policy
|
||||
# fingerprint, or a per-deployment dedup signature) without depending on
|
||||
# TagRateLimitScope's own hashability.
|
||||
# a (tag_id, values) pair mirroring TagRateLimitScope, used to fold enabled_for/disabled_for
|
||||
# into a hashable form without depending on TagRateLimitScope's own hashability
|
||||
ScopeSignature: TypeAlias = tuple[str, tuple[str, ...]] | None
|
||||
|
||||
|
||||
def scope_signature(scope: TagRateLimitScope | None) -> ScopeSignature:
|
||||
"""Normalizes a `TagRateLimitScope` into a plain, hashable tuple, so
|
||||
`enabled_for`/`disabled_for` can be folded into a policy fingerprint (or
|
||||
a per-deployment dedup signature) -- two entries disagreeing on either
|
||||
field must be treated as genuinely different policies rather than
|
||||
merged into one bucket."""
|
||||
"""Normalizes a TagRateLimitScope into a hashable tuple for policy fingerprinting/dedup."""
|
||||
return None if scope is None else (scope.tag_id, scope.values)
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
"""First tag matching `f"{tag_id}:"`, value after the colon. `!`-prefixed tag-routing
|
||||
negation markers are skipped so they're never misread as an identity value."""
|
||||
prefix: Final = f"{tag_id}:"
|
||||
for tag in tags:
|
||||
if tag.startswith("!"):
|
||||
|
|
@ -180,32 +113,10 @@ def extract_identity(tags: Sequence[str], tag_id: str) -> str | None:
|
|||
|
||||
|
||||
def entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str | None, model: str | None) -> bool:
|
||||
"""
|
||||
Applies `entry`'s own scoping fields (`enabled_for`/`disabled_for`/
|
||||
`apply_to_key_alias`/`apply_to_models`), evaluated in this order -- deny
|
||||
overrides allow, checked before any allowlist:
|
||||
|
||||
1. `disabled_for`: the gate tag (often a SECOND, independent tag, but
|
||||
`disabled_for.tag_id` can equally be set to this entry's own
|
||||
`tag_id` to gate on a subset of its own resolved identity) is
|
||||
present and its value is in `disabled_for.values` -> doesn't apply.
|
||||
Absent gate tag never triggers this -- nothing to match against a
|
||||
denylist.
|
||||
2. `enabled_for`: the gate tag is absent, or present but its value is
|
||||
NOT in `enabled_for.values` -> doesn't apply. Unlike `disabled_for`,
|
||||
absence DOES fail this check -- an allowlist gate requires an
|
||||
explicit match, so "not tagged at all" means "not in scope".
|
||||
3. `apply_to_models`: `model` is absent, or present but not in the
|
||||
list -> doesn't apply. Same allowlist semantics as `enabled_for` --
|
||||
a request with no `model` never satisfies this gate.
|
||||
4. `apply_to_key_alias`: the calling key's own alias is absent, or
|
||||
present but not in the list -> doesn't apply. Same allowlist
|
||||
semantics as `enabled_for` -- a key with no alias set never
|
||||
satisfies this gate.
|
||||
|
||||
An entry with none of these fields set always applies -- this is the
|
||||
unscoped behavior every existing entry has today, unchanged.
|
||||
"""
|
||||
"""Applies entry's own scoping fields in order (deny before allow): disabled_for excludes a
|
||||
matching value; enabled_for requires a matching value (absence fails, unlike disabled_for);
|
||||
apply_to_models and apply_to_key_alias are allowlists an absent model/alias never satisfies.
|
||||
An entry with none of these set always applies."""
|
||||
if entry.disabled_for is not None:
|
||||
disabled_gate_value: Final = extract_identity(tags, entry.disabled_for.tag_id)
|
||||
if disabled_gate_value is not None and disabled_gate_value in entry.disabled_for.values:
|
||||
|
|
@ -224,63 +135,18 @@ def entry_applies(entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str
|
|||
def resolve_authoritative_metadata_variable_name(
|
||||
metadata_source: Mapping[str, object],
|
||||
) -> Literal["metadata", "litellm_metadata"]:
|
||||
"""`get_metadata_variable_name_from_kwargs` only checks key presence, which
|
||||
misresolves both at admission time and at `async_log_success_event` time:
|
||||
a caller can forge an empty (or merely present-but-`None`) `litellm_metadata`
|
||||
on an ordinary request -- `kwargs["litellm_params"]` also always carries a
|
||||
`litellm_metadata` key (typically `None`) alongside the real, populated
|
||||
`metadata` dict for a standard (non LITELLM_METADATA_ROUTES) request -- and
|
||||
the key-presence check always picks `litellm_metadata` in both cases,
|
||||
silently reading no tags/identity at all and admitting the request against
|
||||
every configured limit.
|
||||
|
||||
Merely requiring the value to be a non-empty dict is not enough either: a
|
||||
caller can populate its own, unrelated keys on the non-authoritative
|
||||
bucket (e.g. `{"litellm_metadata": {"x": 1}}` on an ordinary route), which
|
||||
is non-empty but still not the field the proxy wrote identity into.
|
||||
`add_litellm_data_to_request` unconditionally stamps `user_api_key_auth`
|
||||
into whichever bucket the route actually resolved as authoritative, and
|
||||
strips any `user_api_key_`-prefixed key a caller pre-populates on the
|
||||
other bucket -- so requiring that marker's presence, not mere
|
||||
truthiness, can't be forged onto the wrong side."""
|
||||
"""Picks metadata vs litellm_metadata by the unforgeable `user_api_key_auth` marker
|
||||
`add_litellm_data_to_request` stamps into whichever bucket is actually authoritative --
|
||||
key presence or truthiness alone can be forged by a caller onto the wrong bucket."""
|
||||
litellm_metadata: Final = metadata_source.get("litellm_metadata")
|
||||
if isinstance(litellm_metadata, Mapping) and "user_api_key_auth" in litellm_metadata:
|
||||
return "litellm_metadata"
|
||||
return "metadata"
|
||||
|
||||
|
||||
def extract_key_hash(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None:
|
||||
"""Same single-authoritative-field lookup as
|
||||
`model_based_tag_rate_limits_hook._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") if isinstance(active, Mapping) else None
|
||||
return key_hash if isinstance(key_hash, str) else None
|
||||
|
||||
|
||||
def extract_key_alias(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None:
|
||||
"""Same single-authoritative-field lookup as
|
||||
`model_based_tag_rate_limits_hook._extract_team_id`, but for the calling
|
||||
virtual key's own `key_alias`: `LiteLLMProxyRequestSetup` sets
|
||||
`metadata["user_api_key_alias"]` to `user_api_key_dict.key_alias`
|
||||
(see `litellm_pre_call_utils.py`)."""
|
||||
active: Final = request_kwargs.get(metadata_variable_name) or EMPTY_MAPPING
|
||||
key_alias: Final = active.get("user_api_key_alias") if isinstance(active, Mapping) else None
|
||||
return key_alias if isinstance(key_alias, str) else None
|
||||
|
||||
|
||||
def _active_metadata_bucket(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> Mapping[str, object]:
|
||||
"""Same fallback `_get_tags_from_request_kwargs` (tag_based_routing.py)
|
||||
already relies on: `request_kwargs` is a flat, top-level-metadata dict at
|
||||
admission time, but `Logging.model_call_details` (what `kwargs` actually
|
||||
is by `async_log_success_event`/`async_log_failure_event` time) never
|
||||
carries `metadata`/`litellm_metadata` at its own top level, only nested
|
||||
under `request_kwargs["litellm_params"]`. Checking only the top level
|
||||
silently finds nothing at success/failure time, exactly the same failure
|
||||
mode that lookup already had to handle."""
|
||||
"""request_kwargs carries metadata at its own top level at admission time, but only nested
|
||||
under litellm_params by async_log_success_event/async_log_failure_event time; checks both."""
|
||||
top_level: Final = request_kwargs.get(metadata_variable_name)
|
||||
if isinstance(top_level, Mapping):
|
||||
return top_level
|
||||
|
|
@ -292,21 +158,28 @@ def _active_metadata_bucket(request_kwargs: Mapping[str, object], metadata_varia
|
|||
return EMPTY_MAPPING
|
||||
|
||||
|
||||
def extract_key_hash(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None:
|
||||
"""Reads the calling virtual key's hash from the authoritative metadata bucket (nested under
|
||||
litellm_params by log time, same as _active_metadata_bucket's other callers); `metadata["user_api_key"]`
|
||||
is already the hashed token despite the plain name."""
|
||||
active: Final = _active_metadata_bucket(request_kwargs, metadata_variable_name)
|
||||
key_hash: Final = active.get("user_api_key")
|
||||
return key_hash if isinstance(key_hash, str) else None
|
||||
|
||||
|
||||
def extract_key_alias(request_kwargs: Mapping[str, object], metadata_variable_name: str) -> str | None:
|
||||
"""Reads the calling virtual key's own key_alias from the authoritative metadata bucket."""
|
||||
active: Final = _active_metadata_bucket(request_kwargs, metadata_variable_name)
|
||||
key_alias: Final = active.get("user_api_key_alias")
|
||||
return key_alias if isinstance(key_alias, str) else None
|
||||
|
||||
|
||||
def order_tags_for_identity_resolution(
|
||||
tags: Sequence[str], request_kwargs: Mapping[str, object], metadata_variable_name: str
|
||||
) -> tuple[str, ...]:
|
||||
"""`extract_identity`/`entry_applies` both resolve a `tag_id` via
|
||||
first-match-by-prefix. `_merge_tags` (litellm_pre_call_utils.py) appends
|
||||
key/team/project tags only if not already present, keeping caller-supplied
|
||||
tags first in the merged `tags` list -- so an authenticated caller could
|
||||
submit e.g. `company_id:attacker-chosen` ahead of the calling key's real
|
||||
`company_id:real-company` tag and have every entry scoped to `company_id`
|
||||
resolve to the caller's own value instead of the key's. `metadata.inherited_tags`
|
||||
is a separate, server-computed snapshot of only the tags the calling
|
||||
key/team/project's own config contributed (see that field's docstring in
|
||||
litellm_pre_call_utils.py), so putting it first makes a policy-backed tag
|
||||
win over a same-prefix caller-supplied one.
|
||||
"""
|
||||
"""Puts server-computed `metadata.inherited_tags` ahead of caller-supplied tags, so a caller
|
||||
can't submit e.g. `company_id:attacker-chosen` and shadow the calling key's real, same-prefix
|
||||
tag; extract_identity/entry_applies both resolve tag_id via first-match-by-prefix."""
|
||||
active: Final = _active_metadata_bucket(request_kwargs, metadata_variable_name)
|
||||
inherited_tags: Final = active.get("inherited_tags")
|
||||
if not isinstance(inherited_tags, (list, tuple)) or not inherited_tags:
|
||||
|
|
@ -315,37 +188,16 @@ def order_tags_for_identity_resolution(
|
|||
|
||||
|
||||
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 a 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 a 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).
|
||||
"""
|
||||
"""Hashes a caller-controlled, unbounded-length tag value to a fixed-length digest, bounding
|
||||
a hook's own contribution to an in-memory cache key or Redis key regardless of input size."""
|
||||
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`, `scope_by_key_hash`, or any of the scoping
|
||||
fields -- both hooks already treat that as two distinct policies
|
||||
elsewhere (model_based_tag_rate_limits_hook.py's own per-deployment
|
||||
dedup signature is one example), 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. `scope_by_key_hash` specifically
|
||||
needs its own slot here rather than relying on a hook's own `_hash_tag`
|
||||
`key_hash`-derived suffix to carry it: that suffix is empty whenever
|
||||
`key_hash` resolves to `None` (no virtual key on the call), which would
|
||||
otherwise collide an unscoped entry with a key-hash-scoped one that
|
||||
agrees on every other field. Hashed to a fixed-length digest for the
|
||||
same reason `fixed_length_identity` hashes `tag_value`: an operator's
|
||||
own `enabled_for`/`disabled_for`/`apply_to_key_alias`/`apply_to_models`
|
||||
list has no length bound.
|
||||
"""
|
||||
"""Fixed-length digest folding every policy-distinguishing field of an entry (limit,
|
||||
period_seconds, scope_by_key_hash, enabled_for/disabled_for, apply_to_key_alias/models) into
|
||||
one bucket-key signature, so two differently-configured entries sharing a name never share
|
||||
a counter."""
|
||||
fingerprint_source: Final = (
|
||||
entry.limit,
|
||||
entry.period_seconds,
|
||||
|
|
@ -359,34 +211,16 @@ def policy_fingerprint(entry: TagRateLimitEntry) -> str:
|
|||
|
||||
|
||||
def bucket_ttl_seconds(entry: TagRateLimitEntry) -> int:
|
||||
"""Redis (and in-memory fallback) TTL for a non-concurrency bucket key.
|
||||
`entry.key_ttl_seconds` overrides the default of period_seconds + 3600
|
||||
when set -- see TagRateLimitEntry.key_ttl_seconds."""
|
||||
"""Redis (and in-memory fallback) ttl for a non-concurrency bucket key: entry.key_ttl_seconds
|
||||
overrides the default of period_seconds + 3600 when set."""
|
||||
return entry.key_ttl_seconds if entry.key_ttl_seconds is not None else entry.period_seconds + 3600
|
||||
|
||||
|
||||
# None => this entry shares the hook's single default cache partition
|
||||
# (matching every entry's behavior before this override existed). Otherwise
|
||||
# a value-stable signature -- not the override int alone -- so two different
|
||||
# entries that happen to choose the identical max_in_memory_cache_size don't
|
||||
# get merged into one shared partition; the same entry (same config content)
|
||||
# always resolves to the same signature across config/index rebuilds, which
|
||||
# is what keeps each hook's own `_partitions` cache from leaking a fresh
|
||||
# partition every time its configuration is re-resolved.
|
||||
# The str before the trailing int is `policy_fingerprint(entry)`: two
|
||||
# entries can share tag_id/name while genuinely disagreeing on
|
||||
# limit/period_seconds/scope_by_key_hash/enabled_for/disabled_for/
|
||||
# apply_to_key_alias/apply_to_models -- both hooks already treat that as two
|
||||
# distinct policies for bucket-key purposes, so a shared
|
||||
# max_in_memory_cache_size must not route them onto the same partition
|
||||
# either, or one entry's high-cardinality traffic can evict the other's
|
||||
# active counters from a cache neither entry asked to share.
|
||||
# max_in_memory_cache_size stays the trailing element: `partition_key[-1]`
|
||||
# reads it directly to size the partition's cache.
|
||||
# None => this entry shares the hook's single default cache partition. Otherwise a value-stable
|
||||
# signature (policy_fingerprint(entry), not the override int alone) so two entries sharing a
|
||||
# max_in_memory_cache_size don't merge into one partition; max_in_memory_cache_size stays the
|
||||
# trailing element since `partition_key[-1]` reads it directly to size the partition's cache.
|
||||
PartitionKey: TypeAlias = tuple[str, str, str, int] | None
|
||||
# Grouping type for async_log_success_event's per-partition tokens/dollars
|
||||
# pipeline dispatch -- named only so the declaration fits on one line; see
|
||||
# that method for why the grouping is needed.
|
||||
PartitionOperations: TypeAlias = dict[PartitionKey, list[RedisPipelineIncrementOperation]]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,17 @@ hook's own test file covers everything specific to how it wires these
|
|||
primitives into its own admission/accounting engine.
|
||||
"""
|
||||
|
||||
import shutil
|
||||
import socket
|
||||
import subprocess
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.hooks.tag_rate_limits_shared import (
|
||||
TAG_RL_CHECK_AND_INCR_SCRIPT,
|
||||
TAG_RL_DECR_FLOOR_ZERO_SCRIPT,
|
||||
bucket_ttl_seconds,
|
||||
entry_applies,
|
||||
extract_identity,
|
||||
|
|
@ -139,6 +149,21 @@ def test_extract_key_alias_ignores_a_non_mapping_authoritative_field():
|
|||
assert extract_key_alias(request_kwargs, "metadata") is None
|
||||
|
||||
|
||||
def test_extract_key_hash_finds_metadata_nested_under_litellm_params_at_log_time():
|
||||
"""Bugbot finding: by async_log_success_event/async_log_failure_event time,
|
||||
kwargs only carries metadata nested under litellm_params (see
|
||||
_active_metadata_bucket's own docstring), not at request_kwargs' own top
|
||||
level -- extract_key_hash must find it there too, or a key-hash-scoped
|
||||
bucket's accounting silently reads a different key than admission did."""
|
||||
request_kwargs = {"litellm_params": {"metadata": {"user_api_key": "real-hash"}}}
|
||||
assert extract_key_hash(request_kwargs, "metadata") == "real-hash"
|
||||
|
||||
|
||||
def test_extract_key_alias_finds_metadata_nested_under_litellm_params_at_log_time():
|
||||
request_kwargs = {"litellm_params": {"metadata": {"user_api_key_alias": "real-alias"}}}
|
||||
assert extract_key_alias(request_kwargs, "metadata") == "real-alias"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# resolve_authoritative_metadata_variable_name -- Veria AI finding: a
|
||||
# caller-supplied, non-empty litellm_metadata must not be selected over
|
||||
|
|
@ -429,3 +454,124 @@ def test_bucket_ttl_seconds_defaults_to_period_plus_one_hour_when_unset():
|
|||
def test_bucket_ttl_seconds_honors_key_ttl_seconds_override():
|
||||
entry = TagRateLimitEntry(name="per_minute", tag_id="end_user_id", limit=1, period_seconds=60, key_ttl_seconds=120)
|
||||
assert bucket_ttl_seconds(entry) == 120
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TAG_RL_CHECK_AND_INCR_SCRIPT / TAG_RL_DECR_FLOOR_ZERO_SCRIPT -- fakeredis has
|
||||
# no EVALSHA support without the optional lupa dependency, so these run
|
||||
# against a throwaway local redis-server instead (same idiom as
|
||||
# test_batch_enqueued_tokens.py's test_redis_lua_path_full_lifecycle).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||
sock.bind(("127.0.0.1", 0))
|
||||
return sock.getsockname()[1]
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def redis_port() -> Iterator[int]:
|
||||
if shutil.which("redis-server") is None:
|
||||
pytest.skip("requires a local redis-server binary to exercise the Lua script path")
|
||||
port = _free_port()
|
||||
proc = subprocess.Popen(
|
||||
["redis-server", "--port", str(port), "--save", "", "--appendonly", "no"],
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
try:
|
||||
deadline = time.monotonic() + 5
|
||||
while time.monotonic() < deadline:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||
sock.settimeout(0.2)
|
||||
if sock.connect_ex(("127.0.0.1", port)) == 0:
|
||||
break
|
||||
else:
|
||||
proc.terminate()
|
||||
pytest.skip("local redis-server did not become ready in time")
|
||||
yield port
|
||||
finally:
|
||||
proc.terminate()
|
||||
proc.wait(timeout=5)
|
||||
|
||||
|
||||
def _register(port: int, script: str):
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
return RedisCache(host="127.0.0.1", port=port).async_register_script(script)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_incr_admits_under_limit_and_sets_ttl(redis_port):
|
||||
run = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
admitted, new_value = await run(keys=["bucket1"], args=[5, 1, 60, 0])
|
||||
assert (admitted, new_value) == (1, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_incr_rejects_over_limit_without_incrementing(redis_port):
|
||||
run = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
await run(keys=["bucket2"], args=[1, 1, 60, 0])
|
||||
rejected, current = await run(keys=["bucket2"], args=[1, 1, 60, 0])
|
||||
assert (rejected, current) == (0, 1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_incr_requests_ttl_is_set_once_and_not_refreshed(redis_port):
|
||||
"""refresh_ttl=0 (the `requests` fixed-window semantics): the epoch-bucketed
|
||||
TTL must be set on the first write and left alone after, or the bucket
|
||||
outlives the epoch it's meant to reset at."""
|
||||
run = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
await run(keys=["bucket3"], args=[5, 1, 100, 0])
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
raw = RedisCache(host="127.0.0.1", port=redis_port)
|
||||
client = raw.init_async_client()
|
||||
first_ttl = await client.ttl("bucket3")
|
||||
await run(keys=["bucket3"], args=[5, 1, 5, 0])
|
||||
second_ttl = await client.ttl("bucket3")
|
||||
assert first_ttl > 5
|
||||
assert second_ttl > 5 # unchanged by the second call's much shorter ttl arg
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_and_incr_concurrency_ttl_refreshes_on_every_admission(redis_port):
|
||||
"""refresh_ttl=1 (the `concurrency` reservation semantics): every admission
|
||||
must push the crash-safety TTL back out, or a long-lived burst of traffic
|
||||
expires the whole counter mid-flight."""
|
||||
run = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
await run(keys=["bucket4"], args=[5, 1, 5, 1])
|
||||
await run(keys=["bucket4"], args=[5, 1, 100, 1])
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
raw = RedisCache(host="127.0.0.1", port=redis_port)
|
||||
client = raw.init_async_client()
|
||||
ttl = await client.ttl("bucket4")
|
||||
assert ttl > 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decr_floor_zero_floors_at_zero_and_deletes_the_key(redis_port):
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
run_incr = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
run_decr = _register(redis_port, TAG_RL_DECR_FLOOR_ZERO_SCRIPT)
|
||||
await run_incr(keys=["bucket5"], args=[5, 1, 60, 1])
|
||||
|
||||
floored = await run_decr(keys=["bucket5"], args=[-5])
|
||||
assert floored == 0
|
||||
|
||||
raw = RedisCache(host="127.0.0.1", port=redis_port)
|
||||
client = raw.init_async_client()
|
||||
assert await client.get("bucket5") is None # floored via DEL, not a TTL-less `SET 0`
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decr_floor_zero_decrements_normally_when_result_stays_non_negative(redis_port):
|
||||
run_incr = _register(redis_port, TAG_RL_CHECK_AND_INCR_SCRIPT)
|
||||
run_decr = _register(redis_port, TAG_RL_DECR_FLOOR_ZERO_SCRIPT)
|
||||
await run_incr(keys=["bucket6"], args=[5, 3, 60, 1])
|
||||
|
||||
remaining = await run_decr(keys=["bucket6"], args=[-1])
|
||||
assert remaining == 2
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue