Skip to content
Merged
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
64 changes: 50 additions & 14 deletions langfuse/_client/observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import contextvars
import inspect
import os
import sys
from functools import wraps
from typing import (
Any,
Expand Down Expand Up @@ -48,6 +49,8 @@
P = ParamSpec("P")
R = TypeVar("R")

_ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT = sys.version_info >= (3, 11)


class LangfuseDecorator:
"""Implementation of the @observe decorator for seamless Langfuse tracing integration.
Expand Down Expand Up @@ -616,6 +619,8 @@ def _finalize_with_error(self, error: BaseException) -> None:

def close(self) -> None:
if self._span_ended:
# Still close the generator so cleanup runs in the preserved context, not at GC time.
self.context.run(self.generator.close)
return

try:
Expand Down Expand Up @@ -678,10 +683,23 @@ def __init__(
self.capture_output = capture_output
self.transform_fn = transform_fn
self._span_ended = False
self._pending_error: Optional[BaseException] = None

def __aiter__(self) -> "_ContextPreservedAsyncGeneratorWrapper":
return self

def _generator_never_resumed(self) -> bool:
try:
state = inspect.getasyncgenstate(self.generator)
except (AttributeError, TypeError):
# getasyncgenstate is Python 3.12+; fall back to the attributes it reads.
frame = getattr(self.generator, "ag_frame", None)
return frame is not None and not getattr(
self.generator, "ag_running", False
)

return state in ("AGEN_CREATED", "AGEN_SUSPENDED")

def _finalize(self) -> None:
if self._span_ended:
return
Expand Down Expand Up @@ -711,39 +729,50 @@ def _finalize_with_error(self, error: BaseException) -> None:

async def aclose(self) -> None:
if self._span_ended:
# Still close the generator so cleanup runs in the preserved context, not at GC time.
await self._close_generator()
return

try:
try:
await asyncio.create_task(
self.generator.aclose(),
context=self.context,
) # type: ignore
except TypeError:
await self.context.run(asyncio.create_task, self.generator.aclose())
await self._close_generator()
except (Exception, asyncio.CancelledError) as error:
self._finalize_with_error(error)
raise
else:
self._finalize()
if self._pending_error is not None:
self._finalize_with_error(self._pending_error)
else:
self._finalize()

async def _close_generator(self) -> None:
if _ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT:
close_task = asyncio.create_task(
self.generator.aclose(),
context=self.context,
) # type: ignore
else:
close_task = self.context.run(asyncio.create_task, self.generator.aclose())

await close_task

async def close(self) -> None:
await self.aclose()

def __del__(self) -> None:
self._finalize()
if self._pending_error is not None:
self._finalize_with_error(self._pending_error)
else:
self._finalize()

async def __anext__(self) -> Any:
try:
# Run the generator's __anext__ in the preserved context
try:
# Python 3.11+ approach with explicit task context
if _ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT:
item = await asyncio.create_task(
self.generator.__anext__(), # type: ignore
context=self.context,
) # type: ignore
except TypeError:
# Python 3.10 fallback - create the task inside the preserved context.
else:
item = await self.context.run(
asyncio.create_task,
self.generator.__anext__(), # type: ignore
Expand All @@ -757,6 +786,13 @@ async def __anext__(self) -> Any:
except StopAsyncIteration:
self._finalize()
raise # Re-raise StopAsyncIteration
except (Exception, asyncio.CancelledError) as e:
except asyncio.CancelledError as e:
if self._generator_never_resumed():
# Defer span end so aclose() can run the generator's cleanup first.
self._pending_error = e
raise
self._finalize_with_error(e)
raise
except Exception as e:
self._finalize_with_error(e)
raise
217 changes: 209 additions & 8 deletions tests/unit/test_observe.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
import asyncio
import contextvars
import gc
import inspect
import json
import sys
from typing import Any, AsyncGenerator, Generator, cast

import pytest

from langfuse import observe
from langfuse._client import observe as observe_module
from langfuse._client.attributes import LangfuseOtelSpanAttributes
from langfuse._client.observe import (
_ContextPreservedAsyncGeneratorWrapper,
Expand Down Expand Up @@ -270,21 +272,220 @@ async def generator() -> AsyncGenerator[str, None]:
assert span.ended == 1


def test_sync_generator_wrapper_close_closes_generator_after_span_ended() -> None:
marker = contextvars.ContextVar("marker", default="ambient")
seen: list[str] = []

def generator() -> Generator[str, None, None]:
try:
yield "item_0"
yield "item_1"
finally:
seen.append(marker.get())

span = SpanRecorder()
context = contextvars.copy_context()
context.run(marker.set, "preserved")
wrapper = _ContextPreservedSyncGeneratorWrapper(
generator(),
context,
cast(Any, span),
False,
None,
)

assert next(wrapper) == "item_0"

# An error from __next__ that never resumed the generator ends the span.
with pytest.raises(RuntimeError):
context.run(lambda: next(wrapper))

assert span.ended == 1
assert seen == []

marker.set("ambient-now")
wrapper.close()

assert seen == ["preserved"]
assert span.ended == 1


@pytest.mark.asyncio
async def test_async_generator_wrapper_fallback_preserves_context(
async def test_async_generator_wrapper_aclose_closes_generator_after_span_ended() -> (
None
):
marker = contextvars.ContextVar("marker", default="ambient")
seen: list[str] = []

async def generator() -> AsyncGenerator[str, None]:
try:
yield "item_0"
yield "item_1"
finally:
seen.append(marker.get())

span = SpanRecorder()
context = contextvars.copy_context()
context.run(marker.set, "preserved")
wrapper = _ContextPreservedAsyncGeneratorWrapper(
generator(),
context,
cast(Any, span),
False,
None,
)

assert await wrapper.__anext__() == "item_0"

# Span ends while the generator is still suspended.
wrapper._finalize_with_error(RuntimeError("ended early"))
assert span.ended == 1
assert seen == []

marker.set("ambient-now")
await wrapper.aclose()

assert seen == ["preserved"]
assert span.ended == 1


@pytest.mark.asyncio
async def test_async_generator_wrapper_defers_span_end_for_unresumed_cancel() -> None:
marker = contextvars.ContextVar("marker", default="ambient")
seen: list[str] = []
cleanup_span_states: list[int] = []

async def generator() -> AsyncGenerator[str, None]:
try:
yield "item_0"
yield "item_1"
finally:
seen.append(marker.get())
cleanup_span_states.append(span.ended)
span.update(cleanup=True)

span = SpanRecorder()
context = contextvars.copy_context()
context.run(marker.set, "preserved")
raw = generator()
wrapper = _ContextPreservedAsyncGeneratorWrapper(
raw,
context,
cast(Any, span),
False,
None,
)

async def consume() -> None:
async for _ in wrapper:
await asyncio.sleep(0)

consumer = asyncio.create_task(consume())
for _ in range(4):
await asyncio.sleep(0)
# Cancel lands before the inner __anext__ task's first step.
asyncio.get_running_loop().call_soon(consumer.cancel)
with pytest.raises(asyncio.CancelledError):
await consumer

assert raw.ag_frame is not None # still suspended, never resumed
# Span end is deferred so the generator's cleanup can still update it.
assert span.ended == 0

marker.set("ambient-now")
await wrapper.aclose()

assert raw.ag_frame is None # closed
assert seen == ["preserved"]
assert cleanup_span_states == [0]
assert span.ended == 1
assert span.updates == [
{"cleanup": True},
{"level": "ERROR", "status_message": "CancelledError"},
]


@pytest.mark.asyncio
async def test_async_generator_wrapper_defers_unresumed_cancel_without_inspect_api(
monkeypatch: pytest.MonkeyPatch,
) -> None:
marker = contextvars.ContextVar("marker", default="ambient")
# Python < 3.12 has no inspect.getasyncgenstate; the fallback must still defer.
monkeypatch.delattr(inspect, "getasyncgenstate", raising=False)

seen: list[str] = []
original_create_task = asyncio.create_task

def create_task_with_type_error(*args: Any, **kwargs: Any) -> asyncio.Task[Any]:
if "context" in kwargs:
raise TypeError("context argument unsupported")
async def generator() -> AsyncGenerator[str, None]:
try:
yield "item_0"
yield "item_1"
finally:
seen.append("closed")

span = SpanRecorder()
wrapper = _ContextPreservedAsyncGeneratorWrapper(
generator(),
contextvars.copy_context(),
cast(Any, span),
False,
None,
)

async def consume() -> None:
async for _ in wrapper:
await asyncio.sleep(0)

consumer = asyncio.create_task(consume())
for _ in range(4):
await asyncio.sleep(0)
asyncio.get_running_loop().call_soon(consumer.cancel)
with pytest.raises(asyncio.CancelledError):
await consumer

assert span.ended == 0

await wrapper.aclose()

assert seen == ["closed"]
assert span.ended == 1
assert span.updates[-1] == {
"level": "ERROR",
"status_message": "CancelledError",
}

return original_create_task(*args, **kwargs)

monkeypatch.setattr(asyncio, "create_task", create_task_with_type_error)
@pytest.mark.asyncio
async def test_async_generator_wrapper_aclose_propagates_cleanup_type_error() -> None:
async def generator() -> AsyncGenerator[str, None]:
try:
yield "item_0"
finally:
raise TypeError("cleanup failed")

span = SpanRecorder()
wrapper = _ContextPreservedAsyncGeneratorWrapper(
generator(),
contextvars.copy_context(),
cast(Any, span),
False,
None,
)

assert await wrapper.__anext__() == "item_0"

with pytest.raises(TypeError, match="cleanup failed"):
await wrapper.aclose()

assert span.ended == 1
assert span.updates[-1] == {"level": "ERROR", "status_message": "cleanup failed"}


@pytest.mark.asyncio
async def test_async_generator_wrapper_fallback_preserves_context(
monkeypatch: pytest.MonkeyPatch,
) -> None:
marker = contextvars.ContextVar("marker", default="ambient")
seen: list[str] = []
monkeypatch.setattr(observe_module, "_ASYNCIO_CREATE_TASK_SUPPORTS_CONTEXT", False)

async def generator() -> AsyncGenerator[str, None]:
try:
Expand Down
Loading