fix(otel): honor per-team Arize sampling rates in OTel v2 fan-out (#44595)

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-06 09:23:11 -05:00 • committed by GitHub
parent 44d5dacbaa
commit 46d2c2a6ea
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 695 additions and 34 deletions

View file

@ -34,6 +34,36 @@ _SUCCESS_SAMPLING_RATE_VAR: Final = "arize_success_sampling_rate"
_ERROR_SAMPLING_RATE_VAR: Final = "arize_error_sampling_rate"
def parse_sampling_rate(value: object, var: str) -> float | None:
"""The rate a key or team set for ``var``, or ``None`` when nothing usable was set.
Unset, empty and ``"None"`` mean "export everything". A value that is not a number
or lies outside ``0.0..1.0`` is logged and treated the same way, so a typo never
silences a tenant's traces. ``0.0`` is a real rate and comes back as ``0.0``.
"""
if value is None or value in ("", "None"):
return None
try:
if not isinstance(value, (str, int, float)):
raise TypeError(type(value).__name__)
rate: Final = float(value)
except (TypeError, ValueError):
verbose_logger.warning(
"ArizeLogger: %s value %r is not a number; exporting the request",
var,
value,
)
return None
if not math.isfinite(rate) or not 0.0 <= rate <= 1.0:
verbose_logger.warning(
"ArizeLogger: %s value %r is outside 0.0..1.0; exporting the request",
var,
value,
)
return None
return rate
class ArizeLogger(OpenTelemetry):
"""
Arize logger that sends traces to an Arize endpoint.
@ -108,26 +138,7 @@ class ArizeLogger(OpenTelemetry):
dynamic_params: Final = kwargs.get("standard_callback_dynamic_params")
if not isinstance(dynamic_params, Mapping):
return None
value: Final = dynamic_params.get(var)
if value is None or value in ("", "None"):
return None
try:
rate: Final = float(value)
except (TypeError, ValueError):
verbose_logger.warning(
"ArizeLogger: %s value %r is not a number; exporting the request",
var,
value,
)
return None
if not math.isfinite(rate) or not 0.0 <= rate <= 1.0:
verbose_logger.warning(
"ArizeLogger: %s value %r is outside 0.0..1.0; exporting the request",
var,
value,
)
return None
return rate
return parse_sampling_rate(dynamic_params.get(var), var)
def _should_export(self, kwargs: dict[str, object], var: str) -> bool:
rate: Final = self._sampling_rate_for_request(kwargs, var)

View file

@ -39,6 +39,17 @@ class OtelDestination(LiteLLMBaseModel):
default="full",
description="``llm_only`` keeps just the model-call spans; the rest of the request tree is not forwarded.",
)
success_sampling_rate: float | None = Field(
default=None,
description=(
"Share of the request trees forwarded, 0.0..1.0, drawn once per request when its root span ends; "
"``None`` forwards every one."
),
)
error_sampling_rate: float | None = Field(
default=None,
description="Same, for the requests with a failed span in their tree; ``None`` forwards every one.",
)
def header_string(self) -> str:
"""Render headers as the ``k=v,k2=v2`` form an ``ExporterSpec`` expects.
@ -53,9 +64,9 @@ class OtelDestination(LiteLLMBaseModel):
def cache_key(self) -> tuple[str, tuple[tuple[str, str], ...], tuple[tuple[str, str], ...], str | None]:
"""Identity for processor reuse, so one destination means one exporter.
``span_scope`` is left out on purpose: the scope decides which spans reach the
processor, not how the processor exports them, so a full and an ``llm_only``
view of the same account share one exporter.
``span_scope`` and the sampling rates are left out on purpose: they decide which
spans reach the processor, not how the processor exports them, so a full and an
``llm_only`` view of the same account share one exporter.
"""
return (
self.endpoint,

View file

@ -1,6 +1,7 @@
"""Provider / exporter factory + the Baggage span processor."""
import queue
import random
import threading
import time
from collections import OrderedDict
@ -33,7 +34,7 @@ from opentelemetry.sdk.trace.export import (
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.trace import Span, SpanContext, SpanKind, Status, Tracer
from opentelemetry.trace import Span, SpanContext, SpanKind, Status, StatusCode, Tracer
from opentelemetry.util.re import parse_env_headers
from opentelemetry.util.types import Attributes, AttributeValue
@ -441,6 +442,53 @@ def is_llm_call_span(span: ReadableSpan) -> bool:
return GenAI.OPERATION_NAME in attributes and MCP.METHOD_NAME not in attributes
#: Request trees held back for a sampled destination until the last span of the trace
#: still open ends, the verdicts kept for the spans that start after that (a post-call
#: database write), and the spans in flight per trace. All bounded, so a span that never
#: ends, or a flood of requests, cannot hold spans for ever: the oldest tree is decided
#: on what it has.
_MAX_PENDING_TREES: Final = 1024
_MAX_PENDING_SPANS_PER_TREE: Final = 512
_MAX_REMEMBERED_VERDICTS: Final = 4096
_MAX_OPEN_TRACES: Final = 16384
class _PendingTree:
"""The spans of one request held for its sampled destinations, and whether any
span of the tree has failed so far."""
__slots__ = ("failed", "held")
def __init__(self) -> None:
self.failed = False
self.held: list[tuple[ReadableSpan, OtelDestination]] = [] # mutable-ok: bounded per tree
def _is_sampled(destination: "OtelDestination") -> bool:
return destination.success_sampling_rate is not None or destination.error_sampling_rate is not None
#: A verdict is one request tree's draw for one destination at its rates: the exporter
#: ``cache_key`` leaves the rates out, so two destinations sharing an exporter with
#: different rates would otherwise share one draw.
_VerdictKey = tuple[int, object, float | None, float | None]
def _verdict_key(trace_id: int, destination: "OtelDestination") -> _VerdictKey:
return (trace_id, destination.cache_key(), destination.success_sampling_rate, destination.error_sampling_rate)
def _keeps_tree(destination: "OtelDestination", failed: bool, draw: Callable[[], float]) -> bool:
"""Whether one draw against the destination's rate keeps a request tree, read the
way the legacy Arize callback reads ``arize_success_sampling_rate`` /
``arize_error_sampling_rate``: a tree with a failed span answers to the error
rate, an unset rate keeps everything and ``0.0`` keeps nothing."""
rate: Final = destination.error_sampling_rate if failed else destination.success_sampling_rate
if rate is None:
return True
return rate > 0.0 and draw() <= rate
def _in_scope(span: ReadableSpan, scope: "OtelSpanScope") -> bool:
return scope == "full" or is_llm_call_span(span)
@ -566,24 +614,38 @@ class TenantFanOutSpanProcessor(SpanProcessor):
excluded_db_systems: frozenset[str] = frozenset(),
pending_drains: int = _MAX_PENDING_DRAINS,
drain_pool: _DrainPool | None = None,
sampling_draw: Callable[[], float] | None = None,
) -> None:
self._operator_sinks: Final = operator_sinks
self._draw: Final = sampling_draw if sampling_draw is not None else random.random
self._excluded_db_systems: Final = excluded_db_systems
self._drain_seconds: Final = shutdown_drain_seconds
self._lock: Final = threading.Condition()
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
self._draining = False # guarded by ``_lock``: shutdown is flushing, nothing is held any more
self._build: Final = processor_factory if processor_factory is not None else _destination_processor
self._processors: OrderedDict[object, SpanProcessor] = OrderedDict() # mutable-ok: bounded LRU
self._retired: OrderedDict[int, SpanProcessor] = OrderedDict() # mutable-ok: drains as exports finish
self._exporting: dict[int, int] = {} # mutable-ok: per-processor in-flight export count
self._pending: OrderedDict[int, _PendingTree] = OrderedDict() # mutable-ok: bounded, by trace id
self._open: OrderedDict[int, int] = OrderedDict() # mutable-ok: bounded, spans in flight by trace id
self._verdicts: OrderedDict[_VerdictKey, bool] = OrderedDict() # mutable-ok: bounded LRU
self._drain: Final = drain_pool if drain_pool is not None else _DrainPool(capacity=pending_drains)
def on_start(self, span: SDKSpan, parent_context: Context | None = None) -> None:
return None
context: Final = span.context
trace_id: Final = context.trace_id if context is not None else 0
with self._lock:
self._open[trace_id] = self._open.get(trace_id, 0) + 1
while len(self._open) > _MAX_OPEN_TRACES:
self._open.popitem(last=False)
def on_end(self, span: ReadableSpan) -> None:
suppressed: Final = suppressed_backends()
attributes: Final = span.attributes or _NO_ATTRIBUTES
context: Final = span.context
trace_id: Final = context.trace_id if context is not None else 0
failed: Final = span.status.status_code is StatusCode.ERROR
for destination in request_destinations():
if (
self._operator_already_writes(span, destination, suppressed)
@ -591,15 +653,94 @@ class TenantFanOutSpanProcessor(SpanProcessor):
or _is_excluded_database_span(attributes, self._excluded_db_systems)
):
continue
processor = self._acquire(destination)
if processor is None:
if not _is_sampled(destination):
self._forward(span, destination)
continue
try:
processor.on_end(_scoped(_for_destination(span, destination), destination.span_scope))
except Exception as exc: # noqa: BLE001 # one destination's failure must not cost the others their span
verbose_logger.debug("OTel V2 fan-out: forwarding to %s failed: %s", destination.endpoint, exc)
finally:
self._release(processor)
for held_span, held_destination in self._route(trace_id, span, destination, failed):
self._forward(held_span, held_destination)
for held_span, held_destination in self._settle(trace_id, failed):
self._forward(held_span, held_destination)
def _forward(self, span: ReadableSpan, destination: "OtelDestination") -> None:
processor: Final = self._acquire(destination)
if processor is None:
return
try:
processor.on_end(_scoped(_for_destination(span, destination), destination.span_scope))
except Exception as exc: # noqa: BLE001 # one destination's failure must not cost the others their span
verbose_logger.debug("OTel V2 fan-out: forwarding to %s failed: %s", destination.endpoint, exc)
finally:
self._release(processor)
def _route(
self, trace_id: int, span: ReadableSpan, destination: "OtelDestination", failed: bool
) -> "tuple[tuple[ReadableSpan, OtelDestination], ...]":
"""The spans to forward now for ``span`` on a sampled ``destination``: the span
itself when its tree's verdict is already known and keeps it, nothing while the
tree is held for the rest of its trace, or what the bounds made the fan-out decide
early (a tree that outgrew its cap is decided on what it has, and so is the oldest
tree once too many are waiting, so a span that never ends holds nothing back for
ever). A failed span marks its tree before any such early draw. Once shutdown
has started flushing, nothing is held: a span ending then is decided at once,
since no later flush would reach it.
The verdict lookup and the hold share one lock acquisition: a tree deciding on
another thread between the two would leave this span in a fresh tree.
"""
with self._lock:
verdict: Final = self._verdicts.get(_verdict_key(trace_id, destination))
if verdict is not None:
return ((span, destination),) if verdict else ()
tree: Final = self._pending.setdefault(trace_id, _PendingTree())
tree.held.append((span, destination))
if failed:
tree.failed = True
if self._draining or len(tree.held) >= _MAX_PENDING_SPANS_PER_TREE:
return self._decide_locked(trace_id)
if len(self._pending) > _MAX_PENDING_TREES:
oldest: Final = next(iter(self._pending))
return self._decide_locked(oldest)
return ()
def _decide(self, trace_id: int) -> "tuple[tuple[ReadableSpan, OtelDestination], ...]":
with self._lock:
return self._decide_locked(trace_id)
def _settle(self, trace_id: int, failed: bool) -> "tuple[tuple[ReadableSpan, OtelDestination], ...]":
"""What the trace's held tree exports once the span that just ended was the last
of its trace still open. The model-call span ends in the post-call callback,
after the server span, so a tree is decided when its trace goes quiet rather than
when its root ends, and a model call that fails late still answers to the error
rate. A trace whose starts this fan-out never saw (attached mid-flight, or forgotten
by the bound) is decided as each of its spans ends.
"""
with self._lock:
tree: Final = self._pending.get(trace_id)
if tree is not None and failed:
tree.failed = True
still_open: Final = self._open.pop(trace_id, 1) - 1
if still_open > 0:
self._open[trace_id] = still_open
return ()
return self._decide_locked(trace_id)
def _decide_locked(self, trace_id: int) -> "tuple[tuple[ReadableSpan, OtelDestination], ...]":
"""Draw once per destination for the held tree; the spans the draws keep, and a
verdict remembered for the spans of this trace that start later."""
tree: Final = self._pending.pop(trace_id, None)
if tree is None:
return ()
verdicts: dict[_VerdictKey, bool] = {} # mutable-ok: one draw per destination of this tree
for _, destination in tree.held:
key = _verdict_key(trace_id, destination)
if key not in verdicts:
verdicts[key] = _keeps_tree(destination, tree.failed, self._draw)
self._verdicts[key] = verdicts[key]
while len(self._verdicts) > _MAX_REMEMBERED_VERDICTS:
self._verdicts.popitem(last=False)
return tuple(
(span, destination) for span, destination in tree.held if verdicts[_verdict_key(trace_id, destination)]
)
def _operator_already_writes(
self, span: ReadableSpan, destination: "OtelDestination", suppressed: frozenset[str]
@ -636,6 +777,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
long as it likes. The drain's workers are daemons, and the whole teardown
shares one deadline.
"""
with self._lock:
self._draining = True
undecided: Final = tuple(self._pending)
for trace_id in undecided:
for held_span, held_destination in self._decide(trace_id):
self._forward(held_span, held_destination)
deadline: Final = time.monotonic() + self._drain_seconds
with self._lock:
self._closed = True

View file

@ -165,6 +165,24 @@ def destination_capable_backends() -> frozenset[str]:
return frozenset(_DESTINATION_BY_CALLBACK) & frozenset(DYNAMIC_HEADERS_BY_CALLBACK)
_SAMPLING_RATE_VARS_BY_CALLBACK: Final[Mapping[str, tuple[str, str]]] = MappingProxyType(
{"arize": ("arize_success_sampling_rate", "arize_error_sampling_rate")}
)
def _sampling_rates(callback_name: str, params: StandardCallbackDynamicParams) -> tuple[float | None, float | None]:
from litellm.integrations.arize.arize import parse_sampling_rate
rate_vars: Final = _SAMPLING_RATE_VARS_BY_CALLBACK.get(callback_name)
if rate_vars is None:
return (None, None)
success_var, error_var = rate_vars
return (
parse_sampling_rate(params.get(success_var), success_var),
parse_sampling_rate(params.get(error_var), error_var),
)
def destination_for(
callback_name: str,
params: StandardCallbackDynamicParams,
@ -190,6 +208,7 @@ def destination_for(
if resolved is None:
return None
endpoint, protocol = resolved
success_sampling_rate, error_sampling_rate = _sampling_rates(callback_name, params)
return OtelDestination(
endpoint=endpoint,
headers=MappingProxyType(dict(headers)),
@ -198,4 +217,6 @@ def destination_for(
callback_name=callback_name,
protocol=protocol,
span_scope=_span_scope(callback_name, params),
success_sampling_rate=success_sampling_rate,
error_sampling_rate=error_sampling_rate,
)

View file

@ -184,7 +184,7 @@ def _request_spans(exporter: InMemorySpanExporter) -> int:
return sum(1 for span in exporter.get_finished_spans() if span.name == "litellm_request")
def _arize_kwargs(callback_vars: dict[str, str] | None = None) -> dict[str, object]:
def _arize_kwargs(callback_vars: dict[str, object] | None = None) -> dict[str, object]:
kwargs: dict[str, object] = {
"model": "gpt-4",
"litellm_params": {"metadata": {}},
@ -276,3 +276,11 @@ async def test_out_of_range_sampling_rate_exports_rather_than_dropping(bad: str)
logger, exporter = _sampled_arize_logger(random_draw=lambda: 0.99)
await logger.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": bad}), None, _START, _END)
assert _request_spans(exporter) == 1
@pytest.mark.asyncio
async def test_a_sampling_rate_that_is_not_a_scalar_exports_rather_than_dropping():
logger, exporter = _sampled_arize_logger(random_draw=lambda: 0.99)
kwargs = _arize_kwargs({"arize_success_sampling_rate": ["0.0"]})
await logger.async_log_success_event(kwargs, None, _START, _END)
assert _request_spans(exporter) == 1

View file

@ -11,6 +11,7 @@ from types import MappingProxyType
from typing import Final
import pytest
from opentelemetry import trace as trace_api
from opentelemetry.sdk.resources import Resource
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
@ -1043,6 +1044,468 @@ class TestFanOut:
assert len(built) == 1
ARIZE_TEAM_PARAMS = {"arize_space_id": "space-team", "arize_api_key": "key-team"}
def arize_destination(**rates: str) -> OtelDestination:
destination = destination_for("arize", {**ARIZE_TEAM_PARAMS, **rates})
assert destination is not None, "the fixture must resolve for the test to mean anything"
return destination
class TestTenantSampling:
"""A team's ``arize_success_sampling_rate`` and ``arize_error_sampling_rate`` decide how
much of its traffic the fan-out delivers to its space, the way they already decide what
the legacy Arize callback exports."""
NAMES = {"POST /v1/chat/completions", "auth /v1/chat/completions", "chat gpt-4"}
@staticmethod
def _fan_out(dest_exporter, draw=None, global_exporter=None):
provider = TracerProvider()
if global_exporter is not None:
provider.add_span_processor(_OverriddenBackendFilter(SimpleSpanProcessor(global_exporter), "arize"))
kwargs = {"processor_factory": lambda _d: SimpleSpanProcessor(dest_exporter)}
if draw is not None:
kwargs["sampling_draw"] = draw
provider.add_span_processor(TenantFanOutSpanProcessor(**kwargs))
return provider
@staticmethod
def _tree(provider, *, failed=False):
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions"):
with tracer.start_as_current_span("auth /v1/chat/completions"):
pass
with tracer.start_as_current_span("chat gpt-4") as span:
if failed:
span.set_status(Status(StatusCode.ERROR, "upstream 500"))
def _run(self, provider, destinations, *, failed=False, trees=1):
def run():
set_request_destinations(destinations)
for _ in range(trees):
self._tree(provider, failed=failed)
in_fresh_context(run)
def test_a_zero_rate_keeps_the_whole_tree_out_of_the_teams_space(self):
"""The reported case: override mode, both rates 0.0, so the team's traffic goes nowhere."""
global_exporter, dest_exporter = InMemorySpanExporter(), InMemorySpanExporter()
provider = self._fan_out(dest_exporter, global_exporter=global_exporter)
self._run(
provider,
(arize_destination(arize_success_sampling_rate="0.0", arize_error_sampling_rate="0.0"),),
)
assert dest_exporter.get_finished_spans() == ()
assert global_exporter.get_finished_spans() == ()
def test_a_rate_of_one_exports_the_whole_tree(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(provider, (arize_destination(arize_success_sampling_rate="1.0"),))
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_an_unset_rate_exports_everything(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.99)
self._run(provider, (arize_destination(),))
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_a_failed_request_is_kept_whole_by_the_error_rate_when_the_success_rate_drops_the_rest(self):
"""Mirrors the legacy callback, where a failed call answers to the error rate, at the
size of a request tree: the failed model call arrives with its parents."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(
provider,
(arize_destination(arize_success_sampling_rate="0.0", arize_error_sampling_rate="1.0"),),
failed=True,
)
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_a_zero_error_rate_drops_the_whole_failed_request_the_success_rate_would_keep(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(
provider,
(arize_destination(arize_success_sampling_rate="1.0", arize_error_sampling_rate="0.0"),),
failed=True,
)
assert dest_exporter.get_finished_spans() == ()
def test_a_model_call_that_fails_after_the_server_span_ended_answers_to_the_error_rate(self):
"""The model-call span is closed in the post-call callback, after the FastAPI server
span has ended, so a request whose only failed span ends late is still a failed
request: it is kept by an error rate of 1.0 that a success rate of 0.0 would drop."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
destinations = (arize_destination(arize_success_sampling_rate="0.0", arize_error_sampling_rate="1.0"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions"):
call = tracer.start_span("chat gpt-4")
assert dest_exporter.get_finished_spans() == (), "undecided while the model call is open"
call.set_status(Status(StatusCode.ERROR, "upstream 500"))
call.end()
in_fresh_context(run)
assert {s.name for s in dest_exporter.get_finished_spans()} == {"POST /v1/chat/completions", "chat gpt-4"}
def test_the_span_that_fills_a_tree_to_its_bound_still_counts_as_failed(self):
"""A tree decided early because it hit the per-tree cap answers to the error rate
when the span that tripped the cap is the one that failed."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
destinations = (arize_destination(arize_success_sampling_rate="1.0", arize_error_sampling_rate="0.0"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions"):
for index in range(otel_providers._MAX_PENDING_SPANS_PER_TREE - 1):
with tracer.start_as_current_span(f"chat {index}"):
pass
with tracer.start_as_current_span("chat gpt-4") as call:
call.set_status(Status(StatusCode.ERROR, "upstream 500"))
in_fresh_context(run)
assert dest_exporter.get_finished_spans() == ()
def test_a_span_that_ends_after_the_root_follows_the_requests_verdict(self):
"""A post-call database write that starts after the whole request tree has ended
goes where the rest of the tree went, without a second draw."""
kept_exporter, dropped_exporter = InMemorySpanExporter(), InMemorySpanExporter()
kept = self._fan_out(kept_exporter, draw=lambda: 0.1)
dropped = self._fan_out(dropped_exporter, draw=lambda: 0.9)
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def late_tree(provider):
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions") as root:
root_context = trace_api.set_span_in_context(root)
with tracer.start_as_current_span("postgres INSERT LiteLLM_SpendLogs", context=root_context):
pass
for provider in (kept, dropped):
def run(provider=provider):
set_request_destinations(destinations)
late_tree(provider)
in_fresh_context(run)
assert {s.name for s in kept_exporter.get_finished_spans()} == {
"POST /v1/chat/completions",
"postgres INSERT LiteLLM_SpendLogs",
}
assert dropped_exporter.get_finished_spans() == ()
def test_a_late_span_whose_verdict_was_forgotten_is_decided_on_its_own_not_stranded(self):
"""Once enough other requests have been decided to evict a request's verdict, a span
of it that starts late (a post-call database write) opens and closes its own count,
so it is decided on a draw of its own as soon as it ends instead of waiting in a tree
nothing would ever close."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.1)
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions") as root:
root_context = trace_api.set_span_in_context(root)
for index in range(otel_providers._MAX_REMEMBERED_VERDICTS):
with tracer.start_as_current_span(f"POST {index}", context=trace_api.Context()):
pass
with tracer.start_as_current_span("postgres INSERT LiteLLM_SpendLogs", context=root_context):
pass
in_fresh_context(run)
names = [s.name for s in dest_exporter.get_finished_spans()]
assert names[0] == "POST /v1/chat/completions"
assert names[-1] == "postgres INSERT LiteLLM_SpendLogs", "exported as it ended, not held for a root"
def test_a_callers_traceparent_neither_steers_the_draw_nor_hides_the_root(self):
"""The draw is random, not read from the trace id a caller may have chosen, and the
server span under a remote parent still closes the tree."""
dest_exporter = InMemorySpanExporter()
draws = iter([0.9, 0.1])
provider = self._fan_out(dest_exporter, draw=lambda: next(draws))
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
for trace_id in (1, 2):
remote = trace_api.SpanContext(
trace_id=trace_id, span_id=1, is_remote=True, trace_flags=trace_api.TraceFlags(0x01)
)
parent = trace_api.set_span_in_context(trace_api.NonRecordingSpan(remote))
with tracer.start_as_current_span("POST /v1/chat/completions", context=parent):
with tracer.start_as_current_span("chat gpt-4"):
pass
in_fresh_context(run)
kept = dest_exporter.get_finished_spans()
assert [s.name for s in kept] == ["chat gpt-4", "POST /v1/chat/completions"]
assert all(s.context.trace_id == 2 for s in kept), "the second request, whose draw of 0.1 passed"
def test_a_root_that_never_ends_holds_the_tree_only_up_to_the_bound(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.0)
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions"):
for index in range(otel_providers._MAX_PENDING_SPANS_PER_TREE):
with tracer.start_as_current_span(f"chat {index}"):
pass
assert len(dest_exporter.get_finished_spans()) == otel_providers._MAX_PENDING_SPANS_PER_TREE
in_fresh_context(run)
def test_a_flood_of_waiting_trees_decides_the_oldest_on_what_it_has(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.0)
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
for index in range(otel_providers._MAX_PENDING_TREES + 1):
root = tracer.start_span(f"POST {index}", context=trace_api.Context())
tracer.start_span(f"chat {index}", context=trace_api.set_span_in_context(root)).end()
if index < otel_providers._MAX_PENDING_TREES:
assert dest_exporter.get_finished_spans() == (), f"tree {index} is held while its root is open"
in_fresh_context(run)
assert [s.name for s in dest_exporter.get_finished_spans()] == ["chat 0"]
def test_a_flood_of_open_traces_forgets_the_oldest_and_decides_it_as_its_spans_end(self):
"""The fan-out counts open spans for a bounded number of traces, a bound far above
what one process keeps in flight, since every trace of the provider is counted, not
only the sampled ones. A trace pushed out of that count is decided on what it has as
its next span ends, root still open, rather than held until a bound gets to it."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.0)
destinations = (arize_destination(arize_success_sampling_rate="0.5"),)
def run():
set_request_destinations(destinations)
tracer = get_tracer(provider, "litellm")
root = tracer.start_span("POST /v1/chat/completions", context=trace_api.Context())
tracer.start_span("auth /v1/chat/completions", context=trace_api.set_span_in_context(root)).end()
for index in range(otel_providers._MAX_OPEN_TRACES):
tracer.start_span(f"POST {index}", context=trace_api.Context())
assert dest_exporter.get_finished_spans() == (), "still held before the count is forgotten"
tracer.start_span("chat gpt-4", context=trace_api.set_span_in_context(root)).end()
assert {s.name for s in dest_exporter.get_finished_spans()} == {
"auth /v1/chat/completions",
"chat gpt-4",
}
root.end()
in_fresh_context(run)
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_shutdown_decides_what_is_still_held(self):
dest_exporter = InMemorySpanExporter()
fan_out = TenantFanOutSpanProcessor(
processor_factory=lambda _d: SimpleSpanProcessor(dest_exporter), sampling_draw=lambda: 0.0
)
provider = TracerProvider()
provider.add_span_processor(fan_out)
def run():
set_request_destinations((arize_destination(arize_success_sampling_rate="0.5"),))
tracer = get_tracer(provider, "litellm")
with tracer.start_as_current_span("POST /v1/chat/completions"):
with tracer.start_as_current_span("chat gpt-4"):
pass
assert dest_exporter.get_finished_spans() == ()
fan_out.shutdown()
in_fresh_context(run)
assert [s.name for s in dest_exporter.get_finished_spans()] == ["chat gpt-4"]
def test_a_span_that_ends_while_shutdown_flushes_is_decided_at_once_not_dropped(self):
"""Shutdown decides the trees it finds, then closes. A sampled span of a trace it
did not find, ending while it flushes with its root still open, is decided on the
spot rather than held by a fan-out that will never forward again."""
dest_exporter = InMemorySpanExporter()
stragglers = [] # mutable-ok: the span the first export ends, mid-shutdown
class _EndsAStragglerOnExport(SimpleSpanProcessor):
def on_end(self, span):
for straggler in stragglers:
if straggler.is_recording():
straggler.end()
super().on_end(span)
fan_out = TenantFanOutSpanProcessor(
processor_factory=lambda _d: _EndsAStragglerOnExport(dest_exporter), sampling_draw=lambda: 0.0
)
provider = TracerProvider()
provider.add_span_processor(fan_out)
def run():
set_request_destinations((arize_destination(arize_success_sampling_rate="0.5"),))
tracer = get_tracer(provider, "litellm")
root = tracer.start_span("POST /v1/chat/completions", context=trace_api.Context())
tracer.start_span("chat gpt-4", context=trace_api.set_span_in_context(root)).end()
late_root = tracer.start_span("POST /v1/embeddings", context=trace_api.Context())
stragglers.append(tracer.start_span("embed ada", context=trace_api.set_span_in_context(late_root)))
fan_out.shutdown()
in_fresh_context(run)
assert {s.name for s in dest_exporter.get_finished_spans()} == {"chat gpt-4", "embed ada"}
def test_a_zero_rate_drops_even_a_draw_of_zero(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.0)
self._run(provider, (arize_destination(arize_success_sampling_rate="0.0"),))
assert dest_exporter.get_finished_spans() == ()
def test_the_draw_is_compared_to_the_rate_the_way_the_legacy_callback_compares_it(self):
dropped_exporter, kept_exporter = InMemorySpanExporter(), InMemorySpanExporter()
destinations = (arize_destination(arize_success_sampling_rate="0.2"),)
self._run(self._fan_out(dropped_exporter, draw=lambda: 0.3), destinations)
self._run(self._fan_out(kept_exporter, draw=lambda: 0.2), destinations)
assert dropped_exporter.get_finished_spans() == ()
assert {s.name for s in kept_exporter.get_finished_spans()} == self.NAMES
def test_a_trace_is_kept_or_dropped_whole(self):
"""One decision per request tree, not one per span, so the team never sees a trace with
its root missing."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(provider, (arize_destination(arize_success_sampling_rate="0.5"),), trees=64)
by_trace = {}
for span in dest_exporter.get_finished_spans():
by_trace.setdefault(span.context.trace_id, set()).add(span.name)
assert all(names == self.NAMES for names in by_trace.values())
assert 0 < len(by_trace) < 64
def test_in_additive_mode_the_operators_own_copy_is_not_sampled(self, monkeypatch):
"""The rates are the team's setting for the team's space; the operator's backbone keeps
every span."""
monkeypatch.setattr(litellm, "otel_tenant_destination_mode", "additive", raising=False)
global_exporter, dest_exporter = InMemorySpanExporter(), InMemorySpanExporter()
provider = self._fan_out(dest_exporter, global_exporter=global_exporter)
self._run(provider, (arize_destination(arize_success_sampling_rate="0.0"),))
assert dest_exporter.get_finished_spans() == ()
assert {s.name for s in global_exporter.get_finished_spans()} == self.NAMES
def test_a_rate_applies_only_to_the_destination_that_carries_it(self):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(provider, (arize_destination(arize_success_sampling_rate="0.0"), LANGFUSE_DEST))
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_two_views_of_one_exporter_each_draw_at_their_own_rate(self):
"""Two destinations to the same space with different rates share an exporter, not a
verdict: the 1.0 view gets the tree once and the 0.0 view adds nothing."""
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(
provider,
(
arize_destination(arize_success_sampling_rate="1.0"),
arize_destination(arize_success_sampling_rate="0.0"),
),
)
assert sorted(s.name for s in dest_exporter.get_finished_spans()) == sorted(self.NAMES)
def test_a_rate_does_not_split_the_teams_exporter(self):
"""Sampling decides which spans reach the processor, not how it exports, so two views of
one account with different rates share one exporter."""
built = []
def factory(destination):
built.append(destination)
return SimpleSpanProcessor(InMemorySpanExporter())
provider = TracerProvider()
provider.add_span_processor(TenantFanOutSpanProcessor(processor_factory=factory))
self._run(provider, (arize_destination(arize_success_sampling_rate="1.0"), arize_destination()))
assert len(built) == 1
@pytest.mark.parametrize("bad", ["abc", "-0.1", "1.5", "nan"])
def test_an_unusable_rate_exports_rather_than_dropping(self, bad):
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter, draw=lambda: 0.99)
self._run(provider, (arize_destination(arize_success_sampling_rate=bad),))
assert {s.name for s in dest_exporter.get_finished_spans()} == self.NAMES
def test_a_team_entry_with_rates_is_sampled_from_auth_to_the_fan_out(self, monkeypatch):
"""The whole path the ticket names: the team's callback vars resolve at auth, and the
fan-out honours them."""
monkeypatch.setenv("LITELLM_OTEL_V2", "true")
is_otel_v2_enabled.cache_clear()
auth = UserAPIKeyAuth(
team_metadata={
"logging": [
{
"callback_name": "arize",
"callback_type": "success",
"callback_vars": {
**ARIZE_TEAM_PARAMS,
"arize_success_sampling_rate": "0.0",
"arize_error_sampling_rate": "0.0",
},
}
]
}
)
destinations = resolve_tenant_otel_destinations(auth)
assert destinations, "the fixture must resolve to a destination for the test to mean anything"
dest_exporter = InMemorySpanExporter()
provider = self._fan_out(dest_exporter)
self._run(provider, destinations)
assert dest_exporter.get_finished_spans() == ()
class TestProviderWiring:
def test_build_tracer_provider_only_filters_when_asked(self):
config = OpenTelemetryV2Config(exporters=[ExporterSpec(kind="in_memory", owner=ExporterOwner.LANGFUSE_OTEL)])