From a1dcd802adf89a0425881bd42be9107f69232a54 Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 15:48:29 -0300 Subject: [PATCH 01/12] fix(proxy): traceparent/baggage fallback must not override caller metadata The W3C traceparent/baggage fallback in add_litellm_metadata_from_request_headers documents itself as last-resort: Lower priority than everything above - only fires when neither the explicit litellm headers nor the Anthropic-metadata path found anything But it guards on the top-level `litellm_trace_id` / `litellm_session_id` body keys and never checks `metadata`, which is the documented way callers set trace_id / session_id (`metadata: {"trace_id": ...}` on /chat/completions). A caller that explicitly sets metadata.trace_id has it silently replaced by the header value, so the implementation contradicts its own stated precedence. This is not a corner case on managed platforms: GCP's front end injects a traceparent into every inbound request, so the fallback fires on traffic whose caller never sent the header. The request still returns 200 and the trace still reaches the logging backend, just under an id the caller never chose, so any caller correlating by its own id silently fails to find its trace. Guard both fallbacks on the caller's request-body metadata as well. Deliberately narrow: - x-litellm-trace-id still outranks the body (documented priority #1) - a request that steers neither field still adopts traceparent/baggage exactly as before - steering is per-field: setting only trace_id still lets session_id come from baggage - litellm_metadata is checked too, for LITELLM_METADATA_ROUTES (/v1/responses, /v1/messages, batches, files) Also corrects the comment, which understated the guard. 4 new tests; each fails without the source change. The existing traceparent and baggage tests are unchanged and still pass. --- litellm/proxy/litellm_pre_call_utils.py | 24 ++++++-- .../proxy/test_litellm_pre_call_utils.py | 55 +++++++++++++++++++ 2 files changed, 75 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 8fc5faee2c9..60748234873 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,6 +150,21 @@ 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 _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: """Only proxy-validated keys are stamped, proven by the unforgeable via_virtual_key marker AND a known non-secret shape: the sha256 hex digest @@ -1574,12 +1589,13 @@ class LiteLLMProxyRequestSetup: # Last-resort fallback: the W3C standards for trace/session propagation # (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 nor the Anthropic-metadata path found - # anything - but lets a caller's existing traceparent/baggage headers + # 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. normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)}) - if "litellm_trace_id" not in data: + if "litellm_trace_id" not in data and not _client_steered_trace_field(data, "trace_id"): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) @@ -1589,7 +1605,7 @@ 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: + if "litellm_session_id" not in data and not _client_steered_trace_field(data, "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 b2241191ced..bd350f56410 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3648,6 +3648,61 @@ def test_add_litellm_metadata_from_request_headers_explicit_trace_id_beats_trace assert data["litellm_session_id"] == "explicit-trace-id-value" +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( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["trace_id"] == "caller-chosen-trace-id" + assert "litellm_trace_id" not in data + + +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( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["session_id"] == "caller-chosen-session-id" + assert "litellm_session_id" not in data + + +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", + } + data = {"metadata": {"trace_id": "caller-chosen-trace-id"}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["trace_id"] == "caller-chosen-trace-id" + assert data["litellm_session_id"] == "baggage-session-42" + + +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( + headers=headers, data=data, _metadata_variable_name="litellm_metadata" + ) + assert data["litellm_metadata"]["trace_id"] == "caller-chosen-trace-id" + assert "litellm_trace_id" not in data + + def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False)) From bae13227fa46b9565c728582a66d3f4e62c58986 Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 16:10:12 -0300 Subject: [PATCH 02/12] fix(proxy): check only the active metadata container and ignore empty values Addresses review. The first version treated any value in either `metadata` or `litellm_metadata` as caller steering. On LITELLM_METADATA_ROUTES the body `metadata` is provider-facing and is only promoted later, so a session_id there suppressed baggage before apply_missing_session_id_policy ran, and a request with a usable baggage session id got a 400 under `missing_session_id: reject`. An empty session_id did the same, since the policy treats "" as absent Now only the active metadata container counts, and only a truthy value, which matches how apply_missing_session_id_policy decides a session id is present. The helper takes a typed `object` rather than a bare dict Adds regression tests for both cases, including the end to end reject path, and drops the test docstrings --- litellm/proxy/litellm_pre_call_utils.py | 31 ++++------ .../proxy/test_litellm_pre_call_utils.py | 57 +++++++++++++++---- 2 files changed, 57 insertions(+), 31 deletions(-) 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", From c2eb718ddd377fb7d4ad3e35422c0296319a1fc3 Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 16:29:48 -0300 Subject: [PATCH 03/12] fix(proxy): count promoted caller trace ids on litellm_metadata routes Addresses review. On LITELLM_METADATA_ROUTES the proxy promotes the caller's trace control fields (trace_id, session_id, ...) from `metadata` into `litellm_metadata`, but only after the header fallback runs. Checking only `litellm_metadata` let traceparent and baggage fill those fields first, and the promotion then skipped them because they were already set, so the trace was recorded under the header ids The check now also reads `metadata` for fields in LITELLM_TRACE_CONTROL_METADATA_FIELDS on those routes, and apply_missing_session_id_policy uses the same check, so `reject` no longer refuses a request whose session id is about to be promoted The earlier test asserting that a session_id in `metadata` must not block baggage on /v1/responses had the premise wrong, since that value is promoted. It is replaced by tests that assert the promoted caller ids win on /v1/responses and /v1/messages, with no policy, reject, and generate --- litellm/proxy/litellm_pre_call_utils.py | 19 ++++++---- .../proxy/test_litellm_pre_call_utils.py | 37 ++++++++++++++++--- 2 files changed, 44 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 0a80f47ce59..8397a904190 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,8 +150,13 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None -def _client_steered_trace_field(metadata: object, field: str) -> bool: - return isinstance(metadata, Mapping) and bool(cast(Mapping[str, object], metadata).get(field)) +def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool: + 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 + ) def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -829,7 +834,7 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or metadata.get("session_id"): + if data.get("litellm_session_id") or _caller_set_trace_field(data, _metadata_variable_name, "session_id"): return match policy: case "generate": @@ -1584,8 +1589,8 @@ 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 _client_steered_trace_field( - cast(object, data.get(_metadata_variable_name)), "trace_id" + if "litellm_trace_id" not in data and not _caller_set_trace_field( + cast(Mapping[str, object], data), _metadata_variable_name, "trace_id" ): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): @@ -1596,8 +1601,8 @@ 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( - cast(object, data.get(_metadata_variable_name)), "session_id" + if "litellm_session_id" not in data and not _caller_set_trace_field( + cast(Mapping[str, object], data), _metadata_variable_name, "session_id" ): baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): 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 bbd8992928d..f99eb5c9b30 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3704,13 +3704,18 @@ def test_add_litellm_metadata_from_request_headers_empty_body_session_id_falls_b 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": {}} +@pytest.mark.parametrize("field", ["trace_id", "session_id"]) +def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_headers(field: str): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + data = {"metadata": {field: "caller-chosen"}, "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" + assert field not in data["litellm_metadata"] + assert f"litellm_{field}" not in data def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: @@ -8200,7 +8205,7 @@ async def test_missing_session_id_omit_keeps_client_supplied_session_id(): ("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"}}), + ("/v1/responses", {"model": "gpt-4o", "input": "hi"}), ], ) async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, client_body: dict[str, object]): @@ -8218,6 +8223,28 @@ async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, c assert updated["litellm_session_id"] == "baggage-session-42" +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"]) +@pytest.mark.parametrize("policy", [None, "reject", "generate"]) +async def test_promoted_caller_trace_ids_beat_traceparent_and_baggage(path: str, policy: str | None): + request = _request_for(path) + request.headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=baggage-session-42", + } + + updated = await add_litellm_data_to_request( + data={"model": "gpt-4o", "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}}, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy} if policy else {}, + ) + + assert updated["litellm_metadata"]["trace_id"] == "caller-trace" + assert updated["litellm_metadata"]["session_id"] == "caller-session" + + @pytest.mark.asyncio @pytest.mark.parametrize( "client_body", From 36f07c87fc285fcc5567322297d849c5e90bf6ac Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 16:41:37 -0300 Subject: [PATCH 04/12] fix(proxy): an empty trace field in litellm_metadata shadows the promoted one Addresses review. Promotion copies a trace control field from `metadata` into `litellm_metadata` only when the key is absent, so a key that is present but empty in `litellm_metadata` wins and the `metadata` value is never used. The check still looked at `metadata` in that case, which let `missing_session_id: reject` accept a request that ends up with an empty session id, and let an empty trace_id or session_id block the headers The check now follows the same rule as promotion. If the active container has the key, its value decides. Only when the key is absent does the promoted `metadata` value count. This makes the conflicting-bucket cases behave exactly as on main --- litellm/proxy/litellm_pre_call_utils.py | 11 +++--- .../proxy/test_litellm_pre_call_utils.py | 34 +++++++++++++++++++ 2 files changed, 40 insertions(+), 5 deletions(-) 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", From c6fce25652821cb65e4bb2a88c1d50ff4ef98316 Mon Sep 17 00:00:00 2001 From: Filipe Andujar Date: Thu, 24 Sep 2026 17:04:59 -0300 Subject: [PATCH 05/12] fix(proxy): stay within the LIT006 cast budget The type-discipline gate failed: this branch added 4 unsuppressed `cast()` calls and LIT006 was already at its limit. Each cast now carries a `# cast-ok` reason. The two in the helper follow an `isinstance` check that proves the Mapping. The two at the call sites stay because the method's `data` parameter is a bare `dict`, and dropping those casts adds basedpyright unknown-type errors instead. No behavior change --- litellm/proxy/litellm_pre_call_utils.py | 19 +++++++++++++------ 1 file changed, 13 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index fab3dab5978..6e785fcd3f9 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -152,12 +152,15 @@ 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)) + 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]) promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS requester: Final = data.get("metadata") - return promoted and isinstance(requester, Mapping) and bool(cast(Mapping[str, object], requester).get(field)) + if not promoted or not isinstance(requester, Mapping): + return False + requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values + return bool(requester_map.get(field)) def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: @@ -1591,7 +1594,9 @@ class LiteLLMProxyRequestSetup: # 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), _metadata_variable_name, "trace_id" + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "trace_id", ): traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): @@ -1603,7 +1608,9 @@ class LiteLLMProxyRequestSetup: "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), _metadata_variable_name, "session_id" + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "session_id", ): baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): From 5baeb9e80b6e6ccef4d206d71d84705a1cc0dd20 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 08:19:28 +0000 Subject: [PATCH 06/12] test(proxy): cover caller metadata precedence over W3C trace headers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../observability/test_langfuse_delivery.py | 243 +++++++++++++++++- 1 file changed, 239 insertions(+), 4 deletions(-) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 13be5a95887..18dc21f80d7 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -2,10 +2,11 @@ import base64 import json import time import uuid -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from pathlib import Path from typing import Final +import pytest import yaml from integration._support.client import Gateway, eventually, object_value, string_value from integration._support.database import read_rows, scratch_database @@ -68,15 +69,20 @@ def _text_prompt(name: str) -> Reply: ) -def _langfuse_config(tmp_path: Path) -> Path: +def _langfuse_config(tmp_path: Path, general_settings: Mapping[str, JsonValue] | None = None) -> Path: config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) settings: Final = { **_SETTINGS.validate_python(config["litellm_settings"]), "success_callback": ["langfuse"], "failure_callback": ["langfuse"], } - path: Final = tmp_path / "langfuse.yaml" - path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + general: Final = { + **_SETTINGS.validate_python(config["general_settings"]), + **(general_settings or {}), + } + name: Final = "langfuse.yaml" if general_settings is None else "langfuse-merged.yaml" + path: Final = tmp_path / name + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general})) return path @@ -367,3 +373,232 @@ def test_prompt_fetch_encodes_the_name_retries_a_5xx_once_and_keeps_langfuse_hea assert leak not in json.dumps(dict(failure.headers)) assert "set-cookie" not in failure.headers and "x-upstream-internal" not in failure.headers assert sum(1 for target in seen_prompt_gets if target.startswith(PROMPTS_PATH + missing_prompt)) == 1 + + +def _responses_result(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + identity, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "integration answer", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + + +def _trace_body(kind: str, model: str, marker: str, metadata: Mapping[str, str] | None) -> dict[str, JsonValue]: + metadata_field: Final[dict[str, JsonValue]] = {} if metadata is None else {"metadata": dict(metadata)} + if kind == "responses": + return {"model": model, "input": marker + "-question", **metadata_field} + if kind == "messages": + return { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": marker + "-question"}], + **metadata_field, + } + return { + "model": model, + "messages": [{"role": "user", "content": marker + "-question"}], + "cache": {"no-cache": True}, + **metadata_field, + } + + +def _w3c_headers(header_trace: str, baggage_session: str | None) -> dict[str, str]: + baggage_field: Final = {} if baggage_session is None else {"baggage": f"session.id={baggage_session}"} + return {"traceparent": f"00-{header_trace}-00f067aa0ba902b7-01", **baggage_field} + + +def _await_span(received: list[Request], destination: Wire, call_id: str) -> Span: + def exported() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") == call_id + ) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=20) + return spans[0] + + +@pytest.mark.parametrize( + ("endpoint", "kind", "metadata_mode", "expected_trace", "expected_session", "expected_target"), + ( + pytest.param( + "/v1/chat/completions", "chat", "both", "caller", "caller", "/v1/chat/completions", id="chat_caller_ids" + ), + pytest.param( + "/v1/responses", "responses", "both", "caller", "caller", "/v1/responses", id="responses_caller_ids" + ), + pytest.param( + "/v1/messages", "messages", "both", "caller", "caller", "/v1/responses", id="messages_caller_ids" + ), + pytest.param( + "/v1/chat/completions", "chat", "none", "header", "baggage", "/v1/chat/completions", id="chat_header_ids" + ), + pytest.param( + "/v1/chat/completions", + "chat", + "trace", + "caller", + "baggage", + "/v1/chat/completions", + id="chat_caller_trace_header_session", + ), + pytest.param( + "/v1/chat/completions", + "chat", + "empty_session", + "header", + "baggage", + "/v1/chat/completions", + id="chat_empty_session_header_ids", + ), + ), +) +def test_langfuse_trace_and_session_prefer_caller_metadata_over_w3c_headers( + gateway: Gateway, + tmp_path: Path, + endpoint: str, + kind: str, + metadata_mode: str, + expected_trace: str, + expected_session: str, + expected_target: str, +) -> None: + marker: Final = "w3c" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = { + "both": {"trace_id": caller_trace, "session_id": caller_session}, + "trace": {"trace_id": caller_trace}, + "empty_session": {"session_id": ""}, + "none": None, + }[metadata_mode] + expected_trace_value: Final = {"caller": caller_trace, "header": header_trace}[expected_trace] + expected_session_value: Final = {"caller": caller_session, "baggage": baggage_session}[expected_session] + 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) + ) 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, metadata), + headers=_w3c_headers(header_trace, baggage_session), + ) + 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 span.trace_id.hex() == expected_trace_value + assert _attribute(span.attributes, "session.id") == expected_session_value + + +def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = "reject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + 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": "reject"}), + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + caller_session: Final = f"my-session-id-{marker}-r1" + header_trace: Final = uuid.uuid4().hex + first: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": marker + "-r1", "metadata": {"session_id": caller_session}}, + headers=_w3c_headers(header_trace, None), + ) + assert first.status_code == 200, first.text + first_span: Final = _await_span(received, destination, first.headers["x-litellm-call-id"]) + assert _attribute(first_span.attributes, "session.id") == caller_session + + baggage_session: Final = "baggage-" + marker + "-r2" + second: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker + "-r2"}], + "metadata": {"session_id": ""}, + }, + headers=_w3c_headers(uuid.uuid4().hex, baggage_session), + ) + assert second.status_code == 200, second.text + second_span: Final = _await_span(received, destination, second.headers["x-litellm-call-id"]) + assert _attribute(second_span.attributes, "session.id") == baggage_session + + third: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker + "-r3"}]}, + headers=_w3c_headers(uuid.uuid4().hex, None), + ) + assert third.status_code == 400, third.text + assert upstream_targets == ["/v1/responses", "/v1/chat/completions"], upstream_targets From d656a4f18bdc142a913571205e5cd08690249ba4 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 09:27:30 +0000 Subject: [PATCH 07/12] 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): From 6f27b562cff4693cc528969b6106c9fe453c3f9e Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 09:42:56 +0000 Subject: [PATCH 08/12] test(proxy): move test context into assertion messages Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/test_litellm_pre_call_utils.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) 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 722caad23cf..8a10f67dbfa 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -8434,8 +8434,6 @@ async def test_missing_session_id_generate_reuses_traceparent_trace_id(): 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 @@ -8447,10 +8445,15 @@ async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path: 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" + caller_trace_msg: Final = "generate must derive the session from the caller's un-promoted metadata.trace_id" + assert updated["litellm_session_id"] == "caller-trace", caller_trace_msg + assert updated["litellm_metadata"]["session_id"] == "caller-trace", caller_trace_msg + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True, ( + "the derived session id must still be marked generated" + ) + assert _spend_log_session_id(updated, "litellm_metadata") == "caller-trace", ( + "spend log and callback session ids must agree on the caller trace id" + ) @pytest.mark.asyncio From 46eac14a6704867a14d3aefbb1233c5f1e54eefb Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 15:11:58 +0000 Subject: [PATCH 09/12] test(integration): audit cells for w3c fallback precedence Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../observability/test_langfuse_delivery.py | 1247 ++++++++++++++++- 1 file changed, 1243 insertions(+), 4 deletions(-) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 6f71070451d..7b460daef5c 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -1,16 +1,33 @@ +import asyncio import base64 import json +import signal +import threading import time import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass from pathlib import Path from typing import Final +import anthropic +import httpx +import openai +import psutil import pytest import yaml -from integration._support.client import Gateway, eventually, object_value, string_value +from _s3_v2_support import _chat_stream_frames, _responses_stream_frames +from integration._support.client import ( + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) from integration._support.database import read_rows, scratch_database -from integration._support.process import owned_proxy +from integration._support.process import group_members, owned_proxy, owned_proxy_process from integration._support.wire import Reply, Request, Wire, wire_server from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest from opentelemetry.proto.common.v1.common_pb2 import KeyValue @@ -399,7 +416,7 @@ def _responses_result(identity: str) -> Reply: ) -def _trace_body(kind: str, model: str, marker: str, metadata: Mapping[str, str] | None) -> dict[str, JsonValue]: +def _trace_body(kind: str, model: str, marker: str, metadata: Mapping[str, JsonValue] | None) -> dict[str, JsonValue]: metadata_field: Final[dict[str, JsonValue]] = {} if metadata is None else {"metadata": dict(metadata)} if kind == "responses": return {"model": model, "input": marker + "-question", **metadata_field} @@ -655,3 +672,1225 @@ def test_missing_session_id_generate_derives_session_from_caller_trace( 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 + + +_AUDIT_ENDPOINTS: Final = ( + pytest.param("/v1/chat/completions", "chat", "/v1/chat/completions", id="chat"), + pytest.param("/v1/responses", "responses", "/v1/responses", id="responses"), + pytest.param("/v1/messages", "messages", "/v1/responses", id="messages"), +) + + +def _audit_upstream(provider_secret: str, marker: str, targets: list[str]): + def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) + assert request.headers["authorization"] == f"Bearer {provider_secret}" + targets.append(request.target) + body: Final = TypeAdapter(dict[str, JsonValue]).validate_json(request.body) + index: Final = len(targets) + if request.target == "/v1/responses": + if body.get("stream"): + return Reply( + content_type="text/event-stream", chunks=_responses_stream_frames(f"resp-{marker}-{index}") + ) + return _responses_result(f"resp-{marker}-{index}") + assert request.target == "/v1/chat/completions", request.target + if body.get("stream"): + return Reply( + content_type="text/event-stream", chunks=_chat_stream_frames(f"chatcmpl-{marker}-{index}") + ) + return _completion(f"{marker}-{index}-answer") + + return upstream + + +def _audit_sink(): + 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") + + return langfuse + + +@dataclass(frozen=True, slots=True) +class _AuditRig: + candidate: Gateway + scenario: Scenario + destination: Wire + + +@pytest.fixture(scope="module") +def audit_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, directory, _langfuse_environment(destination), config=_langfuse_config(directory) + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_generate_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-generate") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "generate"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_reject_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-reject") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "reject"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_omit_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-omit") + with ( + gateway_from_environment() as gateway, + wire_server(_audit_sink()) as destination, + owned_proxy( + gateway, + directory, + _langfuse_environment(destination), + config=_langfuse_config(directory, {"missing_session_id": "omit"}), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +@pytest.fixture(scope="module") +def audit_otel_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_AuditRig]: + directory: Final = tmp_path_factory.mktemp("audit-otel") + + def otlp(request: Request) -> Reply: + return Reply() + + with ( + gateway_from_environment() as gateway, + wire_server(otlp) as destination, + owned_proxy( + gateway, + directory, + {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"}, + config=_otel_config(directory, destination.url), + ) as candidate, + candidate.scenario() as scenario, + ): + yield _AuditRig(candidate, scenario, destination) + + +def _await_spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id, status, metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (call_id,), + ), + lambda values: len(values) == 1, + seconds=240, + ) + return rows[0] + + +def _assert_call( + response: httpx.Response, + received: list[Request], + destination: Wire, + targets: list[str], + expected_target: str, + expected_trace: str | None, + expected_session: str | None, +) -> None: + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + _assert_span_spend(response, received, destination, expected_trace, expected_session) + + +def _assert_span_spend( + response: httpx.Response, + received: list[Request], + destination: Wire, + expected_trace: str | None, + expected_session: str | None, +) -> None: + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, destination, call_id) + if expected_trace is not None: + assert span.trace_id.hex() == expected_trace, f"call {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == expected_session, f"call {call_id}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == expected_session, f"call {call_id}: spend session {row}" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize("metadata_mode", ("both", "none", "trace", "session")) +def test_audit_caller_metadata_wins_over_w3c_per_field( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata_mode: str +) -> None: + marker: Final = "audit" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = { + "both": {"trace_id": caller_trace, "session_id": caller_session}, + "trace": {"trace_id": caller_trace}, + "session": {"session_id": caller_session}, + "none": None, + }[metadata_mode] + expected_trace: Final = caller_trace if metadata_mode in ("both", "trace") else header_trace + expected_session: Final = caller_session if metadata_mode in ("both", "session") else baggage_session + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, metadata), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call( + response, received, audit_rig.destination, targets, expected_target, expected_trace, expected_session + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_caller_metadata_wins_on_streamed_calls( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditstream" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + { + **_trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + "stream": True, + }, + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_rig.destination, targets, expected_target, caller_trace, caller_session) + + +def test_audit_caller_metadata_wins_through_official_sdk_clients(audit_rig: _AuditRig) -> None: + marker: Final = "auditsdk" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + destination: Final = audit_rig.destination + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + base_url: Final = str(candidate.client.base_url).rstrip("/") + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def drive(client_kind: str) -> tuple[str, str, httpx.Headers]: + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"{client_kind}-session-{marker}" + body_metadata: Final = { + "trace_id": caller_trace, + "session_id": caller_session, + "generation_name": marker + "-" + client_kind, + } + headers: Final = _w3c_headers(uuid.uuid4().hex, "baggage-" + marker + "-" + client_kind) + if client_kind == "chat_openai_sync": + return caller_trace, caller_session, ( + openai.OpenAI(base_url=f"{base_url}/v1", api_key=candidate.key) + .chat.completions.with_raw_response.create( + model=model, + messages=[{"role": "user", "content": marker + "-chat"}], + extra_body={"metadata": body_metadata, "cache": {"no-cache": True}}, + extra_headers=headers, + ) + .headers + ) + if client_kind == "responses_openai_async": + + async def responses_call() -> httpx.Headers: + answer: Final = await openai.AsyncOpenAI( + base_url=f"{base_url}/v1", api_key=candidate.key + ).responses.with_raw_response.create( + model=model, + input=marker + "-responses", + extra_body={"metadata": body_metadata, "cache": {"no-cache": True}}, + extra_headers=headers, + ) + return answer.headers + + return caller_trace, caller_session, asyncio.run(responses_call()) + if client_kind == "messages_anthropic_sync": + return caller_trace, caller_session, ( + anthropic.Anthropic(base_url=base_url, api_key=candidate.key) + .messages.with_raw_response.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker + "-messages"}], + extra_body={"metadata": body_metadata}, + extra_headers=headers, + ) + .headers + ) + + async def messages_call() -> httpx.Headers: + answer: Final = await anthropic.AsyncAnthropic( + base_url=base_url, api_key=candidate.key + ).messages.with_raw_response.create( + model=model, + max_tokens=16, + messages=[{"role": "user", "content": marker + "-messages-async"}], + extra_body={"metadata": body_metadata}, + extra_headers=headers, + ) + return answer.headers + + return caller_trace, caller_session, asyncio.run(messages_call()) + + expected: Final = { + call_headers["x-litellm-call-id"]: (caller_trace, caller_session, client_kind) + for client_kind, (caller_trace, caller_session, call_headers) in ( + (kind, drive(kind)) + for kind in ( + "chat_openai_sync", + "responses_openai_async", + "messages_anthropic_sync", + "messages_anthropic_async", + ) + ) + } + + def spans_named() -> tuple[Span, ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") in expected + ) + + spans: Final = eventually(spans_named, lambda values: len(values) == len(expected), seconds=60) + for span in spans: + call_id: Final = str( + _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + ) + caller_trace, caller_session, client_kind = expected[call_id] + assert span.trace_id.hex() == caller_trace, f"{client_kind}: trace id" + assert _attribute(span.attributes, "session.id") == caller_session, f"{client_kind}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == caller_session, f"{client_kind} {call_id}: spend session {row}" + + +def _otel_config(tmp_path: Path, sink_url: str) -> Path: + config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text())) + settings: Final = {**_SETTINGS.validate_python(config["litellm_settings"]), "callbacks": ["otel"]} + path: Final = tmp_path / "otel.yaml" + path.write_text( + yaml.safe_dump( + { + **config, + "litellm_settings": settings, + "callback_settings": { + "otel": {"exporter": "http/json", "endpoint": sink_url, "mapper_names": ["genai"]} + }, + } + ) + ) + return path + + +_OTEL_SPAN: Final = TypeAdapter(dict[str, JsonValue]) + + +def _otel_spans(batches: Sequence[Request]) -> tuple[dict[str, JsonValue], ...]: + spans: list[dict[str, JsonValue]] = [] # mutable-ok: flattens nested OTLP batches into a tuple + for batch in batches: + if not batch.target.endswith("/v1/traces"): + continue + payload: Final = TypeAdapter(JsonValue).validate_json(batch.body) + envelopes: Final = payload if isinstance(payload, list) else [payload] + for envelope in envelopes: + for resource in TypeAdapter(list[JsonValue]).validate_python( + object_value(envelope)["resourceSpans"] + ): + for scope in object_value(resource)["scopeSpans"]: + spans.extend(TypeAdapter(list[JsonValue]).validate_python(object_value(scope)["spans"])) + return tuple(_OTEL_SPAN.validate_python(span) for span in spans) + + +def _otel_attribute(span: Mapping[str, JsonValue], key: str) -> str | None: + for attribute in TypeAdapter(list[JsonValue]).validate_python(span.get("attributes") or []): + entry: Final = object_value(attribute) + if entry["key"] == key: + value: Final = object_value(entry["value"]) + raw: Final = value.get("stringValue") + return str(raw) if raw is not None else None + return None + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_otel_span_carries_caller_ids( + audit_otel_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditotel" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + destination: Final = audit_otel_rig.destination + model: Final = audit_otel_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_otel_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + assert targets == [expected_target], targets + response_id: Final = string_value(object_value(response.json())["id"]) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + + def exported() -> tuple[dict[str, JsonValue], ...]: + received.extend(destination.drain()) + return tuple( + span + for span in _otel_spans(received) + if _otel_attribute(span, "gen_ai.response.id") == response_id + ) + + spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=60) + span: Final = spans[0] + print( + f"H7 record: otel span trace={span['traceId']} header={header_trace} " + f"session.id={_otel_attribute(span, 'session.id')} " + f"gen_ai.conversation.id={_otel_attribute(span, 'gen_ai.conversation.id')}" + ) + assert span["traceId"] == header_trace, ( + f"otel span for {response_id}: trace id must be the ambient W3C header trace" + ) + assert _otel_attribute(span, "gen_ai.conversation.id") == caller_session, ( + f"otel span conversation id for {response_id}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + ("case", "metadata_mode", "headers_mode", "expected_session"), + ( + pytest.param("g1", "trace", "traceparent", "caller_trace", id="g1_caller_trace_and_header"), + pytest.param("g3", "none", "traceparent", "header_trace", id="g3_header_only"), + pytest.param("g5", "session", "traceparent", "caller_session", id="g5_caller_session"), + ), +) +def test_audit_generate_policy_derives_session_from_caller_trace( + audit_generate_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + case: str, + metadata_mode: str, + headers_mode: str, + expected_session: str, +) -> None: + marker: Final = "auditgen" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = {"trace": {"trace_id": caller_trace}, "session": {"session_id": caller_session}, "none": None}[ + metadata_mode + ] + expected: Final = {"caller_trace": caller_trace, "header_trace": header_trace, "caller_session": caller_session}[ + expected_session + ] + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + response: Final = audit_generate_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, metadata), + headers={"traceparent": f"00-{header_trace}-00f067aa0ba902b7-01"}, + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + span: Final = _await_span(received, audit_generate_rig.destination, call_id) + assert _attribute(span.attributes, "session.id") == expected, f"call {call_id}: session.id" + row: Final = _await_spend_row(call_id) + assert row["session_id"] == expected, f"call {call_id}: spend session {row}" + + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_generate_policy_is_stable_across_repeated_caller_trace( + audit_generate_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditgen2" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + responses: Final = tuple( + audit_generate_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, f"{marker}-{attempt}", {"trace_id": caller_trace}), + ) + for attempt in ("first", "second") + ) + sessions: Final[list[str | None]] = [] # mutable-ok: collects the two observed sessions in order + for attempt, response in zip(("first", "second"), responses): + assert response.status_code == 200, f"{attempt}: {response.text}" + row: Final = _await_spend_row(response.headers["x-litellm-call-id"]) + sessions.append(string_value(row["session_id"])) + assert sessions == [caller_trace, caller_trace], ( + f"generate must derive both sessions from the caller trace {caller_trace}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_generate_policy_fresh_session_without_any_ids( + audit_generate_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditgen4" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_generate_rig.scenario.model( + api_base=provider.url + "/v1", api_key=provider_secret + ) + sessions: Final[list[str | None]] = [] # mutable-ok: collects the two observed sessions in order + for attempt in ("first", "second"): + response: Final = audit_generate_rig.candidate.request( + "POST", endpoint, _trace_body(kind, model, f"{marker}-{attempt}", None) + ) + assert response.status_code == 200, response.text + session: Final = _await_spend_row(response.headers["x-litellm-call-id"])["session_id"] + assert session, f"{attempt}: generated session id must be non-empty" + sessions.append(string_value(session)) + assert sessions[0] != sessions[1], f"two id-less calls must not share a session: {sessions}" + assert targets and set(targets) == {expected_target}, targets + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + ("case", "metadata", "baggage", "expected_status", "expected_session"), + ( + pytest.param( + "r1", "caller_session", None, 200, "caller_session", id="r1_caller_session_no_baggage" + ), + pytest.param("r2", "empty_session", "baggage", 200, "baggage", id="r2_empty_session_baggage"), + pytest.param("r3", "none", None, 400, None, id="r3_nothing_rejected"), + ), +) +def test_audit_reject_policy( + audit_reject_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + case: str, + metadata: str, + baggage: str, + expected_status: int, + expected_session: str | None, +) -> None: + marker: Final = "auditreject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + baggage_session: Final = "baggage-" + marker + caller_session: Final = f"my-session-id-{marker}" + body_metadata: Final = { + "caller_session": {"session_id": caller_session}, + "empty_session": {"session_id": ""}, + "none": None, + }[metadata] + expected: Final = {"caller_session": caller_session, "baggage": baggage_session}[expected_session] if expected_session else None + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_reject_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_reject_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, body_metadata), + headers=_w3c_headers(uuid.uuid4().hex, baggage_session if baggage else None), + ) + assert response.status_code == expected_status, response.text + if expected_status != 200: + assert targets == [], targets + rejected_rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (response.headers["x-litellm-call-id"],), + ), + lambda values: len(values) == 1, + seconds=240, + ) + assert rejected_rows[0]["status"] == "failure", ( + f"rejected call {response.headers['x-litellm-call-id']} must write exactly one failure spend row: {rejected_rows}" + ) + return + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_reject_rig.destination, targets, expected_target, None, expected) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_omit_policy_records_no_session(audit_omit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditomit" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_omit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_omit_rig.candidate.request( + "POST", endpoint, _trace_body(kind, model, marker, None) + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_omit_rig.destination, call_id) + row: Final = _await_spend_row(call_id) + span_session: Final = _attribute(span.attributes, "session.id") + assert row["session_id"] == span_session, ( + f"call {call_id}: spend session {row['session_id']!r} must match span session {span_session!r}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize("bad_value", (123, ["x"]), ids=["int", "list"]) +def test_audit_non_string_caller_ids_are_ignored_consistently( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, bad_value: JsonValue +) -> None: + marker: Final = "auditbad" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + destination: Final = audit_rig.destination + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + outcomes: Final[list[tuple[str | None, str | None, str | None]]] = ( + [] + ) # mutable-ok: collects (span trace, span session, spend session) per leg + for attempt, headers in ( + ("with_headers", _w3c_headers(uuid.uuid4().hex, "baggage-" + marker)), + ("no_headers", {}), + ): + response: Final = candidate.request( + "POST", + endpoint, + _trace_body( + kind, + model, + f"{marker}-{attempt}", + {"trace_id": bad_value, "session_id": bad_value}, + ), + headers=headers, + ) + assert response.status_code == 200, f"{attempt}: {response.text}" + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, destination, call_id) + row: Final = _await_spend_row(call_id) + outcomes.append( + ( + span.trace_id.hex(), + str(_attribute(span.attributes, "session.id")), + str(row["session_id"]), + ) + ) + assert outcomes[0] == outcomes[1], ( + f"W3C headers must not change the outcome for caller {bad_value!r}: {outcomes}" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +@pytest.mark.parametrize( + "metadata", + ({"trace_id": "", "session_id": ""}, {"trace_id": None, "session_id": None}), + ids=["empty", "null"], +) +def test_audit_empty_and_null_caller_ids_fall_back_to_w3c( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata: dict[str, object] +) -> None: + marker: Final = "auditempty" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, dict(metadata)), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + _assert_call(response, received, audit_rig.destination, targets, expected_target, header_trace, baggage_session) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_five_kilobyte_caller_ids_win_verbatim( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditbig" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = "T" * 5120 + caller_session: Final = "S" * 5120 + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, call_id) + assert span.trace_id.hex() != header_trace, f"call {call_id}: caller trace must beat the W3C header" + assert _attribute(span.attributes, "session.id") == caller_session, f"call {call_id}: session.id" + assert _await_spend_row(call_id)["session_id"] == caller_session, f"call {call_id}: spend session" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS) +def test_audit_identical_requests_twice_log_per_call( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str +) -> None: + marker: Final = "auditdup" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for attempt in ("first", "second"): + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body( + kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session} + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + assert response.status_code == 200, f"{attempt}: {response.text}" + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, audit_rig.destination, call_id) + assert span.trace_id.hex() == caller_trace, f"{attempt} {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == caller_session, f"{attempt} {call_id}: session.id" + assert _await_spend_row(call_id)["session_id"] == caller_session, f"{attempt} {call_id}: spend" + assert targets and set(targets) == {expected_target}, targets + + +def test_audit_malformed_w3c_headers_are_ignored(audit_rig: _AuditRig) -> None: + marker: Final = "auditmal" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + "/v1/chat/completions", + _trace_body("chat", model, marker, None), + headers={"traceparent": "00-zz-00f067aa0ba902b7-01", "baggage": "not-a-session-key"}, + ) + assert response.status_code == 200, response.text + assert targets == ["/v1/chat/completions"], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, response.headers["x-litellm-call-id"]) + assert len(span.trace_id.hex()) == 32 and "zz" not in span.trace_id.hex(), span.trace_id.hex() + row: Final = _await_spend_row(response.headers["x-litellm-call-id"]) + span_session: Final = _attribute(span.attributes, "session.id") + print(f"S7 record: span session.id={span_session!r} spend session={row['session_id']!r}") + assert span_session in (None, row["session_id"]), ( + f"span session {span_session!r} diverges from spend session {row['session_id']!r}" + ) + assert row["session_id"], row + + +def test_audit_unauthenticated_call_leaves_no_spend_row(audit_rig: _AuditRig) -> None: + marker: Final = "auditunauth" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + denied: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + marker, + {"trace_id": uuid.uuid4().hex, "session_id": f"my-session-id-{marker}"}, + ), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + key="sk-wrong-key", + ) + assert denied.status_code == 401, denied.text + assert targets == [], targets + control: Final = candidate.request( + "POST", "/v1/chat/completions", _trace_body("chat", model, marker + "-control", None) + ) + assert control.status_code == 200, control.text + _await_spend_row(control.headers["x-litellm-call-id"]) + assert ( + read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id=%s', + (denied.headers.get("x-litellm-call-id") or "",), + ) + == [] + ), "an unauthenticated call must not write a spend row" + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +@pytest.mark.parametrize("upstream_status", (500, 401), ids=["upstream_500", "upstream_401"]) +def test_audit_upstream_error_still_logs_caller_session( + audit_rig: _AuditRig, + endpoint: str, + kind: str, + expected_target: str, + upstream_status: int, +) -> None: + marker: Final = "auditerr" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + def upstream(request: Request) -> Reply: + if request.body: + targets.append(request.target) + return Reply(status=upstream_status, body=b'{"error": {"message": "scripted upstream failure"}}') + + with wire_server(upstream) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, {"trace_id": caller_trace, "session_id": caller_session}), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + ) + assert response.status_code == upstream_status, response.text + assert targets and set(targets) == {expected_target}, targets + call_id: Final = response.headers["x-litellm-call-id"] + row: Final = _await_spend_row(call_id) + assert row["session_id"] == caller_session, f"call {call_id}: failure spend session {row}" + + +def test_audit_sink_rejection_does_not_break_the_caller(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditreject" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + attempts: Final[list[int]] = [] # mutable-ok: sink status sequence counter + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + if request.target == TRACES_PATH: + attempts.append(1) + if len(attempts) == 1: + return Reply(status=403) + if len(attempts) == 2: + return Reply(status=404) + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for attempt in ("first", "second", "third"): + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + f"{marker}-{attempt}", + {"trace_id": uuid.uuid4().hex, "session_id": f"my-session-id-{marker}-{attempt}"}, + ), + ) + assert response.status_code == 200, f"{attempt}: {response.text}" + _await_spend_row(response.headers["x-litellm-call-id"]) + assert targets == ["/v1/chat/completions"] * 3, targets + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_string_metadata_body_does_not_crash(audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditstr" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + body: Final = {**_trace_body(kind, model, marker, None), "metadata": "x"} + response: Final = candidate.request( + "POST", endpoint, body, headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker) + ) + assert response.status_code < 500, f"metadata string must not crash the proxy: {response.status_code} {response.text}" + follow_up: Final = candidate.request( + "POST", "/v1/chat/completions", _trace_body("chat", model, marker + "-follow", None) + ) + assert follow_up.status_code == 200, follow_up.text + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +def test_audit_key_metadata_session_still_honoured(audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str) -> None: + marker: Final = "auditkey" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + key_session: Final = f"key-session-{marker}" + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + token: Final = audit_rig.scenario.key(metadata={"session_id": key_session}) + response: Final = audit_rig.candidate.request( + "POST", + endpoint, + _trace_body(kind, model, marker, None), + headers=_w3c_headers(uuid.uuid4().hex, "baggage-" + marker), + key=token, + ) + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + span: Final = _await_span(received, audit_rig.destination, call_id) + row: Final = _await_spend_row(call_id) + assert _attribute(span.attributes, "session.id") == row["session_id"], ( + f"call {call_id}: spend session {row['session_id']!r} must match span session" + ) + + +@pytest.mark.parametrize(("endpoint", "kind", "expected_target"), _AUDIT_ENDPOINTS[:2]) +@pytest.mark.parametrize("metadata_mode", ("both", "none"), ids=["caller_ids", "no_metadata"]) +def test_audit_cache_hit_call_keeps_winning_ids( + audit_rig: _AuditRig, endpoint: str, kind: str, expected_target: str, metadata_mode: str +) -> None: + marker: Final = "auditcache" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}" + metadata: Final = ( + {"trace_id": caller_trace, "session_id": caller_session} if metadata_mode == "both" else None + ) + expected_trace: Final = caller_trace if metadata_mode == "both" else header_trace + expected_session: Final = caller_session if metadata_mode == "both" else baggage_session + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + body: Final = {key: value for key, value in _trace_body(kind, model, marker, metadata).items() if key != "cache"} + headers: Final = _w3c_headers(header_trace, baggage_session) + first: Final = candidate.request("POST", endpoint, body, headers=headers) + assert first.status_code == 200, first.text + second: Final = candidate.request("POST", endpoint, body, headers=headers) + assert second.status_code == 200, second.text + assert second.headers.get("x-litellm-cache-key"), ( + f"second identical call must be a cache hit: {dict(second.headers)}" + ) + call_id: Final = second.headers["x-litellm-call-id"] + assert targets == [expected_target], targets + span: Final = _await_span(received, audit_rig.destination, call_id) + if metadata_mode == "both": + assert span.trace_id.hex() == expected_trace, f"cache hit {call_id}: trace id" + assert _attribute(span.attributes, "session.id") == expected_session, ( + f"cache hit {call_id}: session.id" + ) + else: + span_session: Final = _attribute(span.attributes, "session.id") + print( + f"E2 record: cache-hit span trace={span.trace_id.hex()} session={span_session!r} " + f"header trace={header_trace} baggage session={baggage_session!r}" + ) + assert len(span.trace_id.hex()) == 32, span.trace_id.hex() + assert _await_spend_row(call_id)["session_id"] == expected_session, f"cache hit {call_id}: spend" + + +def test_audit_concurrent_requests_each_keep_their_caller_ids(audit_rig: _AuditRig) -> None: + marker: Final = "auditconc" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with wire_server(_audit_upstream(provider_secret, marker, targets)) as provider: + candidate: Final = audit_rig.candidate + model: Final = audit_rig.scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + jobs: Final = tuple( + (index, endpoint, kind, uuid.uuid4().hex, f"my-session-id-{marker}-{index}") + for index, (endpoint, kind, _) in tuple( + enumerate(tuple(row.values for row in _AUDIT_ENDPOINTS) * 4) + )[:10] + ) + + def fire(job: tuple[int, str, str, str, str]) -> tuple[str, str, httpx.Response]: + index, endpoint, kind, caller_trace, caller_session = job + response: Final = candidate.request( + "POST", + endpoint, + _trace_body( + kind, model, f"{marker}-{index}", {"trace_id": caller_trace, "session_id": caller_session} + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_trace, caller_session, response + + with ThreadPoolExecutor(max_workers=10) as pool: + answered: Final = tuple(pool.map(fire, jobs)) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + assert len(targets) == 10, targets + for caller_trace, caller_session, response in answered: + assert response.status_code == 200, response.text + _assert_span_spend(response, received, audit_rig.destination, caller_trace, caller_session) + + +def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditoutage" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + outage: Final = threading.Event() + + def langfuse(request: Request) -> Reply: + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() + if outage.is_set(): + return Reply(status=503) + return Reply(body=b"", content_type="application/x-protobuf") + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(langfuse) as destination, + owned_proxy( + gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path) + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + jobs: Final = tuple( + (index, endpoint, kind, uuid.uuid4().hex, f"my-session-id-{marker}-{index}") + for index, (endpoint, kind, _) in enumerate(tuple(row.values for row in _AUDIT_ENDPOINTS) * 10) + ) + + def fire(job: tuple[int, str, str, str, str]) -> tuple[str, str, httpx.Response]: + index, endpoint, kind, caller_trace, caller_session = job + stream: Final = index % 3 == 0 + response: Final = candidate.request( + "POST", + endpoint, + { + **_trace_body( + kind, + model, + f"{marker}-{index}", + {"trace_id": caller_trace, "session_id": caller_session}, + ), + "stream": stream, + }, + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_trace, caller_session, response + + outage.set() + with ThreadPoolExecutor(max_workers=30) as pool: + first_wave: Final = tuple(pool.map(fire, jobs[:10])) + outage.clear() + with ThreadPoolExecutor(max_workers=30) as pool: + answered: Final = first_wave + tuple(pool.map(fire, jobs[10:])) + for _, _, response in answered: + assert response.status_code == 200, response.text + expected_by_call: Final = { + response.headers["x-litellm-call-id"]: caller_session + for caller_session, response in ((session, res) for _, session, res in answered) + } + call_ids: Final = sorted(expected_by_call) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + delivered: Final = eventually( + lambda: ( + received.extend(destination.drain()) or tuple( + span + for span in _spans(received) + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + in expected_by_call + ) + ), + lambda spans: len(spans) >= len(answered), + seconds=45, + return_last_on_timeout=True, + ) + delivered_calls: Final = frozenset( + str(_attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id")) + for span in delivered + ) + lost: Final = frozenset(call_ids) - delivered_calls + print(f"sink outage lost {len(lost)} of {len(answered)} spans") + for span in delivered: + span_call: Final = str( + _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") + ) + assert _attribute(span.attributes, "session.id") == expected_by_call[span_call], ( + f"call {span_call}: session.id" + ) + spend_rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + (call_ids,), + ), + lambda values: len(values) == len(answered), + seconds=150, + return_last_on_timeout=True, + ) + print(f"C1 record: {len(spend_rows)} of {len(answered)} spend rows written") + for row in spend_rows: + assert row["session_id"] in frozenset(expected_by_call.values()), row + + +def test_audit_surviving_worker_keeps_serving_after_kill(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = "auditworker" + uuid.uuid4().hex + provider_secret: Final = "synthetic-provider-secret-" + marker + header_trace: Final = uuid.uuid4().hex + baggage_session: Final = "baggage-" + marker + targets: Final[list[str]] = [] # mutable-ok: the upstream records every hit + + with ( + wire_server(_audit_upstream(provider_secret, marker, targets)) as provider, + wire_server(_audit_sink()) as destination, + owned_proxy_process( + gateway, + tmp_path, + _langfuse_environment(destination), + config=_langfuse_config(tmp_path), + workers=2, + ) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret) + workers: Final = tuple( + member for member in group_members(owned.process.pid) if member.pid != owned.process.pid + ) + assert workers, "expected worker processes under the owned proxy" + workers[0].send_signal(signal.SIGKILL) + psutil.wait_procs([workers[0]], timeout=5) + + def fire(index: int) -> tuple[str, httpx.Response]: + caller_trace: Final = uuid.uuid4().hex + caller_session: Final = f"my-session-id-{marker}-{index}" + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + _trace_body( + "chat", + model, + f"{marker}-{index}", + {"trace_id": caller_trace, "session_id": caller_session}, + ), + headers=_w3c_headers(header_trace, baggage_session), + ) + return caller_session, response + + with ThreadPoolExecutor(max_workers=10) as pool: + answered: Final = tuple(pool.map(fire, range(10))) + received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones + for caller_session, response in answered: + assert response.status_code == 200, response.text + call_id: Final = response.headers["x-litellm-call-id"] + span: Final = _await_span(received, destination, call_id) + assert _attribute(span.attributes, "session.id") == caller_session, f"{call_id}: session.id" + expected_by_call: Final = { + response.headers["x-litellm-call-id"]: caller_session + for caller_session, response in answered + } + spend_rows: Final = eventually( + lambda: read_rows( + 'SELECT session_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + (list(expected_by_call),), + ), + lambda values: len(values) == len(answered), + seconds=150, + return_last_on_timeout=True, + ) + print(f"C2 record: {len(spend_rows)} of {len(answered)} spend rows written") + assert spend_rows, "no spend rows survived the worker kill" + for row in spend_rows: + assert row["session_id"] == expected_by_call[string_value(row["litellm_call_id"])], row From fb94af87bf770c0863a15cbc2affdb95c3803d53 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 16:00:42 +0000 Subject: [PATCH 10/12] test(integration): tighten audit chaos cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../observability/test_langfuse_delivery.py | 61 +++++++++++++------ 1 file changed, 41 insertions(+), 20 deletions(-) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 7b460daef5c..93399a173bc 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -67,7 +67,20 @@ def _completion(text: str) -> Reply: def _projects() -> Reply: - return Reply(body=json.dumps({"data": [{"id": "integration-project", "name": "integration"}]}).encode()) + return Reply( + body=json.dumps( + { + "data": [ + { + "id": "integration-project", + "name": "integration", + "organization": {"id": "integration-org", "name": "integration"}, + "metadata": {}, + } + ] + } + ).encode() + ) def _text_prompt(name: str) -> Reply: @@ -828,7 +841,6 @@ def _assert_call( expected_trace: str | None, expected_session: str | None, ) -> None: - call_id: Final = response.headers["x-litellm-call-id"] assert targets == [expected_target], targets _assert_span_spend(response, received, destination, expected_trace, expected_session) @@ -1467,13 +1479,15 @@ def test_audit_malformed_w3c_headers_are_ignored(audit_rig: _AuditRig) -> None: assert len(span.trace_id.hex()) == 32 and "zz" not in span.trace_id.hex(), span.trace_id.hex() row: Final = _await_spend_row(response.headers["x-litellm-call-id"]) span_session: Final = _attribute(span.attributes, "session.id") - print(f"S7 record: span session.id={span_session!r} spend session={row['session_id']!r}") assert span_session in (None, row["session_id"]), ( f"span session {span_session!r} diverges from spend session {row['session_id']!r}" ) assert row["session_id"], row + + + def test_audit_unauthenticated_call_leaves_no_spend_row(audit_rig: _AuditRig) -> None: marker: Final = "auditunauth" + uuid.uuid4().hex provider_secret: Final = "synthetic-provider-secret-" + marker @@ -1732,10 +1746,10 @@ def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_p outage: Final = threading.Event() def langfuse(request: Request) -> Reply: - if request.method == "GET" and request.target.startswith(PROJECTS_PATH): - return _projects() if outage.is_set(): return Reply(status=503) + if request.method == "GET" and request.target.startswith(PROJECTS_PATH): + return _projects() return Reply(body=b"", content_type="application/x-protobuf") with ( @@ -1774,7 +1788,13 @@ def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_p outage.set() with ThreadPoolExecutor(max_workers=30) as pool: first_wave: Final = tuple(pool.map(fire, jobs[:10])) + unhealthy: Final = candidate.request("GET", "/health/services?service=langfuse") + assert unhealthy.status_code != 200 or "unhealthy" in unhealthy.text, ( + f"langfuse must report unhealthy while the sink 503s: {unhealthy.status_code} {unhealthy.text}" + ) outage.clear() + healthy: Final = candidate.request("GET", "/health/services?service=langfuse") + assert healthy.status_code == 200, healthy.text with ThreadPoolExecutor(max_workers=30) as pool: answered: Final = first_wave + tuple(pool.map(fire, jobs[10:])) for _, _, response in answered: @@ -1795,15 +1815,19 @@ def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_p ) ), lambda spans: len(spans) >= len(answered), - seconds=45, - return_last_on_timeout=True, + seconds=60, ) - delivered_calls: Final = frozenset( - str(_attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id")) - for span in delivered + spans_per_call: Final = { + call_id: sum( + 1 + for span in delivered + if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") == call_id + ) + for call_id in call_ids + } + assert sorted(spans_per_call.values()) == [1] * len(answered), ( + f"each call id must arrive on exactly one span: {spans_per_call}" ) - lost: Final = frozenset(call_ids) - delivered_calls - print(f"sink outage lost {len(lost)} of {len(answered)} spans") for span in delivered: span_call: Final = str( _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") @@ -1813,16 +1837,15 @@ def test_audit_sink_outage_mid_burst_loses_no_spend_rows(gateway: Gateway, tmp_p ) spend_rows: Final = eventually( lambda: read_rows( - 'SELECT session_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + 'SELECT session_id, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', (call_ids,), ), lambda values: len(values) == len(answered), seconds=150, - return_last_on_timeout=True, ) - print(f"C1 record: {len(spend_rows)} of {len(answered)} spend rows written") for row in spend_rows: - assert row["session_id"] in frozenset(expected_by_call.values()), row + row_call: Final = string_value(row["litellm_call_id"]) + assert row["session_id"] == expected_by_call[row_call], f"call {row_call}: spend {row}" def test_audit_surviving_worker_keeps_serving_after_kill(gateway: Gateway, tmp_path: Path) -> None: @@ -1888,9 +1911,7 @@ def test_audit_surviving_worker_keeps_serving_after_kill(gateway: Gateway, tmp_p ), lambda values: len(values) == len(answered), seconds=150, - return_last_on_timeout=True, ) - print(f"C2 record: {len(spend_rows)} of {len(answered)} spend rows written") - assert spend_rows, "no spend rows survived the worker kill" for row in spend_rows: - assert row["session_id"] == expected_by_call[string_value(row["litellm_call_id"])], row + row_call: Final = string_value(row["litellm_call_id"]) + assert row["session_id"] == expected_by_call[row_call], f"call {row_call}: spend {row}" From b481774f2373bec9b85cf754e068bc793cda1038 Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 16:03:06 +0000 Subject: [PATCH 11/12] style(test): drop stray blank lines Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/observability/test_langfuse_delivery.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 93399a173bc..0a27561e731 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -1485,9 +1485,6 @@ def test_audit_malformed_w3c_headers_are_ignored(audit_rig: _AuditRig) -> None: assert row["session_id"], row - - - def test_audit_unauthenticated_call_leaves_no_spend_row(audit_rig: _AuditRig) -> None: marker: Final = "auditunauth" + uuid.uuid4().hex provider_secret: Final = "synthetic-provider-secret-" + marker From 4d3940cdb62c455e4024709191ec76422b84de0d Mon Sep 17 00:00:00 2001 From: yucheng Date: Wed, 30 Sep 2026 01:47:34 +0000 Subject: [PATCH 12/12] test(integration): ignore body-less model-info probes in langfuse precedence upstreams Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/observability/test_langfuse_delivery.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 0a27561e731..d2709e833ca 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -528,6 +528,8 @@ def test_langfuse_trace_and_session_prefer_caller_metadata_over_w3c_headers( upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) assert request.headers["authorization"] == f"Bearer {provider_secret}" upstream_targets.append(request.target) if request.target == "/v1/responses": @@ -571,6 +573,8 @@ def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback( upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) assert request.headers["authorization"] == f"Bearer {provider_secret}" upstream_targets.append(request.target) if request.target == "/v1/responses": @@ -650,6 +654,8 @@ def test_missing_session_id_generate_derives_session_from_caller_trace( upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit def upstream(request: Request) -> Reply: + if not request.body: + return Reply(status=404) assert request.headers["authorization"] == f"Bearer {provider_secret}" upstream_targets.append(request.target) if request.target == "/v1/responses":