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 @@ -52,6 +52,8 @@ to include examples, links to docs, or any other relevant information.

### Fixed

- OpenTelemetry trace and span IDs propagated by concurrent workers no longer
interfere with each other, preserving the correct parent-child hierarchy.
- The `google-adk` extra now depends on `mcp`, so fresh installs of
`temporalio[google-adk]` can import `temporalio.contrib.google_adk_agents`
without separately installing `mcp`. Previously the import failed with an
Expand Down
25 changes: 17 additions & 8 deletions temporalio/contrib/opentelemetry/_id_generator.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import random
from contextvars import ContextVar

from opentelemetry.sdk.trace.id_generator import IdGenerator
from opentelemetry.trace import (
Expand Down Expand Up @@ -47,8 +48,12 @@ class TemporalIdGenerator(IdGenerator):
def __init__(self, id_generator: IdGenerator):
"""Initialize a TemporalIdGenerator."""
self._id_generator = id_generator
self.traces: list[int] = []
self.spans: list[int] = []
self._traces: ContextVar[tuple[int, ...]] = ContextVar(
"temporalio_otel_trace_id_seeds", default=()
)
self._spans: ContextVar[tuple[int, ...]] = ContextVar(
"temporalio_otel_span_id_seeds", default=()
)

def seed_span_id(self, span_id: int) -> None:
"""Seed the generator with a span ID to use as the first result.
Expand All @@ -59,24 +64,26 @@ def seed_span_id(self, span_id: int) -> None:
Args:
span_id: The span ID to use as the first generated span ID.
"""
self.spans.append(span_id)
self._spans.set((*self._spans.get(), span_id))

def seed_trace_id(self, trace_id: int) -> None:
"""Seed the generator with a trace ID to use as the first result.

Args:
trace_id: The trace ID to use as the first generated trace ID.
"""
self.traces.append(trace_id)
self._traces.set((*self._traces.get(), trace_id))

def generate_span_id(self) -> int:
"""Generate a span ID using Temporal's deterministic random when in workflow.

Returns:
A 64-bit span ID.
"""
if len(self.spans) > 0:
return self.spans.pop()
spans = self._spans.get()
if spans:
self._spans.set(spans[:-1])
return spans[-1]

if workflow_random := _get_workflow_random():
span_id = workflow_random.getrandbits(64)
Expand All @@ -91,8 +98,10 @@ def generate_trace_id(self) -> int:
Returns:
A 128-bit trace ID.
"""
if len(self.traces) > 0:
return self.traces.pop()
traces = self._traces.get()
if traces:
self._traces.set(traces[:-1])
return traces[-1]

if workflow_random := _get_workflow_random():
trace_id = workflow_random.getrandbits(128)
Expand Down
41 changes: 41 additions & 0 deletions tests/contrib/opentelemetry/test_opentelemetry_plugin.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
import logging
import threading
import uuid
from concurrent.futures import ThreadPoolExecutor
from datetime import timedelta
from typing import Any

Expand All @@ -10,6 +12,7 @@
from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.sdk.trace.id_generator import RandomIdGenerator
from opentelemetry.trace import (
get_tracer,
)
Expand All @@ -18,6 +21,7 @@
from temporalio import activity, nexus, workflow
from temporalio.client import Client, WorkflowFailureError
from temporalio.contrib.opentelemetry import OpenTelemetryPlugin, create_tracer_provider
from temporalio.contrib.opentelemetry._id_generator import TemporalIdGenerator
from temporalio.exceptions import ApplicationError
from temporalio.testing import WorkflowEnvironment

Expand All @@ -29,6 +33,43 @@
logger = logging.getLogger(__name__)


@pytest.mark.parametrize(
("seed_method", "generate_method"),
[
("seed_span_id", "generate_span_id"),
("seed_trace_id", "generate_trace_id"),
],
)
def test_temporal_id_generator_seeds_are_context_local(
seed_method: str, generate_method: str
) -> None:
generator = TemporalIdGenerator(RandomIdGenerator())
first_seeded = threading.Event()
second_seeded = threading.Event()
first_generated = threading.Event()
seed_values = (123, 456)

def generate_first() -> int:
getattr(generator, seed_method)(seed_values[0])
first_seeded.set()
assert second_seeded.wait(timeout=5)
generated = getattr(generator, generate_method)()
first_generated.set()
return generated

def generate_second() -> int:
assert first_seeded.wait(timeout=5)
getattr(generator, seed_method)(seed_values[1])
second_seeded.set()
assert first_generated.wait(timeout=5)
return getattr(generator, generate_method)()

with ThreadPoolExecutor(max_workers=2) as executor:
first = executor.submit(generate_first)
second = executor.submit(generate_second)
assert (first.result(), second.result()) == seed_values


@activity.defn
async def simple_no_context_activity() -> str:
with get_tracer(__name__).start_as_current_span("Activity"):
Expand Down
Loading