From d656a4f18bdc142a913571205e5cd08690249ba4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 09:27:30 +0000 Subject: [PATCH] fix(proxy): generated session uses promoted caller trace id Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/litellm_pre_call_utils.py | 40 +++++++++----- .../observability/test_langfuse_delivery.py | 53 +++++++++++++++++++ .../proxy/test_litellm_pre_call_utils.py | 30 +++++++++++ 3 files changed, 109 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 660bfa7ccd3..ae7320350de 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,17 +150,17 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None -def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool: +def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> object | None: 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 bool(active_map[field]) + return active_map[field] or 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 False + return None requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values - return bool(requester_map.get(field)) + return requester_map.get(field) or None def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -841,11 +841,15 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or _caller_set_trace_field(data, _metadata_variable_name, "session_id"): + if data.get("litellm_session_id") or _caller_trace_field(data, _metadata_variable_name, "session_id") is not None: return match policy: case "generate": - session_id: Final = str(data.get("litellm_trace_id") or metadata.get("trace_id") or uuid.uuid4()) + session_id: Final = str( + data.get("litellm_trace_id") + or _caller_trace_field(data, _metadata_variable_name, "trace_id") + or uuid.uuid4() + ) data["litellm_session_id"] = session_id # rebind-ok: data is an out-param data.setdefault("litellm_trace_id", session_id) metadata["session_id"] = session_id @@ -1598,10 +1602,14 @@ 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 _caller_set_trace_field( - cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object - _metadata_variable_name, - "trace_id", + 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 ): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): @@ -1612,10 +1620,14 @@ 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 _caller_set_trace_field( - cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object - _metadata_variable_name, - "session_id", + 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 ): baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 18dc21f80d7..6f71070451d 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -602,3 +602,56 @@ def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback( ) assert third.status_code == 400, third.text assert upstream_targets == ["/v1/responses", "/v1/chat/completions"], upstream_targets + + +@pytest.mark.parametrize( + ("endpoint", "kind", "expected_target"), + ( + pytest.param("/v1/responses", "responses", "/v1/responses", id="responses_caller_trace"), + pytest.param("/v1/messages", "messages", "/v1/responses", id="messages_caller_trace"), + ), +) +def test_missing_session_id_generate_derives_session_from_caller_trace( + gateway: Gateway, tmp_path: Path, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "gen" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit + + def upstream(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {provider_secret}" + upstream_targets.append(request.target) + if request.target == "/v1/responses": + return _responses_result("resp-" + marker) + assert request.target == "/v1/chat/completions", request.target + return _completion(marker + "-answer") + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(upstream) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, + tmp_path, + _langfuse_environment(destination), + config=_langfuse_config(tmp_path, {"missing_session_id": "generate"}), + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace}), + headers=_w3c_headers(uuid.uuid4().hex, None), + ) + assert response.status_code == 200, response.text + assert upstream_targets == [expected_target], upstream_targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, destination, response.headers["x-litellm-call-id"]) + assert _attribute(span.attributes, "session.id") == caller_trace 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 6b572ef823b..722caad23cf 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8423,6 +8423,36 @@ async def test_missing_session_id_generate_reuses_traceparent_trace_id(): assert _spend_log_session_id(updated) == "4bf92f3577b34da6a3ce929d0e0e4736" +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"]) +@pytest.mark.parametrize( + "headers", + [ + {"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"}, + {}, + ], + ids=["with_traceparent", "no_traceparent"], +) +async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path: str, headers: dict[str, str]): + """On litellm_metadata routes the caller's metadata.trace_id still counts as caller set even though + it has not been promoted yet, so the generated session id derives from it instead of a fresh uuid.""" + request = _request_for(path) + request.headers = headers + + updated = await add_litellm_data_to_request( + data={"model": "gpt-4o", "input": "hi", "metadata": {"trace_id": "caller-trace"}}, + 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"] == "caller-trace" + assert updated["litellm_metadata"]["session_id"] == "caller-trace" + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True + assert _spend_log_session_id(updated, "litellm_metadata") == "caller-trace" + + @pytest.mark.asyncio @pytest.mark.parametrize("policy", ["generate", "reject"]) async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):