diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index fc83c1ddeed..78de702f2fb 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -179,6 +179,8 @@ from litellm.proxy.anthropic_endpoints.streaming_model_restamp import ( AnthropicStreamModelRestamper, ) from litellm.proxy.litellm_pre_call_utils import ( + LiteLLMProxyRequestSetup, + _get_metadata_variable_name, add_litellm_data_to_request, refresh_proxy_server_request_body_snapshot, reject_url_valued_destination, @@ -1849,8 +1851,6 @@ class ProxyBaseLLMRequestProcessing: # Store queue time in metadata after add_litellm_data_to_request to ensure it's preserved if queue_time_seconds is not None: - from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name - _metadata_variable_name: Final = _get_metadata_variable_name(request) if _metadata_variable_name not in self.data: self.data[_metadata_variable_name] = {} @@ -1980,6 +1980,13 @@ class ProxyBaseLLMRequestProcessing: # have mutated `self.data` in place, and the audit-trail snapshot taken in # add_litellm_data_to_request predates that mutation. refresh_proxy_server_request_body_snapshot(self.data) + # Same reason: a pre_call hook that rewrites `spend_logs_metadata` would + # otherwise leave the upstream proxy recording the pre-hook attribution + LiteLLMProxyRequestSetup.add_spend_logs_metadata_to_llm_call_headers( + data=self.data, + _metadata_variable_name=_get_metadata_variable_name(request), + general_settings=general_settings, + ) verbose_proxy_logger.debug("receiving data: %s", self.data) if "messages" in self.data and self.data["messages"]: diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 2b0630e4602..9e5e32f5dce 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1216,51 +1216,13 @@ class LiteLLMProxyRequestSetup: return returned_headers @staticmethod - def add_spend_logs_metadata_to_llm_call_headers( - data: MutableMapping[str, object], # mutable-ok: this helper writes the outbound header into it - _metadata_variable_name: str, - general_settings: Mapping[str, object] | None, - ) -> None: - """ - Emit the request's resolved ``spend_logs_metadata`` as the - ``x-litellm-spend-logs-metadata`` header on the outbound LLM call. - - Proxy-to-proxy attribution: an upstream LiteLLM proxy reads that header in - ``_get_spend_logs_metadata_from_request_headers`` and stores the values in its - own SpendLogs row. ``forward_client_headers_to_llm_api`` only relays headers - the client itself sent, so ``spend_logs_metadata`` the downstream resolved from - the virtual key or the team never reached the upstream. - - Must run after every key/team ``spend_logs_metadata`` merge so the header - carries the same values the downstream writes to its own SpendLogs. The - resolved dict already merges caller, key and team values with key/team losing - to the caller, so the upstream sees one deterministic namespace. - - Opt-in via ``general_settings.forward_spend_logs_metadata_to_llm_api``: the - values are customer identifiers and the header is sent to every configured - provider, not only to LiteLLM upstreams. - """ - 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, so drop any copy the caller put - # in the request body first. `extra_headers` beats `headers` in every provider - # handler, and the emission below can still be skipped (no values, oversized, - # unserializable) - leaving the caller's copy on those paths would let a caller - # forge the attribution the upstream records by making the resolved value too big. - caller_extra_headers: Final = data.get("extra_headers") - if isinstance(caller_extra_headers, dict): - for key in [ - k for k in caller_extra_headers if isinstance(k, str) and k.lower() == SPEND_LOGS_METADATA_HEADER_NAME - ]: - del caller_extra_headers[key] - - metadata: Final = data.get(_metadata_variable_name) + def _encode_spend_logs_metadata_header(metadata: object) -> str | None: + """The header value for ``metadata``'s ``spend_logs_metadata``, or None to send nothing.""" if not isinstance(metadata, dict): - return + return None spend_logs_metadata: Final = metadata.get("spend_logs_metadata") if not isinstance(spend_logs_metadata, dict) or not spend_logs_metadata: - return + return None try: encoded: Final = json.dumps(spend_logs_metadata) @@ -1268,7 +1230,7 @@ class LiteLLMProxyRequestSetup: verbose_proxy_logger.warning( "spend_logs_metadata is not JSON-serializable, not forwarding it to the LLM API" ) - return + return None encoded_size: Final = len(encoded.encode("utf-8")) if encoded_size > MAX_SPEND_LOGS_METADATA_HEADER_BYTES: verbose_proxy_logger.warning( @@ -1276,15 +1238,52 @@ class LiteLLMProxyRequestSetup: encoded_size, MAX_SPEND_LOGS_METADATA_HEADER_BYTES, ) + return None + return encoded + + @staticmethod + def add_spend_logs_metadata_to_llm_call_headers( + data: MutableMapping[str, object], # mutable-ok: this helper writes the outbound header into it + _metadata_variable_name: str, + general_settings: Mapping[str, object] | None, + ) -> None: + """ + Set the outbound ``x-litellm-spend-logs-metadata`` header from the request's + resolved ``spend_logs_metadata``, so an upstream LiteLLM proxy records the same + attribution this proxy does. + + Call it after every step that can change ``spend_logs_metadata``: the key and + team merges, and the pre-call hooks. Re-running it overwrites or removes the + header, so a later call never leaves a stale value behind. + + Opt-in via ``general_settings.forward_spend_logs_metadata_to_llm_api``: the + values are customer identifiers and the header goes to every configured + provider, not only to LiteLLM upstreams. + """ + 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): + for key in [ + k for k in caller_extra_headers if isinstance(k, str) and k.lower() == SPEND_LOGS_METADATA_HEADER_NAME + ]: + del caller_extra_headers[key] + + encoded: Final = LiteLLMProxyRequestSetup._encode_spend_logs_metadata_header(data.get(_metadata_variable_name)) + existing_headers: Final = data.get("headers") if isinstance(existing_headers, dict): - existing_headers[SPEND_LOGS_METADATA_HEADER_NAME] = encoded - else: - # Every provider reads `headers or litellm.headers`, replacing rather than - # merging, so creating this dict from scratch would drop the operator's - # `litellm_settings.headers` from every request the flag applies to. + if encoded is None: + existing_headers.pop(SPEND_LOGS_METADATA_HEADER_NAME, None) + else: + existing_headers[SPEND_LOGS_METADATA_HEADER_NAME] = encoded + elif encoded is not None: + # Providers read `headers or litellm.headers`, replacing rather than merging, + # so building this dict from scratch would drop `litellm_settings.headers` emitted: Final = dict(litellm.headers or {}) emitted[SPEND_LOGS_METADATA_HEADER_NAME] = encoded data["headers"] = emitted # rebind-ok: emitting this header is what this helper is for @@ -2343,8 +2342,6 @@ async def add_litellm_data_to_request( user_api_key_dict=user_api_key_dict, ) - # Runs after the key/team spend_logs_metadata merges above so the forwarded header - # carries the same values this proxy writes to its own SpendLogs. LiteLLMProxyRequestSetup.add_spend_logs_metadata_to_llm_call_headers( data=data, _metadata_variable_name=_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 e7c35ab27ba..b0c07013832 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8027,6 +8027,40 @@ async def test_forward_spend_logs_metadata_leaves_caller_header_alone_when_flag_ assert json.loads(updated["extra_headers"][SPEND_LOGS_METADATA_HEADER_NAME]) == {"user_id": "caller"} +@pytest.mark.parametrize( + "hook_metadata, expected", + [ + ({"user_id": "hook-overwrote", "hook_added": "yes"}, {"user_id": "hook-overwrote", "hook_added": "yes"}), + ({"blob": "x" * 5000}, None), + ({}, None), + ], +) +@pytest.mark.asyncio +async def test_forward_spend_logs_metadata_reruns_after_a_pre_call_hook_edit( + hook_metadata: dict[str, str], expected: dict[str, str] | None +): + """A pre_call hook that rewrites spend_logs_metadata must not leave a stale header behind.""" + general_settings = {"forward_spend_logs_metadata_to_llm_api": True} + data = await add_litellm_data_to_request( + data={"model": "gpt-4o", "messages": []}, + request=_spend_logs_metadata_request(), + user_api_key_dict=_proxy_chain_auth(), + proxy_config=MagicMock(), + general_settings=general_settings, + ) + assert SPEND_LOGS_METADATA_HEADER_NAME in data["headers"] + + data["metadata"]["spend_logs_metadata"] = hook_metadata + LiteLLMProxyRequestSetup.add_spend_logs_metadata_to_llm_call_headers( + data=data, + _metadata_variable_name="metadata", + general_settings=general_settings, + ) + + emitted = data["headers"].get(SPEND_LOGS_METADATA_HEADER_NAME) + assert (json.loads(emitted) if emitted is not None else None) == expected + + @pytest.mark.asyncio async def test_forward_spend_logs_metadata_keeps_globally_configured_headers(): """ @@ -8073,7 +8107,6 @@ async def test_forward_spend_logs_metadata_emits_resolved_key_and_team_values(): "username": "jdoe", "cost_center": "CC-42", } - # the header must carry exactly what this proxy logs for itself assert forwarded == updated["metadata"]["spend_logs_metadata"]