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",