From 648a89eb7e4ce5ff5424336c5413375c6c17c91e Mon Sep 17 00:00:00 2001 From: Yucheng He Date: Fri, 4 Sep 2026 19:03:14 -0700 Subject: [PATCH] fix(proxy): drop a forged spend logs metadata header from provider_specific_header The provider handlers merge provider_specific_header over data["headers"] last of all, so a caller copy there made the upstream proxy record an attribution this proxy never logged. The strip now walks every request-body header dict the handlers merge, not just extra_headers. --- litellm/proxy/litellm_pre_call_utils.py | 30 +++++++--- .../proxy/test_litellm_pre_call_utils.py | 55 ++++++++++++++++++- 2 files changed, 74 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b49a86e35e2..9d03d20baab 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -4,7 +4,7 @@ import json import re import time from collections import OrderedDict -from collections.abc import Mapping, MutableMapping, Sequence +from collections.abc import Iterator, Mapping, MutableMapping, Sequence from datetime import datetime from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, cast @@ -1266,6 +1266,20 @@ class LiteLLMProxyRequestSetup: return None return encoded + @staticmethod + def _caller_supplied_header_dicts( + data: Mapping[str, object], + ) -> Iterator[dict]: # mutable-ok: the caller deletes the forged key out of each one + """Every request-body dict a provider handler merges over ``data["headers"]``.""" + extra_headers: Final = data.get("extra_headers") + if isinstance(extra_headers, dict): + yield extra_headers + provider_specific: Final = data.get("provider_specific_header") + for entry in provider_specific if isinstance(provider_specific, list) else [provider_specific]: + scoped = entry.get("extra_headers") if isinstance(entry, dict) else None # rebind-ok: loop-local + if isinstance(scoped, dict): + yield scoped + @staticmethod def add_spend_logs_metadata_to_llm_call_headers( data: MutableMapping[str, object], # mutable-ok: this helper writes the outbound header into it @@ -1288,15 +1302,15 @@ class LiteLLMProxyRequestSetup: if not general_settings or general_settings.get("forward_spend_logs_metadata_to_llm_api") is not True: return - # Past the opt-in gate the proxy owns this header, and `extra_headers` beats - # `headers` in every provider handler, so a caller copy left in the request body - # would forge the attribution the upstream records whenever nothing is emitted - caller_extra_headers: Final = data.get("extra_headers") - if isinstance(caller_extra_headers, dict): + # Past the opt-in gate the proxy owns this header, and both `extra_headers` and + # `provider_specific_header` are merged over `headers` by the provider handlers, + # so a caller copy left in the request body would forge the attribution the + # upstream records even when this proxy resolved and logged a different one + for caller_headers in LiteLLMProxyRequestSetup._caller_supplied_header_dicts(data): for key in [ - k for k in caller_extra_headers if isinstance(k, str) and k.lower() == SPEND_LOGS_METADATA_HEADER_NAME + k for k in caller_headers if isinstance(k, str) and k.lower() == SPEND_LOGS_METADATA_HEADER_NAME ]: - del caller_extra_headers[key] # rebind-ok: dropping the forged copy is the point + del caller_headers[key] # rebind-ok: dropping the forged copy is the point encoded: Final = LiteLLMProxyRequestSetup._encode_spend_logs_metadata_header(data.get(_metadata_variable_name)) diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 80ba80fa11e..ddffbe2fc23 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -4387,6 +4387,7 @@ async def test_bearer_token_not_in_debug_logs(): # =====================================================================# Tests for credential overrides from model_config (team/project metadata) # ===================================================================== + @pytest.fixture() def setup_test_credentials(): """Populate litellm.credential_list with test credentials and enable feature flag, clean up after.""" @@ -5612,6 +5613,7 @@ class TestApplyKeyTagsPreAuth: # user-facing model name has no provider prefix. # ===================================================================== + def test_resolve_provider_from_deployment_uses_litellm_params_model(): """When custom_llm_provider is unset, fall back to the prefix of model.""" router = MagicMock() @@ -7741,9 +7743,7 @@ def _proxy_chain_auth( return UserAPIKeyAuth( api_key="hashed-key", metadata=( - metadata - if metadata is not None - else {"spend_logs_metadata": {"user_id": "U0099887", "username": "jdoe"}} + metadata if metadata is not None else {"spend_logs_metadata": {"user_id": "U0099887", "username": "jdoe"}} ), team_metadata=( team_metadata @@ -7802,6 +7802,55 @@ async def test_forward_spend_logs_metadata_overrides_caller_supplied_extra_heade ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "provider_specific_header", + [ + { + "custom_llm_provider": "openai", + "extra_headers": {SPEND_LOGS_METADATA_HEADER_NAME: json.dumps({"user_id": "forged"})}, + }, + [ + {"custom_llm_provider": "anthropic", "extra_headers": {"anthropic-beta": "v1"}}, + { + "custom_llm_provider": "openai", + "extra_headers": {SPEND_LOGS_METADATA_HEADER_NAME: json.dumps({"user_id": "forged"})}, + }, + ], + ], + ids=["single", "list"], +) +async def test_forward_spend_logs_metadata_overrides_caller_supplied_provider_specific_header( + provider_specific_header: object, +): + """ + `provider_specific_header` is merged over `headers` last of all, so a caller copy there + forged the attribution the upstream recorded while this proxy logged its own resolved value. + """ + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "messages": [], + "provider_specific_header": provider_specific_header, + }, + request=_spend_logs_metadata_request(), + user_api_key_dict=_proxy_chain_auth(), + proxy_config=MagicMock(), + general_settings={"forward_spend_logs_metadata_to_llm_api": True}, + ) + + assert json.loads(updated["headers"][SPEND_LOGS_METADATA_HEADER_NAME])["user_id"] == "U0099887" + scoped = updated["provider_specific_header"] + entries = scoped if isinstance(scoped, list) else [scoped] + for entry in entries: + assert SPEND_LOGS_METADATA_HEADER_NAME not in entry["extra_headers"], ( + "the caller's copy must be dropped so the proxy's resolved value is what ships" + ) + assert [e["extra_headers"] for e in entries if e["extra_headers"]] == ( + [{"anthropic-beta": "v1"}] if len(entries) > 1 else [] + ), "only the forged key goes, every other scoped header the caller sent survives" + + @pytest.mark.asyncio @pytest.mark.parametrize( "key_metadata",