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:
Deepanshu 2026-09-05 09:05:16 -04:00
parent 74c2dd3e22
commit 38ec960064
2 changed files with 209 additions and 229 deletions

View file

@ -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]]

View file

@ -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