diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e0a4184291e..4e758721d2c 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -100,6 +100,7 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, _get_dynamic_logging_metadata, # pyright: ignore[reportPrivateUsage] # shared proxy helper, same import style as _read_request_body above + get_chain_id_from_headers, ) from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.utils import normalize_route_for_root_path @@ -643,17 +644,32 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils): if isinstance(deployment_model_info, Mapping): _metadata["model_info"] = dict(deployment_model_info) - kwargs: Final = { - "litellm_params": { - **litellm_params_in_body, - "metadata": _metadata, - "proxy_server_request": { - "url": str(request.url), - "method": request.method, - "body": copy.copy(_parsed_body), # use copy instead of deepcopy - "headers": request.headers, - }, + # Match /v1/chat/completions: honor x-litellm-session-id / x-litellm-trace-id + # so spend logs group pass-through calls with the rest of the session (#43540). + # get_standard_logging_payload_trace_id reads litellm_params["litellm_session_id"]. + _request_headers = getattr(request, "headers", None) or {} + chain_id = get_chain_id_from_headers(dict(_request_headers)) + if chain_id: + # Header wins over any client-supplied body metadata (same as chat). + _metadata["session_id"] = chain_id + _metadata["trace_id"] = chain_id + + litellm_params: dict = { + **litellm_params_in_body, + "metadata": _metadata, + "proxy_server_request": { + "url": str(request.url), + "method": request.method, + "body": copy.copy(_parsed_body), # use copy instead of deepcopy + "headers": request.headers, }, + } + if chain_id: + litellm_params["litellm_session_id"] = chain_id + litellm_params["litellm_trace_id"] = chain_id + + kwargs: Final = { + "litellm_params": litellm_params, "call_type": "pass_through_endpoint", "litellm_call_id": litellm_call_id, "passthrough_logging_payload": passthrough_logging_payload, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 3469df082e0..c2fb322b8fd 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -7305,6 +7305,71 @@ def test_passthrough_client_cannot_forge_session_id_omission(client_metadata_key ) +@pytest.mark.parametrize( + "header_name,header_value", + [ + ("x-litellm-session-id", "S1"), + ("x-litellm-trace-id", "T1"), + ], +) +def test_passthrough_honors_session_header_for_spend_logs(header_name: str, header_value: str): + """Pass-through routes must pick up x-litellm-session-id / x-litellm-trace-id + the same way /v1/chat/completions does (#43540). Without this, every + /gemini/... spend row gets a fresh uuid4 session.""" + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.headers = Headers({header_name: header_value}) + mock_request.scope = {} + + kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body={}, + litellm_call_id="lit-43540-call-id", + ) + + litellm_params = kwargs["litellm_params"] + assert litellm_params["litellm_session_id"] == header_value + assert litellm_params["litellm_trace_id"] == header_value + assert litellm_params["metadata"]["session_id"] == header_value + assert litellm_params["metadata"]["trace_id"] == header_value + assert ( + StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id( + logging_obj=MagicMock(litellm_trace_id="fallback-uuid"), + litellm_params=litellm_params, + ) + == header_value + ) + + +def test_passthrough_without_session_header_does_not_invent_one(): + """No session header → leave litellm_session_id unset so spend logging + keeps its existing uuid4 / omit fallback.""" + mock_request = MagicMock(spec=Request) + mock_request.method = "POST" + mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent" + mock_request.headers = Headers({}) + mock_request.scope = {} + + kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint( + request=mock_request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + passthrough_logging_payload=MagicMock(), + logging_obj=MagicMock(), + _parsed_body={}, + litellm_call_id="lit-43540-no-header", + ) + + assert "litellm_session_id" not in kwargs["litellm_params"] + assert "litellm_trace_id" not in kwargs["litellm_params"] + assert "session_id" not in kwargs["litellm_params"]["metadata"] + + @pytest.mark.parametrize("client_metadata_key", ["litellm_metadata", "metadata"]) def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_body(client_metadata_key: str): """A provider route that resolved a router deployment stashes its model_info on request.state. That