diff --git a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py index 010996e29f9..66ee3fccb44 100644 --- a/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py +++ b/litellm/proxy/hooks/model_based_tag_rate_limits_hook.py @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py index 1ad4ec94ff8..0f19ede858c 100644 --- a/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py +++ b/tests/test_litellm/proxy/hooks/test_model_based_tag_rate_limits_hook.py @@ -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): """