diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ae7320350de..107c3dd73d1 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,17 +150,25 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None -def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> object | None: +def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> str | None: + """The caller's value for a trace-control field, counted only when it is a + usable id: a non-empty string. An explicitly empty/unusable value on the + active metadata container still shadows the promoted requester value, but + neither ever counts as "the caller supplied this field" on its own, so a + numeric session id or an empty string cannot suppress the W3C header + fallback or satisfy a missing-session-id policy.""" active: Final = data.get(metadata_variable_name) if isinstance(active, Mapping) and field in active: active_map: Final = cast(Mapping[str, object], active) # cast-ok: isinstance above, free-form JSON values - return active_map[field] or None + active_value: Final = active_map[field] + return active_value if isinstance(active_value, str) and active_value else None promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS requester: Final = data.get("metadata") if not promoted or not isinstance(requester, Mapping): return None requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values - return requester_map.get(field) or None + requester_value: Final = requester_map.get(field) + return requester_value if isinstance(requester_value, str) and requester_value else None def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -841,7 +849,16 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or _caller_trace_field(data, _metadata_variable_name, "session_id") is not None: + caller_session_id: Final = _caller_trace_field(data, _metadata_variable_name, "session_id") + if caller_session_id is not None: + # The caller supplied a usable session id, so generate/reject must not + # fire. Surface it on the root field as well: consumers that read + # ``litellm_session_id`` (router fallbacks, spend logs, sandbox reuse) + # otherwise see no session at all and mint a fresh uuid4 per request. + if not data.get("litellm_session_id"): + data["litellm_session_id"] = caller_session_id # rebind-ok: data is an out-param + return + if data.get("litellm_session_id"): return match policy: case "generate": @@ -1597,42 +1614,44 @@ 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 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 to a DIFFERENT usable id + # - 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 _caller_trace_field( - cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object - _metadata_variable_name, - "trace_id", - ) - is None - ): + if "litellm_trace_id" not in data: traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) - if trace_id_from_traceparent: + # The caller's metadata wins over the header fallback unless + # both carry the same id: stamping the root field then claims + # nothing the caller did not already ask for, and keeps the + # W3C-correlated root trace id instead of a generated uuid4. + caller_trace_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "trace_id", + ) + if trace_id_from_traceparent and ( + caller_trace_id is None or caller_trace_id == trace_id_from_traceparent + ): metadata_from_headers["trace_id"] = trace_id_from_traceparent data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param verbose_proxy_logger.debug( "Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent ) - if ( - "litellm_session_id" not in data - and _caller_trace_field( - cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object - _metadata_variable_name, - "session_id", - ) - is None - ): + if "litellm_session_id" not in data: baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): session_id_from_baggage: Final = _session_id_from_baggage(baggage) - if session_id_from_baggage: + caller_session_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "session_id", + ) + if session_id_from_baggage and ( + caller_session_id is None or caller_session_id == session_id_from_baggage + ): metadata_from_headers["session_id"] = session_id_from_baggage data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param verbose_proxy_logger.debug("Extracted session_id from W3C baggage header") 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 8a10f67dbfa..c54319247c2 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3743,6 +3743,36 @@ def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_heade assert f"litellm_{field}" not in data +def test_add_litellm_metadata_from_request_headers_equal_ids_still_stamp_root_fields(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=matching-session-42", + } + data = { + "metadata": {"trace_id": "4bf92f3577b34da6a3ce929d0e0e4736", "session_id": "matching-session-42"}, + } + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["litellm_session_id"] == "matching-session-42" + assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["metadata"]["session_id"] == "matching-session-42" + + +@pytest.mark.parametrize("non_string_session_id", [4815162342, True, {"session": "nested"}]) +def test_add_litellm_metadata_from_request_headers_non_string_body_session_id_falls_back_to_baggage( + non_string_session_id: object, +): + headers = {"baggage": "session.id=header-session-42"} + data = {"metadata": {"session_id": non_string_session_id}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "header-session-42" + assert data["metadata"]["session_id"] == "header-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)) @@ -8456,6 +8486,58 @@ async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path: ) +@pytest.mark.asyncio +@pytest.mark.parametrize("policy", ["generate", "reject"]) +async def test_missing_session_id_policy_promotes_caller_session_to_root_field(policy: str): + """A caller-supplied usable session id satisfies the missing_session_id policies on + litellm_metadata routes (where it is not yet the managed metadata field) and must also + land on the root ``litellm_session_id`` field: consumers that read the root field + (router fallbacks, spend logs, sandbox reuse) otherwise mint a fresh uuid4 per request.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": "caller-session-42"}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy}, + ) + + assert updated["litellm_trace_id"] == "root-trace-42" + assert updated["litellm_session_id"] == "caller-session-42" + assert updated["litellm_metadata"]["session_id"] == "caller-session-42" + assert SESSION_ID_GENERATED_METADATA_KEY not in updated["litellm_metadata"] + + +@pytest.mark.asyncio +async def test_missing_session_id_generate_ignores_non_string_caller_session_id(): + """A non-string session id is not a usable session: the generate policy must fall through + to generation instead of letting an unusable value strand the root session field.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": 4815162342}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "generate"}, + ) + + assert updated["litellm_session_id"] == "root-trace-42" + assert updated["litellm_metadata"]["session_id"] == "root-trace-42" + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True + + @pytest.mark.asyncio @pytest.mark.parametrize("policy", ["generate", "reject"]) async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):