diff --git a/CHANGELOG.md b/CHANGELOG.md index 5a8995696..3134fa923 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -28,6 +28,11 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- `temporalio.contrib.opentelemetry`: `TracingInterceptor` and `OpenTelemetryInterceptor` no longer + log `Failed to detach context` when a context is torn down on a different thread while + OpenTelemetry's threading instrumentation (enabled by strands, among others) is active; a + context is now detached exactly when its token is still valid in the current + `contextvars.Context`, which it stays when a workflow resumes on another pool thread. ### Security ## [1.34.0] - 2026-09-30 diff --git a/temporalio/contrib/opentelemetry/_context.py b/temporalio/contrib/opentelemetry/_context.py new file mode 100644 index 000000000..a57f8d48d --- /dev/null +++ b/temporalio/contrib/opentelemetry/_context.py @@ -0,0 +1,51 @@ +"""Attach OpenTelemetry contexts so that the matching detach is always safe.""" + +from __future__ import annotations + +from collections.abc import Iterator +from contextlib import contextmanager +from contextvars import Token + +import opentelemetry.context +from opentelemetry.context import Context + + +@contextmanager +def attached_context(context: Context | None) -> Iterator[None]: + """Attach ``context`` for the block and detach it afterwards where possible. + + ``None`` attaches nothing. The block's ``finally`` can run in a different + ``contextvars.Context`` than the one that attached: a context manager + abandoned by an evicted workflow is finalized wherever garbage collection + happens to run. The token is not valid there, and + ``opentelemetry.context.detach`` would log "Failed to detach context" even + though there is nothing to detach. Checking that the attached context is + still current does not catch every such case, because OpenTelemetry's + threading instrumentation (enabled by strands, among others) propagates + the same ``Context`` object into new threads. The thread is no test + either: workflow activations move between pool threads while the asyncio + task keeps its ``contextvars.Context``, and those detaches must happen. + Only ``contextvars`` knows which ``Context`` a token belongs to, so + :func:`_detach` performs the reset ``detach`` performs and ignores the + ``ValueError`` raised for a token from another ``Context``. + """ + if context is None: + yield + return + token = opentelemetry.context.attach(context) + try: + yield + finally: + _detach(context, token) + + +def _detach(context: Context, token: Token[Context]) -> bool: + """Detach ``context`` if it is current and ``token`` is valid here.""" + if context is not opentelemetry.context.get_current(): + return False + try: + token.var.reset(token) + except ValueError: + # The token was created in a different contextvars.Context. + return False + return True diff --git a/temporalio/contrib/opentelemetry/_interceptor.py b/temporalio/contrib/opentelemetry/_interceptor.py index 6dca4596e..f503d5ac5 100644 --- a/temporalio/contrib/opentelemetry/_interceptor.py +++ b/temporalio/contrib/opentelemetry/_interceptor.py @@ -36,6 +36,7 @@ import temporalio.nexus.system.workflow_service.models import temporalio.worker import temporalio.workflow +from temporalio.contrib.opentelemetry._context import attached_context from temporalio.exceptions import ApplicationError, ApplicationErrorCategory # OpenTelemetry dynamically, lazily chooses its context implementation at @@ -183,8 +184,7 @@ def _start_as_current_span( kind: opentelemetry.trace.SpanKind, context: Context | None = None, ) -> Iterator[None]: - token = opentelemetry.context.attach(context) if context else None - try: + with attached_context(context): with self.tracer.start_as_current_span( name, attributes=attributes, @@ -219,9 +219,6 @@ def _start_as_current_span( ) ) raise - finally: - if token and context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) def _completed_workflow_span( self, params: _CompletedWorkflowSpanParams @@ -556,8 +553,7 @@ async def handle_query(self, input: temporalio.worker.HandleQueryInput) -> Any: # We need to put this interceptor on the context too context = self._set_on_context(context) # Run under context with new span - token = opentelemetry.context.attach(context) - try: + with attached_context(context): # This won't be created if there was no context header self._completed_span( f"HandleQuery:{input.query}", @@ -567,13 +563,6 @@ async def handle_query(self, input: temporalio.worker.HandleQueryInput) -> Any: kind=opentelemetry.trace.SpanKind.SERVER, ) return await super().handle_query(input) - finally: - # In some exceptional cases this finally is executed with a - # different contextvars.Context than the one the token was created - # on. As such we do a best effort detach to avoid using a mismatched - # token. - if context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) def handle_update_validator( self, input: temporalio.worker.HandleUpdateInput @@ -644,31 +633,23 @@ def _top_level_workflow_context( success = False exception: Exception | None = None # Run under this context - token = opentelemetry.context.attach(context) - - try: - yield None - success = True - except temporalio.exceptions.FailureError as err: - # We only record the failure errors since those are the only ones - # that lead to workflow completions - exception = err - raise - finally: - # Create a completed span before detaching context - if exception or (success and success_is_complete): - self._completed_span( - f"CompleteWorkflow:{temporalio.workflow.info().workflow_type}", - exception=exception, - kind=opentelemetry.trace.SpanKind.INTERNAL, - ) - - # In some exceptional cases this finally is executed with a - # different contextvars.Context than the one the token was created - # on. As such we do a best effort detach to avoid using a mismatched - # token. - if context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) + with attached_context(context): + try: + yield None + success = True + except temporalio.exceptions.FailureError as err: + # We only record the failure errors since those are the only ones + # that lead to workflow completions + exception = err + raise + finally: + # Create a completed span before detaching context + if exception or (success and success_is_complete): + self._completed_span( + f"CompleteWorkflow:{temporalio.workflow.info().workflow_type}", + exception=exception, + kind=opentelemetry.trace.SpanKind.INTERNAL, + ) def _context_to_headers( self, headers: Mapping[str, temporalio.api.common.v1.Payload] diff --git a/temporalio/contrib/opentelemetry/_otel_interceptor.py b/temporalio/contrib/opentelemetry/_otel_interceptor.py index ff07f0f1d..5207095dc 100644 --- a/temporalio/contrib/opentelemetry/_otel_interceptor.py +++ b/temporalio/contrib/opentelemetry/_otel_interceptor.py @@ -35,6 +35,7 @@ import temporalio.nexus.system.workflow_service.models import temporalio.worker import temporalio.workflow +from temporalio.contrib.opentelemetry._context import attached_context from temporalio.contrib.opentelemetry._tracer_provider import ( ReplaySafeTracerProvider, ) @@ -125,8 +126,7 @@ def _maybe_span( yield return - token = opentelemetry.context.attach(context) if context else None - try: + with attached_context(context): with tracer.start_as_current_span( name, attributes=attributes, @@ -148,9 +148,6 @@ def _maybe_span( ) ) raise - finally: - if token and context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) class OpenTelemetryInterceptor( @@ -334,8 +331,7 @@ async def execute_activity( self, input: temporalio.worker.ExecuteActivityInput ) -> Any: context = _headers_to_context(input.headers) - token = opentelemetry.context.attach(context) - try: + with attached_context(context): info = temporalio.activity.info() with _maybe_span( get_tracer(__name__), @@ -349,9 +345,6 @@ async def execute_activity( kind=opentelemetry.trace.SpanKind.SERVER, ): return await super().execute_activity(input) - finally: - if context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) class _TracingNexusOperationInboundInterceptor( @@ -368,12 +361,8 @@ def __init__( @contextmanager def _top_level_context(self, headers: Mapping[str, str]) -> Iterator[None]: context = _nexus_headers_to_context(headers) - token = opentelemetry.context.attach(context) - try: + with attached_context(context): yield - finally: - if context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) async def execute_nexus_operation_start( self, input: temporalio.worker.ExecuteNexusOperationStartInput @@ -501,12 +490,8 @@ async def handle_update_handler( @contextmanager def _top_level_workflow_context(self, input: _InputWithHeaders) -> Iterator[None]: context = _headers_to_context(input.headers) - token = opentelemetry.context.attach(context) - try: + with attached_context(context): yield - finally: - if context is opentelemetry.context.get_current(): - opentelemetry.context.detach(token) class _TracingWorkflowOutboundInterceptor( diff --git a/tests/contrib/opentelemetry/test_opentelemetry.py b/tests/contrib/opentelemetry/test_opentelemetry.py index 0ea9530e6..4c11e95f8 100644 --- a/tests/contrib/opentelemetry/test_opentelemetry.py +++ b/tests/contrib/opentelemetry/test_opentelemetry.py @@ -1,20 +1,22 @@ from __future__ import annotations import asyncio +import contextvars import gc import logging import queue import threading import uuid from collections.abc import Callable, Generator, Iterable -from concurrent.futures import ThreadPoolExecutor +from concurrent.futures import Future, ThreadPoolExecutor from contextlib import contextmanager from dataclasses import dataclass from datetime import timedelta -from typing import Any, cast +from typing import Any, ParamSpec, TypeVar, cast import nexusrpc import opentelemetry.context +import opentelemetry.trace import pytest from opentelemetry import baggage, context from opentelemetry.sdk.trace import ReadableSpan, TracerProvider @@ -26,17 +28,30 @@ from temporalio.client import Client, WithStartWorkflowOperation, WorkflowUpdateStage from temporalio.common import RetryPolicy, WorkflowIDConflictPolicy from temporalio.contrib.opentelemetry import ( + OpenTelemetryInterceptor, TracingInterceptor, TracingWorkflowInboundInterceptor, + create_tracer_provider, ) +from temporalio.contrib.opentelemetry import _context as otel_context from temporalio.contrib.opentelemetry import workflow as otel_workflow +from temporalio.contrib.opentelemetry._otel_interceptor import ( + _TracingWorkflowInboundInterceptor as _OtelTracingWorkflowInboundInterceptor, +) from temporalio.exceptions import ( ApplicationError, ApplicationErrorCategory, NexusOperationError, ) from temporalio.testing import WorkflowEnvironment -from temporalio.worker import UnsandboxedWorkflowRunner, Worker +from temporalio.worker import ( + ExecuteWorkflowInput, + Interceptor, + UnsandboxedWorkflowRunner, + Worker, + WorkflowInboundInterceptor, + WorkflowInterceptorClassInput, +) from tests.helpers import LogCapturer from tests.helpers.nexus import make_nexus_endpoint_name @@ -916,54 +931,45 @@ async def test_opentelemetry_context_restored_after_activity( client_with_tracing: Client, activity: Callable[[], None], expect_failure: bool, + monkeypatch: pytest.MonkeyPatch, ) -> None: - attach_count = 0 - detach_count = 0 - original_attach = context.attach - original_detach = context.detach - - def tracked_attach(ctx): # type:ignore[reportMissingParameterType] - nonlocal attach_count - attach_count += 1 - return original_attach(ctx) - - def tracked_detach(token): # type:ignore[reportMissingParameterType] - nonlocal detach_count - detach_count += 1 - return original_detach(token) - - context.attach = tracked_attach - context.detach = tracked_detach + # Every context the interceptors attach must be detached again, also when + # the activity raises. Counting through the interceptors' own detach keeps + # this independent of other users of opentelemetry.context, such as the + # threading instrumentation's attach/detach around every thread's run(). + detached: list[bool] = [] + original_detach = otel_context._detach - try: - task_queue = f"task_queue_{uuid.uuid4()}" - async with Worker( - client_with_tracing, - task_queue=task_queue, - workflows=[ContextClearWorkflow], - activities=[activity], - ): - with baggage_values({"user.id": "test-123"}): - try: - await client_with_tracing.execute_workflow( - ContextClearWorkflow.run, - id=f"workflow_{uuid.uuid4()}", - task_queue=task_queue, - ) - assert not expect_failure, ( - "This test should have raised an exception" - ) - except Exception: - assert expect_failure, "This test is not expeced to raise" + def tracked_detach(ctx: Any, token: Any) -> bool: + result = original_detach(ctx, token) + detached.append(result) + return result - assert attach_count == detach_count, ( - f"Context leak detected: {attach_count} attaches vs {detach_count} detaches. " - ) - assert attach_count > 0, "Expected at least one context attach/detach" + monkeypatch.setattr(otel_context, "_detach", tracked_detach) - finally: - context.attach = original_attach - context.detach = original_detach + task_queue = f"task_queue_{uuid.uuid4()}" + async with Worker( + client_with_tracing, + task_queue=task_queue, + workflows=[ContextClearWorkflow], + activities=[activity], + ): + with baggage_values({"user.id": "test-123"}): + try: + await client_with_tracing.execute_workflow( + ContextClearWorkflow.run, + id=f"workflow_{uuid.uuid4()}", + task_queue=task_queue, + ) + assert not expect_failure, "This test should have raised an exception" + except Exception: + assert expect_failure, "This test is not expeced to raise" + + assert detached, "Expected at least one context attach/detach" + assert all(detached), ( + f"Context leak detected: {detached.count(False)} of {len(detached)} " + "attached contexts were not detached" + ) @activity.defn @@ -1052,7 +1058,7 @@ async def test_opentelemetry_standalone_activity_tracing( assert start_activity_span.attributes["temporalActivityType"] == "tracing_activity" -def test_opentelemetry_safe_detach(): +def _v1_workflow_context(success_is_complete: bool = True) -> Any: class _fake_self: def _load_workflow_context_carrier(*_args): return None @@ -1063,11 +1069,25 @@ def _set_on_context(self, ctx: Any): def _completed_span(*args: Any, **_kwargs: Any): pass - # create a context manager and force enter to happen on this thread - context_manager = TracingWorkflowInboundInterceptor._top_level_workflow_context( + return TracingWorkflowInboundInterceptor._top_level_workflow_context( _fake_self(), # type: ignore - success_is_complete=True, + success_is_complete=success_is_complete, + ) + + +def _v2_workflow_context() -> Any: + class _fake_input: + headers: dict[str, Any] = {} + + return _OtelTracingWorkflowInboundInterceptor._top_level_workflow_context( + object(), # type: ignore + _fake_input(), # type: ignore ) + + +def _assert_context_detach_is_safe(make_context_manager: Callable[[], Any]) -> None: + # create a context manager and force enter to happen on this thread + context_manager = make_context_manager() context_manager.__enter__() # move reference to context manager into queue @@ -1097,3 +1117,167 @@ def otel_context_error(record: logging.LogRecord) -> bool: assert capturer.find(otel_context_error) is None, ( "Detach from context message should not be logged" ) + + +@pytest.mark.parametrize( + "make_context_manager", + [_v1_workflow_context, _v2_workflow_context], + ids=["TracingInterceptor", "OpenTelemetryInterceptor"], +) +@pytest.mark.parametrize( + "threading_instrumented", [False, True], ids=["plain", "threading-instrumented"] +) +def test_opentelemetry_safe_detach( + make_context_manager: Callable[[], Any], threading_instrumented: bool +): + if not threading_instrumented: + _assert_context_detach_is_safe(make_context_manager) + return + # OpenTelemetry's threading instrumentation (strands turns it on when an + # Agent is created) propagates the current Context object into new + # threads, so a context-identity check alone would detach a token minted + # on another thread. + threading_instrumentation = pytest.importorskip( + "opentelemetry.instrumentation.threading" + ) + instrumentor = threading_instrumentation.ThreadingInstrumentor() + already_instrumented = instrumentor.is_instrumented_by_opentelemetry + if not already_instrumented: + instrumentor.instrument() + try: + _assert_context_detach_is_safe(make_context_manager) + finally: + if not already_instrumented: + instrumentor.uninstrument() + + +def _run_in_context_on_new_thread( + ctx: contextvars.Context, fn: Callable[[], Any] +) -> Any: + with ThreadPoolExecutor(max_workers=1) as pool: + return pool.submit(ctx.run, fn).result(timeout=5) + + +@pytest.mark.parametrize( + "make_context_manager", + [lambda: _v1_workflow_context(success_is_complete=False), _v2_workflow_context], + ids=["TracingInterceptor", "OpenTelemetryInterceptor"], +) +def test_opentelemetry_detach_after_thread_change( + make_context_manager: Callable[[], Any], +): + # Workflow activations run on a thread pool, so the asyncio task that + # entered one of these context managers can leave it on another thread + # while keeping the same contextvars.Context. The token is valid there and + # the detach must restore the outer context. + task_context = contextvars.copy_context() + outer = opentelemetry.context.set_value("outer", True) + task_context.run(opentelemetry.context.attach, outer) + context_manager = make_context_manager() + + _run_in_context_on_new_thread(task_context, context_manager.__enter__) + assert task_context.run(opentelemetry.context.get_current) is not outer + _run_in_context_on_new_thread( + task_context, lambda: context_manager.__exit__(None, None, None) + ) + assert task_context.run(opentelemetry.context.get_current) is outer + + +_P = ParamSpec("_P") +_T = TypeVar("_T") + + +class _AlternatingThreadExecutor(ThreadPoolExecutor): + """Runs each submission on a different thread than the one before it, the + way successive activations of a workflow can land on different pool + threads under load.""" + + def __init__(self) -> None: + super().__init__(max_workers=1) + self._pools = [ThreadPoolExecutor(max_workers=1) for _ in range(2)] + self._submissions = 0 + + def submit( + self, fn: Callable[_P, _T], /, *args: _P.args, **kwargs: _P.kwargs + ) -> Future[_T]: + pool = self._pools[self._submissions % len(self._pools)] + self._submissions += 1 + return pool.submit(fn, *args, **kwargs) + + def shutdown(self, wait: bool = True, *, cancel_futures: bool = False) -> None: + for pool in self._pools: + pool.shutdown(wait, cancel_futures=cancel_futures) + super().shutdown(wait, cancel_futures=cancel_futures) + + +@workflow.defn +class TimerWorkflow: + @workflow.run + async def run(self) -> None: + # Two activations: the start, and the timer firing. + await asyncio.sleep(0.01) + + +@pytest.mark.parametrize( + "make_interceptor", + [lambda: TracingInterceptor(get_tracer(__name__)), OpenTelemetryInterceptor], + ids=["TracingInterceptor", "OpenTelemetryInterceptor"], +) +async def test_opentelemetry_context_restored_after_activation_thread_change( + client: Client, + make_interceptor: Callable[[], Interceptor], + reset_otel_tracer_provider: Any, # type: ignore[reportUnusedParameter] +): + # OpenTelemetryInterceptor insists on a replay-safe global provider. + opentelemetry.trace.set_tracer_provider(create_tracer_provider()) + records: list[tuple[threading.Thread, threading.Thread, bool]] = [] + + class ContextRestoredWorkflowInboundInterceptor(WorkflowInboundInterceptor): + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + entered_on = threading.current_thread() + before = opentelemetry.context.get_current() + result = await super().execute_workflow(input) + records.append( + ( + entered_on, + threading.current_thread(), + opentelemetry.context.get_current() is before, + ) + ) + return result + + class ContextRestoredInterceptor(Interceptor): + def workflow_interceptor_class( + self, input: WorkflowInterceptorClassInput + ) -> type[WorkflowInboundInterceptor] | None: + return ContextRestoredWorkflowInboundInterceptor + + executor = _AlternatingThreadExecutor() + try: + async with Worker( + client, + task_queue=f"task_queue_{uuid.uuid4()}", + workflows=[TimerWorkflow], + # The first interceptor is the outermost, so it sees whatever the + # tracing interceptor leaves attached. + interceptors=[ContextRestoredInterceptor(), make_interceptor()], + workflow_task_executor=executor, + workflow_runner=UnsandboxedWorkflowRunner(), + ) as worker: + await client.execute_workflow( + TimerWorkflow.run, + id=f"workflow_{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + finally: + executor.shutdown() + + assert len(records) == 1 + entered_on, left_on, context_restored = records[0] + assert entered_on is not left_on, ( + "The workflow was expected to change threads between its activations" + ) + assert context_restored, ( + "The tracing interceptor must detach its context even though the " + "workflow left it on a different thread than it entered on" + )