diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index d48451de6b1..ae7320350de 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -150,6 +150,19 @@ 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) -> 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 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 None + requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values + return requester_map.get(field) or 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 @@ -828,11 +841,15 @@ 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_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 @@ -1579,12 +1596,21 @@ 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 - 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 _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): trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent) @@ -1594,7 +1620,15 @@ 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 _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): session_id_from_baggage: Final = _session_id_from_baggage(baggage) diff --git a/tests/integration/observability/test_langfuse_delivery.py b/tests/integration/observability/test_langfuse_delivery.py index 13be5a95887..d2709e833ca 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,1518 @@ 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_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") + 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/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 84266325226..8a10f67dbfa 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -3673,6 +3673,76 @@ 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 _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan: return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False)) @@ -8155,6 +8225,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", @@ -8274,6 +8423,39 @@ 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_keeps_client_supplied_session_id(policy: str):