diff --git a/CHANGELOG.md b/CHANGELOG.md index 19f761152..ab0b4e69a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -56,6 +56,8 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- OpenAI Agents tracing preserves caller spans across workers and restores context + after activity, local-activity, and child-workflow completion callbacks. - `contrib.deepagents`: summarization middleware configured with a model name string now routes its LLM calls through Activities instead of running them in the Workflow. - **Experimental**: External storage metrics now report the wall-clock time storage was in flight. diff --git a/temporalio/contrib/openai_agents/_otel_trace_interceptor.py b/temporalio/contrib/openai_agents/_otel_trace_interceptor.py index 63f8f9d83..c329cbd89 100644 --- a/temporalio/contrib/openai_agents/_otel_trace_interceptor.py +++ b/temporalio/contrib/openai_agents/_otel_trace_interceptor.py @@ -2,13 +2,14 @@ from __future__ import annotations +from collections.abc import Iterator +from contextlib import contextmanager from typing import Any import opentelemetry.trace import temporalio.converter -from ..opentelemetry._id_generator import TemporalIdGenerator from ._trace_interceptor import ( OpenAIAgentsContextPropagationInterceptor, _InputWithHeaders, @@ -22,19 +23,19 @@ class OTelOpenAIAgentsContextPropagationInterceptor( def __init__( self, - otel_id_generator: TemporalIdGenerator, payload_converter: temporalio.converter.PayloadConverter = temporalio.converter.default().payload_converter, add_temporal_spans: bool = True, ) -> None: """Initialize OTEL-aware context propagation interceptor. Args: - otel_id_generator: Generator for OTEL-compatible IDs. payload_converter: Converter for serializing trace context. add_temporal_spans: Whether to add Temporal-specific spans. """ - super().__init__(payload_converter, add_temporal_spans, start_traces=True) - self._otel_id_generator = otel_id_generator + super().__init__( + payload_converter=payload_converter, + add_temporal_spans=add_temporal_spans, + ) def header_contents(self) -> dict[str, Any]: """Get header contents enhanced with OpenTelemetry span context. @@ -50,39 +51,46 @@ def header_contents(self) -> dict[str, Any]: **super().header_contents(), "otelSpanId": span_context.span_id, "otelTraceId": span_context.trace_id, + "otelTraceFlags": int(span_context.trace_flags), + "otelTraceState": span_context.trace_state.to_header(), } else: return super().header_contents() + @contextmanager def context_from_header( self, input: _InputWithHeaders, - ): - """Extracts and initializes trace information the input header.""" + ) -> Iterator[None]: + """Use propagated IDs as a remote parent without recording replicas.""" span_info = self.get_header_contents(input) - - if span_info is None: - return - otel_span_id = span_info.get("otelSpanId") - otel_trace_id = span_info.get("otelTraceId") - - # Seed the trace id before the trace is reconstructed so the workflow's root - # OTEL span shares the caller's trace id rather than generating a new one. - if otel_trace_id and self._otel_id_generator: - self._otel_id_generator.seed_trace_id(otel_trace_id) - - # If only a trace was propagated from the caller, we need to seed for trace context - if otel_span_id and self._otel_id_generator and span_info.get("spanId") is None: - self._otel_id_generator.seed_span_id(otel_span_id) - - super().trace_context_from_header_contents(span_info) - - # If a span was propagated from the caller, we need to seed for span context - if ( - otel_span_id - and self._otel_id_generator - and span_info.get("spanId") is not None - ): - self._otel_id_generator.seed_span_id(otel_span_id) - - super().span_context_from_header_contents(span_info) + with super().context_from_header(input=input): + if ( + span_info is not None + and span_info.get("otelSpanId") + and span_info.get("otelTraceId") + ): + span_context: opentelemetry.trace.SpanContext = ( + opentelemetry.trace.SpanContext( + trace_id=span_info["otelTraceId"], + span_id=span_info["otelSpanId"], + is_remote=True, + trace_flags=opentelemetry.trace.TraceFlags( + span_info.get( + "otelTraceFlags", opentelemetry.trace.TraceFlags.SAMPLED + ) + ), + trace_state=opentelemetry.trace.TraceState.from_header( + [span_info["otelTraceState"]] + if span_info.get("otelTraceState") + else [] + ), + ) + ) + with opentelemetry.trace.use_span( + opentelemetry.trace.NonRecordingSpan(span_context), + end_on_exit=False, + ): + yield + else: + yield diff --git a/temporalio/contrib/openai_agents/_temporal_openai_agents.py b/temporalio/contrib/openai_agents/_temporal_openai_agents.py index 6023ad090..045c9618d 100644 --- a/temporalio/contrib/openai_agents/_temporal_openai_agents.py +++ b/temporalio/contrib/openai_agents/_temporal_openai_agents.py @@ -6,7 +6,9 @@ import typing from collections.abc import AsyncIterator, Callable, Collection, Iterator, Sequence from contextlib import asynccontextmanager, contextmanager +from contextvars import Token from datetime import timedelta +from weakref import WeakKeyDictionary import pydantic from agents import ModelProvider, Trace, set_trace_provider @@ -52,6 +54,10 @@ from temporalio.worker.workflow_sandbox import SandboxedWorkflowRunner if typing.TYPE_CHECKING: + from openinference.instrumentation.openai_agents._processor import ( + OpenInferenceTracingProcessor, + ) + from temporalio.contrib.openai_agents import ( SandboxClientProvider, StatefulMCPServerProvider, @@ -62,6 +68,9 @@ _otel_trace_start_patch_lock = threading.RLock() _otel_trace_start_patch_ref_count = 0 _otel_trace_start_original: Callable[..., typing.Any] | None = None +_otel_trace_end_original: ( + Callable[["OpenInferenceTracingProcessor", Trace], None] | None +) = None _otel_trace_start_instrumentor: typing.Any | None = None @@ -69,25 +78,41 @@ def _install_otel_instrumentation(tracer_provider: typing.Any) -> None: """Configure OpenInference while at least one tracing context is active.""" global _otel_trace_start_instrumentor global _otel_trace_start_original + global _otel_trace_end_original global _otel_trace_start_patch_ref_count from openinference.instrumentation.openai_agents import OpenAIAgentsInstrumentor from openinference.instrumentation.openai_agents._processor import ( OpenInferenceTracingProcessor, ) - from opentelemetry.context import attach + from opentelemetry.context import Context, attach, detach from opentelemetry.trace import set_span_in_context with _otel_trace_start_patch_lock: if _otel_trace_start_patch_ref_count == 0: original_on_trace_start = OpenInferenceTracingProcessor.on_trace_start + original_on_trace_end = OpenInferenceTracingProcessor.on_trace_end _otel_trace_start_original = original_on_trace_start + _otel_trace_end_original = original_on_trace_end + trace_tokens: WeakKeyDictionary[Trace, Token[Context]] = WeakKeyDictionary() + + def on_trace_start( + self: OpenInferenceTracingProcessor, trace: Trace + ) -> None: + original_on_trace_start(self=self, trace=trace) + trace_tokens[trace] = attach( + set_span_in_context(self._root_spans[trace.trace_id]) + ) - def on_trace_start(self: typing.Any, trace: Trace) -> None: # type: ignore[reportUnusedFunction] - original_on_trace_start(self, trace) - attach(set_span_in_context(self._root_spans[trace.trace_id])) + def on_trace_end(self: OpenInferenceTracingProcessor, trace: Trace) -> None: + try: + original_on_trace_end(self=self, trace=trace) + finally: + if token := trace_tokens.pop(trace, None): + detach(token) setattr(OpenInferenceTracingProcessor, "on_trace_start", on_trace_start) + setattr(OpenInferenceTracingProcessor, "on_trace_end", on_trace_end) try: _otel_trace_start_instrumentor = OpenAIAgentsInstrumentor() _otel_trace_start_instrumentor.instrument( @@ -99,7 +124,13 @@ def on_trace_start(self: typing.Any, trace: Trace) -> None: # type: ignore[repo "on_trace_start", _otel_trace_start_original, ) + setattr( + OpenInferenceTracingProcessor, + "on_trace_end", + _otel_trace_end_original, + ) _otel_trace_start_original = None + _otel_trace_end_original = None _otel_trace_start_instrumentor = None raise _otel_trace_start_patch_ref_count += 1 @@ -109,6 +140,7 @@ def _uninstall_otel_instrumentation() -> None: """Tear down OpenInference after the final tracing context exits.""" global _otel_trace_start_instrumentor global _otel_trace_start_original + global _otel_trace_end_original global _otel_trace_start_patch_ref_count from openinference.instrumentation.openai_agents._processor import ( @@ -130,7 +162,14 @@ def _uninstall_otel_instrumentation() -> None: "on_trace_start", _otel_trace_start_original, ) + if _otel_trace_end_original is not None: + setattr( + OpenInferenceTracingProcessor, + "on_trace_end", + _otel_trace_end_original, + ) _otel_trace_start_original = None + _otel_trace_end_original = None _otel_trace_start_instrumentor = None @@ -426,7 +465,6 @@ def workflow_runner(runner: WorkflowRunner | None) -> WorkflowRunner: interceptor = OTelOpenAIAgentsContextPropagationInterceptor( add_temporal_spans=add_temporal_spans, - otel_id_generator=provider.id_generator(), ) @asynccontextmanager diff --git a/temporalio/contrib/openai_agents/_trace_interceptor.py b/temporalio/contrib/openai_agents/_trace_interceptor.py index 66297e20b..dfa172df6 100644 --- a/temporalio/contrib/openai_agents/_trace_interceptor.py +++ b/temporalio/contrib/openai_agents/_trace_interceptor.py @@ -3,16 +3,17 @@ from __future__ import annotations import abc -from collections.abc import Mapping -from contextlib import contextmanager +import asyncio +import contextvars +from collections.abc import Callable, Iterator, Mapping +from contextlib import AbstractContextManager, ExitStack, contextmanager from typing import Any, Protocol -from agents import CustomSpanData, custom_span, get_current_span, trace +from agents import CustomSpanData, Span, Trace, custom_span, get_current_span, trace from agents.tracing import ( get_trace_provider, ) from agents.tracing.scope import Scope -from agents.tracing.spans import Span import temporalio.api.common.v1 import temporalio.client @@ -32,7 +33,7 @@ class _InputWithHeaders(Protocol): def temporal_span( add_temporal_spans: bool, span_name: str, -): +) -> Iterator[None]: """Create a temporal span context manager. Args: @@ -42,7 +43,7 @@ def temporal_span( Yields: A span context with temporal metadata if enabled. """ - if add_temporal_spans: + if add_temporal_spans and get_trace_provider().get_current_trace() is not None: """Extracts and initializes trace information the input header.""" data = ( { @@ -92,8 +93,9 @@ def __init__( payload_converter: The payload converter to use for serializing/deserializing trace context. Defaults to the default Temporal payload converter. add_temporal_spans: Whether to add temporal-specific spans to traces. - start_traces: Whether to start new traces if none exist. This will cause duplication if the underlying - trace provider actually process start events. Primarily designed for use with Open Telemetry integration. + start_traces: Whether to emit start events for reconstructed context. + Keep disabled when the processor records spans, to avoid duplicating + the caller's spans. """ super().__init__() self._payload_converter = payload_converter @@ -210,17 +212,23 @@ def span_context_from_header_contents(self, span_info: dict[str, Any]): else: Scope.set_current_span(current_span) + @contextmanager def context_from_header( self, input: _InputWithHeaders, - ): - """Extracts and initializes trace information the input header.""" - span_info = self.get_header_contents(input) - if span_info is None: - return - - self.trace_context_from_header_contents(span_info) - self.span_context_from_header_contents(span_info) + ) -> Iterator[None]: + """Restore the caller's trace context for an inbound operation.""" + previous_trace: Trace | None = Scope.get_current_trace() + previous_span: Span[Any] | None = Scope.get_current_span() + try: + span_info = self.get_header_contents(input) + if span_info is not None: + self.trace_context_from_header_contents(span_info=span_info) + self.span_context_from_header_contents(span_info=span_info) + yield + finally: + Scope.set_current_span(previous_span) + Scope.set_current_trace(previous_trace) @contextmanager def maybe_span(self, span_name: str, data: dict[str, Any] | None): @@ -318,9 +326,11 @@ def __init__( async def execute_activity( self, input: temporalio.worker.ExecuteActivityInput ) -> Any: - self._root.context_from_header(input) - with temporal_span(self._root._add_temporal_spans, "temporal:executeActivity"): - return await self.next.execute_activity(input) + with self._root.context_from_header(input=input): + with temporal_span( + self._root._add_temporal_spans, "temporal:executeActivity" + ): + return await self.next.execute_activity(input) class _ContextPropagationWorkflowInboundInterceptor( @@ -342,30 +352,35 @@ def root(self): async def execute_workflow( self, input: temporalio.worker.ExecuteWorkflowInput ) -> Any: - self.root().context_from_header(input) - with temporal_span(self.root()._add_temporal_spans, "temporal:executeWorkflow"): - return await self.next.execute_workflow(input) + with self.root().context_from_header(input=input): + with temporal_span( + self.root()._add_temporal_spans, "temporal:executeWorkflow" + ): + return await self.next.execute_workflow(input) async def handle_signal(self, input: temporalio.worker.HandleSignalInput) -> None: - self.root().context_from_header(input) - with temporal_span(self.root()._add_temporal_spans, "temporal:handleSignal"): - return await self.next.handle_signal(input) + with self.root().context_from_header(input=input): + with temporal_span( + self.root()._add_temporal_spans, "temporal:handleSignal" + ): + return await self.next.handle_signal(input) async def handle_query(self, input: temporalio.worker.HandleQueryInput) -> Any: - with temporal_span(self.root()._add_temporal_spans, "temporal:handleQuery"): - return await self.next.handle_query(input) + with self.root().context_from_header(input=input): + with temporal_span(self.root()._add_temporal_spans, "temporal:handleQuery"): + return await self.next.handle_query(input) def handle_update_validator( self, input: temporalio.worker.HandleUpdateInput ) -> None: - self.root().context_from_header(input) - self.next.handle_update_validator(input) + with self.root().context_from_header(input=input): + self.next.handle_update_validator(input) async def handle_update_handler( self, input: temporalio.worker.HandleUpdateInput ) -> Any: - self.root().context_from_header(input) - return await self.next.handle_update_handler(input) + with self.root().context_from_header(input=input): + return await self.next.handle_update_handler(input) class _ContextPropagationWorkflowOutboundInterceptor( @@ -398,48 +413,60 @@ async def signal_external_workflow( def start_activity( self, input: temporalio.worker.StartActivityInput ) -> temporalio.workflow.ActivityHandle: - trace = get_trace_provider().get_current_trace() - span: Span | None = None - if trace and self.root()._add_temporal_spans: - span = custom_span( - name="temporal:startActivity", data={"activity": input.activity} - ) - span.start(mark_as_current=True) - - self.root().set_header_from_context(input) - handle = self.next.start_activity(input) - if span: - handle.add_done_callback(lambda _: span.finish()) # type: ignore - return handle + with self._span_for_start( + input=input, + span_name="temporal:startActivity", + data={"activity": input.activity}, + ) as (context, finish_on_completion): + handle = context.run(self.next.start_activity, input=input) + finish_on_completion(handle) + return handle async def start_child_workflow( self, input: temporalio.worker.StartChildWorkflowInput ) -> temporalio.workflow.ChildWorkflowHandle: - trace = get_trace_provider().get_current_trace() - span: Span | None = None - if trace and self.root()._add_temporal_spans: - span = custom_span( - name="temporal:startChildWorkflow", data={"workflow": input.workflow} - ) - span.start(mark_as_current=True) - self.root().set_header_from_context(input) - handle = await self.next.start_child_workflow(input) - if span: - handle.add_done_callback(lambda _: span.finish()) # type: ignore - return handle + with self._span_for_start( + input=input, + span_name="temporal:startChildWorkflow", + data={"workflow": input.workflow}, + ) as (_, finish_on_completion): + handle = await self.next.start_child_workflow(input=input) + finish_on_completion(handle) + return handle def start_local_activity( self, input: temporalio.worker.StartLocalActivityInput ) -> temporalio.workflow.ActivityHandle: - trace = get_trace_provider().get_current_trace() - span: Span | None = None - if trace and self.root()._add_temporal_spans: - span = custom_span( - name="temporal:startLocalActivity", data={"activity": input.activity} - ) - span.start(mark_as_current=True) - self.root().set_header_from_context(input) - handle = self.next.start_local_activity(input) - if span: - handle.add_done_callback(lambda _: span.finish()) # type: ignore - return handle + with self._span_for_start( + input=input, + span_name="temporal:startLocalActivity", + data={"activity": input.activity}, + ) as (context, finish_on_completion): + handle = context.run(self.next.start_local_activity, input=input) + finish_on_completion(handle) + return handle + + @contextmanager + def _span_for_start( + self, + input: _InputWithHeaders, + span_name: str, + data: dict[str, Any], + ) -> Iterator[tuple[contextvars.Context, Callable[[asyncio.Future[Any]], None]]]: + context: contextvars.Context = contextvars.copy_context() + span_scope: AbstractContextManager[None] = self.root().maybe_span( + span_name=span_name, data=data + ) + with ExitStack() as cleanup: + context.run(span_scope.__enter__) + cleanup.callback(context.run, span_scope.__exit__, None, None, None) + context.run(self.root().set_header_from_context, input=input) + + def finish_on_completion(handle: asyncio.Future[Any]) -> None: + # Tokens cannot be detached in a copy of their original Context. + handle.add_done_callback( + lambda _: span_scope.__exit__(None, None, None), context=context + ) + + yield context, finish_on_completion + cleanup.pop_all() diff --git a/tests/contrib/openai_agents/test_openai_tracing.py b/tests/contrib/openai_agents/test_openai_tracing.py index 28b804cc1..93adb7374 100644 --- a/tests/contrib/openai_agents/test_openai_tracing.py +++ b/tests/contrib/openai_agents/test_openai_tracing.py @@ -1,21 +1,53 @@ +import asyncio import uuid from datetime import timedelta from typing import Any +import opentelemetry.context import opentelemetry.trace -from agents import Span, Trace, TracingProcessor, custom_span, trace +import pytest +from agents import ( + Agent, + Runner, + Span, + Trace, + TracingProcessor, + custom_span, + function_tool, + trace, +) from agents.tracing import get_trace_provider +from agents.tracing.provider import DefaultTraceProvider +from openinference.instrumentation.openai_agents._processor import ( + OpenInferenceTracingProcessor, +) +from openinference.semconv.trace import SpanAttributes from opentelemetry.sdk.trace import ReadableSpan from opentelemetry.sdk.trace.export import SimpleSpanProcessor from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter +from opentelemetry.sdk.trace.id_generator import RandomIdGenerator from temporalio import activity, workflow +from temporalio.api.enums.v1 import EventType from temporalio.client import Client +from temporalio.common import RetryPolicy from temporalio.contrib.openai_agents import _temporal_openai_agents +from temporalio.contrib.openai_agents._temporal_trace_provider import ( + TemporalTraceProvider, +) from temporalio.contrib.openai_agents.testing import ( AgentEnvironment, + ResponseBuilders, + TestModel, ) from temporalio.contrib.opentelemetry import create_tracer_provider +from temporalio.exceptions import ( + ActivityError, + ApplicationError, + CancelledError, + ChildWorkflowError, +) +from temporalio.worker import Replayer from temporalio.worker.workflow_sandbox import ( SandboxedWorkflowRunner, SandboxRestrictions, @@ -51,31 +83,50 @@ def force_flush(self) -> None: pass +@pytest.mark.usefixtures("reset_otel_tracer_provider") def test_otel_instrumentation_lifecycle_does_not_nest() -> None: - from openinference.instrumentation.openai_agents._processor import ( - OpenInferenceTracingProcessor, - ) - from opentelemetry import trace - + exporter = set_test_tracer_provider() original = OpenInferenceTracingProcessor.on_trace_start - _temporal_openai_agents._install_otel_instrumentation(trace.get_tracer_provider()) + original_end = OpenInferenceTracingProcessor.on_trace_end + _temporal_openai_agents._install_otel_instrumentation( + opentelemetry.trace.get_tracer_provider() + ) try: installed_patch = OpenInferenceTracingProcessor.on_trace_start + installed_end = OpenInferenceTracingProcessor.on_trace_end assert installed_patch is not original + assert installed_end is not original_end _temporal_openai_agents._install_otel_instrumentation( - trace.get_tracer_provider() + opentelemetry.trace.get_tracer_provider() ) try: assert OpenInferenceTracingProcessor.on_trace_start is installed_patch + assert OpenInferenceTracingProcessor.on_trace_end is installed_end finally: _temporal_openai_agents._uninstall_otel_instrumentation() assert OpenInferenceTracingProcessor.on_trace_start is installed_patch + assert OpenInferenceTracingProcessor.on_trace_end is installed_end + previous_context = opentelemetry.context.get_current() + with pytest.raises(ValueError, match="trace body failed"): + with trace("Standalone trace"): + with custom_span("Standalone child"): + raise ValueError("trace body failed") + assert opentelemetry.context.get_current() is previous_context + processor = openinference_processor() + assert not processor._root_spans + assert not processor._otel_spans + assert not processor._tokens + assert {span.name for span in exporter.get_finished_spans()} == { + "Standalone trace", + "Standalone child", + } finally: _temporal_openai_agents._uninstall_otel_instrumentation() assert OpenInferenceTracingProcessor.on_trace_start is original + assert OpenInferenceTracingProcessor.on_trace_end is original_end async def test_tracing(client: Client): @@ -521,9 +572,11 @@ async def ready() -> bool: ) +@pytest.mark.parametrize("add_temporal_spans", [False, True]) async def test_workflow_only_trace_to_spans( client: Client, reset_otel_tracer_provider: Any, # type: ignore[reportUnusedParameter] + add_temporal_spans: bool, ): """Test: Workflow-only trace -> spans (with worker restart).""" exporter = set_test_tracer_provider() @@ -533,7 +586,7 @@ async def test_workflow_only_trace_to_spans( # First worker: Start workflow (no external trace context) async with AgentEnvironment( model=research_mock_model(), - add_temporal_spans=False, + add_temporal_spans=add_temporal_spans, use_otel_instrumentation=True, ) as env: new_client = env.applied_on_client(client) @@ -563,7 +616,7 @@ async def ready() -> bool: # Second worker: Complete the workflow with fresh objects (new instrumentation) async with AgentEnvironment( model=research_mock_model(), - add_temporal_spans=False, + add_temporal_spans=add_temporal_spans, use_otel_instrumentation=True, ) as env: new_client = env.applied_on_client(client) @@ -579,8 +632,16 @@ async def ready() -> bool: await workflow_handle.signal(SelfTracingWorkflow.proceed) result = await workflow_handle.result() assert result == "done" + processor = openinference_processor() spans = exporter.get_finished_spans() + print_otel_spans(spans) + assert not processor._root_spans + assert not processor._otel_spans + assert not processor._tokens + span_ids = {span.context.span_id for span in spans if span.context} + assert len(span_ids) == len(spans) + assert all(not span.parent or span.parent.span_id in span_ids for span in spans) assert len(spans) >= 2 # Workflow trace + workflow span @@ -677,7 +738,7 @@ async def test_otel_tracing_in_runner( ResearchWorkflow, max_cached_workflows=0, ) as worker: - with trace("Research workflow"): + with env.openai_agents_plugin.tracing_context(), trace("Research workflow"): workflow_handle = await client.start_workflow( ResearchWorkflow.run, "Caribbean vacation spots in April, optimizing for surfing, hiking and water sports", @@ -859,7 +920,7 @@ async def test_sdk_trace_to_otel_span_parenting( ), ) as worker: # Start SDK trace in client, then start workflow within that trace - with trace("Client SDK trace"): + with env.openai_agents_plugin.tracing_context(), trace("Client SDK trace"): workflow_handle = await new_client.start_workflow( OtelSpanWorkflow.run, id=f"sdk-trace-otel-span-workflow-{uuid.uuid4()}", @@ -953,3 +1014,533 @@ async def ready() -> bool: assert len(span_ids) == len(set(span_ids)), ( f"All spans should have unique IDs, got: {span_ids}" ) + + +def openinference_processor() -> OpenInferenceTracingProcessor: + provider = get_trace_provider() + if isinstance(provider, TemporalTraceProvider): + provider = provider._original_provider + assert isinstance(provider, DefaultTraceProvider) + return next( + processor + for processor in provider._multi_processor._processors + if isinstance(processor, OpenInferenceTracingProcessor) + ) + + +@function_tool +async def trace_test_tool() -> str: + return "tool result" + + +@workflow.defn +class AgentToolTraceWorkflow: + @workflow.run + async def run(self) -> str: + result = await Runner.run( + starting_agent=Agent(name="Trace agent", tools=[trace_test_tool]), + input="Call the tool.", + ) + return str(result.final_output) + + +@pytest.mark.parametrize("max_cached_workflows", [1000, 0]) +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_otel_remote_context_preserves_root( + client: Client, + caplog: pytest.LogCaptureFixture, + max_cached_workflows: int, +) -> None: + exporter = set_test_tracer_provider() + roots: list[tuple[int, int, str]] = [] + retained: list[tuple[int, int, int]] = [] + restored: list[bool] = [] + async with AgentEnvironment( + model=TestModel.returning_responses( + [ + response + for _ in range(3) + for response in ( + ResponseBuilders.tool_call(arguments="{}", name="trace_test_tool"), + ResponseBuilders.output_message("done"), + ) + ] + ), + use_otel_instrumentation=True, + ) as env: + client = env.applied_on_client(client) + async with new_worker( + client, + AgentToolTraceWorkflow, + max_cached_workflows=max_cached_workflows, + ) as worker: + with opentelemetry.trace.get_tracer(__name__).start_as_current_span( + "Outer span" + ): + parent_context = opentelemetry.context.get_current() + for index in range(3): + with env.openai_agents_plugin.tracing_context(): + processor = openinference_processor() + with trace(f"Client trace {index}"): + root = opentelemetry.trace.get_current_span() + session_id = f"session-{index}" + root.set_attribute(SpanAttributes.SESSION_ID, session_id) + context = root.get_span_context() + roots.append( + (context.trace_id, context.span_id, session_id) + ) + assert ( + await client.execute_workflow( + AgentToolTraceWorkflow.run, + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + == "done" + ) + restored.append( + opentelemetry.context.get_current() is parent_context + ) + retained.append( + ( + len(processor._root_spans), + len(processor._otel_spans), + len(processor._tokens), + ) + ) + + spans = exporter.get_finished_spans() + spans_by_id = {span.context.span_id: span for span in spans if span.context} + missing_parents = [ + span.name + for span in spans + if span.parent and span.parent.span_id not in spans_by_id + ] + print_otel_spans(spans) + print(f"Client roots: {roots}") + print(f"Missing parents: {missing_parents}") + print(f"Retained roots/spans/tokens: {retained}") + print(f"Caller context restored: {restored}") + for trace_id, span_id, session_id in roots: + assert span_id in spans_by_id, "The original client root was not exported" + root_span = spans_by_id[span_id] + assert root_span.context and root_span.context.trace_id == trace_id + assert root_span.attributes + assert root_span.attributes[SpanAttributes.SESSION_ID] == session_id + assert len(spans_by_id) == len(spans) + assert not missing_parents + assert retained == [(0, 0, 0)] * 3 + assert all(restored) + for span in spans: + if span.name == "trace_test_tool": + assert span.parent + assert spans_by_id[span.parent.span_id].name == "turn" + assert "Failed to detach context" not in caplog.text + + +@workflow.defn +class CallbackChildWorkflow: + @workflow.run + async def run(self, outcome: str = "success") -> str: + if outcome == "failure": + raise ApplicationError("Child failed", non_retryable=True) + if outcome == "cancel": + await workflow.wait_condition(lambda: False) + return "success" + + +@activity.defn +async def callback_outcome_activity(outcome: str) -> str: + if outcome == "failure": + raise ApplicationError("Activity failed", non_retryable=True) + if outcome == "cancel": + await asyncio.Future() + return "success" + + +@workflow.defn +class CallbackSpanWorkflow: + @workflow.run + async def run(self, outcome: str = "success") -> dict[str, Any]: + observations: list[dict[str, Any]] = [] + errors: list[str] = [] + with custom_span("Callback caller") as caller: + otel_caller = opentelemetry.trace.get_current_span().get_span_context() + for operation in ("activity", "local_activity", "child_workflow"): + handle: asyncio.Future[Any] + if operation == "activity" and outcome == "success": + handle = workflow.start_activity( + simple_no_context_activity, + start_to_close_timeout=timedelta(seconds=30), + ) + elif operation == "local_activity" and outcome == "success": + handle = workflow.start_local_activity( + simple_no_context_activity, + start_to_close_timeout=timedelta(seconds=30), + ) + elif operation == "child_workflow": + handle = await workflow.start_child_workflow( + CallbackChildWorkflow.run, + arg=outcome, + ) + else: + start = ( + workflow.start_activity + if operation == "activity" + else workflow.start_local_activity + ) + handle = start( + callback_outcome_activity, + arg=outcome, + start_to_close_timeout=timedelta(seconds=30), + retry_policy=RetryPolicy(maximum_attempts=1), + ) + for phase in ("scheduled", "completed"): + if phase == "completed": + if outcome == "cancel": + handle.cancel() + try: + await handle + except ( + ActivityError, + ApplicationError, + ChildWorkflowError, + CancelledError, + asyncio.CancelledError, + ): + errors.append(operation) + current = get_trace_provider().get_current_span() + observations.append( + { + "operation": operation, + "phase": phase, + "agent_span": current.span_id if current else None, + "otel_span": opentelemetry.trace.get_current_span() + .get_span_context() + .span_id, + } + ) + with custom_span(f"After {operation}"): + pass + return { + "agent_span": caller.span_id, + "otel_span": otel_caller.span_id, + "observations": observations, + "errors": errors, + } + + +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_otel_callback_spans_restore_context( + client: Client, + caplog: pytest.LogCaptureFixture, + max_cached_workflows: int = 1000, +) -> None: + exporter = set_test_tracer_provider() + async with AgentEnvironment( + model=research_mock_model(), + use_otel_instrumentation=True, + ) as env: + client = env.applied_on_client(client) + async with new_worker( + client, + CallbackSpanWorkflow, + CallbackChildWorkflow, + activities=[simple_no_context_activity], + max_cached_workflows=max_cached_workflows, + workflow_runner=SandboxedWorkflowRunner( + SandboxRestrictions.default.with_passthrough_modules("opentelemetry") + ), + ) as worker: + with env.openai_agents_plugin.tracing_context(): + processor = openinference_processor() + with trace("Callback trace"): + result = await client.execute_workflow( + CallbackSpanWorkflow.run, + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + spans = exporter.get_finished_spans() + print_otel_spans(spans) + print(f"Callback contexts: {result}") + print( + "Retained roots/spans/tokens:", + len(processor._root_spans), + len(processor._otel_spans), + len(processor._tokens), + ) + assert "Failed to detach context" not in caplog.text + for observation in result["observations"]: + assert observation["agent_span"] == result["agent_span"] + assert observation["otel_span"] == result["otel_span"] + assert not processor._root_spans + assert not processor._otel_spans + assert not processor._tokens + assert len(spans) == len({span.context.span_id for span in spans if span.context}) + for span in spans: + if span.name.startswith("After "): + assert span.parent and span.parent.span_id == result["otel_span"] + + +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_otel_callback_replay( + client: Client, + caplog: pytest.LogCaptureFixture, +) -> None: + await test_otel_callback_spans_restore_context( + client=client, + caplog=caplog, + max_cached_workflows=0, + ) + + +@pytest.mark.parametrize("outcome", ["failure", "cancel"]) +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_otel_callback_failure_and_cancellation( + client: Client, + caplog: pytest.LogCaptureFixture, + outcome: str, +) -> None: + exporter = set_test_tracer_provider() + async with AgentEnvironment( + model=research_mock_model(), + use_otel_instrumentation=True, + ) as env: + client = env.applied_on_client(client) + async with new_worker( + client, + CallbackSpanWorkflow, + CallbackChildWorkflow, + activities=[simple_no_context_activity, callback_outcome_activity], + workflow_runner=SandboxedWorkflowRunner( + SandboxRestrictions.default.with_passthrough_modules("opentelemetry") + ), + ) as worker: + with env.openai_agents_plugin.tracing_context(): + processor = openinference_processor() + with trace("Callback outcome trace"): + result = await client.execute_workflow( + CallbackSpanWorkflow.run, + arg=outcome, + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + print(f"Callback outcome {outcome}: {result}") + print_otel_spans(exporter.get_finished_spans()) + assert result["errors"] == ["activity", "local_activity", "child_workflow"] + for observation in result["observations"]: + assert observation["agent_span"] == result["agent_span"] + assert observation["otel_span"] == result["otel_span"] + assert not processor._root_spans + assert not processor._otel_spans + assert not processor._tokens + assert "Failed to detach context" not in caplog.text + + +@activity.defn +async def inspect_otel_context_activity() -> tuple[int, int, str]: + context = opentelemetry.trace.get_current_span().get_span_context() + return context.trace_id, int(context.trace_flags), context.trace_state.to_header() + + +@workflow.defn +class InspectOtelContextWorkflow: + @workflow.run + async def run(self) -> tuple[int, int, str]: + return await workflow.execute_activity( + inspect_otel_context_activity, + start_to_close_timeout=timedelta(seconds=30), + ) + + +@pytest.mark.parametrize("sampled", [True, False]) +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_otel_remote_sampling_context( + client: Client, + caplog: pytest.LogCaptureFixture, + sampled: bool, +) -> None: + exporter = set_test_tracer_provider() + id_generator = RandomIdGenerator() + remote_context = opentelemetry.trace.SpanContext( + trace_id=id_generator.generate_trace_id(), + span_id=id_generator.generate_span_id(), + is_remote=True, + trace_flags=opentelemetry.trace.TraceFlags( + opentelemetry.trace.TraceFlags.SAMPLED + if sampled + else opentelemetry.trace.TraceFlags.DEFAULT + ), + trace_state=opentelemetry.trace.TraceState([("vendor", "state")]), + ) + async with AgentEnvironment( + model=research_mock_model(), + use_otel_instrumentation=True, + ) as env: + client = env.applied_on_client(client) + async with new_worker( + client, + InspectOtelContextWorkflow, + activities=[inspect_otel_context_activity], + ) as worker: + with env.openai_agents_plugin.tracing_context(): + processor = openinference_processor() + with opentelemetry.trace.use_span( + opentelemetry.trace.NonRecordingSpan(remote_context) + ): + with trace("Remote sampling"): + result = await client.execute_workflow( + InspectOtelContextWorkflow.run, + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + print(f"Remote sampling {sampled}: {result}") + assert result == ( + remote_context.trace_id, + int(remote_context.trace_flags), + remote_context.trace_state.to_header(), + ) + assert bool(exporter.get_finished_spans()) == sampled + assert not processor._root_spans + assert not processor._otel_spans + assert not processor._tokens + assert "Failed to detach context" not in caplog.text + + +@pytest.mark.parametrize("add_temporal_spans", [True, False]) +async def test_callback_context_without_otel( + client: Client, add_temporal_spans: bool +) -> None: + processor = MemoryTracingProcessor() + processor.trace_events = [] + processor.span_events = [] + get_trace_provider().set_processors([processor]) + async with AgentEnvironment( + model=research_mock_model(), + add_temporal_spans=add_temporal_spans, + ) as env: + client = env.applied_on_client(client) + async with new_worker( + client, + CallbackSpanWorkflow, + CallbackChildWorkflow, + activities=[simple_no_context_activity], + workflow_runner=SandboxedWorkflowRunner( + SandboxRestrictions.default.with_passthrough_modules("opentelemetry") + ), + ) as worker: + with trace("Native callback trace"): + result = await client.execute_workflow( + CallbackSpanWorkflow.run, + id=str(uuid.uuid4()), + task_queue=worker.task_queue, + ) + print(f"Native tracing, temporal spans {add_temporal_spans}: {result}") + assert result["agent_span"] != "no-op" + for observation in result["observations"]: + assert observation["agent_span"] == result["agent_span"] + started = {span.span_id for span, start in processor.span_events if start} + ended = {span.span_id for span, start in processor.span_events if not start} + assert started == ended + assert ( + any( + span.span_data.export().get("name") == "temporal:startActivity" + for span, _ in processor.span_events + ) + == add_temporal_spans + ) + + +@workflow.defn +class ConcurrentTraceCommandsWorkflow: + @workflow.run + async def run(self) -> None: + child_start = asyncio.create_task( + workflow.start_child_workflow(CallbackChildWorkflow.run) + ) + await asyncio.sleep(0) + activity_handle = workflow.start_activity( + simple_no_context_activity, + start_to_close_timeout=timedelta(seconds=30), + ) + child_handle = await child_start + await activity_handle + await child_handle + + +@pytest.mark.parametrize("use_otel_instrumentation", [False, True]) +@pytest.mark.parametrize("replay", [False, True]) +@pytest.mark.usefixtures("reset_otel_tracer_provider") +async def test_tracing_preserves_child_command_order( + client: Client, use_otel_instrumentation: bool, replay: bool +) -> None: + set_test_tracer_provider() + runner = SandboxedWorkflowRunner( + SandboxRestrictions.default.with_passthrough_modules( + "agents", "openai", "mcp", "opentelemetry" + ) + ) + async with new_worker( + client, + ConcurrentTraceCommandsWorkflow, + CallbackChildWorkflow, + activities=[simple_no_context_activity], + workflow_runner=runner, + ) as worker: + original_handle = await client.start_workflow( + ConcurrentTraceCommandsWorkflow.run, + id=f"command-order-original-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + await original_handle.result() + history = await original_handle.fetch_history() + command_events = ( + EventType.EVENT_TYPE_START_CHILD_WORKFLOW_EXECUTION_INITIATED, + EventType.EVENT_TYPE_ACTIVITY_TASK_SCHEDULED, + ) + original_order = [ + event.event_type + for event in history.events + if event.event_type in command_events + ] + print(f"Command history: {original_handle.id}") + print(f"Original command order: {[EventType.Name(e) for e in original_order]}") + + async with AgentEnvironment( + model=research_mock_model(), + use_otel_instrumentation=use_otel_instrumentation, + ) as env: + if replay: + await Replayer( + workflows=[ConcurrentTraceCommandsWorkflow, CallbackChildWorkflow], + workflow_runner=runner, + namespace=client.namespace, + plugins=[env.openai_agents_plugin], + ).replay_workflow(history=history) + print("History replay with tracing: passed") + else: + client = env.applied_on_client(client) + if not use_otel_instrumentation: + get_trace_provider().set_processors([MemoryTracingProcessor()]) + async with new_worker( + client, + ConcurrentTraceCommandsWorkflow, + CallbackChildWorkflow, + activities=[simple_no_context_activity], + workflow_runner=runner, + ) as worker: + with env.openai_agents_plugin.tracing_context(), trace("Commands"): + traced_handle = await client.start_workflow( + ConcurrentTraceCommandsWorkflow.run, + id=f"command-order-traced-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + await traced_handle.result() + traced_history = await traced_handle.fetch_history() + traced_order = [ + event.event_type + for event in traced_history.events + if event.event_type in command_events + ] + print(f"Command history: {traced_handle.id}") + print(f"Traced command order: {[EventType.Name(e) for e in traced_order]}") + assert traced_order == original_order