From 02b8b5cc8036a6a2d507e7f8e31a1757285d9f25 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 7 Oct 2026 19:24:50 -0700 Subject: [PATCH] fix(proxy): traceparent/baggage fallback must not override caller metadata (#43688) * 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. * 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 * 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 * 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 * 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 * test(proxy): cover caller metadata precedence over W3C trace headers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): generated session uses promoted caller trace id Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): move test context into assertion messages Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit cells for w3c fallback precedence Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): tighten audit chaos cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(test): drop stray blank lines Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * 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> * fix(proxy): keep root trace/session fields on equal ids and usable caller values The _caller_trace_field gating added for caller-metadata precedence was presence-aware, which regressed three pre-call behaviors versus main: - Equal ids in W3C headers and body metadata suppressed the traceparent/baggage branch entirely, leaving the root litellm_trace_id/litellm_session_id unset so router fallbacks minted a fresh uuid4 per request. The header branch now also fires when the caller value equals the header-derived id, so equal ids stamp the root fields exactly like main. - A truthy but non-string metadata value (e.g. session_id 4815162342) counted as caller-supplied and suppressed the baggage fallback, so downstream str-only consumers (code interpreter sandbox reuse) got a new sandbox per turn. _caller_trace_field now counts only non-empty string values, and the generate policy falls through to generation when the caller value is not usable. - On litellm_metadata routes a usable caller session id suppressed the missing_session_id policies while never landing on the root field, so the root session stayed unset. The policy now promotes the caller's usable session id to litellm_session_id instead of leaving it stranded. * test(integration): non-string caller ids fall back to W3C like empty ids The audit cell compared a headers leg against a headerless leg for equality, which only held while a truthy non-string id suppressed the W3C fallback on both legs. With unusable values ignored again, the headers leg resolves to the header ids while the headerless leg cannot, so assert the headers leg's concrete outcome instead, mirroring the empty/null cell. * style(proxy): trim the session promotion comment to the non-obvious why --------- Co-authored-by: Filipe Andujar Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/litellm_pre_call_utils.py | 67 +- .../observability/test_langfuse_delivery.py | 1541 ++++++++++++++++- .../unit/proxy/test_litellm_pre_call_utils.py | 264 +++ 3 files changed, 1857 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 03a93016177..46cdf328af6 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -154,6 +154,27 @@ def _session_id_from_baggage(baggage: str) -> str | None: return None +def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> str | None: + """The caller's value for a trace-control field, counted only when it is a + usable id: a non-empty string. An explicitly empty/unusable value on the + active metadata container still shadows the promoted requester value, but + neither ever counts as "the caller supplied this field" on its own, so a + numeric session id or an empty string cannot suppress the W3C header + fallback or satisfy a missing-session-id policy.""" + active: Final = data.get(metadata_variable_name) + if isinstance(active, Mapping) and field in active: + active_map: Final = cast(Mapping[str, object], active) # cast-ok: isinstance above, free-form JSON values + active_value: Final = active_map[field] + return active_value if isinstance(active_value, str) and active_value else None + promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS + requester: Final = data.get("metadata") + if not promoted or not isinstance(requester, Mapping): + return None + requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values + requester_value: Final = requester_map.get(field) + return requester_value if isinstance(requester_value, str) and requester_value else None + + def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None: """Only proxy-validated keys are stamped, proven by the unforgeable via_virtual_key marker AND a known non-secret shape: the sha256 hex digest @@ -839,11 +860,22 @@ def apply_missing_session_id_policy( ): metadata["session_id"] = body_session_id return - if data.get("litellm_session_id") or metadata.get("session_id"): + caller_session_id: Final = _caller_trace_field(data, _metadata_variable_name, "session_id") + if caller_session_id is not None: + # Consumers that read the root field (router fallbacks, spend logs, + # sandbox reuse) otherwise see no session and mint a uuid4 per request. + if not data.get("litellm_session_id"): + data["litellm_session_id"] = caller_session_id # rebind-ok: data is an out-param + return + if data.get("litellm_session_id"): return match policy: case "generate": - 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 @@ -1596,16 +1628,28 @@ 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 - # (from real OTel instrumentation) correlate with litellm's own logs - # instead of generating an unrelated trace_id. + # explicit litellm headers, the Anthropic-metadata path, nor the + # caller's own request metadata set the field to a DIFFERENT usable id + # - but lets a caller's existing traceparent/baggage headers (from + # real OTel instrumentation) correlate with litellm's own logs instead + # of generating an unrelated trace_id. normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)}) if "litellm_trace_id" not in data: traceparent: Final = normalized_headers.get("traceparent") if isinstance(traceparent, str): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) - if trace_id_from_traceparent: + # The caller's metadata wins over the header fallback unless + # both carry the same id: stamping the root field then claims + # nothing the caller did not already ask for, and keeps the + # W3C-correlated root trace id instead of a generated uuid4. + caller_trace_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "trace_id", + ) + if trace_id_from_traceparent and ( + caller_trace_id is None or caller_trace_id == trace_id_from_traceparent + ): metadata_from_headers["trace_id"] = trace_id_from_traceparent data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param verbose_proxy_logger.debug( @@ -1615,7 +1659,14 @@ class LiteLLMProxyRequestSetup: baggage: Final = normalized_headers.get("baggage") if isinstance(baggage, str): session_id_from_baggage: Final = _session_id_from_baggage(baggage) - if session_id_from_baggage: + caller_session_id: Final = _caller_trace_field( + cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object + _metadata_variable_name, + "session_id", + ) + if session_id_from_baggage and ( + caller_session_id is None or caller_session_id == session_id_from_baggage + ): metadata_from_headers["session_id"] = session_id_from_baggage data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param verbose_proxy_logger.debug("Extracted session_id from W3C baggage header") diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 13be5a95887..00eee082c33 100644 --- a/tests/integration/observability/test_langfuse_delivery.py +++ b/tests/integration/observability/test_langfuse_delivery.py @@ -1,15 +1,33 @@ +import asyncio import base64 import json +import signal +import threading import time import uuid -from collections.abc import 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 @@ -49,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: @@ -68,15 +99,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 +403,1494 @@ 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, 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} + 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: + 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": + 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: + 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": + 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 + + +@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: + 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": + 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 + + +_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: + 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_fall_back_to_w3c( + 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 + 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, {"trace_id": bad_value, "session_id": bad_value}), + 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) +@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") + 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 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 ( + 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])) + 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: + 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=60, + ) + 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}" + ) + 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, litellm_call_id FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(%s)', + (call_ids,), + ), + lambda values: len(values) == len(answered), + seconds=150, + ) + for row in spend_rows: + 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: + 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, + ) + for row in spend_rows: + row_call: Final = string_value(row["litellm_call_id"]) + assert row["session_id"] == expected_by_call[row_call], f"call {row_call}: spend {row}" diff --git a/tests/unit/proxy/test_litellm_pre_call_utils.py b/tests/unit/proxy/test_litellm_pre_call_utils.py index 86492537fd3..c6c473ae71e 100644 --- a/tests/unit/proxy/test_litellm_pre_call_utils.py +++ b/tests/unit/proxy/test_litellm_pre_call_utils.py @@ -3674,6 +3674,106 @@ 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(): + 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(): + 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(): + 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(): + 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 + + +@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" + + +@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 field not in data["litellm_metadata"] + assert f"litellm_{field}" not in data + + +def test_add_litellm_metadata_from_request_headers_equal_ids_still_stamp_root_fields(): + headers = { + "traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01", + "baggage": "session.id=matching-session-42", + } + data = { + "metadata": {"trace_id": "4bf92f3577b34da6a3ce929d0e0e4736", "session_id": "matching-session-42"}, + } + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["litellm_session_id"] == "matching-session-42" + assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736" + assert data["metadata"]["session_id"] == "matching-session-42" + + +@pytest.mark.parametrize("non_string_session_id", [4815162342, True, {"session": "nested"}]) +def test_add_litellm_metadata_from_request_headers_non_string_body_session_id_falls_back_to_baggage( + non_string_session_id: object, +): + headers = {"baggage": "session.id=header-session-42"} + data = {"metadata": {"session_id": non_string_session_id}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "header-session-42" + assert data["metadata"]["session_id"] == "header-session-42" + + def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False)) @@ -8222,6 +8322,85 @@ 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"}), + ], +) +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("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 +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", @@ -8341,6 +8520,91 @@ 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]): + 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"}, + ) + + 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 +@pytest.mark.parametrize("policy", ["generate", "reject"]) +async def test_missing_session_id_policy_promotes_caller_session_to_root_field(policy: str): + """A caller-supplied usable session id satisfies the missing_session_id policies on + litellm_metadata routes (where it is not yet the managed metadata field) and must also + land on the root ``litellm_session_id`` field: consumers that read the root field + (router fallbacks, spend logs, sandbox reuse) otherwise mint a fresh uuid4 per request.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": "caller-session-42"}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": policy}, + ) + + assert updated["litellm_trace_id"] == "root-trace-42" + assert updated["litellm_session_id"] == "caller-session-42" + assert updated["litellm_metadata"]["session_id"] == "caller-session-42" + assert SESSION_ID_GENERATED_METADATA_KEY not in updated["litellm_metadata"] + + +@pytest.mark.asyncio +async def test_missing_session_id_generate_ignores_non_string_caller_session_id(): + """A non-string session id is not a usable session: the generate policy must fall through + to generation instead of letting an unusable value strand the root session field.""" + request = _request_for("/v1/responses") + + updated = await add_litellm_data_to_request( + data={ + "model": "gpt-4o", + "input": "hi", + "litellm_trace_id": "root-trace-42", + "metadata": {"session_id": 4815162342}, + }, + request=request, + user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"), + proxy_config=MagicMock(), + general_settings={"missing_session_id": "generate"}, + ) + + assert updated["litellm_session_id"] == "root-trace-42" + assert updated["litellm_metadata"]["session_id"] == "root-trace-42" + assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True + + @pytest.mark.asyncio @pytest.mark.parametrize("policy", ["generate", "reject"]) async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):