Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
65 changes: 53 additions & 12 deletions langfuse/_client/observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
Any,
AsyncGenerator,
Callable,
Coroutine,
Dict,
Generator,
Iterable,
Expand Down Expand Up @@ -561,7 +562,7 @@ def _handle_observe_result(


class _ContextPreservedSyncGeneratorWrapper:
"""Sync generator wrapper that ensures each iteration runs in preserved context."""
"""Preserve tracing context across synchronous generator operations."""

def __init__(
self,
Expand Down Expand Up @@ -640,9 +641,19 @@ def __del__(self) -> None:
pass

def __next__(self) -> Any:
return self._advance(method=self.generator.__next__)

def send(self, value: Any) -> Any:
return self._advance(method=self.generator.send, args=(value,))

def throw(self, *args: Any) -> Any:
return self._advance(method=self.generator.throw, args=args)
Comment thread
1fanwang marked this conversation as resolved.

def _advance(
self, *, method: Callable[..., Any], args: Tuple[Any, ...] = ()
) -> Any:
try:
# Run the generator's __next__ in the preserved context
item = self.context.run(next, self.generator)
item: Any = self.context.run(method, *args)
if self.capture_output:
self.items.append(item)

Expand All @@ -652,13 +663,13 @@ def __next__(self) -> Any:
self._finalize()
raise # Re-raise StopIteration

except (Exception, asyncio.CancelledError) as e:
except BaseException as e:
self._finalize_with_error(e)
raise


class _ContextPreservedAsyncGeneratorWrapper:
"""Async generator wrapper that ensures each iteration runs in preserved context."""
"""Preserve tracing context across asynchronous generator operations."""

def __init__(
self,
Expand Down Expand Up @@ -767,19 +778,49 @@ def __del__(self) -> None:
self._finalize()

async def __anext__(self) -> Any:
return await self._advance(method=self.generator.__anext__)

async def asend(self, value: Any) -> Any:
return await self._advance(method=self.generator.asend, args=(value,))

async def athrow(self, *args: Any) -> Any:
return await self._advance(method=self.generator.athrow, args=args)

async def _run_operation(
self, method: Callable[..., Coroutine[Any, Any, Any]], args: Tuple[Any, ...]
) -> Tuple[Any, Optional[BaseException]]:
try:
# Run the generator's __anext__ in the preserved context
return await method(*args), None
except (KeyboardInterrupt, SystemExit) as error:
# Tasks otherwise re-raise these before their awaiter can handle them.
return None, error

async def _advance(
self,
*,
method: Callable[..., Coroutine[Any, Any, Any]],
args: Tuple[Any, ...] = (),
) -> Any:
try:
operation: Coroutine[Any, Any, Tuple[Any, Optional[BaseException]]] = (
self._run_operation(method=method, args=args)
)
item: Any
error: Optional[BaseException]
if _ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT:
item = await asyncio.create_task(
self.generator.__anext__(), # type: ignore
item, error = await asyncio.create_task(
coro=operation,
context=self.context,
) # type: ignore
)
else:
item = await self.context.run(
item, error = await self.context.run(
asyncio.create_task,
self.generator.__anext__(), # type: ignore
operation,
)

if error is not None:
raise error

if self.capture_output:
self.items.append(item)

Expand All @@ -795,6 +836,6 @@ async def __anext__(self) -> Any:
raise
self._finalize_with_error(e)
raise
except Exception as e:
except BaseException as e:
self._finalize_with_error(e)
raise
224 changes: 222 additions & 2 deletions tests/unit/test_observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,18 +3,24 @@
import gc
import inspect
import json
import logging
import sys
from typing import Any, AsyncGenerator, Generator, cast
from contextlib import asynccontextmanager, contextmanager
from tempfile import TemporaryFile
from typing import Any, AsyncGenerator, BinaryIO, Generator, cast

import pytest
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import StatusCode, get_current_span

from langfuse import observe
from langfuse import Langfuse, observe
from langfuse._client import observe as observe_module
from langfuse._client.attributes import LangfuseOtelSpanAttributes
from langfuse._client.observe import (
_ContextPreservedAsyncGeneratorWrapper,
_ContextPreservedSyncGeneratorWrapper,
)
from tests.conftest import InMemorySpanExporter


class SpanRecorder:
Expand All @@ -35,6 +41,220 @@ def _finished_spans_by_name(memory_exporter: Any, name: str) -> list[Any]:
return [span for span in memory_exporter.get_finished_spans() if span.name == name]


@pytest.fixture
def native_memory_client(
monkeypatch: pytest.MonkeyPatch, memory_exporter: InMemorySpanExporter
) -> Generator[Langfuse, None, None]:
monkeypatch.setenv(name="LANGFUSE_PUBLIC_KEY", value="test-generator-protocol")
client = Langfuse(
public_key="test-generator-protocol",
secret_key="test-secret-key",
base_url="http://127.0.0.1:9",
tracer_provider=TracerProvider(),
span_exporter=memory_exporter,
tracing_enabled=True,
sample_rate=1.0,
)
try:
yield client
finally:
client.shutdown()


@pytest.mark.parametrize("suppress", [False, True])
@pytest.mark.parametrize(
"error_type", [ValueError, BaseException, KeyboardInterrupt, SystemExit]
)
def test_observed_context_manager_preserves_exception_handling(
native_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
suppress: bool,
error_type: type[BaseException],
) -> None:
opened_file: BinaryIO | None = None
caught_error: BaseException | None = None
application_error = error_type("application failed")

@contextmanager
@observe(capture_output=False)
def resource() -> Generator[BinaryIO, None, None]:
with TemporaryFile() as handle:
try:
yield handle
except error_type:
if not suppress:
raise

manager = resource()
try:
try:
with manager as handle:
opened_file = handle
raise application_error
except BaseException as error:
caught_error = error

assert opened_file is not None
native_memory_client.flush()
spans = memory_exporter.get_finished_spans()
logging.getLogger(__name__).info(
"sync suppress=%s error=%r resource_closed=%s finished_spans=%s",
suppress,
caught_error,
opened_file.closed,
len(spans),
)
if suppress:
assert caught_error is None
else:
assert caught_error is application_error
assert opened_file.closed
assert len(spans) == 1
assert spans[0].status.status_code == (
StatusCode.UNSET if suppress else StatusCode.ERROR
)
finally:
manager.gen.close()


@pytest.mark.asyncio
@pytest.mark.parametrize("suppress", [False, True])
@pytest.mark.parametrize(
"error_type", [ValueError, BaseException, KeyboardInterrupt, SystemExit]
)
async def test_observed_async_context_manager_preserves_exception_handling(
native_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
suppress: bool,
error_type: type[BaseException],
) -> None:
opened_file: BinaryIO | None = None
caught_error: BaseException | None = None
application_error = error_type("application failed")

@asynccontextmanager
@observe(capture_output=False)
async def resource() -> AsyncGenerator[BinaryIO, None]:
with TemporaryFile() as handle:
try:
yield handle
except error_type:
if not suppress:
raise

manager = resource()
try:
try:
async with manager as handle:
opened_file = handle
raise application_error
except BaseException as error:
caught_error = error

assert opened_file is not None
native_memory_client.flush()
spans = memory_exporter.get_finished_spans()
logging.getLogger(__name__).info(
"async suppress=%s error=%r resource_closed=%s finished_spans=%s",
suppress,
caught_error,
opened_file.closed,
len(spans),
)
if suppress:
assert caught_error is None
else:
assert caught_error is application_error
assert opened_file.closed
assert len(spans) == 1
assert spans[0].status.status_code == (
StatusCode.UNSET if suppress else StatusCode.ERROR
)
finally:
await manager.gen.aclose()


def test_observed_generator_send_and_throw_preserve_context_and_output(
langfuse_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
) -> None:
observation_ids: list[str | None] = []

@observe()
def stream() -> Generator[str, str, None]:
observation_ids.append(langfuse_memory_client.get_current_observation_id())
try:
value = yield "ready"
observation_ids.append(langfuse_memory_client.get_current_observation_id())
yield value
except ValueError:
observation_ids.append(langfuse_memory_client.get_current_observation_id())
yield "recovered"

generator = stream()
try:
assert next(generator) == "ready"
assert generator.send("sent") == "sent"
assert generator.throw(ValueError("recover")) == "recovered"
with pytest.raises(StopIteration):
next(generator)

langfuse_memory_client.flush()
spans = memory_exporter.get_finished_spans()
assert len(spans) == 1
assert len(observation_ids) == 3
assert observation_ids[0] is not None
assert len(set(observation_ids)) == 1
assert not get_current_span().get_span_context().is_valid
assert (
(spans[0].attributes[LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT])
== "readysentrecovered"
)
finally:
generator.close()


@pytest.mark.asyncio
async def test_observed_async_generator_send_and_throw_preserve_context_and_output(
langfuse_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
) -> None:
observation_ids: list[str | None] = []

@observe()
async def stream() -> AsyncGenerator[str, str]:
observation_ids.append(langfuse_memory_client.get_current_observation_id())
try:
value = yield "ready"
observation_ids.append(langfuse_memory_client.get_current_observation_id())
yield value
except ValueError:
observation_ids.append(langfuse_memory_client.get_current_observation_id())
yield "recovered"

generator = stream()
try:
assert await generator.__anext__() == "ready"
assert await generator.asend("sent") == "sent"
assert await generator.athrow(ValueError("recover")) == "recovered"
with pytest.raises(StopAsyncIteration):
await generator.__anext__()

langfuse_memory_client.flush()
spans = memory_exporter.get_finished_spans()
assert len(spans) == 1
assert len(observation_ids) == 3
assert observation_ids[0] is not None
assert len(set(observation_ids)) == 1
assert not get_current_span().get_span_context().is_valid
assert (
(spans[0].attributes[LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT])
== "readysentrecovered"
)
finally:
await generator.aclose()


@pytest.mark.asyncio
async def test_capture_output_false_preserves_type_when_current_span_is_updated(
langfuse_memory_client: Any, memory_exporter: Any
Expand Down