Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
Next Next commit
fix(observe): preserve generator control methods
Signed-off-by: 1fanwang <1fannnw@gmail.com>
  • Loading branch information
1fanwang committed Sep 19, 2026
commit 78c12d60eeb437c4628a00ecd5956cb686a92934
43 changes: 34 additions & 9 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 @@ -658,7 +669,7 @@ def __next__(self) -> Any:


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,17 +778,31 @@ 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 _advance(
self,
*,
method: Callable[..., Coroutine[Any, Any, Any]],
args: Tuple[Any, ...] = (),
) -> Any:
try:
# Run the generator's __anext__ in the preserved context
operation: Coroutine[Any, Any, Any] = method(*args)
if _ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT:
item = await asyncio.create_task(
self.generator.__anext__(), # type: ignore
item: Any = await asyncio.create_task(
coro=operation,
context=self.context,
) # type: ignore
)
else:
item = await self.context.run(
asyncio.create_task,
self.generator.__anext__(), # type: ignore
operation,
)

if self.capture_output:
Expand Down
167 changes: 166 additions & 1 deletion tests/unit/test_observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,17 +4,20 @@
import inspect
import json
import sys
from contextlib import asynccontextmanager, contextmanager, nullcontext
from typing import Any, AsyncGenerator, Generator, cast

import pytest
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 +38,168 @@ 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.mark.parametrize("suppress", [False, True])
def test_observed_context_manager_preserves_exception_handling(
langfuse_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
suppress: bool,
) -> None:
closed: list[bool] = []

@contextmanager
@observe(capture_output=False)
def resource() -> Generator[None, None, None]:
try:
yield
except ValueError:
if not suppress:
raise
finally:
closed.append(True)

manager = resource()
try:
with (
nullcontext()
if suppress
else pytest.raises(ValueError, match="application failed")
):
with manager:
raise ValueError("application failed")

assert closed == [True]
langfuse_memory_client.flush()
spans = memory_exporter.get_finished_spans()
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])
async def test_observed_async_context_manager_preserves_exception_handling(
langfuse_memory_client: Langfuse,
memory_exporter: InMemorySpanExporter,
suppress: bool,
) -> None:
closed: list[bool] = []

@asynccontextmanager
@observe(capture_output=False)
async def resource() -> AsyncGenerator[None, None]:
try:
yield
except ValueError:
if not suppress:
raise
finally:
closed.append(True)

manager = resource()
try:
with (
nullcontext()
if suppress
else pytest.raises(ValueError, match="application failed")
):
async with manager:
raise ValueError("application failed")

assert closed == [True]
langfuse_memory_client.flush()
spans = memory_exporter.get_finished_spans()
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