mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
ec529f4b94
commit
f24d362d8c
2 changed files with 439 additions and 31 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue