Skip to content
Closed
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
24 changes: 17 additions & 7 deletions langfuse/_client/observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -616,6 +616,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 @@ -711,22 +713,30 @@ 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()

async def _close_generator(self) -> None:
try:
close_task = asyncio.create_task(
self.generator.aclose(),
context=self.context,
) # type: ignore
except TypeError:
# Python 3.10 create_task has no context param; guard only the call itself.
close_task = self.context.run(asyncio.create_task, self.generator.aclose())

await close_task

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

Expand Down
151 changes: 151 additions & 0 deletions tests/unit/test_observe.py
Original file line number Diff line number Diff line change
Expand Up @@ -270,6 +270,157 @@ 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_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 (e.g. the cancel race below).
wrapper._finalize_with_error(asyncio.CancelledError())
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_closes_generator_cancel_never_resumed() -> 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")
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
assert span.ended == 1
assert seen == []

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

assert raw.ag_frame is None # closed
assert seen == ["preserved"]
assert span.ended == 1


@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,
Expand Down