Skip to content
Merged
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
3 changes: 3 additions & 0 deletions cadence/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,16 @@

# Import main client functionality
from .client import Client
from .context import ContextPropagator, ContextVarPropagator
from .worker import Registry
from . import workflow

__version__ = "0.1.0"

__all__ = [
"Client",
"ContextPropagator",
"ContextVarPropagator",
"Registry",
"workflow",
]
17 changes: 15 additions & 2 deletions cadence/_internal/activity/_activity_executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -19,6 +20,7 @@
RespondActivityTaskCompletedRequest,
)
from cadence.client import Client
from cadence.context import ContextPropagator
from cadence.metrics import (
duration_between,
duration_from_nanoseconds,
Expand Down Expand Up @@ -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
Expand All @@ -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-"
)
Expand Down Expand Up @@ -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(
Expand Down
31 changes: 26 additions & 5 deletions cadence/_internal/activity/_context.py
Original file line number Diff line number Diff line change
@@ -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):
Expand All @@ -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
Expand All @@ -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:
Expand All @@ -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()
Expand Down Expand Up @@ -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
Expand All @@ -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")
Expand Down
72 changes: 72 additions & 0 deletions cadence/_internal/context.py
Original file line number Diff line number Diff line change
@@ -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
Comment thread
gitar-bot[bot] marked this conversation as resolved.


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
10 changes: 9 additions & 1 deletion cadence/_internal/workflow/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading