mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
fix(langfuse): sample on a hash of the full trace id and tolerate bad LANGFUSE_SAMPLE_RATE
TraceIdRatioBased reads the low 64 bits of the trace id. litellm trace ids are UUIDs, whose variant bits sit at the top of that word, so every fractional rate up to 0.5 dropped all traces. A SHA-256 of the full id gives an unbiased, deterministic decision. Values outside [0, 1] or non numeric now warn and export everything instead of raising during callback construction, which surfaced as a 500 on the first request of each worker Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
fc180f9c28
commit
78bc59b9b7
2 changed files with 142 additions and 28 deletions
|
|
@ -25,7 +25,9 @@ from opentelemetry.sdk.resources import Resource
|
|||
from opentelemetry.sdk.trace import ReadableSpan, TracerProvider
|
||||
from opentelemetry.sdk.trace.export import SpanExporter, SpanExportResult
|
||||
from opentelemetry.sdk.trace.id_generator import RandomIdGenerator
|
||||
from opentelemetry.sdk.trace.sampling import TraceIdRatioBased
|
||||
from opentelemetry.sdk.trace.sampling import Decision, Sampler, SamplingResult
|
||||
from opentelemetry.trace import Link, SpanKind, TraceState
|
||||
from opentelemetry.util.types import Attributes
|
||||
from requests import RequestException
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -36,8 +38,10 @@ __all__ = (
|
|||
"RELEASE_ATTRIBUTE",
|
||||
"DiscardingSpanExporter",
|
||||
"RetryingSpanExporter",
|
||||
"TraceIdHashSampler",
|
||||
"acquire_langfuse_client",
|
||||
"build_isolated_tracer_provider",
|
||||
"configured_sample_rate",
|
||||
"evict_stale_langfuse_resources",
|
||||
"lease_langfuse_client",
|
||||
"open_trace_context",
|
||||
|
|
@ -201,7 +205,64 @@ class _RequestedSpanIdGenerator(RandomIdGenerator):
|
|||
_litellm_built_providers: Final[WeakSet] = WeakSet()
|
||||
|
||||
|
||||
def build_isolated_tracer_provider(*, environment: str | None, release: str | None) -> TracerProvider:
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TraceIdHashSampler(Sampler):
|
||||
"""Sample on a SHA-256 of the trace id rather than its low 64 bits.
|
||||
|
||||
litellm trace ids are UUIDs, whose variant bits pin the top of that low word, so
|
||||
``TraceIdRatioBased`` drops every trace at rates up to 0.5 and skews above it.
|
||||
"""
|
||||
|
||||
rate: float
|
||||
|
||||
def should_sample(
|
||||
self,
|
||||
parent_context: Context | None,
|
||||
trace_id: int,
|
||||
name: str,
|
||||
kind: SpanKind | None = None,
|
||||
attributes: Attributes = None,
|
||||
links: Sequence[Link] | None = None,
|
||||
trace_state: TraceState | None = None,
|
||||
) -> SamplingResult:
|
||||
digest: Final = sha256(trace_id.to_bytes(16, "big")).digest()
|
||||
sampled: Final = int.from_bytes(digest[:8], "big") < round(self.rate * 2**64)
|
||||
parent: Final = otel_trace.get_current_span(parent_context).get_span_context()
|
||||
return SamplingResult(
|
||||
Decision.RECORD_AND_SAMPLE if sampled else Decision.DROP,
|
||||
attributes if sampled else None,
|
||||
parent.trace_state if parent.is_valid else None,
|
||||
)
|
||||
|
||||
def get_description(self) -> str:
|
||||
return f"TraceIdHashSampler{{{self.rate}}}"
|
||||
|
||||
|
||||
def _parse_sample_rate(raw: str) -> float | None:
|
||||
try:
|
||||
rate: Final = float(raw)
|
||||
except ValueError:
|
||||
return None
|
||||
return rate if 0.0 <= rate <= 1.0 else None
|
||||
|
||||
|
||||
def configured_sample_rate() -> float:
|
||||
"""``LANGFUSE_SAMPLE_RATE`` as a fraction, exporting everything when it is unset or unusable."""
|
||||
raw: Final = os.environ.get("LANGFUSE_SAMPLE_RATE")
|
||||
if raw is None:
|
||||
return 1.0
|
||||
parsed: Final = _parse_sample_rate(raw)
|
||||
if parsed is None:
|
||||
verbose_logger.warning(
|
||||
"LANGFUSE_SAMPLE_RATE=%r is not a number between 0.0 and 1.0; ignoring it and exporting every trace", raw
|
||||
)
|
||||
return 1.0
|
||||
return parsed
|
||||
|
||||
|
||||
def build_isolated_tracer_provider(
|
||||
*, environment: str | None, release: str | None, sample_rate: float = 1.0
|
||||
) -> TracerProvider:
|
||||
"""Give the langfuse client a provider of its own instead of the process-wide one.
|
||||
|
||||
v4 is built on OpenTelemetry and otherwise either claims the global tracer
|
||||
|
|
@ -211,13 +272,9 @@ def build_isolated_tracer_provider(*, environment: str | None, release: str | No
|
|||
|
||||
The resource is rebuilt here because langfuse only applies ``environment``
|
||||
and ``release`` when it constructs the provider itself, and the sampler is
|
||||
rebuilt for the same reason: ``LANGFUSE_SAMPLE_RATE`` is otherwise silently
|
||||
installed for the same reason: ``sample_rate`` is otherwise silently
|
||||
ignored and every trace exports.
|
||||
"""
|
||||
raw_sample_rate: Final = os.environ.get("LANGFUSE_SAMPLE_RATE")
|
||||
sample_rate: Final = float(raw_sample_rate) if raw_sample_rate is not None else 1.0
|
||||
if not 0.0 <= sample_rate <= 1.0:
|
||||
raise ValueError(f"Sample rate must be between 0.0 and 1.0, got {sample_rate}")
|
||||
attributes: Final = MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
|
|
@ -227,7 +284,7 @@ def build_isolated_tracer_provider(*, environment: str | None, release: str | No
|
|||
)
|
||||
provider: Final = TracerProvider(
|
||||
resource=Resource.create(dict(attributes)),
|
||||
sampler=TraceIdRatioBased(sample_rate) if sample_rate < 1 else None,
|
||||
sampler=TraceIdHashSampler(sample_rate) if sample_rate < 1 else None,
|
||||
id_generator=_RequestedSpanIdGenerator(),
|
||||
)
|
||||
with _LIVE_CLIENTS_LOCK:
|
||||
|
|
@ -549,6 +606,7 @@ def acquire_langfuse_client(
|
|||
base_url=parameters.get("base_url"),
|
||||
)
|
||||
)
|
||||
sample_rate: Final = configured_sample_rate()
|
||||
with LangfuseResourceManager._lock: # pyright: ignore[reportPrivateUsage] # registry has no public accessor
|
||||
cached: Final = _evict_if_stale_locked(
|
||||
public_key=public_key,
|
||||
|
|
@ -557,9 +615,10 @@ def acquire_langfuse_client(
|
|||
)
|
||||
client: Final = Langfuse(
|
||||
**parameters, # pyright: ignore[reportArgumentType] # kwargs-ok: dict mirrors the typed ctor, values resolved by the callers
|
||||
sample_rate=sample_rate,
|
||||
tracer_provider=None
|
||||
if cached is not None
|
||||
else build_isolated_tracer_provider(environment=environment, release=release),
|
||||
else build_isolated_tracer_provider(environment=environment, release=release, sample_rate=sample_rate),
|
||||
span_exporter=span_exporter,
|
||||
)
|
||||
register_langfuse_client(client)
|
||||
|
|
|
|||
|
|
@ -6,8 +6,11 @@ call would otherwise record its own duration instead of the call's.
|
|||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
|
||||
import opentelemetry.trace as otel_trace
|
||||
import pytest
|
||||
|
|
@ -29,7 +32,9 @@ from litellm.integrations.langfuse.langfuse_sdk import (
|
|||
_lifecycle_state,
|
||||
_litellm_built_providers,
|
||||
_teardown_langfuse_client,
|
||||
acquire_langfuse_client,
|
||||
build_isolated_tracer_provider,
|
||||
configured_sample_rate,
|
||||
evict_stale_langfuse_resources,
|
||||
lease_langfuse_client,
|
||||
open_trace_context,
|
||||
|
|
@ -371,27 +376,77 @@ def test_isolated_provider_carries_environment_and_release():
|
|||
assert attributes["langfuse.release"] == "v9"
|
||||
|
||||
|
||||
def test_langfuse_sample_rate_drops_spans_on_the_isolated_provider(monkeypatch):
|
||||
"""The SDK only installs its sampler on providers it builds itself; v2 sampled via the same env var."""
|
||||
monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "0")
|
||||
dropped_exporter = InMemorySpanExporter()
|
||||
dropping_provider = build_isolated_tracer_provider(environment=None, release=None)
|
||||
dropping_provider.add_span_processor(SimpleSpanProcessor(dropped_exporter))
|
||||
dropping_provider.get_tracer("test").start_span("dropped").end()
|
||||
assert not dropped_exporter.get_finished_spans()
|
||||
|
||||
monkeypatch.delenv("LANGFUSE_SAMPLE_RATE")
|
||||
kept_exporter = InMemorySpanExporter()
|
||||
keeping_provider = build_isolated_tracer_provider(environment=None, release=None)
|
||||
keeping_provider.add_span_processor(SimpleSpanProcessor(kept_exporter))
|
||||
keeping_provider.get_tracer("test").start_span("kept").end()
|
||||
assert [span.name for span in kept_exporter.get_finished_spans()] == ["kept"]
|
||||
def _generations_exported_at(sample_rate: float, trace_ids: tuple[str, ...]) -> frozenset[str]:
|
||||
exporter: Final = InMemorySpanExporter()
|
||||
provider: Final = build_isolated_tracer_provider(environment=None, release=None, sample_rate=sample_rate)
|
||||
provider.add_span_processor(SimpleSpanProcessor(exporter))
|
||||
pk: Final = f"pk-sample-{sample_rate}"
|
||||
LangfuseResourceManager._instances.pop(pk, None)
|
||||
client: Final = Langfuse(
|
||||
public_key=pk,
|
||||
secret_key="sk-sample",
|
||||
host="http://127.0.0.1:1",
|
||||
tracer_provider=provider,
|
||||
span_exporter=exporter,
|
||||
)
|
||||
try:
|
||||
for trace_id in trace_ids:
|
||||
context, claim_root = open_trace_context(client=client, trace_id=trace_id, parent_observation_id=None)
|
||||
start_generation(
|
||||
client=client,
|
||||
context=context,
|
||||
name="sampled",
|
||||
start_time=CALL_START,
|
||||
claim_trace_root=claim_root,
|
||||
attributes={},
|
||||
).end()
|
||||
finally:
|
||||
LangfuseResourceManager._instances.pop(pk, None)
|
||||
return frozenset(format(span.context.trace_id, "032x") for span in exporter.get_finished_spans())
|
||||
|
||||
|
||||
def test_invalid_sample_rate_fails_at_construction_like_the_sdk(monkeypatch):
|
||||
monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "1.5")
|
||||
with pytest.raises(ValueError, match=r"between 0\.0 and 1\.0"):
|
||||
build_isolated_tracer_provider(environment=None, release=None)
|
||||
def test_sample_rate_zero_drops_and_one_keeps_every_trace():
|
||||
trace_ids: Final = tuple(resolve_trace_id(uuid.uuid4()) for _ in range(20))
|
||||
assert _generations_exported_at(0, trace_ids) == frozenset()
|
||||
assert _generations_exported_at(1, trace_ids) == frozenset(trace_ids)
|
||||
|
||||
|
||||
def test_fractional_sample_rate_keeps_a_deterministic_share_of_uuid_trace_ids():
|
||||
trace_ids: Final = tuple(resolve_trace_id(uuid.uuid4()) for _ in range(400))
|
||||
kept: Final = _generations_exported_at(0.5, trace_ids)
|
||||
assert 140 <= len(kept) <= 260
|
||||
assert _generations_exported_at(0.5, trace_ids) == kept
|
||||
assert kept < _generations_exported_at(0.9, trace_ids)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raw", ["1.5", "-0.5", "abc"])
|
||||
def test_unusable_sample_rate_warns_and_exports_everything(
|
||||
monkeypatch: pytest.MonkeyPatch, raw: str, caplog: pytest.LogCaptureFixture
|
||||
):
|
||||
monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", raw)
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
|
||||
assert configured_sample_rate() == 1.0
|
||||
assert "LANGFUSE_SAMPLE_RATE" in caplog.text
|
||||
|
||||
pk: Final = f"pk-unusable-rate-{raw}"
|
||||
LangfuseResourceManager._instances.pop(pk, None)
|
||||
try:
|
||||
client: Final = acquire_langfuse_client(
|
||||
parameters={"public_key": pk, "secret_key": "sk", "base_url": "http://127.0.0.1:1"},
|
||||
environment=None,
|
||||
release=None,
|
||||
mock_mode=True,
|
||||
)
|
||||
assert client._resources.sample_rate == 1.0
|
||||
finally:
|
||||
LangfuseResourceManager._instances.pop(pk, None)
|
||||
|
||||
|
||||
def test_configured_sample_rate_reads_the_env_var(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.delenv("LANGFUSE_SAMPLE_RATE", raising=False)
|
||||
assert configured_sample_rate() == 1.0
|
||||
monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "0.25")
|
||||
assert configured_sample_rate() == 0.25
|
||||
|
||||
|
||||
def _isolated_client_with_exporter():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue