diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 60748234873..0a80f47ce59 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,19 +150,8 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None -def _client_steered_trace_field(data: dict, field: str) -> bool: - """Did the caller explicitly steer ``field`` in the request body? - - ``metadata`` (``litellm_metadata`` on ``LITELLM_METADATA_ROUTES``) is the - documented way to set ``trace_id`` / ``session_id``, so a value there is an - explicit caller choice and must outrank the traceparent/baggage fallback - below, which only guards on the top-level ``litellm_*`` keys. - """ - for variable_name in ("metadata", "litellm_metadata"): - container = data.get(variable_name) - if isinstance(container, dict) and container.get(field) is not None: - return True - return False +def _client_steered_trace_field(metadata: object, field: str) -> bool: + return isinstance(metadata, Mapping) and bool(cast(Mapping[str, object], metadata).get(field)) def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -1590,12 +1579,14 @@ class LiteLLMProxyRequestSetup: # (https://www.w3.org/TR/trace-context/, https://www.w3.org/TR/baggage/). # Lower priority than everything above - only fires when neither the # explicit litellm headers, the Anthropic-metadata path, nor the - # caller's own request-body metadata set the field - but lets a - # caller's existing traceparent/baggage headers - # (from real OTel instrumentation) correlate with litellm's own logs - # instead of generating an unrelated trace_id. + # caller's own request metadata set the field - but lets a caller's + # existing traceparent/baggage headers (from real OTel instrumentation) + # 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(data, "trace_id"): + if "litellm_trace_id" not in data and not _client_steered_trace_field( + cast(object, data.get(_metadata_variable_name)), "trace_id" + ): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) @@ -1605,7 +1596,9 @@ 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(data, "session_id"): + if "litellm_session_id" not in data and not _client_steered_trace_field( + cast(object, data.get(_metadata_variable_name)), "session_id" + ): baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): session_id_from_baggage: Final = _session_id_from_baggage(baggage) 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 bd350f56410..bbd8992928d 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3649,12 +3649,6 @@ def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_trace def test_add_litellm_metadata_from_request_headers_body_trace_id_beats_traceparent(): - """A caller that set metadata.trace_id keeps it: the traceparent fallback is - documented as last-resort, so it must not overwrite an explicit choice. - - This matters in practice because some platforms (e.g. GCP) inject a - traceparent into every inbound request, so the fallback would otherwise fire - on traffic whose caller never sent the header at all.""" headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"} data = {"metadata": {"trace_id": "caller-chosen-trace-id"}} LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( @@ -3665,7 +3659,6 @@ def test_add_litellm_metadata_from_request_headers_body_trace_id_beats_tracepare def test_add_litellm_metadata_from_request_headers_body_session_id_beats_baggage(): - """Same for session_id and the baggage header.""" headers = {"baggage": "session.id=baggage-session-42"} data = {"metadata": {"session_id": "caller-chosen-session-id"}} LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( @@ -3676,8 +3669,6 @@ def test_add_litellm_metadata_from_request_headers_body_session_id_beats_baggage def test_add_litellm_metadata_from_request_headers_body_steering_is_per_field(): - """Steering one field must not suppress the fallback for the other: - a caller setting only trace_id still gets session_id from baggage.""" headers = { "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", "baggage": "session.id=baggage-session-42", @@ -3691,9 +3682,6 @@ def test_add_litellm_metadata_from_request_headers_body_steering_is_per_field(): def test_add_litellm_metadata_from_request_headers_litellm_metadata_steering_honoured(): - """Routes in LITELLM_METADATA_ROUTES (/v1/responses, /v1/messages, batches, - files) carry their metadata in litellm_metadata, so steering there counts - the same as steering in metadata.""" headers = {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"} data = {"litellm_metadata": {"trace_id": "caller-chosen-trace-id"}} LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( @@ -3703,6 +3691,28 @@ def test_add_litellm_metadata_from_request_headers_litellm_metadata_steering_hon assert "litellm_trace_id" not in data +@pytest.mark.parametrize("empty_session_id", ["", None]) +def test_add_litellm_metadata_from_request_headers_empty_body_session_id_falls_back_to_baggage( + empty_session_id: str | None, +): + headers = {"baggage": "session.id=baggage-session-42"} + data = {"metadata": {"session_id": empty_session_id}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "baggage-session-42" + 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": {}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert data["litellm_session_id"] == "baggage-session-42" + + def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False)) @@ -8185,6 +8195,29 @@ async def test_missing_session_id_omit_keeps_client_supplied_session_id(): assert _spend_log_session_id(updated) == "client-session-1" +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("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"}}), + ], +) +async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, client_body: dict[str, object]): + request = _request_for(path) + request.headers = {"baggage": "session.id=baggage-session-42"} + + updated = await add_litellm_data_to_request( + data=client_body, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "reject"}, + ) + + assert updated["litellm_session_id"] == "baggage-session-42" + + @pytest.mark.asyncio @pytest.mark.parametrize( "client_body",