fix(rate-limiting): stop a caller-supplied tag from shadowing a policy-backed identity tag (ported from #38289/#38292)

_extract_identity/_entry_applies resolve a tag_id by first-match-by-prefix
over the merged tags list, but _merge_tags (litellm_pre_call_utils.py)
keeps caller-supplied tags ahead of key/team/project tags in that list.
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
rate-limit entry scoped to company_id resolve to the caller's own value
instead of the key's.

Adds _order_tags_for_identity_resolution, which puts metadata.inherited_tags
(the server-computed snapshot of only the tags the calling key/team/
project's own config contributed) ahead of the full tags list before
either lookup runs, checking both the flat admission-time shape and the
litellm_params-nested success-event shape. Wired into both admission and
success-event tag resolution in this hook; the sibling global hook reuses
this same function (it already imports helpers from this module).
This commit is contained in:
Deepanshu 2026-08-27 07:22:58 -04:00
parent 01b75cdfd9
commit ec529f4b94
2 changed files with 159 additions and 3 deletions

View file

@ -316,6 +316,44 @@ def _extract_key_alias(request_kwargs: Mapping[str, object], metadata_variable_n
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` time) never carries
`metadata`/`litellm_metadata` at its own top level, only nested under
`request_kwargs["litellm_params"]`. Checking only the top level would
silently find nothing at success time, exactly the same failure mode
that lookup already had to handle."""
top_level: Final = request_kwargs.get(metadata_variable_name)
if isinstance(top_level, Mapping):
return top_level
litellm_params: Final = request_kwargs.get("litellm_params")
nested: Final = litellm_params.get(metadata_variable_name) if isinstance(litellm_params, Mapping) else None
return nested if isinstance(nested, Mapping) else _EMPTY_MAPPING
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.
"""
inherited_tags: Final = _active_metadata_bucket(request_kwargs, metadata_variable_name).get("inherited_tags")
if not isinstance(inherited_tags, (list, tuple)) or not inherited_tags:
return tuple(tags)
return tuple(dict.fromkeys((*inherited_tags, *tags)))
def _entries_for_unit(deployment: Mapping[str, object], unit: _LimitUnit) -> tuple[TagRateLimitEntry, ...]:
raw_tag_rate_limits: Final = (deployment.get("model_info") or _EMPTY_MAPPING).get("tag_rate_limits")
if not raw_tag_rate_limits:
@ -1303,8 +1341,10 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
if not configured:
return healthy_deployments
tags: Final = _get_tags_from_request_kwargs(
resolved_request_kwargs, metadata_variable_name=metadata_variable_name
tags: Final = _order_tags_for_identity_resolution(
_get_tags_from_request_kwargs(resolved_request_kwargs, metadata_variable_name=metadata_variable_name),
resolved_request_kwargs,
metadata_variable_name,
)
present_deployment_ids: Final[frozenset[str]] = frozenset(
@ -1816,7 +1856,11 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
if not configured:
return
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

View file

@ -912,6 +912,61 @@ async def test_filter_deployments_allows_under_limit_and_rejects_at_limit(time_c
assert exc_info.value.detail["limit_name"] == "per_minute"
@pytest.mark.asyncio
async def test_filter_deployments_a_caller_supplied_tag_cannot_shadow_the_policy_backed_identity_tag(time_controller):
"""
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 real
`company_id:real-company` tag (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. `_order_tags_for_identity_resolution`
must put `inherited_tags` first so admission charges the real bucket.
"""
limiter = _make_limiter(time_controller)
router = litellm.Router(
model_list=[
_deployment(
"grp",
"dep-1",
{
"request_limits": {
"limits": [{"name": "per-company", "tag_id": "company_id", "limit": 1, "period_seconds": 86400}]
}
},
)
]
)
limiter.update_variables(llm_router=router)
healthy = router.model_list
poisoned_request_kwargs = {
"metadata": {
"tags": ["company_id:attacker-chosen", "company_id:real-company"],
"inherited_tags": ["company_id:real-company"],
}
}
result = await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=poisoned_request_kwargs
)
assert result == healthy
# 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_request_kwargs = {
"metadata": {"tags": ["company_id:real-company"], "inherited_tags": ["company_id:real-company"]}
}
with pytest.raises(ProxyRateLimitError):
await limiter.async_filter_deployments(
model="grp", healthy_deployments=healthy, messages=None, request_kwargs=victim_request_kwargs
)
@pytest.mark.asyncio
async def test_filter_deployments_falls_back_to_deployment_model_name_for_routing_group_calls(time_controller):
"""
@ -1617,6 +1672,63 @@ async def test_log_success_event_accounts_when_litellm_params_carries_a_null_lit
)
@pytest.mark.asyncio
async def test_log_success_event_accounts_the_key_backed_tag_not_a_caller_forged_one(time_controller):
"""
Security regression: admission (async_filter_deployments) sees a flat
request_kwargs where metadata.inherited_tags sits at the top level, but
kwargs at async_log_success_event time is Logging.model_call_details,
which only ever nests metadata under kwargs["litellm_params"] (see
test_log_success_event_accounts_when_litellm_params_carries_a_null_litellm_metadata_key).
_order_tags_for_identity_resolution's own inherited_tags lookup must
check both shapes, or a caller-forged company_id tag that admission
correctly ignored could still get accounted against at success time,
charging a different bucket than the one admission actually checked.
"""
limiter = _make_limiter(time_controller)
router = litellm.Router(
model_list=[
_deployment(
"grp",
"dep-1",
{
"token_limits": {
"limits": [{"name": "daily", "tag_id": "company_id", "limit": 500000, "period_seconds": 86400}]
}
},
)
]
)
limiter.update_variables(llm_router=router)
kwargs = {
"litellm_params": {
"metadata": {
"tags": ["company_id:attacker-chosen"],
"inherited_tags": ["company_id:real-company"],
},
},
"standard_logging_object": {
"model_group": "grp",
"model_id": "dep-1",
"total_tokens": 42,
"response_cost": 0.01,
},
}
await limiter.async_log_success_event(kwargs=kwargs, response_obj=None, start_time=0, end_time=0)
await asyncio.sleep(0)
now = time_controller.now().timestamp()
real_key = _expected_bucket_key("grp", "tokens", "daily", "company_id", "real-company", 86400, now, limit=500000)
forged_key = _expected_bucket_key(
"grp", "tokens", "daily", "company_id", "attacker-chosen", 86400, now, limit=500000
)
assert (
float(await limiter.internal_usage_cache.async_get_cache(key=real_key, litellm_parent_otel_span=None)) == 42.0
)
assert await limiter.internal_usage_cache.async_get_cache(key=forged_key, litellm_parent_otel_span=None) is None
@pytest.mark.asyncio
async def test_log_success_event_reads_nested_litellm_metadata_when_that_is_authoritative(time_controller):
"""