fix(rate-limiting): stop pre-call fallback retries from double-charging

veria-ai finding on PR #36541: ProxyBaseLLMRequestProcessing._pre_call_with_fallbacks
reruns the whole pre-call pipeline (this hook included) once per fallback
model on any ProxyRateLimitError, not only one this hook itself raised, but
keeps the same litellm_call_id across every attempt. Without this fix, a
call admitted once by this hook and then rejected by a different, later
check in the same pass would get charged again on every fallback retry for
what is still one logical client request, letting a caller consume one
shared tag-quota unit per fallback attempt.

A "requests" or "concurrency" check whose key was already charged/reserved
for this call_id now renews at zero net cost as part of the same
all-or-nothing atomic batch, mirroring model_based_tag_rate_limits_hook's
identical fix for its own per-hop retries.

Renewal requires the repeat admission to carry the same authenticated
key_hash as whichever admission first claimed this call_id's stash:
litellm_call_id is caller-controlled via the x-litellm-call-id header (the
same forgery vector that hook's pending-reservations mirror was hardened
against earlier in this PR), so two unrelated requests choosing an
identical call_id must not be able to renew each other's charge.

Not live-verified: reliably reproducing this specific path needs a second,
different check to reject after this hook has already admitted in the same
pre-call pass, which depends on callback registration order this session
didn't have a clean way to control from proxy config alone. Verified
instead with regression tests against the real hook class covering both the
legitimate renewal case and the forged-call-id case, plus four existing
tests updated to use distinct call_ids for what they model as separate
logical requests now that repeat-call_id admission has real behavior tied
to it.
This commit is contained in:
Deepanshu 2026-08-25 20:28:03 -04:00
parent 334de9031b
commit 68405241a0
2 changed files with 181 additions and 14 deletions

View file

