diff --git a/langfuse/_client/observe.py b/langfuse/_client/observe.py index bd0a3edee..3e3837d8c 100644 --- a/langfuse/_client/observe.py +++ b/langfuse/_client/observe.py @@ -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: @@ -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() diff --git a/tests/unit/test_observe.py b/tests/unit/test_observe.py index 5527be9b9..4830601fc 100644 --- a/tests/unit/test_observe.py +++ b/tests/unit/test_observe.py @@ -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,