From 46d2c2a6ea79584226f77c81eef10c350d01c270 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 6 Oct 2026 09:23:11 -0500 Subject: [PATCH] fix(otel): honor per-team Arize sampling rates in OTel v2 fan-out (#44595) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/integrations/arize/arize.py | 51 +- .../integrations/otel/model/destination.py | 17 +- .../integrations/otel/plumbing/providers.py | 167 ++++++- .../integrations/otel/presets/destinations.py | 21 + tests/unit/integrations/arize/test_arize.py | 10 +- .../otel/test_otel_v2_destinations.py | 463 ++++++++++++++++++ 6 files changed, 695 insertions(+), 34 deletions(-) diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py index 57e60fea759..58024cb2e1e 100644 --- a/litellm/integrations/arize/arize.py +++ b/litellm/integrations/arize/arize.py @@ -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) diff --git a/litellm/integrations/otel/model/destination.py b/litellm/integrations/otel/model/destination.py index 4c8b3b1118d..56b3123660c 100644 --- a/litellm/integrations/otel/model/destination.py +++ b/litellm/integrations/otel/model/destination.py @@ -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, diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index 531cd528739..13bfa5a20ea 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -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 diff --git a/litellm/integrations/otel/presets/destinations.py b/litellm/integrations/otel/presets/destinations.py index f3cbe0c8f9c..67450684658 100644 --- a/litellm/integrations/otel/presets/destinations.py +++ b/litellm/integrations/otel/presets/destinations.py @@ -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, ) diff --git a/tests/unit/integrations/arize/test_arize.py b/tests/unit/integrations/arize/test_arize.py index 5fde627680d..e9cab65c545 100644 --- a/tests/unit/integrations/arize/test_arize.py +++ b/tests/unit/integrations/arize/test_arize.py @@ -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 diff --git a/tests/unit/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py index 39c1bb07483..b80d75adc0b 100644 --- a/tests/unit/integrations/otel/test_otel_v2_destinations.py +++ b/tests/unit/integrations/otel/test_otel_v2_destinations.py @@ -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)])