diff --git a/litellm/integrations/langfuse/langfuse_sdk.py b/litellm/integrations/langfuse/langfuse_sdk.py index 8dbc2142c41..5c2f12cb1d4 100644 --- a/litellm/integrations/langfuse/langfuse_sdk.py +++ b/litellm/integrations/langfuse/langfuse_sdk.py @@ -5,14 +5,13 @@ 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 from hashlib import sha256 from importlib.metadata import version from itertools import chain -from time import sleep +from time import monotonic, sleep from types import MappingProxyType from typing import Final, Literal @@ -605,6 +604,19 @@ def acquire_langfuse_tracing( return created +class _FlushWorker(threading.Thread): + """Daemon, so a channel still blocked at the deadline cannot hold up interpreter exit.""" + + def __init__(self, channel: LangfuseTracing, timeout_millis: int) -> None: + super().__init__(name="langfuse-flush", daemon=True) + self.channel: Final = channel + self.timeout_millis: Final = timeout_millis + self.flushed = False + + def run(self) -> None: + self.flushed = self.channel.flush(self.timeout_millis) + + def flush_langfuse_tracing(timeout_millis: int = 30_000) -> bool: """Force-flush every export channel this process acquired, all within one ``timeout_millis`` deadline. @@ -613,13 +625,13 @@ def flush_langfuse_tracing(timeout_millis: int = 30_000) -> bool: """ with _TRACING_LOCK: channels: Final = tuple(_TRACING.values()) - 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) + workers: Final = tuple(_FlushWorker(channel, timeout_millis) for channel in channels) + deadline: Final = monotonic() + timeout_millis / 1000 + for worker in workers: + worker.start() + for worker in workers: + worker.join(max(0.0, deadline - monotonic())) + return all(not worker.is_alive() and worker.flushed for worker in workers) 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 9d5005a7876..b22a3f8be8e 100644 --- a/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py +++ b/tests/test_litellm/integrations/langfuse/test_langfuse_sdk.py @@ -19,7 +19,7 @@ import opentelemetry.trace as otel_trace import pytest from langfuse import LangfuseOtelSpanAttributes as A from langfuse.api import UnauthorizedError -from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace import SpanProcessor, TracerProvider from opentelemetry.sdk.trace.export import SpanExporter, SpanExportResult from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter @@ -558,6 +558,26 @@ def test_flush_langfuse_tracing_flushes_channels_concurrently_under_one_deadline assert second_exported.is_set() +def test_flush_langfuse_tracing_leaves_an_overrunning_channel_on_a_daemon_thread(): + """A channel whose flush outlives the deadline is reported as failed and must not be able to + hold up interpreter exit, so the thread still flushing it has to be a daemon.""" + release = threading.Event() + + class BlocksUntilReleased(SpanProcessor): + def force_flush(self, timeout_millis: int = 30_000) -> bool: + return release.wait(timeout=10.0) + + _acquire(public_key="pk-overrunning-flush", mock_mode=True, flush_interval=600.0).provider.add_span_processor( + BlocksUntilReleased() + ) + try: + assert flush_langfuse_tracing(timeout_millis=200) is False + stuck = [thread for thread in threading.enumerate() if thread.name.startswith("langfuse-flush")] + assert stuck and all(thread.daemon for thread in stuck) + finally: + release.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")