fix(rate-limiting): close identity-shadowing, tag disclosure, and cross-attempt accounting gaps in the global tag hook (ported from #38347)

- Wire _order_tags_for_identity_resolution into this hook's own admission
  and success-event tag resolution, matching the sibling model-based hook.
  Without it a caller could put a forged tag ahead of the policy-backed
  inherited one and dodge or mis-bucket every global limit.
- Drop tag_value from the client-facing rejection detail: it can resolve
  from inherited_tags (server-assigned key/team/project metadata), and
  echoing it back would disclose that identity to the rejected caller.
- Track every model an admission attempt saw for a call_id (admitted_models,
  a frozenset) instead of overwriting a single field, only recording once an
  attempt clears every check. Otherwise an apply_to_models-scoped entry that
  matched an earlier _pre_call_with_fallbacks attempt lost its success-time
  accounting once a later attempt admitted with a different model, and a
  rejected attempt's model could wrongly drive later accounting.
- Add async_post_call_failure_hook, releasing a reservation from an earlier
  admission attempt when _pre_call_with_fallbacks exhausts every fallback
  and re-raises without ever running the real LLM call -- the only other
  release paths (success/failure/disconnect) are tied to that call, which
  never happens.
This commit is contained in:
Deepanshu 2026-08-27 07:23:10 -04:00
parent ec529f4b94
commit f24d362d8c
2 changed files with 439 additions and 31 deletions

View file

@ -99,6 +99,7 @@ from litellm.proxy.hooks.model_based_tag_rate_limits_hook import (
_extract_key_hash, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_fixed_length_identity, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_LimitUnit, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_order_tags_for_identity_resolution, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_partition_key, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_PartitionKey, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
_PartitionOperations, # pyright: ignore[reportPrivateUsage] # reused across module boundaries, see module docstring
@ -124,6 +125,20 @@ else:
Span: TypeAlias = object
def _entry_applies_any_admitted_model(
entry: TagRateLimitEntry, tags: Sequence[str], key_alias: str | None, admitted_models: frozenset[str]
) -> bool:
"""Same as `_entry_applies`, except an `apply_to_models`-scoped entry
counts as applying if ANY model an admission attempt for this call_id
saw was in scope -- not just whichever model the call ultimately served.
A `_pre_call_with_fallbacks` retry re-admits with a different model for
the same call_id, and an entry that matched an earlier attempt must
still get its success-time accounting."""
if not admitted_models:
return _entry_applies(entry, tags, key_alias, None)
return any(_entry_applies(entry, tags, key_alias, model) for model in admitted_models)
def _hash_tag(entry: TagRateLimitEntry, unit: _LimitUnit, tag_value: str, key_hash: str | None) -> str:
"""
Global-hook equivalent of `model_based_tag_rate_limits_hook._hash_tag`,
@ -183,11 +198,13 @@ class _GlobalTagRateLimitStash:
"""
admission_time: float | None = None
# The caller-facing `model` admission read from `data.get("model")`, so
# async_log_success_event's tokens/dollars accounting gates
# apply_to_models against the same, originally-requested model admission
# decided on -- not whatever model a later fallback actually served.
model: str | None = None
# Every model any admission attempt for this call_id has classified
# entries against, accumulated rather than overwritten: a
# _pre_call_with_fallbacks retry re-runs admission with a *different*
# model for the same call_id, and an apply_to_models entry that matched
# an earlier attempt must still get its accounting at success time even
# though the request ultimately serves from a later attempt's model.
admitted_models: frozenset[str] = field(default_factory=frozenset)
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
@ -400,6 +417,16 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
"global_tag_rate_limits_hook: failed to release concurrency slot %s: %s", key, e
)
@staticmethod
def _record_admitted_model(stash: "_GlobalTagRateLimitStash", model: str | None, renewal_allowed: bool) -> None:
"""Only called once this admission attempt has cleared every check
without raising -- a rejected attempt's model must never join
admitted_models, or a later successful attempt's accounting could
wrongly credit an apply_to_models entry that never actually admitted
this request under that model."""
if renewal_allowed and model is not None:
stash.admitted_models = stash.admitted_models | frozenset((model,))
@staticmethod
def _ttl_for(unit: _LimitUnit, entry: TagRateLimitEntry) -> int:
if unit == "concurrency":
@ -500,7 +527,11 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
"error": "tag_rate_limit_exceeded",
"type": unit,
"tag_id": entry.tag_id,
"tag_value": tag_value,
# tag_value deliberately excluded: it can resolve from
# inherited_tags (server-assigned key/team/project metadata),
# and echoing it back would disclose that identity to the
# caller. verbose_proxy_logger.debug above still logs it
# server-side for observability.
"limit_name": entry.name,
"limit": entry.limit,
"period_seconds": entry.period_seconds,
@ -536,7 +567,11 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
stash: Final = _claim_stash_for_data(data)
metadata_variable_name: Final = get_metadata_variable_name_from_kwargs(data)
tags: Final = _get_tags_from_request_kwargs(data, metadata_variable_name=metadata_variable_name)
tags: Final = _order_tags_for_identity_resolution(
_get_tags_from_request_kwargs(data, metadata_variable_name=metadata_variable_name),
data,
metadata_variable_name,
)
key_alias: Final = user_api_key_dict.key_alias
key_hash: Final = user_api_key_dict.api_key
model: Final = data.get("model") if isinstance(data.get("model"), str) else None
@ -551,9 +586,9 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
now: Final = self._time_provider().timestamp()
stash.admission_time = now
stash.model = model
classified: Final = self._classify(config, tags, key_alias, key_hash, now, model)
if not classified:
self._record_admitted_model(stash, model, renewal_allowed)
return data
read_only_checks: Final = tuple(c for c in classified if not c.is_atomic)
@ -637,39 +672,52 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
if request_keys:
stash.charged_request_keys.extend(request_keys) # mutable-ok: see field's own docstring
self._record_admitted_model(stash, model, renewal_allowed)
return data
async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None:
stash: Final = _stash_for_call(_call_id_from_kwargs(request_data))
async def _release_pending_for_call_id(self, request_kwargs: Mapping[str, object]) -> None:
stash: Final = _stash_for_call(_call_id_from_kwargs(request_kwargs))
if stash is None or not stash.pending_concurrency_keys:
return
release_keys: Final = tuple(stash.pending_concurrency_keys)
stash.pending_concurrency_keys.clear()
await self._release_keys(release_keys)
async def async_release_disconnect_state_hook(self, request_data: Mapping[str, object]) -> None:
await self._release_pending_for_call_id(request_data)
async def async_post_call_failure_hook(
self,
request_data: dict, # mutable-ok: must match CustomLogger.async_post_call_failure_hook's own base signature exactly
original_exception: Exception,
user_api_key_dict: UserAPIKeyAuth,
traceback_str: str | None = None,
) -> None:
"""
A request that never reaches Router (every fallback model also
rejected, or none configured) never runs the actual LLM call, so
neither async_log_success_event nor async_log_failure_event -- both
tied to that call's own wrapper -- ever fires for it. This is the
only remaining release path for a reservation from an earlier,
successful admission attempt in the same _pre_call_with_fallbacks
chain. litellm_call_id survives proxy/utils.py's own stripping here
(only litellm_logging_obj is popped), so the same ContextVar-based
stash lookup as the other release hooks still works.
"""
await self._release_pending_for_call_id(request_data)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None:
# No special-case skip for this hook's own tag_rate_limit_exceeded
# rejection: that rejection never reaches the point where a
# concurrency reservation is queued (see async_pre_call_hook), so
# stash.pending_concurrency_keys is already empty in that case and
# the check below naturally no-ops. Skipping release based on the
# exception's error marker alone would be wrong here, since
# model_based_tag_rate_limits_hook raises the identical marker --
# that rejection can land after this hook already reserved a slot
# for this same request, and that slot must still be released.
stash: Final = _stash_for_call(_call_id_from_kwargs(kwargs))
if stash is None or not stash.pending_concurrency_keys:
return
release_keys: Final = tuple(stash.pending_concurrency_keys)
stash.pending_concurrency_keys.clear()
await self._release_keys(release_keys)
# Always release regardless of which hook raised: this hook's own
# rejection never reserves a slot, so pending_concurrency_keys is
# already empty in that case and the check below no-ops; a rejection
# from model_based_tag_rate_limits_hook (same error marker) can still
# land after this hook already reserved its own slot.
await self._release_pending_for_call_id(kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None:
stash: Final = _stash_for_call(_call_id_from_kwargs(kwargs))
if stash is not None and stash.pending_concurrency_keys:
release_keys: Final = tuple(stash.pending_concurrency_keys)
stash.pending_concurrency_keys.clear()
release_task: Final = asyncio.create_task(self._release_keys(release_keys))
release_task: Final = asyncio.create_task(self._release_pending_for_call_id(kwargs))
_BACKGROUND_TASKS.add(release_task) # mutable-ok: see _BACKGROUND_TASKS's own docstring
release_task.add_done_callback(_BACKGROUND_TASKS.discard)
@ -690,7 +738,11 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
key_hash: Final = _extract_key_hash(litellm_params_for_metadata, metadata_variable_name)
key_alias: Final = _extract_key_alias(litellm_params_for_metadata, metadata_variable_name)
tags: Final = _get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name)
tags: Final = _order_tags_for_identity_resolution(
_get_tags_from_request_kwargs(kwargs, metadata_variable_name=metadata_variable_name),
kwargs,
metadata_variable_name,
)
if not tags:
return
@ -699,7 +751,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
if stash is not None and stash.admission_time is not None
else self._time_provider().timestamp()
)
model: Final = stash.model if stash is not None else None
admitted_models: Final = stash.admitted_models if stash is not None else frozenset()
increment_by_unit: Final[Mapping[_LimitUnit, float]] = MappingProxyType(
{
"tokens": float(standard_logging_object.get("total_tokens") or 0),
@ -716,7 +768,7 @@ class _PROXY_GlobalTagRateLimitsHook( # pyright: ignore[reportUnusedClass] # o
tag_value = _extract_identity(tags, entry.tag_id)
if tag_value is None:
continue
if not _entry_applies(entry, tags, key_alias, model):
if not _entry_applies_any_admitted_model(entry, tags, key_alias, admitted_models):
continue
increment_value = increment_by_unit[unit]
if increment_value == 0:

View file

@ -517,6 +517,174 @@ async def test_apply_to_models_fallback_does_not_re_narrow_accounting_to_the_ser
)
@pytest.mark.asyncio
async def test_apply_to_models_accounts_when_a_fallback_retry_re_admits_with_a_different_model(
time_controller, monkeypatch
):
"""
Unlike the test above (one admission, a later Router-level fallback
reported only at success time), this simulates _pre_call_with_fallbacks
itself re-running async_pre_call_hook for the SAME call_id with a
DIFFERENT model after some other hook rejected the original one. The
first admission (opus-chain, in apply_to_models scope) must still get
its success-time accounting even though the second admission
(sonnet-chain, out of scope) is the one that actually proceeds --
overwriting a single "last admitted model" field would silently drop
the spend for the in-scope entry.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"dollar_limits": {
"limits": [
{
"name": "chain_spend",
"tag_id": "end_user_id",
"limit": 10.0,
"period_seconds": 86400,
"apply_to_models": ["opus-chain"],
}
]
}
},
)
hook = _make_hook(time_controller)
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-1"), "model": "opus-chain"},
call_type="completion",
)
# _pre_call_with_fallbacks retry: same call_id, fallback model outside apply_to_models.
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-1"), "model": "sonnet-chain"},
call_type="completion",
)
kwargs = {
"litellm_call_id": "call-1",
"metadata": {"tags": ["end_user_id:u1"]},
"model": "sonnet-chain",
"standard_logging_object": {
"total_tokens": 0,
"response_cost": 12.0,
"model": "sonnet-chain",
"model_group": "sonnet-chain",
},
}
await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0)
await asyncio.sleep(0)
# The $12 spend landed in the opus-chain-scoped bucket, so a fresh
# opus-chain request is now over the $10 limit and rejected.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-2"), "model": "opus-chain"},
call_type="completion",
)
@pytest.mark.asyncio
async def test_a_rejected_admission_attempts_model_does_not_drive_later_accounting(time_controller, monkeypatch):
"""
A fallback retry's FIRST attempt can itself be rejected (by this same
entry, or a different hook) before it ever admits. That rejected
attempt's model must not join the stash's admitted-models history: a
later, successful attempt against a different (out-of-scope) model must
not have its accounting wrongly credited to an apply_to_models entry
that never actually admitted this request.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"concurrency_limits": {
"limits": [
{
"name": "conc-a",
"tag_id": "end_user_id",
"limit": 1,
"period_seconds": 60,
"apply_to_models": ["model-a"],
}
]
},
"dollar_limits": {
"limits": [
{
"name": "chain_spend",
"tag_id": "end_user_id",
"limit": 10.0,
"period_seconds": 86400,
"apply_to_models": ["model-a"],
}
]
},
},
)
hook = _make_hook(time_controller)
# Occupy model-a's only concurrency slot with an unrelated call.
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="occupier"), "model": "model-a"},
call_type="completion",
)
# call-1 attempt #1: model-a, rejected (slot taken) -- never truly admitted.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-1"), "model": "model-a"},
call_type="completion",
)
# call-1 attempt #2 (fallback retry): model-b, not in apply_to_models=[model-a], admits.
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-1"), "model": "model-b"},
call_type="completion",
)
kwargs = {
"litellm_call_id": "call-1",
"metadata": {"tags": ["end_user_id:u1"]},
"model": "model-b",
"standard_logging_object": {"total_tokens": 0, "response_cost": 50.0, "model": "model-b", "model_group": "model-b"},
}
await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0)
await asyncio.sleep(0)
# Release the occupier's concurrency slot so the final check below is
# gated only by chain_spend (dollars), not by conc-a still being full.
await hook.async_log_success_event(
kwargs={"litellm_call_id": "occupier", "metadata": {"tags": ["end_user_id:u1"]}},
response_obj=None,
start_time=0,
end_time=0,
)
await asyncio.sleep(0)
# chain_spend (apply_to_models=[model-a]) must still be empty: model-a's
# own admission attempt was rejected, never admitted, so a fresh
# model-a request is still allowed under the $10 limit.
result = await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-2"), "model": "model-a"},
call_type="completion",
)
assert result is not None
# ---------------------------------------------------------------------------
# _pre_call_with_fallbacks reruns admission for the same logical request:
# a repeat call_id must renew, not double-charge -- veria-ai finding on
@ -798,6 +966,54 @@ async def test_concurrency_reservation_released_on_disconnect(time_controller, m
assert result is not None
@pytest.mark.asyncio
async def test_concurrency_reservation_released_when_every_fallback_is_exhausted(time_controller, monkeypatch):
"""
When _pre_call_with_fallbacks exhausts every fallback model (another
hook rejects each one) and re-raises, the request never reaches the
real LLM call -- neither async_log_success_event nor
async_log_failure_event, both tied to that call's own wrapper, ever
fires. async_post_call_failure_hook is the only remaining release path
for a reservation this hook already admitted earlier in that same
fallback chain.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"concurrency_limits": {
"limits": [{"name": "conc", "tag_id": "end_user_id", "limit": 1, "period_seconds": 60}]
}
},
)
hook = _make_hook(time_controller)
# This hook admits and reserves the only slot.
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-1"), "model": "model-a"},
call_type="completion",
)
# _pre_call_with_fallbacks eventually gives up (every fallback rejected
# by some other hook) and reports the failure via post_call_failure_hook.
await hook.async_post_call_failure_hook(
request_data={"litellm_call_id": "call-1"},
original_exception=ProxyRateLimitError(
detail={"error": "some_other_hooks_limit"}, headers={}, rate_limit_type=None
),
user_api_key_dict=_key(),
)
result = await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={**_data(["end_user_id:u1"], call_id="call-2"), "model": "model-a"},
call_type="completion",
)
assert result is not None
@pytest.mark.asyncio
async def test_concurrent_requests_do_not_share_each_others_reservation_state(time_controller, monkeypatch):
"""Two logically distinct requests running as separate asyncio Tasks must
@ -1003,3 +1219,143 @@ async def test_config_reload_takes_effect_on_next_request(time_controller, monke
data=_data(["end_user_id:u1"], call_id="call-3"),
call_type="completion",
)
# ---------------------------------------------------------------------------
# Identity resolution: policy-backed tags must win over caller-supplied ones
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_a_caller_supplied_tag_cannot_shadow_the_policy_backed_identity_tag(time_controller, monkeypatch):
"""
Security regression: _merge_tags (litellm_pre_call_utils.py) keeps
caller-supplied tags first in the merged tags list, appending key/team/
project-contributed tags only if not already present. Since
_extract_identity/_entry_applies resolve a tag_id by
first-match-by-prefix, an authenticated caller could otherwise submit
e.g. company_id:attacker-chosen ahead of the key's real
company_id:real-company (surfaced via metadata.inherited_tags) and have
every company_id-scoped entry resolve to the forged value instead of the
real one -- letting the caller dodge the limit entirely by rotating
fabricated identities, or evade being charged against their own real
bucket. The hook must order tags so inherited_tags wins.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"request_limits": {
"limits": [{"name": "per-company", "tag_id": "company_id", "limit": 1, "period_seconds": 86400}]
}
},
)
hook = _make_hook(time_controller)
key = _key()
poisoned_data = {
"litellm_call_id": "attack-1",
"metadata": {
"tags": ["company_id:attacker-chosen", "company_id:real-company"],
"inherited_tags": ["company_id:real-company"],
},
}
await hook.async_pre_call_hook(user_api_key_dict=key, cache=DualCache(), data=poisoned_data, call_type="completion")
# The real company's own bucket must have been charged by the attack
# request, not a bucket keyed to the attacker's forged value -- so a
# second, genuine company_id:real-company request is now rejected.
victim_data = {
"litellm_call_id": "victim-1",
"metadata": {"tags": ["company_id:real-company"], "inherited_tags": ["company_id:real-company"]},
}
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(user_api_key_dict=key, cache=DualCache(), data=victim_data, call_type="completion")
@pytest.mark.asyncio
async def test_success_accounting_also_resolves_identity_from_the_policy_backed_tag(time_controller, monkeypatch):
"""Same forged-tag scenario as the admission-time test above, but for
async_log_success_event's own identity resolution (tokens/dollars)."""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"dollar_limits": {
"limits": [{"name": "per-company-spend", "tag_id": "company_id", "limit": 10.0, "period_seconds": 86400}]
}
},
)
hook = _make_hook(time_controller)
# kwargs at async_log_success_event time is Logging.model_call_details:
# metadata/inherited_tags live nested under kwargs["litellm_params"],
# never at the top level.
kwargs = {
"litellm_call_id": "attack-1",
"litellm_params": {
"metadata": {
"tags": ["company_id:attacker-chosen", "company_id:real-company"],
"inherited_tags": ["company_id:real-company"],
},
},
"standard_logging_object": {"total_tokens": 0, "response_cost": 20.0},
}
await hook.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0)
await asyncio.sleep(0)
# The $20 spend must have landed against company_id:real-company, so a
# fresh request under the genuine identity is now over the $10 limit.
with pytest.raises(ProxyRateLimitError):
await hook.async_pre_call_hook(
user_api_key_dict=_key(),
cache=DualCache(),
data={
"litellm_call_id": "victim-1",
"metadata": {"tags": ["company_id:real-company"], "inherited_tags": ["company_id:real-company"]},
},
call_type="completion",
)
@pytest.mark.asyncio
async def test_rejection_detail_does_not_disclose_the_resolved_tag_value(time_controller, monkeypatch):
"""
tag_value can resolve from inherited_tags (server-assigned key/team/
project metadata via _order_tags_for_identity_resolution), so echoing it
back in the client-facing 429 detail would disclose that identity to
the caller who triggered the rejection.
"""
monkeypatch.setattr(
litellm,
"global_tag_rate_limits",
{
"request_limits": {
"limits": [{"name": "per-company", "tag_id": "company_id", "limit": 1, "period_seconds": 86400}]
}
},
)
hook = _make_hook(time_controller)
key = _key()
data = {
"litellm_call_id": "call-1",
"metadata": {"tags": ["company_id:secret-internal-name"], "inherited_tags": ["company_id:secret-internal-name"]},
}
await hook.async_pre_call_hook(user_api_key_dict=key, cache=DualCache(), data=data, call_type="completion")
with pytest.raises(ProxyRateLimitError) as exc_info:
await hook.async_pre_call_hook(
user_api_key_dict=key,
cache=DualCache(),
data={
"litellm_call_id": "call-2",
"metadata": {
"tags": ["company_id:secret-internal-name"],
"inherited_tags": ["company_id:secret-internal-name"],
},
},
call_type="completion",
)
assert "tag_value" not in exc_info.value.detail