@ -189,6 +189,24 @@ class _GlobalTagRateLimitStash:
# decided on -- not whatever model a later fallback actually served.
model: str | None = None
pending_concurrency_keys: list[tuple[str, _PartitionKey]] = field(default_factory=list) # mutable-ok: queue
# "requests" keys already charged for this call_id -- veria-ai finding:
# ProxyBaseLLMRequestProcessing._pre_call_with_fallbacks reruns the whole
# pre-call pipeline (this hook included) once per fallback model on ANY
# ProxyRateLimitError, not only one this hook itself raised, but reuses
# the same litellm_call_id (self.data is mutated in place, only `model`
# changes) across every attempt -- so this stash is the SAME object each
# time. A "requests" check matching an already-charged key here renews
# at zero net cost instead of charging a second unit for the same
# logical request; see async_pre_call_hook's own comment for how.
charged_request_keys: list[str] = field(default_factory=list) # mutable-ok: see comment above
# The server-authenticated key_hash (UserAPIKeyAuth.api_key) of whichever
# call first claimed this stash. litellm_call_id is caller-controlled via
# the x-litellm-call-id header (the exact forgery vector
# model_based_tag_rate_limits_hook's own pending-reservations mirror was
# hardened against earlier), so two unrelated requests sharing a
# caller-chosen id must not be allowed to "renew" each other's charge --
# only a later admission carrying this same, authenticated key_hash may.
owner_key_hash: str | None = None
# Sentinel key for a call with no litellm_call_id at all (claim and lookup
@ -509,12 +527,12 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
if config is None:
return data
# Unlike model_based_tag_rate_limits_hook's async_filter_deployments
# (called once per routing hop, so a still-queued reservation can
# legitimately belong to an earlier, already-failed hop of the same
# request), async_pre_call_hook fires exactly once per request --
# there is no "prior hop" case here, so no stale-reservation release
# is needed at the top of admission.
# async_pre_call_hook fires once per request in the common case, but
# ProxyBaseLLMRequestProcessing._pre_call_with_fallbacks can re-run
# this same pipeline once per fallback model on any ProxyRateLimitError
# (not only one this hook raised) -- see charged_request_keys' own
# docstring for how a repeat run for the same call_id renews rather
# than re-charges both "requests" and "concurrency" checks below.
stash: Final = _claim_stash_for_data(data)
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(data)
@ -523,6 +541,14 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
key_hash: Final = user_api_key_dict.api_key
model: Final = data.get("model") if isinstance(data.get("model"), str) else None
# Only a repeat admission carrying the SAME authenticated key_hash as
# whichever call first claimed this stash may renew its charges --
# see owner_key_hash's own docstring for why a bare call_id match is
# not enough. First admission for this stash claims ownership here.
if stash.owner_key_hash is None:
stash.owner_key_hash = key_hash
renewal_allowed: Final = stash.owner_key_hash == key_hash
now: Final = self._time_provider().timestamp()
stash.admission_time = now
stash.model = model
@ -543,13 +569,32 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
await self._partition_for(_partition_key(check.entry))
) # mutable-ok: see comment above
atomic_partitions: Final = tuple(atomic_partitions_list)
already_reserved_concurrency_keys: Final = frozenset(
key for key, _partition_key in stash.pending_concurrency_keys
)
failing_index, values = await self._atomic_check_and_increment(
tuple(
(
partition.internal_usage_cache,
check.key,
check.entry.limit,
1.0,
# A key already charged/reserved for this call_id (an
# earlier _pre_call_with_fallbacks attempt for the
# same logical request) renews at zero net cost
# instead of charging or reserving a second unit --
# folded into this same all-or-nothing batch so a
# rollback here (some other check in the batch
# rejecting) refunds that zero-cost renewal as a
# genuine no-op, same reasoning as
# model_based_tag_rate_limits_hook's identical fix
# for its own per-hop retries.
0.0
if renewal_allowed
and (
(check.unit == "requests" and check.key in stash.charged_request_keys)
or (check.unit == "concurrency" and check.key in already_reserved_concurrency_keys)
)
else 1.0,
self._ttl_for(check.unit, check.entry),
)
for partition, check in zip(atomic_partitions, atomic_checks)
@ -561,12 +606,37 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
failing_check.unit, failing_check.entry, failing_check.tag_value, model, current=values[0]
)
# Only genuinely new reservations, never a key already in
# already_reserved_concurrency_keys: that key's own check just
# renewed at zero net cost above, so re-adding it here would
# make release (which decrements once per queued entry) decrement
# twice for a counter that was only ever incremented once.
concurrency_reservations: Final = tuple(
(check.key, _partition_key(check.entry)) for check in atomic_checks if check.unit == "concurrency"
(check.key, _partition_key(check.entry))
for check in atomic_checks
if check.unit == "concurrency" and check.key not in already_reserved_concurrency_keys
)
if concurrency_reservations:
stash.pending_concurrency_keys.extend(concurrency_reservations) # mutable-ok: see field's own docstring
# Only recorded when renewal_allowed: an admission that didn't
# own this stash (a call_id collision from a different key_hash)
# must not contaminate the rightful owner's own renewal
# tracking, or a later, genuine fallback retry from the owner
# could wrongly treat the impostor's charge as its own and
# renew for free.
request_keys: Final = (
tuple(
check.key
for check in atomic_checks
if check.unit == "requests" and check.key not in stash.charged_request_keys
)
if renewal_allowed
else ()
)
if request_keys:
stash.charged_request_keys.extend(request_keys) # mutable-ok: see field's own docstring
return data
async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None:

View file

@ -98,13 +98,15 @@ async def test_request_limit_shared_across_keys_by_default(time_controller, monk
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="key-a"), cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion"
)
# A different key, identical tag value: must be rejected too -- proves
# the bucket is genuinely shared, not per-key by default.
# A different key, identical tag value, and a distinct call_id (a
# genuinely separate logical request, not a fallback retry of the same
# one) -- must be rejected too, proving the bucket is genuinely shared,
# not per-key by default.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="key-b"),
cache=DualCache(),
data=_data(["end_user_id:u1"]),
data=_data(["end_user_id:u1"], call_id="call-2"),
call_type="completion",
)
@ -125,7 +127,9 @@ async def test_request_limit_is_independent_of_model(time_controller, monkeypatc
hook = _make_hook(time_controller)
data_model_a = {**_data(["end_user_id:u1"]), "model": "gpt-4o"}
data_model_b = {**_data(["end_user_id:u1"]), "model": "claude-3"}
# A distinct call_id: this is a separate logical request, not the same
# one retrying against a different model via _pre_call_with_fallbacks.
data_model_b = {**_data(["end_user_id:u1"], call_id="call-2"), "model": "claude-3"}
await hook.async_pre_call_hook(
user_api_key_dict=_key(), cache=DualCache(), data=data_model_a, call_type="completion"
)
@ -198,11 +202,13 @@ async def test_apply_to_key_alias_enforces_for_the_listed_key(time_controller, m
data=_data(["end_user_id:u1"]),
call_type="completion",
)
# A distinct call_id: a second, separate request from the same key, not
# a fallback retry of the first.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="premium-key"),
cache=DualCache(),
data=_data(["end_user_id:u1"]),
data=_data(["end_user_id:u1"], call_id="call-2"),
call_type="completion",
)
@ -237,11 +243,13 @@ async def test_apply_to_key_alias_composes_with_scope_by_key_hash(time_controlle
data=_data(["end_user_id:u1"]),
call_type="completion",
)
# A distinct call_id: a second, separate request from the same key, not
# a fallback retry of the first.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="key-a", api_key="hashA"),
cache=DualCache(),
data=_data(["end_user_id:u1"]),
data=_data(["end_user_id:u1"], call_id="call-2"),
call_type="completion",
)
# key-b is unaffected by key-a's exhausted bucket.
@ -509,6 +517,95 @@ async def test_apply_to_models_fallback_does_not_re_narrow_accounting_to_the_ser
)
# ---------------------------------------------------------------------------
# _pre_call_with_fallbacks reruns admission for the same logical request:
# a repeat call_id must renew, not double-charge -- veria-ai finding on
# PR #36541
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_repeat_admission_for_the_same_call_id_and_key_renews_instead_of_double_charging(
time_controller, monkeypatch
):
"""
ProxyBaseLLMRequestProcessing._pre_call_with_fallbacks reruns the whole
pre-call pipeline (this hook included) once per fallback model on ANY
ProxyRateLimitError, not only one this hook itself raised, but keeps the
same litellm_call_id across every attempt (self.data is mutated in
place; only "model" changes). Without this fix, an unrelated rejection
(a different rate limiter, a budget cap) triggering N fallback attempts
would charge this hook's own "requests" cap N times for one logical
client call. A limit of 1 makes a double-charge directly observable: if
the second admission (same call_id, same key) charged again instead of
renewing, this would raise.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"request_limits": {
"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]
}
},
)
hook = _make_hook(time_controller)
key = _key(alias="key-a")
await hook.async_pre_call_hook(
user_api_key_dict=key, cache=DualCache(), data=_data(["end_user_id:u1"]), call_type="completion"
)
# Same call_id, same key, different model -- exactly what
# _pre_call_with_fallbacks produces for a fallback attempt of the same
# logical request.
result = await hook.async_pre_call_hook(
user_api_key_dict=key,
cache=DualCache(),
data={**_data(["end_user_id:u1"]), "model": "fallback-model"},
call_type="completion",
)
assert result is not None
@pytest.mark.asyncio
async def test_a_forged_shared_call_id_from_a_different_key_does_not_get_a_free_renewal(time_controller, monkeypatch):
"""
Security regression: litellm_call_id is caller-controlled via the
x-litellm-call-id header (the same forgery vector
model_based_tag_rate_limits_hook's own pending-reservations mirror was
hardened against earlier in this PR). Two unrelated requests choosing
the identical call_id must not be able to renew each other's charge --
only a second admission carrying the SAME authenticated key_hash as
whichever request first claimed that call_id may. A limit of 1 makes
this observable: if the second, different-key admission wrongly
renewed, it would succeed instead of raising.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"request_limits": {
"limits": [{"name": "daily", "tag_id": "end_user_id", "limit": 1, "period_seconds": 86400}]
}
},
)
hook = _make_hook(time_controller)
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="key-a", api_key="hashA"),
cache=DualCache(),
data=_data(["end_user_id:u1"], call_id="forged-call-id"),
call_type="completion",
)
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(alias="key-b", api_key="hashB"),
cache=DualCache(),
data=_data(["end_user_id:u1"], call_id="forged-call-id"),
call_type="completion",
)
# ---------------------------------------------------------------------------
# Concurrency: reservation at admission, release on success/failure/disconnect
# ---------------------------------------------------------------------------