litellm/tests/e2e/otel_client.py
ryan-crabbe-berri 679f7e636e
test(e2e): record steps for raw transport calls and poll helpers (#45150)
Add @step labels to the HttpTransport methods, the poll and wait helpers and the boot helpers that did real IO without recording a step, so a test that reaches the proxy through them no longer reports an empty or gappy step timeline in the JUnit report.
2026-10-07 15:21:26 -07:00

294 lines
12 KiB
Python

"""Jaeger read-back for the OTEL trace-completeness tests: typed models over the
Jaeger query API (the destination's own API - completeness is judged on what the
backend actually holds, never on "export succeeded" proxy-side).
Traces are fetched server-side by the ``litellm.call_id`` tag the gen-AI span
carries (the request's x-litellm-call-id response header), so read-back is
immune to the query page filling up with unrelated traffic (background jobs,
other suites sharing the stack). Jaeger returns every span of a matching trace,
so the completeness assertions see the whole tree. A failed query is a hard
failure, never an empty result - an unreachable destination must not read as
"the trace never arrived".
Service spans that end after the response (the cost write is one) carry no
call id and land as the root of their own trace with a link back to the
request span, so they are fetched by operation name and matched by that link
to the request root rather than by the tag query.
External reads go through ``e2e_http`` (the only module allowed to call
``requests.*``).
"""
from __future__ import annotations
import json
import time
from collections.abc import Iterator
from dataclasses import dataclass
from typing import Final
import pytest
from pydantic import BaseModel, ConfigDict, Field
from e2e_config import OTEL_QUERY_URL, POLL_INTERVAL, POLL_TIMEOUT
from e2e_http import URL, NetworkError, NoBody, Result, Success, get
from e2e_metadata import step
#: OTEL resource service.name the proxy exports under (OTEL_SERVICE_NAME default).
JAEGER_SERVICE = "litellm"
#: Span tag carrying the request's x-litellm-call-id (stamped on the gen-AI span).
CALL_ID_TAG = "litellm.call_id"
#: The v2 gen-AI span attribute recording time-to-first-token for streamed
#: calls: seconds from the upstream request being issued to the first streamed
#: chunk (stamped only for streaming; added in #32236).
TTFT_TAG = "gen_ai.response.time_to_first_chunk"
#: Jaeger's rendering of a span whose OTEL status is ERROR.
ERROR_STATUS_TAG = "otel.status_code"
class JaegerTag(BaseModel):
model_config = ConfigDict(extra="ignore")
key: str
value: str | int | float | bool | None = None
class JaegerReference(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
ref_type: str = Field(alias="refType")
trace_id: str = Field(alias="traceID")
span_id: str = Field(alias="spanID")
class JaegerSpan(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
span_id: str = Field(alias="spanID")
operation_name: str = Field(alias="operationName")
start_time: int = Field(default=0, alias="startTime")
#: Span duration in microseconds, as reported by the Jaeger query API.
duration: int = 0
references: list[JaegerReference] = []
tags: list[JaegerTag] = []
@property
def kind(self) -> str:
for tag in self.tags:
if tag.key == "span.kind":
return str(tag.value)
return ""
def tag(self, key: str) -> str | int | float | bool | None:
for entry in self.tags:
if entry.key == key:
return entry.value
return None
class JaegerTrace(BaseModel):
model_config = ConfigDict(extra="ignore", populate_by_name=True)
trace_id: str = Field(alias="traceID")
spans: list[JaegerSpan] = []
def span_names(self) -> list[str]:
return sorted(span.operation_name for span in self.spans)
class JaegerTracesPage(BaseModel):
model_config = ConfigDict(extra="ignore")
data: list[JaegerTrace] = []
class _TracesQuery(BaseModel):
service: str
tags: str | None = None
operation: str | None = None
limit: int = 20
lookback: str = "1h"
start: int | None = None
end: int | None = None
def _ticks() -> Iterator[None]:
while True:
yield None
time.sleep(POLL_INTERVAL)
def _settled(trace: JaegerTrace, names: set[str], prefixes: set[str]) -> bool:
present = set(trace.span_names())
return names.issubset(present) and all(any(name.startswith(prefix) for name in present) for prefix in prefixes)
def root_span(trace: JaegerTrace) -> JaegerSpan | None:
"""The single span whose references all point outside the trace (a span
with no references qualifies). None when there is not exactly one."""
in_trace = {span.span_id for span in trace.spans}
roots = [span for span in trace.spans if all(ref.span_id not in in_trace for ref in span.references)]
return roots[0] if len(roots) == 1 else None
def served_genai_spans(trace: JaegerTrace, genai_span: str) -> list[JaegerSpan]:
"""The gen-AI spans for attempts that actually served the request.
The proxy opens one gen-AI span per upstream attempt, so a call the router
retried carries an error span for every failed attempt beside the one that
answered. Only the served attempt streams chunks, so only it records TTFT
or a streaming flag; asserting over the raw span list makes every one of
these tests fail whenever the upstream 429s, 529s, or hands back a stale
credential on the first try."""
return [span for span in trace.spans if span.operation_name == genai_span and span.tag(ERROR_STATUS_TAG) != "ERROR"]
def one_served_genai_span(trace: JaegerTrace, genai_span: str) -> JaegerSpan:
served = served_genai_spans(trace, genai_span)
assert len(served) == 1, (
f"a streamed call must produce exactly ONE served gen-AI span, got {len(served)}; spans: {trace.span_names()}"
)
return served[0]
def _follows(trace: JaegerTrace, parent_trace_id: str, parent_span_id: str) -> bool:
root = root_span(trace)
return root is not None and any(
ref.trace_id == parent_trace_id and ref.span_id == parent_span_id for ref in root.references
)
@dataclass(frozen=True, slots=True)
class CallTraces:
hits: tuple[JaegerTrace, ...]
linked: tuple[JaegerTrace, ...]
@dataclass(frozen=True, slots=True)
class _Observation:
traces: CallTraces
missing: tuple[str, ...]
unreachable: NetworkError | None
def settled(self, names: set[str], prefixes: set[str]) -> bool:
hits: Final = self.traces.hits
return self.unreachable is None and len(hits) == 1 and not self.missing and _settled(hits[0], names, prefixes)
@dataclass(frozen=True, slots=True)
class OtelReader:
query_url: str
def _query_traces(self, call_id: str) -> Result[JaegerTracesPage]:
return get(
URL(f"{self.query_url}/api/traces"),
headers=NoBody(),
params=_TracesQuery(service=JAEGER_SERVICE, tags=json.dumps({CALL_ID_TAG: call_id})),
response_type=JaegerTracesPage,
timeout=30.0,
)
def traces_for_call(self, call_id: str) -> list[JaegerTrace]:
"""Every trace holding a span tagged with this call id. Jaeger matches
spans server-side and returns their full traces; more than one hit for
one call IS the split-trace bug, so this never collapses to one."""
match self._query_traces(call_id):
case Success(data=page):
return page.data
case failure:
pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}")
def _query_operation(self, operation: str, *, start: int) -> Result[JaegerTracesPage]:
return get(
URL(f"{self.query_url}/api/traces"),
headers=NoBody(),
params=_TracesQuery(
service=JAEGER_SERVICE,
operation=operation,
limit=200,
start=start,
end=int(time.time() * 1_000_000),
),
response_type=JaegerTracesPage,
timeout=30.0,
)
def linked_traces(self, *, operation: str, parent: JaegerTrace) -> tuple[JaegerTrace, ...] | NetworkError:
"""Traces whose root span references the parent trace's root span.
Detached post-response work lands as the root of its own trace with a
link back to the request span instead of the call-id tag, so it is
found by operation name, windowed to start at the parent root's start
time (the detached span always starts after it), and matched on that
link. A NetworkError is handed back so the polling caller can tell an
unreachable read-back endpoint from a span that never arrived."""
parent_root: Final = root_span(parent)
if parent_root is None:
return ()
match self._query_operation(operation, start=parent_root.start_time):
case Success(data=page):
return tuple(t for t in page.data if _follows(t, parent.trace_id, parent_root.span_id))
case NetworkError() as failure:
return failure
case failure:
pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}")
@step("Poll Jaeger for the traces of call {call_id}")
def poll_traces_for_call(
self,
*,
call_id: str,
settled_names: set[str],
settled_prefixes: set[str],
linked_names: frozenset[str] = frozenset(),
) -> CallTraces:
"""Poll until exactly one trace holds the call, it carries every span
name in ``settled_names`` plus at least one name per prefix in
``settled_prefixes``, and every name in ``linked_names`` is either in
that trace or is the root of its own trace referencing the request
root (post-response work detaches per #42826). At the deadline the
last observed state is returned as-is so the caller's assertions
report the real final state - on a split trace this never settles and
the orphan comes back. A read-back endpoint still failing at the
deadline (either query) is a hard failure, not a missing span."""
deadline: Final = time.monotonic() + POLL_TIMEOUT
last: Final = self._poll(call_id, settled_names, settled_prefixes, linked_names, deadline)
if last.unreachable is not None:
pytest.fail(
f"Jaeger query API at {self.query_url} stayed unreachable until the "
f"{POLL_TIMEOUT}s poll deadline: {last.unreachable}"
)
return last.traces
def _observe(self, call_id: str, linked_names: frozenset[str]) -> _Observation:
match self._query_traces(call_id):
case NetworkError() as failure:
return _Observation(CallTraces((), ()), tuple(linked_names), failure)
case Success(data=page):
if len(page.data) != 1:
return _Observation(CallTraces(tuple(page.data), ()), tuple(linked_names), None)
hit: Final = page.data[0]
present: Final = frozenset(hit.span_names())
results: Final = {
name: self.linked_traces(operation=name, parent=hit) for name in linked_names if name not in present
}
unreachable: Final = next((r for r in results.values() if isinstance(r, NetworkError)), None)
linked: Final = tuple(t for r in results.values() if not isinstance(r, NetworkError) for t in r)
missing: Final = tuple(name for name, r in results.items() if isinstance(r, NetworkError) or not r)
return _Observation(CallTraces((hit,), linked), missing, unreachable)
case failure:
pytest.fail(f"Jaeger query API at {self.query_url} failed: {failure}")
def _poll(
self,
call_id: str,
names: set[str],
prefixes: set[str],
linked_names: frozenset[str],
deadline: float,
) -> _Observation:
observations: Final = (self._observe(call_id, linked_names) for _ in _ticks())
return next(o for o in observations if o.settled(names, prefixes) or time.monotonic() >= deadline)
def build_otel_reader() -> OtelReader:
return OtelReader(query_url=OTEL_QUERY_URL)