From 38ec9600644696760d3561277982e2872439e499 Mon Sep 17 00:00:00 2001 From: Deepanshu Date: Sat, 5 Sep 2026 09:05:16 -0400 Subject: [PATCH] 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. --- litellm/proxy/hooks/tag_rate_limits_shared.py | 292 ++++-------------- .../hooks/test_tag_rate_limits_shared.py | 146 +++++++++ 2 files changed, 209 insertions(+), 229 deletions(-) diff --git a/litellm/proxy/hooks/tag_rate_limits_shared.py b/litellm/proxy/hooks/tag_rate_limits_shared.py index a21523c7885..75ae67dd120 100644 --- a/litellm/proxy/hooks/tag_rate_limits_shared.py +++ b/litellm/proxy/hooks/tag_rate_limits_shared.py @@ -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]] diff --git a/tests/test_litellm/proxy/hooks/test_tag_rate_limits_shared.py b/tests/test_litellm/proxy/hooks/test_tag_rate_limits_shared.py index 2e4954cd030..efe3b530e8d 100644 --- a/tests/test_litellm/proxy/hooks/test_tag_rate_limits_shared.py +++ b/tests/test_litellm/proxy/hooks/test_tag_rate_limits_shared.py @@ -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