diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 8397a904190..fab3dab5978 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -151,12 +151,13 @@ def _session_id_from_baggage(baggage: str) -> str | None: def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool: + active: Final = data.get(metadata_variable_name) + active_mapping: Final = cast(Mapping[str, object], active) if isinstance(active, Mapping) else None + if active_mapping is not None and field in active_mapping: + return bool(active_mapping.get(field)) 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 - ) + requester: Final = data.get("metadata") + return promoted and isinstance(requester, Mapping) and bool(cast(Mapping[str, object], requester).get(field)) def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: 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 f99eb5c9b30..4bbb6238ddf 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8245,6 +8245,40 @@ async def test_promoted_caller_trace_ids_beat_traceparent_and_baggage(path: str, assert updated["litellm_metadata"]["session_id"] == "caller-session" +@pytest.mark.asyncio +async def test_missing_session_id_reject_ignores_requester_session_id_shadowed_by_empty_litellm_metadata(): + with pytest.raises(ProxyException) as exc_info: + await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "metadata": {"session_id": "caller-session"}, + "litellm_metadata": {"session_id": ""}, + }, + request=_request_for("/v1/responses"), + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "reject"}, + ) + assert exc_info.value.code == "400" + + +def test_add_litellm_metadata_from_request_headers_empty_litellm_metadata_field_falls_back_to_headers(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = { + "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}, + "litellm_metadata": {"trace_id": "", "session_id": ""}, + } + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["litellm_session_id"] == "baggage-session-42" + + @pytest.mark.asyncio @pytest.mark.parametrize( "client_body",