Skip to content
Open
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
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
74 changes: 41 additions & 33 deletions temporalio/contrib/openai_agents/_otel_trace_interceptor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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
48 changes: 43 additions & 5 deletions temporalio/contrib/openai_agents/_temporal_openai_agents.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -62,32 +68,51 @@
_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


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(
Expand All @@ -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
Expand All @@ -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 (
Expand All @@ -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


Expand Down Expand Up @@ -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
Expand Down
Loading