fix(proxy): resolve inherited_tags from litellm_params at success-event time too

order_tags_for_identity_resolution only checked the top level of the metadata
dict, which is correct for admission's flat request_kwargs but never present at
async_log_success_event time -- Logging.model_call_details only ever nests
metadata under kwargs["litellm_params"]. Admission correctly preferred the
key-backed identity tag, but token/dollar accounting fell through to the
caller-forged one instead, charging a different bucket than the one admission
actually checked. Adds the same litellm_params fallback _get_tags_from_request_kwargs
already relies on. bugbot caught this on review.
This commit is contained in:
Deepanshu 2026-08-26 14:01:32 -04:00
parent d1cc806ba4
commit 3a2f893134
2 changed files with 79 additions and 2 deletions

View file

@ -247,6 +247,26 @@ def extract_key_alias(request_kwargs: Mapping[str, object], metadata_variable_na
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`/`async_log_failure_event` time) never
carries `metadata`/`litellm_metadata` at its own top level, only nested
under `request_kwargs["litellm_params"]`. Checking only the top level
silently finds nothing at success/failure 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")
if isinstance(litellm_params, Mapping):
nested: Final = litellm_params.get(metadata_variable_name)
if isinstance(nested, Mapping):
return nested
return EMPTY_MAPPING
def order_tags_for_identity_resolution(
tags: Sequence[str], request_kwargs: Mapping[str, object], metadata_variable_name: str
) -> tuple[str, ...]:
@ -262,8 +282,8 @@ def order_tags_for_identity_resolution(
litellm_pre_call_utils.py), so putting it first makes a policy-backed tag
win over a same-prefix caller-supplied one.
"""
active: Final = request_kwargs.get(metadata_variable_name) or EMPTY_MAPPING
inherited_tags: Final = active.get("inherited_tags") if isinstance(active, Mapping) else None
active: Final = _active_metadata_bucket(request_kwargs, metadata_variable_name)
inherited_tags: Final = active.get("inherited_tags")
if not isinstance(inherited_tags, (list, tuple)) or not inherited_tags:
return tuple(tags)
return tuple(dict.fromkeys((*inherited_tags, *tags)))

View file

@ -1378,6 +1378,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):
"""
Bugbot finding: 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 only
checked the top level, so 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):
"""