mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
test(proxy): cover caller metadata precedence over W3C trace headers
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7e04a23a36
commit
5baeb9e80b
1 changed files with 239 additions and 4 deletions
|
|
@ -2,10 +2,11 @@ import base64
|
|||
import json
|
||||
import time
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows, scratch_database
|
||||
|
|
@ -68,15 +69,20 @@ def _text_prompt(name: str) -> Reply:
|
|||
)
|
||||
|
||||
|
||||
def _langfuse_config(tmp_path: Path) -> Path:
|
||||
def _langfuse_config(tmp_path: Path, general_settings: Mapping[str, JsonValue] | None = None) -> Path:
|
||||
config: Final = _PROXY_CONFIG.validate_python(yaml.safe_load(STOCK_CONFIG.read_text()))
|
||||
settings: Final = {
|
||||
**_SETTINGS.validate_python(config["litellm_settings"]),
|
||||
"success_callback": ["langfuse"],
|
||||
"failure_callback": ["langfuse"],
|
||||
}
|
||||
path: Final = tmp_path / "langfuse.yaml"
|
||||
path.write_text(yaml.safe_dump({**config, "litellm_settings": settings}))
|
||||
general: Final = {
|
||||
**_SETTINGS.validate_python(config["general_settings"]),
|
||||
**(general_settings or {}),
|
||||
}
|
||||
name: Final = "langfuse.yaml" if general_settings is None else "langfuse-merged.yaml"
|
||||
path: Final = tmp_path / name
|
||||
path.write_text(yaml.safe_dump({**config, "litellm_settings": settings, "general_settings": general}))
|
||||
return path
|
||||
|
||||
|
||||
|
|
@ -367,3 +373,232 @@ def test_prompt_fetch_encodes_the_name_retries_a_5xx_once_and_keeps_langfuse_hea
|
|||
assert leak not in json.dumps(dict(failure.headers))
|
||||
assert "set-cookie" not in failure.headers and "x-upstream-internal" not in failure.headers
|
||||
assert sum(1 for target in seen_prompt_gets if target.startswith(PROMPTS_PATH + missing_prompt)) == 1
|
||||
|
||||
|
||||
def _responses_result(identity: str) -> Reply:
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"id": "msg_" + identity,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"status": "completed",
|
||||
"content": [{"type": "output_text", "text": "integration answer", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _trace_body(kind: str, model: str, marker: str, metadata: Mapping[str, str] | None) -> dict[str, JsonValue]:
|
||||
metadata_field: Final[dict[str, JsonValue]] = {} if metadata is None else {"metadata": dict(metadata)}
|
||||
if kind == "responses":
|
||||
return {"model": model, "input": marker + "-question", **metadata_field}
|
||||
if kind == "messages":
|
||||
return {
|
||||
"model": model,
|
||||
"max_tokens": 16,
|
||||
"messages": [{"role": "user", "content": marker + "-question"}],
|
||||
**metadata_field,
|
||||
}
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": marker + "-question"}],
|
||||
"cache": {"no-cache": True},
|
||||
**metadata_field,
|
||||
}
|
||||
|
||||
|
||||
def _w3c_headers(header_trace: str, baggage_session: str | None) -> dict[str, str]:
|
||||
baggage_field: Final = {} if baggage_session is None else {"baggage": f"session.id={baggage_session}"}
|
||||
return {"traceparent": f"00-{header_trace}-00f067aa0ba902b7-01", **baggage_field}
|
||||
|
||||
|
||||
def _await_span(received: list[Request], destination: Wire, call_id: str) -> Span:
|
||||
def exported() -> tuple[Span, ...]:
|
||||
received.extend(destination.drain())
|
||||
return tuple(
|
||||
span
|
||||
for span in _spans(received)
|
||||
if _attribute(span.attributes, "langfuse.observation.metadata.litellm_call_id") == call_id
|
||||
)
|
||||
|
||||
spans: Final = eventually(exported, lambda values: len(values) == 1, seconds=20)
|
||||
return spans[0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("endpoint", "kind", "metadata_mode", "expected_trace", "expected_session", "expected_target"),
|
||||
(
|
||||
pytest.param(
|
||||
"/v1/chat/completions", "chat", "both", "caller", "caller", "/v1/chat/completions", id="chat_caller_ids"
|
||||
),
|
||||
pytest.param(
|
||||
"/v1/responses", "responses", "both", "caller", "caller", "/v1/responses", id="responses_caller_ids"
|
||||
),
|
||||
pytest.param(
|
||||
"/v1/messages", "messages", "both", "caller", "caller", "/v1/responses", id="messages_caller_ids"
|
||||
),
|
||||
pytest.param(
|
||||
"/v1/chat/completions", "chat", "none", "header", "baggage", "/v1/chat/completions", id="chat_header_ids"
|
||||
),
|
||||
pytest.param(
|
||||
"/v1/chat/completions",
|
||||
"chat",
|
||||
"trace",
|
||||
"caller",
|
||||
"baggage",
|
||||
"/v1/chat/completions",
|
||||
id="chat_caller_trace_header_session",
|
||||
),
|
||||
pytest.param(
|
||||
"/v1/chat/completions",
|
||||
"chat",
|
||||
"empty_session",
|
||||
"header",
|
||||
"baggage",
|
||||
"/v1/chat/completions",
|
||||
id="chat_empty_session_header_ids",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_langfuse_trace_and_session_prefer_caller_metadata_over_w3c_headers(
|
||||
gateway: Gateway,
|
||||
tmp_path: Path,
|
||||
endpoint: str,
|
||||
kind: str,
|
||||
metadata_mode: str,
|
||||
expected_trace: str,
|
||||
expected_session: str,
|
||||
expected_target: str,
|
||||
) -> None:
|
||||
marker: Final = "w3c" + uuid.uuid4().hex
|
||||
provider_secret: Final = "synthetic-provider-secret-" + marker
|
||||
header_trace: Final = uuid.uuid4().hex
|
||||
baggage_session: Final = "baggage-" + marker
|
||||
caller_trace: Final = uuid.uuid4().hex
|
||||
caller_session: Final = f"my-session-id-{marker}"
|
||||
metadata: Final = {
|
||||
"both": {"trace_id": caller_trace, "session_id": caller_session},
|
||||
"trace": {"trace_id": caller_trace},
|
||||
"empty_session": {"session_id": ""},
|
||||
"none": None,
|
||||
}[metadata_mode]
|
||||
expected_trace_value: Final = {"caller": caller_trace, "header": header_trace}[expected_trace]
|
||||
expected_session_value: Final = {"caller": caller_session, "baggage": baggage_session}[expected_session]
|
||||
upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit
|
||||
|
||||
def upstream(request: Request) -> Reply:
|
||||
assert request.headers["authorization"] == f"Bearer {provider_secret}"
|
||||
upstream_targets.append(request.target)
|
||||
if request.target == "/v1/responses":
|
||||
return _responses_result("resp-" + marker)
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return _completion(marker + "-answer")
|
||||
|
||||
def langfuse(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.startswith(PROJECTS_PATH):
|
||||
return _projects()
|
||||
return Reply(body=b"", content_type="application/x-protobuf")
|
||||
|
||||
with (
|
||||
wire_server(upstream) as provider,
|
||||
wire_server(langfuse) as destination,
|
||||
owned_proxy(
|
||||
gateway, tmp_path, _langfuse_environment(destination), config=_langfuse_config(tmp_path)
|
||||
) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret)
|
||||
response: Final = candidate.request(
|
||||
"POST",
|
||||
endpoint,
|
||||
_trace_body(kind, model, marker, metadata),
|
||||
headers=_w3c_headers(header_trace, baggage_session),
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert upstream_targets == [expected_target], upstream_targets
|
||||
received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones
|
||||
span: Final = _await_span(received, destination, response.headers["x-litellm-call-id"])
|
||||
assert span.trace_id.hex() == expected_trace_value
|
||||
assert _attribute(span.attributes, "session.id") == expected_session_value
|
||||
|
||||
|
||||
def test_missing_session_id_reject_accepts_caller_metadata_and_baggage_fallback(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
marker: Final = "reject" + uuid.uuid4().hex
|
||||
provider_secret: Final = "synthetic-provider-secret-" + marker
|
||||
upstream_targets: Final[list[str]] = [] # mutable-ok: records which upstream endpoint each call hit
|
||||
|
||||
def upstream(request: Request) -> Reply:
|
||||
assert request.headers["authorization"] == f"Bearer {provider_secret}"
|
||||
upstream_targets.append(request.target)
|
||||
if request.target == "/v1/responses":
|
||||
return _responses_result("resp-" + marker)
|
||||
assert request.target == "/v1/chat/completions", request.target
|
||||
return _completion(marker + "-answer")
|
||||
|
||||
def langfuse(request: Request) -> Reply:
|
||||
if request.method == "GET" and request.target.startswith(PROJECTS_PATH):
|
||||
return _projects()
|
||||
return Reply(body=b"", content_type="application/x-protobuf")
|
||||
|
||||
with (
|
||||
wire_server(upstream) as provider,
|
||||
wire_server(langfuse) as destination,
|
||||
owned_proxy(
|
||||
gateway,
|
||||
tmp_path,
|
||||
_langfuse_environment(destination),
|
||||
config=_langfuse_config(tmp_path, {"missing_session_id": "reject"}),
|
||||
) as candidate,
|
||||
candidate.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(api_base=provider.url + "/v1", api_key=provider_secret)
|
||||
received: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls keep earlier ones
|
||||
|
||||
caller_session: Final = f"my-session-id-{marker}-r1"
|
||||
header_trace: Final = uuid.uuid4().hex
|
||||
first: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{"model": model, "input": marker + "-r1", "metadata": {"session_id": caller_session}},
|
||||
headers=_w3c_headers(header_trace, None),
|
||||
)
|
||||
assert first.status_code == 200, first.text
|
||||
first_span: Final = _await_span(received, destination, first.headers["x-litellm-call-id"])
|
||||
assert _attribute(first_span.attributes, "session.id") == caller_session
|
||||
|
||||
baggage_session: Final = "baggage-" + marker + "-r2"
|
||||
second: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": marker + "-r2"}],
|
||||
"metadata": {"session_id": ""},
|
||||
},
|
||||
headers=_w3c_headers(uuid.uuid4().hex, baggage_session),
|
||||
)
|
||||
assert second.status_code == 200, second.text
|
||||
second_span: Final = _await_span(received, destination, second.headers["x-litellm-call-id"])
|
||||
assert _attribute(second_span.attributes, "session.id") == baggage_session
|
||||
|
||||
third: Final = candidate.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": marker + "-r3"}]},
|
||||
headers=_w3c_headers(uuid.uuid4().hex, None),
|
||||
)
|
||||
assert third.status_code == 400, third.text
|
||||
assert upstream_targets == ["/v1/responses", "/v1/chat/completions"], upstream_targets
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue