diff --git a/langfuse/_client/client.py b/langfuse/_client/client.py index af06e69ef..087172a04 100644 --- a/langfuse/_client/client.py +++ b/langfuse/_client/client.py @@ -1030,6 +1030,9 @@ def start_as_current_observation( The created observation will be the child of the current span in the context. + An asyncio.CancelledError escaping the context marks the observation as ERROR + with the cancellation message and is re-raised to preserve task cancellation. + Args: trace_context: Optional context for connecting to an existing trace name: Name of the observation (e.g., function or operation name) @@ -1389,7 +1392,16 @@ def _start_as_current_otel_span_with_processed_media( ) # For span-like types (span, agent, tool, chain, retriever, evaluator, guardrail), no generation properties needed - yield span_class(**common_args) # type: ignore[arg-type] + observation = span_class(**common_args) # type: ignore[arg-type] + try: + yield observation + except asyncio.CancelledError as error: + if otel_span.is_recording(): + observation.update( + level="ERROR", + status_message=str(error) or type(error).__name__, + ) + raise finally: if baggage_token is not None: diff --git a/tests/unit/test_otel.py b/tests/unit/test_otel.py index 46a085a71..b0cfdba44 100644 --- a/tests/unit/test_otel.py +++ b/tests/unit/test_otel.py @@ -1,9 +1,11 @@ +import asyncio import json from datetime import datetime from hashlib import sha256 from typing import List, Sequence import pytest +from opentelemetry import context as context_api from opentelemetry import trace as trace_api from opentelemetry.sdk.resources import Resource from opentelemetry.sdk.trace import ReadableSpan, TracerProvider @@ -2787,6 +2789,100 @@ def test_blocked_scope_drop_logs_scope_name( ) +@pytest.mark.asyncio +@pytest.mark.parametrize("as_type", ["span", "generation"]) +@pytest.mark.parametrize("cancellation", ["empty", "message", "timeout"]) +async def test_observation_context_records_cancellation( + langfuse_memory_client, get_span, as_type, cancellation +): + client = langfuse_memory_client + ready = asyncio.Event() + restored_context = None + + async def work(): + nonlocal restored_context + try: + with client.start_as_current_observation( + name="cancelled-observation", as_type=as_type + ): + ready.set() + await asyncio.Event().wait() + finally: + restored_context = context_api.get_current() + + with client.start_as_current_observation(name="parent") as parent: + parent_context = context_api.get_current() + task = asyncio.create_task(work()) + await asyncio.wait_for(ready.wait(), timeout=5) + + if cancellation == "timeout": + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(task, timeout=0) + else: + message = "service shutting down" if cancellation == "message" else None + task.cancel(message) + with pytest.raises(asyncio.CancelledError) as exc_info: + await task + assert str(exc_info.value) == (message or "") + + assert task.cancelled() + assert restored_context is parent_context + assert trace_api.get_current_span().get_span_context().span_id == int( + parent.id, 16 + ) + + client.flush() + span = get_span("cancelled-observation") + expected_message = ( + "service shutting down" if cancellation == "message" else "CancelledError" + ) + assert span.attributes.get(LangfuseOtelSpanAttributes.OBSERVATION_LEVEL) == "ERROR" + assert ( + span.attributes[LangfuseOtelSpanAttributes.OBSERVATION_STATUS_MESSAGE] + == expected_message + ) + assert span.status.status_code == trace_api.StatusCode.ERROR + assert span.status.description == expected_message + assert span.end_time is not None + assert span.parent.span_id == int(parent.id, 16) + assert get_span("parent").status.status_code == trace_api.StatusCode.UNSET + + +@pytest.mark.parametrize("as_type", ["span", "generation"]) +def test_observation_cancellation_preserves_manual_end( + langfuse_memory_client, get_span, memory_exporter, as_type +): + client = langfuse_memory_client + cancellation = asyncio.CancelledError("cancelled child") + with client.start_as_current_observation(name="parent") as parent: + parent_context = context_api.get_current() + with pytest.raises(asyncio.CancelledError) as exc_info: + with client.start_as_current_observation( + name="manual-end-child", + as_type=as_type, + trace_context={ + "trace_id": parent.trace_id, + "parent_span_id": parent.id, + }, + end_on_exit=False, + ) as child: + raise cancellation + + assert exc_info.value is cancellation + assert context_api.get_current() is parent_context + assert child._otel_span.is_recording() + client.flush() + assert memory_exporter.get_finished_spans() == [] + child.end() + + client.flush() + span = get_span("manual-end-child") + assert span.attributes.get(LangfuseOtelSpanAttributes.OBSERVATION_LEVEL) == "ERROR" + assert span.status.description == "cancelled child" + assert span.parent.span_id == int(parent.id, 16) + assert len(memory_exporter.get_finished_spans()) == 2 + + class TestConcurrencyAndAsync(TestOTelBase): """Tests for asynchronous and concurrent span operations."""