diff --git a/cadence/__init__.py b/cadence/__init__.py index c1c2a17..626fff9 100644 --- a/cadence/__init__.py +++ b/cadence/__init__.py @@ -6,6 +6,7 @@ # Import main client functionality from .client import Client +from .context import ContextPropagator, ContextVarPropagator from .worker import Registry from . import workflow @@ -13,6 +14,8 @@ __all__ = [ "Client", + "ContextPropagator", + "ContextVarPropagator", "Registry", "workflow", ] diff --git a/cadence/_internal/activity/_activity_executor.py b/cadence/_internal/activity/_activity_executor.py index 1c5a8aa..e213926 100644 --- a/cadence/_internal/activity/_activity_executor.py +++ b/cadence/_internal/activity/_activity_executor.py @@ -4,10 +4,11 @@ from logging import getLogger import time from traceback import format_exception -from typing import Any, Callable, Optional, Union, cast +from typing import Any, Callable, Optional, Sequence, Union, cast from google.protobuf.duration import to_timedelta from google.protobuf.timestamp import to_datetime from cadence._internal.activity._context import _Context, _SyncContext +from cadence._internal.context import header_to_dict from cadence._internal.activity._definition import BaseDefinition, ExecutionStrategy from cadence._internal.activity._heartbeat import _HeartbeatSender from cadence.activity import ActivityInfo, ActivityDefinition @@ -19,6 +20,7 @@ RespondActivityTaskCompletedRequest, ) from cadence.client import Client +from cadence.context import ContextPropagator from cadence.metrics import ( duration_between, duration_from_nanoseconds, @@ -53,6 +55,7 @@ def __init__( max_workers: int, registry: Callable[[str], ActivityDefinition], metrics_emitter: MetricsEmitter | None = None, + context_propagators: Sequence[ContextPropagator] = (), ): self._client = client self._data_converter = client.data_converter @@ -62,6 +65,7 @@ def __init__( self._metrics_emitter: MetricsEmitter = ( metrics_emitter if metrics_emitter is not None else NoOpMetricsEmitter() ) + self._context_propagators = tuple(context_propagators) self._thread_pool = ThreadPoolExecutor( max_workers=max_workers, thread_name_prefix=f"{task_list}-activity-" ) @@ -138,13 +142,22 @@ def _create_context( ) if activity_def.strategy == ExecutionStrategy.ASYNC: - return _Context(self._client, info, activity_def, heartbeat_sender) + return _Context( + self._client, + info, + activity_def, + heartbeat_sender, + self._context_propagators, + header_to_dict(task.header), + ) return _SyncContext( self._client, info, activity_def, self._thread_pool, heartbeat_sender, + self._context_propagators, + header_to_dict(task.header), ) async def _report_failure( diff --git a/cadence/_internal/activity/_context.py b/cadence/_internal/activity/_context.py index 5d574b3..c4900ee 100644 --- a/cadence/_internal/activity/_context.py +++ b/cadence/_internal/activity/_context.py @@ -1,15 +1,18 @@ import asyncio +import contextvars import threading from concurrent.futures import Future as ConcurrentFuture from concurrent.futures.thread import ThreadPoolExecutor from datetime import timedelta -from typing import Any, Type +from typing import Any, Mapping, Sequence, Type from cadence import Client from cadence._internal.activity._definition import BaseDefinition from cadence._internal.activity._heartbeat import _HeartbeatSender +from cadence._internal.context import extract_headers from cadence.activity import ActivityInfo, ActivityContext from cadence.api.v1.common_pb2 import Payload +from cadence.context import ContextPropagator class _Context(ActivityContext): @@ -19,6 +22,8 @@ def __init__( info: ActivityInfo, activity_def: BaseDefinition[[Any], Any], heartbeat_sender: _HeartbeatSender, + context_propagators: Sequence[ContextPropagator] = (), + headers: Mapping[str, bytes] | None = None, ): self._client = client self._info = info @@ -28,10 +33,15 @@ def __init__( self._heartbeat_tasks: set[asyncio.Future[Any]] = set() self._heartbeat_tasks_lock = threading.Lock() self._cancel_event = asyncio.Event() + self._context_propagators = tuple(context_propagators) + self._headers = dict(headers) if headers is not None else {} async def execute(self, payload: Payload) -> Any: params = self._to_params(payload) - self._activity_task = asyncio.create_task(self._run_activity(params)) + # Fresh context: activity vars come from task headers, not worker ambient state. + self._activity_task = asyncio.create_task( + self._run_activity(params), context=contextvars.Context() + ) try: return await self._activity_task except asyncio.CancelledError as e: @@ -46,7 +56,8 @@ async def execute(self, payload: Payload) -> Any: async def _run_activity(self, params: list[Any]) -> Any: with self._activate(): - return await self._activity_def.impl_fn(*params) + with extract_headers(self._context_propagators, self._headers): + return await self._activity_def.impl_fn(*params) async def _wait_pending_heartbeats(self) -> None: tasks = self._pending_heartbeat_tasks() @@ -115,8 +126,17 @@ def __init__( activity_def: BaseDefinition[[Any], Any], executor: ThreadPoolExecutor, heartbeat_sender: _HeartbeatSender, + context_propagators: Sequence[ContextPropagator] = (), + headers: Mapping[str, bytes] | None = None, ): - super().__init__(client, info, activity_def, heartbeat_sender) + super().__init__( + client, + info, + activity_def, + heartbeat_sender, + context_propagators, + headers, + ) self._executor = executor self._sync_cancel_event = threading.Event() self._loop: asyncio.AbstractEventLoop | None = None @@ -134,7 +154,8 @@ async def execute(self, payload: Payload) -> Any: def _run(self, args: list[Any]) -> Any: with self._activate(): - return self._activity_def.impl_fn(*args) + with extract_headers(self._context_propagators, self._headers): + return self._activity_def.impl_fn(*args) def client(self) -> Client: raise RuntimeError("client is only supported in async activities") diff --git a/cadence/_internal/context.py b/cadence/_internal/context.py new file mode 100644 index 0000000..a310b33 --- /dev/null +++ b/cadence/_internal/context.py @@ -0,0 +1,72 @@ +"""Internal helpers for converting and propagating Cadence headers.""" + +from __future__ import annotations + +from collections.abc import Iterator, Mapping, Sequence +from contextlib import ExitStack, contextmanager +from typing import Protocol + +from cadence.api.v1.common_pb2 import Header, Payload +from cadence.context import ContextPropagator + + +class _HeaderCarrier(Protocol): + """Proto message that carries a Cadence ``Header`` field.""" + + header: Header + + +def header_to_dict(header: Header) -> dict[str, bytes]: + """Return a detached copy of a proto header's byte payloads.""" + return {key: bytes(payload.data) for key, payload in header.fields.items()} + + +def header_from_dict(headers: Mapping[str, bytes]) -> Header | None: + """Build a Header, omitting it entirely when no values were injected.""" + if not headers: + return None + return Header(fields={key: Payload(data=value) for key, value in headers.items()}) + + +def inject_headers(propagators: Sequence[ContextPropagator]) -> dict[str, bytes]: + """Inject ordered propagators, with later propagators winning duplicate keys.""" + headers: dict[str, bytes] = {} + for propagator in propagators: + headers.update(propagator.inject()) + return headers + + +def set_header(attrs: _HeaderCarrier, propagators: Sequence[ContextPropagator]) -> None: + """Attach injected context to ``attrs``, leaving the field unset when empty.""" + set_header_from_dict(attrs, inject_headers(propagators)) + + +def set_header_from_dict(attrs: _HeaderCarrier, headers: Mapping[str, bytes]) -> None: + """Attach ``headers`` to ``attrs``, leaving the field unset when empty.""" + header = header_from_dict(headers) + if header is not None: + attrs.header.CopyFrom(header) + + +def validate_propagators(propagators: Sequence[ContextPropagator]) -> None: + """Fail fast when a propagator does not implement the protocol.""" + for index, propagator in enumerate(propagators): + if not callable(getattr(propagator, "inject", None)): + raise TypeError( + f"context_propagators[{index}] is missing a callable inject() method" + ) + if not callable(getattr(propagator, "extract", None)): + raise TypeError( + f"context_propagators[{index}] is missing a callable extract() method" + ) + + +@contextmanager +def extract_headers( + propagators: Sequence[ContextPropagator], headers: Mapping[str, bytes] +) -> Iterator[None]: + """Activate propagators in order and reliably unwind partial extraction.""" + with ExitStack() as stack: + for propagator in propagators: + stack.enter_context(propagator.extract(headers)) + yield diff --git a/cadence/_internal/workflow/context.py b/cadence/_internal/workflow/context.py index 307ac23..21408be 100644 --- a/cadence/_internal/workflow/context.py +++ b/cadence/_internal/workflow/context.py @@ -13,6 +13,7 @@ from cadence._internal.workflow.statemachine.marker_state_machine import ( SIDE_EFFECT_MARKER_NAME, ) +from cadence._internal.context import inject_headers, set_header from cadence.api.v1 import workflow_pb2 from cadence.api.v1.common_pb2 import ( ActivityType, @@ -29,6 +30,7 @@ ) from cadence.api.v1.tasklist_pb2 import TaskList, TaskListKind from cadence.data_converter import DataConverter +from cadence.context import ContextPropagator from cadence.workflow import ( ActivityOptions, ChildWorkflowFuture, @@ -51,11 +53,13 @@ def __init__( self, info: WorkflowInfo, decision_manager: DecisionManager, + context_propagators: tuple[ContextPropagator, ...] = (), ): self._info = info self._replay_mode = True self._replay_current_time: Optional[datetime] = None self._decision_manager = decision_manager + self._context_propagators = context_propagators self._cancellation_info: WorkflowCancellationInfo | None = None def info(self) -> WorkflowInfo: @@ -110,13 +114,13 @@ async def execute_activity( task_list=TaskList(kind=TaskListKind.TASK_LIST_KIND_NORMAL, name=task_list), input=activity_input, retry_policy=retry_policy_to_proto(opts.get("retry_policy")), - header=None, request_local_dispatch=False, schedule_to_close_timeout=_round_to_nearest_second(schedule_to_close), schedule_to_start_timeout=_round_to_nearest_second(schedule_to_start), start_to_close_timeout=_round_to_nearest_second(start_to_close), heartbeat_timeout=_round_to_nearest_second(heartbeat), ) + set_header(schedule_attributes, self._context_propagators) future = self._decision_manager.schedule_activity(schedule_attributes) result_payload = await future @@ -213,6 +217,7 @@ def _build_child_workflow_attrs( ), task_start_to_close_timeout=_round_to_nearest_second(task_timeout), ) + set_header(schedule_attributes, self._context_propagators) cron_schedule = kwargs.get("cron_schedule") if cron_schedule: @@ -358,6 +363,9 @@ def request_cancel( def is_cancel_requested(self) -> bool: return self._cancellation_info is not None + def inject_propagated_headers(self) -> dict[str, bytes]: + return inject_headers(self._context_propagators) + @contextmanager def _activate(self) -> Iterator["Context"]: token = WorkflowContext._var.set(self) diff --git a/cadence/_internal/workflow/workflow_engine.py b/cadence/_internal/workflow/workflow_engine.py index 3417b3e..0a2926e 100644 --- a/cadence/_internal/workflow/workflow_engine.py +++ b/cadence/_internal/workflow/workflow_engine.py @@ -3,8 +3,9 @@ from asyncio import CancelledError, InvalidStateError from dataclasses import dataclass from functools import singledispatchmethod -from typing import List, Optional +from typing import List, Mapping, Optional, Sequence +from cadence._internal.context import extract_headers, set_header_from_dict from cadence._internal.workflow.context import Context from cadence._internal.workflow.decision_events_iterator import DecisionEventsIterator from cadence._internal.workflow.deterministic_event_loop import ( @@ -34,6 +35,7 @@ ) from cadence.api.v1.tasklist_pb2 import TaskList from cadence.error import ContinueAsNewError +from cadence.context import ContextPropagator from cadence.workflow import WorkflowDefinition, WorkflowInfo logger = logging.getLogger(__name__) @@ -46,7 +48,13 @@ class DecisionResult: class WorkflowEngine: - def __init__(self, info: WorkflowInfo, workflow_definition: WorkflowDefinition): + def __init__( + self, + info: WorkflowInfo, + workflow_definition: WorkflowDefinition, + context_propagators: Sequence[ContextPropagator] = (), + headers: Mapping[str, bytes] | None = None, + ): self._event_loop = DeterministicEventLoop() self._decision_manager = DecisionManager(self._event_loop) self._data_converter = info.data_converter @@ -55,7 +63,9 @@ def __init__(self, info: WorkflowInfo, workflow_definition: WorkflowDefinition): self._event_loop, workflow_definition, ) - self._context = Context(info, self._decision_manager) + self._context_propagators = tuple(context_propagators) + self._headers = dict(headers) if headers is not None else {} + self._context = Context(info, self._decision_manager, self._context_propagators) def process_decision( self, @@ -77,30 +87,31 @@ def process_decision( try: # Activate workflow context for the entire decision processing with self._context._activate() as ctx: - # Log decision task processing start with full context (matches Java ReplayDecisionTaskHandler) - logger.info( - "Processing decision task for workflow", - extra={ - "workflow_type": ctx.info().workflow_type, - "workflow_id": ctx.info().workflow_id, - "run_id": ctx.info().workflow_run_id, - "query": query.query_type if query else None, - }, - ) + with extract_headers(self._context_propagators, self._headers): + # Log decision task processing start with full context (matches Java ReplayDecisionTaskHandler) + logger.info( + "Processing decision task for workflow", + extra={ + "workflow_type": ctx.info().workflow_type, + "workflow_id": ctx.info().workflow_id, + "run_id": ctx.info().workflow_run_id, + "query": query.query_type if query else None, + }, + ) - # Create DecisionEventsIterator for structured event processing - events_iterator = DecisionEventsIterator(events) + # Create DecisionEventsIterator for structured event processing + events_iterator = DecisionEventsIterator(events) - # Process decision events using iterator-driven approach - self._process_decision_events(ctx, events_iterator) + # Process decision events using iterator-driven approach + self._process_decision_events(ctx, events_iterator) - if query: - return self._execute_query(query) + if query: + return self._execute_query(query) - # Collect all pending decisions from state machines - decisions = self._decision_manager.collect_pending_decisions() + # Collect all pending decisions from state machines + decisions = self._decision_manager.collect_pending_decisions() - return DecisionResult(decisions=decisions, query_result=None) + return DecisionResult(decisions=decisions, query_result=None) # TODO: reevaluate if this is needed to log error here or in the caller except Exception as e: @@ -252,6 +263,8 @@ def _maybe_complete_workflow(self) -> Optional[Decision]: attrs.task_start_to_close_timeout.FromTimedelta( e.task_start_to_close_timeout ) + if e.headers is not None: + set_header_from_dict(attrs, e.headers) return Decision( continue_as_new_workflow_execution_decision_attributes=attrs, ) diff --git a/cadence/activity.py b/cadence/activity.py index 82d607c..2ad9e6d 100644 --- a/cadence/activity.py +++ b/cadence/activity.py @@ -170,8 +170,10 @@ def wait_for_cancelled(self, timeout: timedelta | None = None) -> bool: ... @contextmanager def _activate(self) -> Iterator[None]: token = ActivityContext._var.set(self) - yield None - ActivityContext._var.reset(token) + try: + yield None + finally: + ActivityContext._var.reset(token) @staticmethod def is_set() -> bool: diff --git a/cadence/client.py b/cadence/client.py index 5d11fb0..2c60195 100644 --- a/cadence/client.py +++ b/cadence/client.py @@ -17,6 +17,7 @@ ) from cadence._internal.workflow.memo import memo_to_proto from cadence._internal.workflow.retry_policy import retry_policy_to_proto +from cadence._internal.context import set_header, validate_propagators from cadence.api.v1 import schedule_pb2 from cadence.api.v1.common_pb2 import ( Memo, @@ -62,6 +63,7 @@ from cadence.api.v1 import workflow_pb2 from cadence.api.v1.tasklist_pb2 import TaskList from cadence.data_converter import DataConverter, DefaultDataConverter +from cadence.context import ContextPropagator from cadence.metrics import MetricsEmitter, NoOpMetricsEmitter from cadence.workflow import ( ActiveClusterSelectionPolicy, @@ -164,6 +166,7 @@ class ClientOptions(TypedDict, total=False): compression: Compression metrics_emitter: MetricsEmitter interceptors: list[ClientInterceptor] + context_propagators: Sequence[ContextPropagator] _DEFAULT_OPTIONS: ClientOptions = { @@ -176,6 +179,7 @@ class ClientOptions(TypedDict, total=False): "compression": Compression.NoCompression, "metrics_emitter": NoOpMetricsEmitter(), "interceptors": [], + "context_propagators": (), } @@ -220,6 +224,10 @@ def schedule_stub(self) -> ScheduleAPIStub: def metrics_emitter(self) -> MetricsEmitter: return self._options["metrics_emitter"] + @property + def context_propagators(self) -> tuple[ContextPropagator, ...]: + return tuple(self._options["context_propagators"]) + async def ready(self) -> None: await self._channel.channel_ready() @@ -316,6 +324,8 @@ def _build_start_workflow_request( if memo_proto is not None: request.memo.CopyFrom(memo_proto) + set_header(request, self.context_propagators) + return request async def start_workflow( @@ -758,6 +768,10 @@ def _validate_and_copy_defaults(options: ClientOptions) -> ClientOptions: for key, value in _DEFAULT_OPTIONS.items(): if key not in options: cast(dict, options)[key] = value + cast(dict, options)["context_propagators"] = tuple( + options.get("context_propagators") or () + ) + validate_propagators(options["context_propagators"]) return options diff --git a/cadence/context.py b/cadence/context.py new file mode 100644 index 0000000..7f7ae4d --- /dev/null +++ b/cadence/context.py @@ -0,0 +1,59 @@ +"""Context propagation between Cadence clients, workflows, and activities.""" + +from __future__ import annotations + +from abc import abstractmethod +from collections.abc import Callable, Iterator, Mapping +from contextlib import contextmanager +from contextvars import ContextVar +from typing import ContextManager, Generic, Protocol, TypeVar, cast + +T = TypeVar("T") +_UNSET = object() + + +class ContextPropagator(Protocol): + """Inject and extract application context in Cadence headers.""" + + @abstractmethod + def inject(self) -> Mapping[str, bytes]: + """Return headers for the context currently active in this execution.""" + raise NotImplementedError() + + @abstractmethod + def extract(self, headers: Mapping[str, bytes]) -> ContextManager[None]: + """Activate context from ``headers`` for the duration of a scope.""" + raise NotImplementedError() + + +class ContextVarPropagator(ContextPropagator, Generic[T]): + """A generic :class:`ContextVar`-backed context propagator.""" + + def __init__( + self, + var: ContextVar[T], + header_key: str, + serialize: Callable[[T], bytes], + deserialize: Callable[[bytes], T], + ) -> None: + self._var = var + self._header_key = header_key + self._serialize = serialize + self._deserialize = deserialize + + def inject(self) -> Mapping[str, bytes]: + value = self._var.get(_UNSET) + if value is _UNSET: + return {} + return {self._header_key: self._serialize(cast(T, value))} + + @contextmanager + def extract(self, headers: Mapping[str, bytes]) -> Iterator[None]: + if self._header_key not in headers: + yield + return + token = self._var.set(self._deserialize(headers[self._header_key])) + try: + yield + finally: + self._var.reset(token) diff --git a/cadence/error.py b/cadence/error.py index c391e12..92d18fc 100644 --- a/cadence/error.py +++ b/cadence/error.py @@ -12,6 +12,7 @@ def __init__( task_list: str | None = None, execution_start_to_close_timeout: timedelta | None = None, task_start_to_close_timeout: timedelta | None = None, + headers: dict[str, bytes] | None = None, ): super().__init__("ContinueAsNew") self.workflow_args = args @@ -19,6 +20,7 @@ def __init__( self.task_list = task_list self.execution_start_to_close_timeout = execution_start_to_close_timeout self.task_start_to_close_timeout = task_start_to_close_timeout + self.headers = headers class ActivityFailure(Exception): diff --git a/cadence/sample/context_propagation_example.py b/cadence/sample/context_propagation_example.py new file mode 100644 index 0000000..69c64d9 --- /dev/null +++ b/cadence/sample/context_propagation_example.py @@ -0,0 +1,74 @@ +"""Run a ContextVar propagation example against a Cadence server.""" + +import asyncio +import os +from contextvars import ContextVar +from datetime import timedelta + +from cadence import Client, ContextVarPropagator, workflow +from cadence.api.v1.history_pb2 import EventFilterType +from cadence.api.v1.service_workflow_pb2 import GetWorkflowExecutionHistoryRequest +from cadence.worker import Registry, Worker + +REQUEST_ID: ContextVar[str] = ContextVar("request_id") +REQUEST_ID_PROPAGATOR = ContextVarPropagator( + REQUEST_ID, + "request-id", + lambda value: value.encode(), + bytes.decode, +) +registry = Registry() + + +@registry.activity() +async def read_request_id() -> str: + return REQUEST_ID.get() + + +@registry.workflow() +class RequestIdWorkflow: + @workflow.run + async def run(self) -> str: + return await read_request_id.with_options( + schedule_to_close_timeout=timedelta(seconds=30) + ).execute() + + +async def main() -> None: + domain = os.environ.get("CADENCE_DOMAIN", "default") + target = os.environ.get("CADENCE_TARGET", "localhost:7833") + task_list = "context-propagation-example" + client = Client( + domain=domain, + target=target, + context_propagators=(REQUEST_ID_PROPAGATOR,), + ) + + async with client, Worker(client, task_list, registry): + token = REQUEST_ID.set("example-request-id") + try: + execution = await client.start_workflow( + "RequestIdWorkflow", + task_list=task_list, + execution_start_to_close_timeout=timedelta(minutes=1), + ) + finally: + REQUEST_ID.reset(token) + + history = await client.workflow_stub.GetWorkflowExecutionHistory( + GetWorkflowExecutionHistoryRequest( + domain=domain, + workflow_execution=execution, + wait_for_new_event=True, + history_event_filter_type=EventFilterType.EVENT_FILTER_TYPE_CLOSE_EVENT, + skip_archival=True, + ) + ) + result = history.history.events[ + -1 + ].workflow_execution_completed_event_attributes.result.data.decode() + print(f"Workflow received request context: {result}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/cadence/testing/_workflow_environment.py b/cadence/testing/_workflow_environment.py index 0f7513f..c60f952 100644 --- a/cadence/testing/_workflow_environment.py +++ b/cadence/testing/_workflow_environment.py @@ -41,8 +41,10 @@ import inspect import logging import uuid +import asyncio +import contextvars from asyncio import Future, get_running_loop -from collections.abc import Iterator +from collections.abc import Iterator, Mapping from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager from datetime import datetime, timedelta, timezone @@ -50,6 +52,7 @@ Any, Callable, Optional, + Sequence, Tuple, Type, Union, @@ -58,12 +61,19 @@ ) from cadence._internal.activity._definition import BaseDefinition +from cadence._internal.context import ( + extract_headers, + header_from_dict, + header_to_dict, + inject_headers, +) from cadence._internal.workflow.deterministic_event_loop import DeterministicEventLoop from cadence._internal.workflow.workflow_instance import WorkflowInstance from cadence.activity import ActivityContext, ActivityInfo from cadence.api.v1.common_pb2 import Payload, WorkflowExecution from cadence.api.v1.query_pb2 import QueryRejectCondition from cadence.client import Client, ClientOptions, StartWorkflowOptions +from cadence.context import ContextPropagator from cadence.data_converter import DataConverter, DefaultDataConverter from cadence.metrics import NoOpMetricsEmitter from cadence.workflow import ( @@ -80,6 +90,24 @@ logger = logging.getLogger(__name__) +def _headers_after_proto_roundtrip(headers: Mapping[str, bytes]) -> dict[str, bytes]: + """Round-trip headers through proto, matching production decision encoding.""" + proto = header_from_dict(headers) + if proto is None: + return {} + return header_to_dict(proto) + + +async def _run_in_fresh_context(awaitable_factory: Callable[[], Any]) -> Any: + """Run an awaitable in a new contextvars context, mirroring activity workers.""" + + async def _wrapper() -> Any: + return await awaitable_factory() + + task = asyncio.create_task(_wrapper(), context=contextvars.Context()) + return await task + + class _Unset: """Sentinel for an unset keyword argument.""" @@ -179,7 +207,10 @@ async def execute_activity( *args: Any, **kwargs: Unpack[ActivityOptions], ) -> ResultType: - return await self._env._invoke_activity(activity, result_type, args, self._info) + outbound = self.inject_propagated_headers() + return await self._env._invoke_activity( + activity, result_type, args, self._info, outbound + ) async def execute_child_workflow( self, @@ -200,8 +231,14 @@ async def start_child_workflow( *args: Any, **kwargs: Unpack[ChildWorkflowOptions], ) -> ChildWorkflowFuture[ResultType]: + outbound = self.inject_propagated_headers() return await self._env._start_child_workflow( - workflow_type, result_type, args, kwargs, self._info + workflow_type, + result_type, + args, + kwargs, + self._info, + headers=outbound, ) async def start_timer(self, duration: timedelta) -> None: @@ -281,6 +318,9 @@ async def signal_external_workflow( def is_cancel_requested(self) -> bool: return False + def inject_propagated_headers(self) -> dict[str, bytes]: + return inject_headers(self._env._context_propagators) + class _Execution: """Holds the in-memory state machine for a single workflow execution.""" @@ -290,6 +330,7 @@ def __init__( env: "TestWorkflowEnvironment", workflow_definition: WorkflowDefinition, info: WorkflowInfo, + headers: dict[str, bytes], ) -> None: self._env = env self.info = info @@ -301,6 +342,7 @@ def __init__( # at the boundary (mirroring the production WorkflowEngine). self._instance = WorkflowInstance(self._loop, workflow_definition) self._context = _InMemoryWorkflowContext(env, info) + self._headers = dict(headers) self.completed: bool = False self.result_payload: Optional[Payload] = None self.error: Optional[BaseException] = None @@ -353,7 +395,8 @@ def _driving(self) -> Iterator[None]: self._env._driving_execution = self try: with self._context._activate(): - yield + with extract_headers(self._env._context_propagators, self._headers): + yield finally: self._env._driving_execution = None @@ -380,15 +423,16 @@ def _run_until_settled(self) -> None: def run_query(self, query_type: str, args: Payload) -> Payload: with self._context._activate(): - query_def = self._definition.queries.get(query_type) - if query_def is None: - raise ValueError( - f"Unknown query type '{query_type}'. " - f"Known types: {list(self._definition.queries.keys())}" - ) - query_args = query_def.params_from_payload(self._data_converter, args) - result = self._instance.handle_query(query_def, query_args) - return self._data_converter.to_data([result]) + with extract_headers(self._env._context_propagators, self._headers): + query_def = self._definition.queries.get(query_type) + if query_def is None: + raise ValueError( + f"Unknown query type '{query_type}'. " + f"Known types: {list(self._definition.queries.keys())}" + ) + query_args = query_def.params_from_payload(self._data_converter, args) + result = self._instance.handle_query(query_def, query_args) + return self._data_converter.to_data([result]) def _collect_outcome(self) -> None: signal_failure = self._instance.get_signal_failure() @@ -424,6 +468,7 @@ def __init__(self, env: "TestWorkflowEnvironment") -> None: data_converter=env._data_converter, identity="test-workflow-environment", metrics_emitter=NoOpMetricsEmitter(), + context_propagators=env._context_propagators, ) async def ready(self) -> None: @@ -515,11 +560,13 @@ def __init__( task_list: str = "test-task-list", data_converter: Optional[DataConverter] = None, start_time: Optional[datetime] = None, + context_propagators: Sequence[ContextPropagator] = (), ) -> None: self._registry = registry self._domain = domain self._default_task_list = task_list self._data_converter = data_converter or DefaultDataConverter() + self._context_propagators = tuple(context_propagators) self._activity_mocks: dict[str, Union[_ValueMock, _FnMock]] = {} self._executions: dict[str, _Execution] = {} self._last_execution: Optional[_Execution] = None @@ -681,7 +728,14 @@ async def _start_workflow( workflow_task_list=task_list, data_converter=self._data_converter, ) - execution = _Execution(self, definition, info) + execution = _Execution( + self, + definition, + info, + _headers_after_proto_roundtrip( + inject_headers(self._client.context_propagators) + ), + ) self._executions[workflow_id] = execution self._last_execution = execution @@ -737,6 +791,7 @@ async def _invoke_activity( result_type: Type[ResultType], args: Tuple[Any, ...], info: WorkflowInfo, + headers: Mapping[str, bytes], ) -> ResultType: name = activity if isinstance(activity, str) else getattr(activity, "name") dc = self._data_converter @@ -763,15 +818,24 @@ async def _invoke_activity( else: call_args = list(args) - if isinstance(mock, _ValueMock): - result: Any = mock.value - elif isinstance(mock, _FnMock): - result = mock.fn(*call_args) - if inspect.iscoroutine(result): - result = await result - else: - assert definition is not None - result = await self._run_real_activity(definition, call_args, name, info) + activity_headers = _headers_after_proto_roundtrip(headers) + + async def _body() -> Any: + with extract_headers(self._context_propagators, activity_headers): + if isinstance(mock, _ValueMock): + result: Any = mock.value + elif isinstance(mock, _FnMock): + result = mock.fn(*call_args) + if inspect.iscoroutine(result): + result = await result + else: + assert definition is not None + result = await self._run_real_activity( + definition, call_args, name, info + ) + return result + + result = await _run_in_fresh_context(_body) result_payload = dc.to_data([result]) return cast(ResultType, dc.from_data(result_payload, [result_type])[0]) @@ -797,6 +861,8 @@ async def _start_child_workflow( args: Tuple[Any, ...], kwargs: ChildWorkflowOptions, parent_info: WorkflowInfo, + *, + headers: Mapping[str, bytes], ) -> ChildWorkflowFuture[ResultType]: definition = self._resolve_workflow(workflow_type) child_id = kwargs.get("workflow_id") or ( @@ -816,7 +882,7 @@ async def _start_child_workflow( input_payload = self._data_converter.to_data(list(args)) result_payload = await self._run_workflow_coro( - definition, child_info, input_payload + definition, child_info, input_payload, headers ) loop = cast(DeterministicEventLoop, get_running_loop()) @@ -835,16 +901,20 @@ async def _run_workflow_coro( definition: WorkflowDefinition, info: WorkflowInfo, input_payload: Payload, + headers: Mapping[str, bytes], ) -> Payload: - instance = definition.cls() - run_method = definition.get_run_method(instance) + child_headers = _headers_after_proto_roundtrip(headers) run_args = definition.run_signature.params_from_payload( self._data_converter, input_payload ) - child_ctx = _InMemoryWorkflowContext(self, info) - token = WorkflowContext._var.set(child_ctx) - try: - result = await run_method(*run_args) - finally: - WorkflowContext._var.reset(token) + + async def _body() -> Any: + instance = definition.cls() + run_method = definition.get_run_method(instance) + child_ctx = _InMemoryWorkflowContext(self, info) + with child_ctx._activate(): + with extract_headers(self._context_propagators, child_headers): + return await run_method(*run_args) + + result = await _run_in_fresh_context(_body) return self._data_converter.to_data([result]) diff --git a/cadence/worker/_activity.py b/cadence/worker/_activity.py index 8b66bac..5746c4a 100644 --- a/cadence/worker/_activity.py +++ b/cadence/worker/_activity.py @@ -60,6 +60,7 @@ def __init__( max_concurrent, registry.get_activity, options["metrics_emitter"], + context_propagators=options.get("context_propagators", ()), ) self._poller = Poller[PollForActivityTaskResponse]( self._num_pollers, diff --git a/cadence/worker/_decision_task_handler.py b/cadence/worker/_decision_task_handler.py index 5c9caec..5db096f 100644 --- a/cadence/worker/_decision_task_handler.py +++ b/cadence/worker/_decision_task_handler.py @@ -6,6 +6,7 @@ from typing import Optional, Sequence from cadence._internal.workflow.history_event_iterator import iterate_history_events +from cadence._internal.context import header_to_dict from cadence._internal.workflow.memo import memo_from_proto from cadence.api.v1.common_pb2 import Payload from cadence.api.v1.decision_pb2 import Decision @@ -91,6 +92,7 @@ def __init__( super().__init__(client, task_list, identity, **options) self._registry = registry self._executor = executor + self._context_propagators = tuple(options.get("context_propagators", ())) async def _handle_task_implementation( self, task: PollForDecisionTaskResponse @@ -194,6 +196,8 @@ async def _handle_task_implementation( workflow_engine = WorkflowEngine( info=workflow_info, workflow_definition=workflow_definition, + context_propagators=self._context_propagators, + headers=header_to_dict(started_attrs.header), ) exec_start_ns = time.monotonic_ns() diff --git a/cadence/worker/_types.py b/cadence/worker/_types.py index d588c7a..138a54a 100644 --- a/cadence/worker/_types.py +++ b/cadence/worker/_types.py @@ -1,9 +1,10 @@ from __future__ import annotations from datetime import timedelta -from typing import TYPE_CHECKING, TypedDict +from typing import TYPE_CHECKING, Sequence, TypedDict if TYPE_CHECKING: + from cadence.context import ContextPropagator from cadence.metrics import MetricsEmitter @@ -18,6 +19,7 @@ class WorkerOptions(TypedDict, total=False): disable_activity_worker: bool identity: str metrics_emitter: MetricsEmitter + context_propagators: Sequence[ContextPropagator] _DEFAULT_WORKER_OPTIONS: WorkerOptions = { @@ -28,6 +30,7 @@ class WorkerOptions(TypedDict, total=False): "decision_task_pollers": 2, "disable_workflow_worker": False, "disable_activity_worker": False, + "context_propagators": (), } _LONG_POLL_TIMEOUT = timedelta(seconds=60) diff --git a/cadence/worker/_worker.py b/cadence/worker/_worker.py index d14115a..9a3cecd 100644 --- a/cadence/worker/_worker.py +++ b/cadence/worker/_worker.py @@ -8,6 +8,7 @@ from cadence.worker._activity import ActivityWorker from cadence.worker._decision import DecisionWorker from cadence.worker._types import WorkerOptions, _DEFAULT_WORKER_OPTIONS +from cadence._internal.context import validate_propagators logger = logging.getLogger(__name__) @@ -70,6 +71,8 @@ def _validate_and_copy_defaults( if "metrics_emitter" not in options: cast(dict, options)["metrics_emitter"] = client.metrics_emitter + if "context_propagators" not in options: + cast(dict, options)["context_propagators"] = client.context_propagators # TODO: More validation @@ -77,3 +80,7 @@ def _validate_and_copy_defaults( for key, value in _DEFAULT_WORKER_OPTIONS.items(): if key not in options: cast(dict, options)[key] = value + cast(dict, options)["context_propagators"] = tuple( + options.get("context_propagators") or () + ) + validate_propagators(options["context_propagators"]) diff --git a/cadence/workflow.py b/cadence/workflow.py index 87ca245..5dab26e 100644 --- a/cadence/workflow.py +++ b/cadence/workflow.py @@ -246,6 +246,7 @@ def continue_as_new( task_list=task_list, execution_start_to_close_timeout=execution_start_to_close_timeout, task_start_to_close_timeout=task_start_to_close_timeout, + headers=WorkflowContext.get().inject_propagated_headers(), ) @@ -652,6 +653,10 @@ def mutable_side_effect( @abstractmethod def is_cancel_requested(self) -> bool: ... + def inject_propagated_headers(self) -> dict[str, bytes]: + """Return headers to attach to outbound workflow decisions.""" + return {} + @contextmanager def _activate(self) -> Iterator["WorkflowContext"]: token = WorkflowContext._var.set(self) diff --git a/tests/cadence/_internal/activity/test_activity_executor.py b/tests/cadence/_internal/activity/test_activity_executor.py index 2ff6743..74805d7 100644 --- a/tests/cadence/_internal/activity/test_activity_executor.py +++ b/tests/cadence/_internal/activity/test_activity_executor.py @@ -11,12 +11,14 @@ from cadence._internal.activity import ActivityExecutor from cadence.activity import ActivityInfo, ActivityDefinition from cadence.api.v1.common_pb2 import ( + Header, WorkflowExecution, ActivityType, Payload, Failure, WorkflowType, ) +from cadence.context import ContextVarPropagator from cadence.api.v1.service_worker_pb2 import ( RespondActivityTaskCompletedResponse, PollForActivityTaskResponse, @@ -152,6 +154,144 @@ def activity_fn(): ) +@pytest.mark.parametrize("is_async", [True, False]) +async def test_activity_without_header_does_not_inherit_worker_context( + client, is_async +): + from contextvars import ContextVar + + worker_stub = client.worker_stub + worker_stub.RespondActivityTaskFailed = AsyncMock( + return_value=RespondActivityTaskFailedResponse() + ) + value: ContextVar[str] = ContextVar("activity-ambient-context") + propagator = ContextVarPropagator( + value, "context", lambda item: item.encode(), bytes.decode + ) + reg = Registry() + + if is_async: + + @reg.activity(name="activity_type") + async def activity_fn(): + value.get() + + else: + + @reg.activity(name="activity_type") + def activity_fn(): + value.get() + + executor = ActivityExecutor( + client, + "task_list", + "identity", + 1, + reg.get_activity, + context_propagators=(propagator,), + ) + value.set("worker-boot-value") + task = fake_task("activity_type", "") + + await executor.execute(task) + + worker_stub.RespondActivityTaskFailed.assert_called_once() + assert value.get() == "worker-boot-value" + + +@pytest.mark.parametrize("is_async", [True, False]) +async def test_activity_context_propagation_does_not_leak(client, is_async): + from contextvars import ContextVar + + worker_stub = client.worker_stub + worker_stub.RespondActivityTaskCompleted = AsyncMock( + return_value=RespondActivityTaskCompletedResponse() + ) + value: ContextVar[str] = ContextVar("activity-context") + propagator = ContextVarPropagator( + value, "context", lambda item: item.encode(), bytes.decode + ) + reg = Registry() + + if is_async: + + @reg.activity(name="activity_type") + async def activity_fn(): + return value.get() + + else: + + @reg.activity(name="activity_type") + def activity_fn(): + return value.get() + + executor = ActivityExecutor( + client, + "task_list", + "identity", + 1, + reg.get_activity, + context_propagators=(propagator,), + ) + task = fake_task("activity_type", "") + task.header.CopyFrom(Header(fields={"context": Payload(data=b"propagated")})) + + await executor.execute(task) + + assert ( + worker_stub.RespondActivityTaskCompleted.call_args.args[0].result.data + == b'"propagated"' + ) + with pytest.raises(LookupError): + value.get() + + +@pytest.mark.parametrize("is_async", [True, False]) +async def test_activity_context_propagation_restores_after_exception(client, is_async): + from contextvars import ContextVar + + worker_stub = client.worker_stub + worker_stub.RespondActivityTaskFailed = AsyncMock( + return_value=RespondActivityTaskFailedResponse() + ) + value: ContextVar[str] = ContextVar("activity-exception-context") + propagator = ContextVarPropagator( + value, "context", lambda item: item.encode(), bytes.decode + ) + reg = Registry() + + if is_async: + + @reg.activity(name="activity_type") + async def activity_fn(): + assert value.get() == "propagated" + raise RuntimeError("activity failed") + + else: + + @reg.activity(name="activity_type") + def activity_fn(): + assert value.get() == "propagated" + raise RuntimeError("activity failed") + + executor = ActivityExecutor( + client, + "task_list", + "identity", + 1, + reg.get_activity, + context_propagators=(propagator,), + ) + task = fake_task("activity_type", "") + task.header.CopyFrom(Header(fields={"context": Payload(data=b"propagated")})) + + await executor.execute(task) + + worker_stub.RespondActivityTaskFailed.assert_called_once() + with pytest.raises(LookupError): + value.get() + + async def test_activity_sync_failure(client): worker_stub = client.worker_stub worker_stub.RespondActivityTaskFailed = AsyncMock( diff --git a/tests/cadence/test_context.py b/tests/cadence/test_context.py new file mode 100644 index 0000000..1a951a3 --- /dev/null +++ b/tests/cadence/test_context.py @@ -0,0 +1,582 @@ +from __future__ import annotations + +from contextlib import contextmanager +from contextvars import ContextVar +from datetime import datetime, timedelta, timezone +from collections.abc import Iterator, Mapping +from unittest.mock import AsyncMock + +import pytest + +from cadence import workflow +from cadence._internal.context import ( + extract_headers, + header_from_dict, + header_to_dict, + inject_headers, +) +from cadence._internal.workflow.workflow_engine import WorkflowEngine +from cadence.api.v1.common_pb2 import ActivityType, Header, Payload +from cadence.api.v1.history_pb2 import ( + ActivityTaskCompletedEventAttributes, + ActivityTaskScheduledEventAttributes, + ActivityTaskStartedEventAttributes, + DecisionTaskCompletedEventAttributes, + DecisionTaskScheduledEventAttributes, + DecisionTaskStartedEventAttributes, + HistoryEvent, + WorkflowExecutionStartedEventAttributes, +) +from cadence.api.v1.service_workflow_pb2 import ( + SignalWithStartWorkflowExecutionResponse, + StartWorkflowExecutionResponse, +) +from cadence.client import Client, ClientOptions +from cadence.client import ( + _validate_and_copy_defaults as _validate_and_copy_client_defaults, +) +from cadence.context import ContextVarPropagator +from cadence.data_converter import DefaultDataConverter +from cadence.metrics import NoOpMetricsEmitter +from cadence.testing import TestWorkflowEnvironment +from cadence.worker._types import WorkerOptions +from cadence.worker._worker import _validate_and_copy_defaults +from cadence.worker import Registry +from cadence.workflow import WorkflowDefinition, WorkflowDefinitionOptions, WorkflowInfo + + +def _string_propagator( + var: ContextVar[str], header_key: str = "context" +) -> ContextVarPropagator[str]: + return ContextVarPropagator( + var, header_key, lambda value: value.encode(), bytes.decode + ) + + +def test_context_var_propagator_unset_restores_and_nests() -> None: + value: ContextVar[str] = ContextVar("value") + propagator = _string_propagator(value) + + assert propagator.inject() == {} + root = value.set("root") + try: + assert propagator.inject() == {"context": b"root"} + with propagator.extract({"context": b"outer"}): + assert value.get() == "outer" + with propagator.extract({"context": b"inner"}): + assert value.get() == "inner" + assert value.get() == "outer" + assert value.get() == "root" + with propagator.extract({}): + assert value.get() == "root" + finally: + value.reset(root) + + +def test_header_conversion_and_ordered_injection() -> None: + assert header_from_dict({}) is None + header = Header(fields={"binary": Payload(data=b"\x00\xff"), "empty": Payload()}) + assert header_to_dict(header) == {"binary": b"\x00\xff", "empty": b""} + + class Propagator: + def __init__(self, result: Mapping[str, bytes]) -> None: + self.result = result + + def inject(self) -> Mapping[str, bytes]: + return self.result + + @contextmanager + def extract(self, headers: Mapping[str, bytes]) -> Iterator[None]: + yield + + assert inject_headers( + (Propagator({"first": b"1", "same": b"old"}), Propagator({"same": b"new"})) + ) == {"first": b"1", "same": b"new"} + + +def test_extract_failure_unwinds_entered_propagators() -> None: + entered: list[str] = [] + + class Propagator: + def __init__(self, name: str, fail: bool = False) -> None: + self.name = name + self.fail = fail + + def inject(self) -> Mapping[str, bytes]: + return {} + + @contextmanager + def extract(self, headers: Mapping[str, bytes]) -> Iterator[None]: + assert headers == {"all": b"headers"} + entered.append(f"enter:{self.name}") + if self.fail: + raise RuntimeError("extract failed") + try: + yield + finally: + entered.append(f"exit:{self.name}") + + with pytest.raises(RuntimeError, match="extract failed"): + with extract_headers( + (Propagator("one"), Propagator("two", fail=True)), {"all": b"headers"} + ): + pass + assert entered == ["enter:one", "enter:two", "exit:one"] + + +@pytest.mark.asyncio +async def test_client_injects_start_and_signal_with_start_headers() -> None: + value: ContextVar[str] = ContextVar("client-context") + propagator = _string_propagator(value) + client = object.__new__(Client) + client._options = { + "domain": "domain", + "target": "target", + "data_converter": DefaultDataConverter(), + "identity": "identity", + "metrics_emitter": NoOpMetricsEmitter(), + "context_propagators": (propagator,), + } + client._workflow_stub = type( + "WorkflowStub", + (), + { + "StartWorkflowExecution": AsyncMock( + return_value=StartWorkflowExecutionResponse(run_id="start-run") + ), + "SignalWithStartWorkflowExecution": AsyncMock( + return_value=SignalWithStartWorkflowExecutionResponse( + run_id="signal-start-run" + ) + ), + }, + )() + + token = value.set("client") + try: + await client.start_workflow( + "workflow", + task_list="task-list", + execution_start_to_close_timeout=timedelta(seconds=1), + ) + await client.signal_with_start_workflow( + "workflow", + "signal", + [], + task_list="task-list", + execution_start_to_close_timeout=timedelta(seconds=1), + ) + finally: + value.reset(token) + + start_request = client.workflow_stub.StartWorkflowExecution.call_args.args[0] + signal_start_request = ( + client.workflow_stub.SignalWithStartWorkflowExecution.call_args.args[0] + ) + assert header_to_dict(start_request.header) == {"context": b"client"} + assert header_to_dict(signal_start_request.start_request.header) == { + "context": b"client" + } + + +def test_client_defaults_to_no_propagators_and_snapshots_mutable_sequence() -> None: + value: ContextVar[str] = ContextVar("client-default-context") + propagator = _string_propagator(value) + + defaulted = _validate_and_copy_client_defaults( + ClientOptions(domain="domain", target="target") + ) + assert defaulted["context_propagators"] == () + + mutable = [propagator] + snapshotted = _validate_and_copy_client_defaults( + ClientOptions(domain="domain", target="target", context_propagators=mutable) + ) + mutable.append(_string_propagator(value, "other")) + assert snapshotted["context_propagators"] == (propagator,) + + +def test_worker_context_propagators_inherit_and_allow_override() -> None: + value: ContextVar[str] = ContextVar("worker-context") + propagator = _string_propagator(value) + override_propagator = _string_propagator(value, "override") + client = object.__new__(Client) + client._options = { + "identity": "client", + "metrics_emitter": NoOpMetricsEmitter(), + "context_propagators": (propagator,), + } + + inherited = WorkerOptions() + _validate_and_copy_defaults(client, "task-list", inherited) + assert inherited["context_propagators"] == (propagator,) + + replaced = WorkerOptions(context_propagators=[override_propagator]) + _validate_and_copy_defaults(client, "task-list", replaced) + assert replaced["context_propagators"] == (override_propagator,) + + disabled = WorkerOptions(context_propagators=[]) + _validate_and_copy_defaults(client, "task-list", disabled) + assert disabled["context_propagators"] == () + + +def test_production_workflow_started_header_reaches_activity_decision() -> None: + value: ContextVar[str] = ContextVar("production-activity-context") + propagator = _string_propagator(value) + observed: list[str] = [] + + class Workflow: + @workflow.run + async def run(self) -> None: + observed.append(value.get()) + await workflow.execute_activity( + "activity", + str, + schedule_to_close_timeout=timedelta(seconds=1), + ) + + events = _started_events({"context": b"from-start"}) + result = _production_engine( + Workflow, + propagator, + header_to_dict(events[0].workflow_execution_started_event_attributes.header), + ).process_decision(events) + + assert observed == ["from-start"] + attrs = result.decisions[0].schedule_activity_task_decision_attributes + assert header_to_dict(attrs.header) == {"context": b"from-start"} + with pytest.raises(LookupError): + value.get() + + +def test_production_workflow_started_header_reaches_child_decision() -> None: + value: ContextVar[str] = ContextVar("production-child-context") + propagator = _string_propagator(value) + observed: list[str] = [] + + class Workflow: + @workflow.run + async def run(self) -> None: + observed.append(value.get()) + await workflow.execute_child_workflow( + "child", + str, + execution_start_to_close_timeout=timedelta(seconds=1), + ) + + events = _started_events({"context": b"from-start"}) + result = _production_engine( + Workflow, + propagator, + header_to_dict(events[0].workflow_execution_started_event_attributes.header), + ).process_decision(events) + + assert observed == ["from-start"] + attrs = result.decisions[0].start_child_workflow_execution_decision_attributes + assert header_to_dict(attrs.header) == {"context": b"from-start"} + with pytest.raises(LookupError): + value.get() + + +def test_production_workflow_continue_as_new_injects_header() -> None: + value: ContextVar[str] = ContextVar("production-continue-context") + propagator = _string_propagator(value) + + class Workflow: + @workflow.run + async def run(self) -> None: + workflow.continue_as_new() + + events = _started_events({"context": b"from-start"}) + result = _production_engine( + Workflow, + propagator, + header_to_dict(events[0].workflow_execution_started_event_attributes.header), + ).process_decision(events) + + attrs = result.decisions[0].continue_as_new_workflow_execution_decision_attributes + assert header_to_dict(attrs.header) == {"context": b"from-start"} + with pytest.raises(LookupError): + value.get() + + +def test_production_workflow_continue_as_new_uses_workflow_context() -> None: + value: ContextVar[str] = ContextVar("production-continue-mutation") + propagator = _string_propagator(value) + + class Workflow: + @workflow.run + async def run(self) -> None: + value.set("updated-in-workflow") + workflow.continue_as_new() + + events = _started_events({"context": b"from-start"}) + result = _production_engine( + Workflow, + propagator, + {"context": b"from-start"}, + ).process_decision(events) + + attrs = result.decisions[0].continue_as_new_workflow_execution_decision_attributes + assert header_to_dict(attrs.header) == {"context": b"updated-in-workflow"} + + +def test_inject_failure_is_propagated() -> None: + class GoodPropagator: + def inject(self) -> Mapping[str, bytes]: + return {"good": b"1"} + + @contextmanager + def extract(self, headers: Mapping[str, bytes]) -> Iterator[None]: + yield + + class BadPropagator: + def inject(self) -> Mapping[str, bytes]: + raise RuntimeError("inject failed") + + @contextmanager + def extract(self, headers: Mapping[str, bytes]) -> Iterator[None]: + yield + + with pytest.raises(RuntimeError, match="inject failed"): + inject_headers((GoodPropagator(), BadPropagator())) + + +def test_replay_restores_context_each_decision_despite_header_drift() -> None: + value: ContextVar[str] = ContextVar("replay-context") + propagator = _string_propagator(value) + observed: list[str] = [] + + class Workflow: + @workflow.run + async def run(self) -> str: + observed.append(value.get()) + result = await workflow.execute_activity( + "activity", + str, + schedule_to_close_timeout=timedelta(seconds=1), + ) + observed.append(value.get()) + return result + + engine = _production_engine(Workflow, propagator, {"context": b"from-start"}) + + first = engine.process_decision(_started_events({"context": b"from-start"})) + scheduled = first.decisions[0].schedule_activity_task_decision_attributes + assert header_to_dict(scheduled.header) == {"context": b"from-start"} + + # The recorded schedule carries a different header than this worker injects now, + # which must not be treated as a nondeterministic replay. + second = engine.process_decision( + _activity_completion_events( + activity_id=scheduled.activity_id, + recorded_headers={"context": b"recorded-elsewhere"}, + result="activity-result", + ) + ) + + assert observed == ["from-start", "from-start"] + completed = second.decisions[0].complete_workflow_execution_decision_attributes + assert DefaultDataConverter().from_data(completed.result, [str]) == [ + "activity-result" + ] + with pytest.raises(LookupError): + value.get() + + +@pytest.mark.asyncio +async def test_test_environment_propagates_client_workflow_activity_and_child() -> None: + value: ContextVar[str] = ContextVar("test-environment-context") + propagator = _string_propagator(value) + registry = Registry() + + @registry.activity(name="read-context") + async def read_context() -> str: + return value.get() + + @registry.workflow + class Child: + @workflow.run + async def run(self) -> str: + return value.get() + + @registry.workflow + class Parent: + @workflow.run + async def run(self) -> str: + activity_value = await workflow.execute_activity( + "read-context", + str, + schedule_to_close_timeout=timedelta(seconds=1), + ) + child_value = await workflow.execute_child_workflow( + "Child", + str, + execution_start_to_close_timeout=timedelta(seconds=1), + ) + return f"{value.get()}:{activity_value}:{child_value}" + + with TestWorkflowEnvironment( + registry, context_propagators=(propagator,) + ) as environment: + token = value.set("client") + try: + await environment.client.start_workflow("Parent", task_list="test") + finally: + value.reset(token) + + assert environment.get_workflow_result(str) == "client:client:client" + with pytest.raises(LookupError): + value.get() + + +@pytest.mark.asyncio +async def test_test_environment_query_sees_propagated_context() -> None: + value: ContextVar[str] = ContextVar("test-query-context") + propagator = _string_propagator(value) + registry = Registry() + + @registry.workflow + class WaitingWorkflow: + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: False) + + @workflow.query(name="context") + def context(self) -> str: + return value.get() + + with TestWorkflowEnvironment( + registry, context_propagators=(propagator,) + ) as environment: + token = value.set("from-start") + try: + execution = await environment.client.start_workflow( + "WaitingWorkflow", task_list="test" + ) + finally: + value.reset(token) + + assert ( + await environment.client.query_workflow( + execution.workflow_id, "", "context", result_type=str + ) + == "from-start" + ) + + +@pytest.mark.asyncio +async def test_test_environment_activity_does_not_inherit_workflow_locals() -> None: + value: ContextVar[str] = ContextVar("local-only") + registry = Registry() + + @registry.activity(name="read-local") + async def read_local() -> str: + return value.get() + + @registry.workflow + class Parent: + @workflow.run + async def run(self) -> str: + value.set("workflow-local") + return await workflow.execute_activity( + "read-local", + str, + schedule_to_close_timeout=timedelta(seconds=1), + ) + + with TestWorkflowEnvironment(registry) as environment: + await environment.client.start_workflow("Parent", task_list="test") + with pytest.raises(LookupError): + environment.get_workflow_result(str) + + +def _started_events(headers: Mapping[str, bytes]) -> list[HistoryEvent]: + started = WorkflowExecutionStartedEventAttributes(header=header_from_dict(headers)) + events = [ + HistoryEvent(workflow_execution_started_event_attributes=started), + HistoryEvent( + decision_task_scheduled_event_attributes=DecisionTaskScheduledEventAttributes() + ), + HistoryEvent( + decision_task_started_event_attributes=DecisionTaskStartedEventAttributes( + scheduled_event_id=2 + ) + ), + ] + for event_id, event in enumerate(events, start=1): + event.event_id = event_id + event.event_time.FromDatetime(datetime.fromtimestamp(event_id, tz=timezone.utc)) + return events + + +def _activity_completion_events( + activity_id: str, recorded_headers: Mapping[str, bytes], result: str +) -> list[HistoryEvent]: + """Replay the first decision's output, then deliver the activity result.""" + events = [ + HistoryEvent( + event_id=4, + decision_task_completed_event_attributes=DecisionTaskCompletedEventAttributes( + scheduled_event_id=2, started_event_id=3 + ), + ), + HistoryEvent( + event_id=5, + activity_task_scheduled_event_attributes=ActivityTaskScheduledEventAttributes( + activity_id=activity_id, + activity_type=ActivityType(name="activity"), + header=header_from_dict(recorded_headers), + ), + ), + HistoryEvent( + event_id=6, + activity_task_started_event_attributes=ActivityTaskStartedEventAttributes( + scheduled_event_id=5 + ), + ), + HistoryEvent( + event_id=7, + activity_task_completed_event_attributes=ActivityTaskCompletedEventAttributes( + scheduled_event_id=5, + result=DefaultDataConverter().to_data([result]), + ), + ), + HistoryEvent( + event_id=8, + decision_task_scheduled_event_attributes=DecisionTaskScheduledEventAttributes(), + ), + HistoryEvent( + event_id=9, + decision_task_started_event_attributes=DecisionTaskStartedEventAttributes( + scheduled_event_id=8 + ), + ), + ] + for event in events: + event.event_time.FromDatetime( + datetime.fromtimestamp(event.event_id, tz=timezone.utc) + ) + return events + + +def _production_engine( + workflow_class: type, + propagator: ContextVarPropagator[str], + headers: Mapping[str, bytes], +) -> WorkflowEngine: + return WorkflowEngine( + info=WorkflowInfo( + workflow_type="workflow", + workflow_domain="domain", + workflow_id="workflow-id", + workflow_run_id="run-id", + workflow_task_list="task-list", + data_converter=DefaultDataConverter(), + ), + workflow_definition=WorkflowDefinition.wrap( + workflow_class, WorkflowDefinitionOptions(name="workflow") + ), + context_propagators=(propagator,), + headers=headers, + ) diff --git a/tests/cadence/worker/test_metrics_propagation.py b/tests/cadence/worker/test_metrics_propagation.py index 702091a..dd9cc79 100644 --- a/tests/cadence/worker/test_metrics_propagation.py +++ b/tests/cadence/worker/test_metrics_propagation.py @@ -17,6 +17,7 @@ def _mock_client(metrics_emitter=None): type(client).domain = PropertyMock(return_value="test-domain") type(client).identity = PropertyMock(return_value="test-identity") client.metrics_emitter = metrics_emitter or NoOpMetricsEmitter() + client.context_propagators = () worker_stub = Mock() worker_stub.PollForDecisionTask = AsyncMock(return_value=Mock(task_token=b"")) worker_stub.PollForActivityTask = AsyncMock(return_value=Mock(task_token=b"")) diff --git a/tests/cadence/worker/test_worker.py b/tests/cadence/worker/test_worker.py index 267217c..9733004 100644 --- a/tests/cadence/worker/test_worker.py +++ b/tests/cadence/worker/test_worker.py @@ -31,6 +31,7 @@ async def poll(_, timeout=0.0): client.worker_stub = worker_stub type(client).domain = PropertyMock(return_value="domain") type(client).identity = PropertyMock(return_value="identity") + type(client).context_propagators = PropertyMock(return_value=()) async with Worker( client, diff --git a/tests/integration_tests/workflow/test_context_propagation.py b/tests/integration_tests/workflow/test_context_propagation.py new file mode 100644 index 0000000..547ec5e --- /dev/null +++ b/tests/integration_tests/workflow/test_context_propagation.py @@ -0,0 +1,110 @@ +from contextvars import ContextVar +from datetime import timedelta + +from typing import cast + +import pytest + +from cadence import ContextVarPropagator, Registry, workflow +from cadence.api.v1.common_pb2 import WorkflowExecution +from cadence.api.v1.history_pb2 import EventFilterType, HistoryEvent +from cadence.api.v1.service_workflow_pb2 import ( + GetWorkflowExecutionHistoryRequest, + GetWorkflowExecutionHistoryResponse, +) +from cadence.worker import Worker +from tests.integration_tests.helper import CadenceHelper, DOMAIN_NAME + +REQUEST_CONTEXT: ContextVar[str] = ContextVar("integration_request_context") +REQUEST_CONTEXT_PROPAGATOR = ContextVarPropagator( + REQUEST_CONTEXT, + "request-context", + lambda value: value.encode(), + bytes.decode, +) +registry = Registry() + + +@registry.activity() +async def read_request_context() -> str: + return REQUEST_CONTEXT.get() + + +@registry.workflow() +class ContextPropagationWorkflow: + """Reads context directly and via an activity, then continues as new once.""" + + @workflow.run + async def run(self, remaining_runs: int) -> str: + activity_value = await read_request_context.with_options( + schedule_to_close_timeout=timedelta(seconds=10) + ).execute() + if remaining_runs > 0: + workflow.continue_as_new(remaining_runs - 1) + return f"{REQUEST_CONTEXT.get()}/{activity_value}" + + +async def _await_close_event( + worker: Worker, execution: WorkflowExecution +) -> HistoryEvent: + """Block until the given run closes and return its final history event.""" + response: GetWorkflowExecutionHistoryResponse = ( + await worker.client.workflow_stub.GetWorkflowExecutionHistory( + GetWorkflowExecutionHistoryRequest( + domain=DOMAIN_NAME, + workflow_execution=execution, + wait_for_new_event=True, + history_event_filter_type=EventFilterType.EVENT_FILTER_TYPE_CLOSE_EVENT, + skip_archival=True, + ) + ) + ) + return cast(HistoryEvent, response.history.events[-1]) + + +async def test_context_propagates_through_activity_and_continue_as_new( + helper: CadenceHelper, +) -> None: + propagation_helper = CadenceHelper( + {**helper.options, "context_propagators": (REQUEST_CONTEXT_PROPAGATOR,)}, + helper.test_name, + helper.fspath, + ) + async with propagation_helper.worker(registry) as worker: + token = REQUEST_CONTEXT.set("integration-value") + try: + execution = await worker.client.start_workflow( + "ContextPropagationWorkflow", + 1, + task_list=worker.task_list, + execution_start_to_close_timeout=timedelta(seconds=30), + ) + finally: + REQUEST_CONTEXT.reset(token) + + # First run: ends by continuing as new, carrying the context on the new run. + first_run_close = await _await_close_event(worker, execution) + continued = first_run_close.workflow_execution_continued_as_new_event_attributes + assert continued.new_execution_run_id, ( + f"expected first run to continue as new, got {first_run_close}" + ) + assert continued.header.fields["request-context"].data == b"integration-value" + + # Second run: sees the context only via the continue-as-new header, and + # passes it on to its own activity. + second_run_close = await _await_close_event( + worker, + WorkflowExecution( + workflow_id=execution.workflow_id, + run_id=continued.new_execution_run_id, + ), + ) + assert ( + second_run_close.workflow_execution_completed_event_attributes.result.data + == b'"integration-value/integration-value"' + ) + + # The worker ran activities and decisions in this process; none of that may leak + # context back into the caller's scope. + with pytest.raises(LookupError): + REQUEST_CONTEXT.get()