diff --git a/litellm/integrations/langfuse/langfuse_sdk.py b/litellm/integrations/langfuse/langfuse_sdk.py index f27d8f5e72c..8dbc2142c41 100644 --- a/litellm/integrations/langfuse/langfuse_sdk.py +++ b/litellm/integrations/langfuse/langfuse_sdk.py @@ -5,6 +5,7 @@ import re import threading from base64 import b64encode from collections.abc import Iterable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor, wait from contextvars import ContextVar from dataclasses import dataclass from datetime import datetime @@ -605,11 +606,20 @@ def acquire_langfuse_tracing( def flush_langfuse_tracing(timeout_millis: int = 30_000) -> bool: - """Force-flush every export channel this process acquired; ``True`` when all of them succeeded.""" + """Force-flush every export channel this process acquired, all within one ``timeout_millis`` deadline. + + ``True`` only when every channel flushed in time; a channel still blocked at the deadline is left to + finish in the background rather than pushing the deadline out for the channels after it. + """ with _TRACING_LOCK: channels: Final = tuple(_TRACING.values()) - results: Final = tuple(channel.flush(timeout_millis) for channel in channels) - return all(results) + if not channels: + return True + pool: Final = ThreadPoolExecutor(max_workers=len(channels), thread_name_prefix="langfuse-flush") + futures: Final = tuple(pool.submit(channel.flush, timeout_millis) for channel in channels) + done, pending = wait(futures, timeout=timeout_millis / 1000) + pool.shutdown(wait=False) + return not pending and all(future.result() for future in done) def build_langfuse_tracing( diff --git a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py b/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py index 9f365097780..9d5005a7876 100644 --- a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py +++ b/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py @@ -7,6 +7,7 @@ otherwise record its own duration instead of the call's. import json import logging +import threading import uuid from base64 import b64encode from datetime import datetime, timedelta, timezone @@ -19,6 +20,7 @@ import pytest from langfuse import LangfuseOtelSpanAttributes as A from langfuse.api import UnauthorizedError from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SpanExporter, SpanExportResult from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from litellm.integrations.langfuse.langfuse import ( @@ -523,6 +525,39 @@ def test_flush_langfuse_tracing_exports_the_queued_spans_of_every_channel(monkey assert [len(exporter.get_finished_spans()) for exporter in exporters] == [1, 1] +def test_flush_langfuse_tracing_flushes_channels_concurrently_under_one_deadline(monkeypatch: pytest.MonkeyPatch): + """A channel stuck on an unreachable host must not spend the whole deadline before the + next channel gets its turn; the first exporter here only returns once the second exported.""" + second_exported = threading.Event() + + class WaitsForTheOther(SpanExporter): + def export(self, spans): + return SpanExportResult.SUCCESS if second_exported.wait(timeout=5.0) else SpanExportResult.FAILURE + + def shutdown(self) -> None: + return None + + class Unblocks(SpanExporter): + def export(self, spans): + second_exported.set() + return SpanExportResult.SUCCESS + + def shutdown(self) -> None: + return None + + exporters = iter((WaitsForTheOther(), Unblocks())) + + def build_next(*, public_key: str, secret_key: str, base_url: str) -> SpanExporter: + return next(exporters) + + monkeypatch.setattr("litellm.integrations.langfuse.langfuse_sdk._build_span_exporter", build_next) + for public_key in ("pk-concurrent-flush-a", "pk-concurrent-flush-b"): + _acquire(public_key=public_key, mock_mode=False, flush_interval=600.0).tracer.start_span("generation").end() + + assert flush_langfuse_tracing(timeout_millis=2_000) is True + assert second_exported.is_set() + + def test_a_changed_sample_rate_rebuilds_the_channel(monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("LANGFUSE_SAMPLE_RATE", "0.25") quarter = _acquire(public_key="pk-resample-test")