From c2eb718ddd377fb7d4ad3e35422c0296319a1fc3 Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 16:29:48 -0300 Subject: [PATCH] fix(proxy): count promoted caller trace ids on litellm_metadata routes Addresses review. On LITELLM_METADATA_ROUTES the proxy promotes the caller's trace control fields (trace_id, session_id, ...) from `metadata` into `litellm_metadata`, but only after the header fallback runs. Checking only `litellm_metadata` let traceparent and baggage fill those fields first, and the promotion then skipped them because they were already set, so the trace was recorded under the header ids The check now also reads `metadata` for fields in LITELLM_TRACE_CONTROL_METADATA_FIELDS on those routes, and apply_missing_session_id_policy uses the same check, so `reject` no longer refuses a request whose session id is about to be promoted The earlier test asserting that a session_id in `metadata` must not block baggage on /v1/responses had the premise wrong, since that value is promoted. It is replaced by tests that assert the promoted caller ids win on /v1/responses and /v1/messages, with no policy, reject, and generate --- litellm/proxy/litellm_pre_call_utils.py | 19 ++++++---- .../proxy/test_litellm_pre_call_utils.py | 37 ++++++++++++++++--- 2 files changed, 44 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 0a80f47ce59..8397a904190 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,8 +150,13 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None -def _client_steered_trace_field(metadata: object, field: str) -> bool: - return isinstance(metadata, Mapping) and bool(cast(Mapping[str, object], metadata).get(field)) +def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool: + promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS + sources: Final = (metadata_variable_name, "metadata") if promoted else (metadata_variable_name,) + return any( + isinstance(container := data.get(name), Mapping) and bool(cast(Mapping[str, object], container).get(field)) + for name in sources + ) def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -829,7 +834,7 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or metadata.get("session_id"): + if data.get("litellm_session_id") or _caller_set_trace_field(data, _metadata_variable_name, "session_id"): return match policy: case "generate": @@ -1584,8 +1589,8 @@ class LiteLLMProxyRequestSetup: # correlate with litellm's own logs instead of generating an unrelated # trace_id. normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)}) - if "litellm_trace_id" not in data and not _client_steered_trace_field( - cast(object, data.get(_metadata_variable_name)), "trace_id" + if "litellm_trace_id" not in data and not _caller_set_trace_field( + cast(Mapping[str, object], data), _metadata_variable_name, "trace_id" ): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): @@ -1596,8 +1601,8 @@ class LiteLLMProxyRequestSetup: verbose_proxy_logger.debug( "Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent ) - if "litellm_session_id" not in data and not _client_steered_trace_field( - cast(object, data.get(_metadata_variable_name)), "session_id" + if "litellm_session_id" not in data and not _caller_set_trace_field( + cast(Mapping[str, object], data), _metadata_variable_name, "session_id" ): baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): 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 bbd8992928d..f99eb5c9b30 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3704,13 +3704,18 @@ def test_add_litellm_metadata_from_request_headers_empty_body_session_id_falls_b assert data["metadata"]["session_id"] == "baggage-session-42" -def test_add_litellm_metadata_from_request_headers_provider_metadata_does_not_block_baggage(): - headers = {"baggage": "session.id=baggage-session-42"} - data = {"metadata": {"session_id": "provider-facing-value"}, "litellm_metadata": {}} +@pytest.mark.parametrize("field", ["trace_id", "session_id"]) +def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_headers(field: str): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = {"metadata": {field: "caller-chosen"}, "litellm_metadata": {}} LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( headers=headers, data=data, _metadata_variable_name="litellm_metadata" ) - assert data["litellm_session_id"] == "baggage-session-42" + assert field not in data["litellm_metadata"] + assert f"litellm_{field}" not in data def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: @@ -8200,7 +8205,7 @@ async def test_missing_session_id_omit_keeps_client_supplied_session_id(): ("path", "client_body"), [ ("/v1/chat/completions", {"model": "gpt-4o", "messages": [], "metadata": {"session_id": ""}}), - ("/v1/responses", {"model": "gpt-4o", "input": "hi", "metadata": {"session_id": "provider-facing"}}), + ("/v1/responses", {"model": "gpt-4o", "input": "hi"}), ], ) async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, client_body: dict[str, object]): @@ -8218,6 +8223,28 @@ async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, c assert updated["litellm_session_id"] == "baggage-session-42" +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"]) +@pytest.mark.parametrize("policy", [None, "reject", "generate"]) +async def test_promoted_caller_trace_ids_beat_traceparent_and_baggage(path: str, policy: str | None): + request = _request_for(path) + request.headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + + updated = await add_litellm_data_to_request( + data={"model": "gpt-4o", "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}}, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy} if policy else {}, + ) + + assert updated["litellm_metadata"]["trace_id"] == "caller-trace" + assert updated["litellm_metadata"]["session_id"] == "caller-session" + + @pytest.mark.asyncio @pytest.mark.parametrize( "client_body",