From 02d494d3b828eb1ac1aed871bd38528cb162d200 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:40:44 -0700 Subject: [PATCH 1/3] Added Execution and ExecutionType to temporalio.common. A linked channel names its owner as an execution, a workflow or a standalone activity, so the client needs a typed value for it. --- temporalio/common.py | 58 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 58 insertions(+) diff --git a/temporalio/common.py b/temporalio/common.py index 04081bf19..cf126ff95 100644 --- a/temporalio/common.py +++ b/temporalio/common.py @@ -18,6 +18,7 @@ Generic, TypeAlias, TypeVar, + cast, get_origin, get_type_hints, overload, @@ -108,6 +109,63 @@ def _validate(self) -> None: raise ValueError("Maximum attempts cannot be negative") +class ExecutionType(IntEnum): + """What kind of execution an :class:`Execution` names. + + .. warning:: + This API is experimental and unstable. + """ + + UNSPECIFIED = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_UNSPECIFIED) + WORKFLOW = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_WORKFLOW) + ACTIVITY = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_ACTIVITY) + """A standalone activity, one started by a client rather than a workflow.""" + + +@dataclass(frozen=True) +class Execution: + """One execution in a namespace: a workflow or a standalone activity. + + ``business_id`` is the id the caller chose, the workflow id or the + activity id. ``run_id`` pins one run of it. Unset, a call reaches the + current run of a workflow chain, as a Signal does. + + .. warning:: + This API is experimental and unstable. + """ + + type: ExecutionType + business_id: str + run_id: str | None = None + + @classmethod + def workflow(cls, workflow_id: str, run_id: str | None = None) -> Execution: + """A workflow execution.""" + return cls(ExecutionType.WORKFLOW, workflow_id, run_id) + + @classmethod + def activity(cls, activity_id: str, run_id: str | None = None) -> Execution: + """A standalone activity execution.""" + return cls(ExecutionType.ACTIVITY, activity_id, run_id) + + def to_proto(self) -> temporalio.api.common.v1.Execution: + """This execution as the API names it.""" + return temporalio.api.common.v1.Execution( + type=cast( + "temporalio.api.enums.v1.ExecutionType.ValueType", int(self.type) + ), + business_id=self.business_id, + run_id=self.run_id or "", + ) + + @staticmethod + def from_proto(proto: temporalio.api.common.v1.Execution) -> Execution: + """From the API's form. An empty run id reads as unset.""" + return Execution( + ExecutionType(proto.type), proto.business_id, proto.run_id or None + ) + + class WorkflowIDReusePolicy(IntEnum): """How already-in-use workflow IDs are handled on start. From f8bf154b17b22ecc4355866ec4bf6135357b402d Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:40:44 -0700 Subject: [PATCH 2/3] Added the notification channel calls to the client. The five calls reach an independent channel or one linked to an execution. Polls return workflow.Notification, so the workflow side shares the type later. --- CHANGELOG.md | 5 + temporalio/client/__init__.py | 20 ++ temporalio/client/_channel.py | 161 ++++++++++++++++ temporalio/client/_client.py | 295 ++++++++++++++++++++++++++++++ temporalio/client/_impl.py | 146 +++++++++++++++ temporalio/client/_interceptor.py | 130 +++++++++++++ temporalio/workflow/__init__.py | 2 + temporalio/workflow/_channels.py | 81 ++++++++ 8 files changed, 840 insertions(+) create mode 100644 temporalio/client/_channel.py create mode 100644 temporalio/workflow/_channels.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 03a2243ed..005a5eb16 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -37,6 +37,11 @@ to include examples, links to docs, or any other relevant information. worker-side factories registered with `StrandsPlugin(sandboxes=...)`. - Added the `temporalio.contrib.gcp.cloud_run.id` module with the `CloudRunIdPlugin` client plugin to set the worker identity on Cloud Run. +- **Experimental**: notification channels. Requires a server that serves notification channels. + - `Client.notify_channel`, `poll_channel`, `describe_channel`, `register_channel_listener` and + `unregister_channel_listener` reach a named channel on the server. The channel is either + independent or linked to an execution, which `execution=` (`temporalio.common.Execution`) or + the `workflow_id=` shorthand names. ### Changed diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 2eef41a39..87e70537c 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -64,6 +64,12 @@ from ._callback import ( Callback, ) +from ._channel import ( + ChannelAddress, + ChannelDescription, + ChannelKind, + ChannelListener, +) from ._client import ( Client, ClientConfig, @@ -106,6 +112,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -120,10 +127,13 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseActivityInput, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -137,6 +147,7 @@ TriggerScheduleInput, UnpauseActivityInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateActivityOptionsInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, @@ -315,6 +326,11 @@ "TerminateNexusOperationInput", "ListNexusOperationsInput", "CountNexusOperationsInput", + "DescribeChannelInput", + "NotifyChannelInput", + "PollChannelInput", + "RegisterChannelListenerInput", + "UnregisterChannelListenerInput", "StartWorkflowUpdateInput", "UpdateWithStartUpdateWorkflowInput", "UpdateWithStartStartWorkflowInput", @@ -351,6 +367,10 @@ "CloudOperationsClient", "Plugin", "Callback", + "ChannelDescription", + "ChannelKind", + "ChannelListener", + "ChannelAddress", "_ClientImpl", "_apply_headers", "_decode_user_metadata", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py new file mode 100644 index 000000000..91ef53e05 --- /dev/null +++ b/temporalio/client/_channel.py @@ -0,0 +1,161 @@ +"""Notification channel descriptions as the client reports them.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +from enum import IntEnum + +import temporalio.api.notification.v1 +import temporalio.common +from temporalio.workflow import Notification + +from ._callback import Callback + +__all__ = [ + "ChannelAddress", + "ChannelDescription", + "ChannelKind", + "ChannelListener", +] + + +class ChannelKind(IntEnum): + """Where a channel lives, which decides how a call addresses it. + + .. warning:: + This API is experimental and unstable. + """ + + UNSPECIFIED = int( + temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_UNSPECIFIED + ) + """The server did not say; an older server answers this.""" + + INDEPENDENT = int( + temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_INDEPENDENT + ) + """Its own execution, keyed by namespace and channel name. + + Any number of workflows subscribe to it and callbacks register on it. + """ + + LINKED = int(temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED) + """Kept in one execution's state, keyed by namespace, execution and name. + + The owning execution, a workflow or a standalone activity, is its listener + by construction. A call reaches it with the ``execution`` argument, or + with ``workflow_id`` when the owner is a workflow. + """ + + +@dataclass(frozen=True) +class ChannelListener: + """One listener of a channel: a workflow or a callback. + + .. warning:: + This API is experimental and unstable. + """ + + listener_id: str + """Assigned by the server when the listener registered.""" + + workflow_id: str | None + """The subscribed workflow, when the listener is one.""" + + run_id: str | None + """The run that subscribed. Delivery follows the chain's current run.""" + + callback: Callback | None + """The callback the server invokes, when the listener is one.""" + + registered_time: datetime | None + """When the listener registered.""" + + @staticmethod + def _from_proto( + proto: temporalio.api.notification.v1.ChannelListener, + ) -> ChannelListener: + callback: Callback | None = None + if proto.HasField("callback") and proto.callback.HasField("nexus"): + callback = Callback( + url=proto.callback.nexus.url, headers=dict(proto.callback.nexus.header) + ) + workflow = proto.workflow if proto.HasField("workflow") else None + return ChannelListener( + listener_id=proto.listener_id, + workflow_id=workflow.workflow_id if workflow else None, + run_id=workflow.run_id if workflow else None, + callback=callback, + registered_time=( + proto.registered_time.ToDatetime(tzinfo=timezone.utc) + if proto.HasField("registered_time") + else None + ), + ) + + +@dataclass(frozen=True) +class ChannelDescription: + """What the server knows about a channel. + + .. warning:: + This API is experimental and unstable. + """ + + listeners: Sequence[ChannelListener] + """Who is listening, workflows and callbacks alike.""" + + latest: Notification | None + """The notification with the highest counter the channel retains.""" + + retained_count: int + """How many notifications the channel keeps for pollers.""" + + kind: ChannelKind = ChannelKind.UNSPECIFIED + """Which kind of channel this is. + + A linked channel of a running execution exists by construction, so a + describe with ``execution`` answers :attr:`ChannelKind.LINKED` with no + listeners and nothing retained for a name nobody has notified yet. + """ + + linked_to: temporalio.common.Execution | None = None + """The owner of a linked channel and the run that holds it. + + ``None`` for an independent channel. + """ + + +@dataclass(frozen=True) +class ChannelAddress: + """Where a channel call reaches a channel: its name and, when linked, its owner. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + """The channel name.""" + + execution: temporalio.common.Execution | None + """The execution the channel is linked to, or ``None`` for an independent one. + + Pass both to :py:meth:`temporalio.client.Client.poll_channel` and the + other channel calls as ``channel`` and ``execution``. + """ + + @property + def workflow_id(self) -> str | None: + """The owning workflow's id, when the owner is a workflow. + + ``None`` for an independent channel and for one a standalone activity + owns, which only ``execution`` reaches. + """ + if ( + self.execution is not None + and self.execution.type == temporalio.common.ExecutionType.WORKFLOW + ): + return self.execution.business_id + return None diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index b5db0c4b7..a0ce23c5a 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -62,22 +62,29 @@ AsyncActivityHandle, AsyncActivityIDReference, ) +from ._callback import Callback +from ._channel import ChannelDescription from ._impl import _ClientImpl from ._interceptor import ( CountActivitiesInput, CountNexusOperationsInput, CountWorkflowsInput, CreateScheduleInput, + DescribeChannelInput, GetWorkerBuildIdCompatibilityInput, GetWorkerTaskReachabilityInput, ListActivitiesInput, ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, + PollChannelInput, + RegisterChannelListenerInput, StartActivityInput, StartWorkflowInput, StartWorkflowUpdateWithStartInput, + UnregisterChannelListenerInput, UpdateWithStartUpdateWorkflowInput, UpdateWorkerBuildIdCompatibilityInput, ) @@ -113,6 +120,9 @@ from ._interceptor import Interceptor from ._plugin import Plugin +DEFAULT_CHANNEL_POLL_WAIT = timedelta(seconds=30) +"""How long :py:meth:`Client.poll_channel` waits for a notification by default.""" + class Client: """Client for accessing Temporal. @@ -2842,6 +2852,270 @@ async def get_worker_task_reachability( ) ) + async def notify_channel( + self, + channel: str, + *, + position: bytes = b"", + counter: int = 0, + metadata: Mapping[str, Any] | None = None, + execution: temporalio.common.Execution | None = None, + workflow_id: str | None = None, + run_id: str | None = None, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> int: + """Notify the listeners of ``channel`` that a source they consume has moved. + + A workflow listening with :py:func:`temporalio.workflow.subscribe_channel` + or :py:func:`temporalio.workflow.linked_channel` runs a Workflow Task + that carries the notification. The server folds notifications per + listener while one is pending, keeping the one with the highest + ``counter``, so a burst of writes costs a listener one task. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel. Scoped to the namespace, or to the + execution when one is given. + position: Where the source stands after the write, in the writer's + own terms. Opaque to the server. + counter: Orders notifications from this channel's writers. Derive it + from ``position``, since only the source can order its positions. + metadata: Details for the listener, such as which topic moved. Each + value is encoded with the client's data converter. + execution: Address the channel linked to this execution, a workflow + or a standalone activity, instead of the independent channel of + that name. Without a run id the call reaches the current run of + a workflow chain, as a Signal does. + workflow_id: Shorthand for ``execution`` naming a workflow. Not + with ``execution``. + run_id: With ``workflow_id``, the run of its chain to address. + rpc_metadata: Headers used on the RPC call. Keys here override + client-level RPC metadata keys. + rpc_timeout: Optional RPC deadline to set for the RPC call. + + Returns: + How many listeners the channel had when the notification arrived. + + Raises: + ValueError: Both ``execution`` and ``workflow_id`` were given, or + ``run_id`` without ``workflow_id``. + """ + return await self._impl.notify_channel( + NotifyChannelInput( + channel=channel, + position=position, + counter=counter, + metadata=metadata, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + execution=_channel_execution(execution, workflow_id, run_id), + ) + ) + + async def poll_channel( + self, + channel: str, + *, + after_counter: int = 0, + wait: bool | timedelta = True, + max_notifications: int = 100, + execution: temporalio.common.Execution | None = None, + workflow_id: str | None = None, + run_id: str | None = None, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> list[temporalio.workflow.Notification]: + """Read the notifications ``channel`` retains above a counter. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel. Scoped to the namespace, or to the + execution when one is given. + after_counter: Only notifications with a counter above this one are + returned. Pass the highest counter seen so far to page. + wait: How long the server holds the call when nothing is retained + above ``after_counter``. ``True`` waits up to + :py:data:`DEFAULT_CHANNEL_POLL_WAIT`, ``False`` returns at once. + max_notifications: Upper bound on the notifications returned. + execution: Address the channel linked to this execution, a workflow + or a standalone activity, instead of the independent channel of + that name. Without a run id the call reaches the current run of + a workflow chain. + workflow_id: Shorthand for ``execution`` naming a workflow. Not + with ``execution``. + run_id: With ``workflow_id``, the run of its chain to address. + rpc_metadata: Headers used on the RPC call. Keys here override + client-level RPC metadata keys. + rpc_timeout: Optional RPC deadline to set for the RPC call. + + Returns: + The notifications, oldest first. Empty when the wait ran out. + + Raises: + ValueError: Both ``execution`` and ``workflow_id`` were given, or + ``run_id`` without ``workflow_id``. + """ + if wait is True: + wait_for: timedelta | None = DEFAULT_CHANNEL_POLL_WAIT + elif wait is False: + wait_for = None + else: + wait_for = wait + return await self._impl.poll_channel( + PollChannelInput( + channel=channel, + after_counter=after_counter, + wait=wait_for, + max_notifications=max_notifications, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + execution=_channel_execution(execution, workflow_id, run_id), + ) + ) + + async def describe_channel( + self, + channel: str, + *, + execution: temporalio.common.Execution | None = None, + workflow_id: str | None = None, + run_id: str | None = None, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> ChannelDescription: + """Describe ``channel``: its kind, its listeners and what it retains. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel. Scoped to the namespace, or to the + execution when one is given. + execution: Describe the channel linked to this execution, a + workflow or a standalone activity, instead of the independent + channel of that name. A linked channel of a running execution + exists by construction, so the answer for a name nobody has + notified yet is a linked channel with no listeners and nothing + retained, not a not-found error. + workflow_id: Shorthand for ``execution`` naming a workflow. Not + with ``execution``. + run_id: With ``workflow_id``, the run of its chain to address. + rpc_metadata: Headers used on the RPC call. Keys here override + client-level RPC metadata keys. + rpc_timeout: Optional RPC deadline to set for the RPC call. + + Raises: + ValueError: Both ``execution`` and ``workflow_id`` were given, or + ``run_id`` without ``workflow_id``. + """ + return await self._impl.describe_channel( + DescribeChannelInput( + channel=channel, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + execution=_channel_execution(execution, workflow_id, run_id), + ) + ) + + async def register_channel_listener( + self, + channel: str, + callback: Callback, + *, + execution: temporalio.common.Execution | None = None, + workflow_id: str | None = None, + run_id: str | None = None, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> str: + """Register ``callback`` as a listener of ``channel``. + + The server invokes the callback with each notification on the channel + until :py:meth:`unregister_channel_listener` removes it. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel. Scoped to the namespace, or to the + execution when one is given. + callback: The callback to invoke. + execution: Listen on the channel linked to this execution, a + workflow or a standalone activity, instead of the independent + channel of that name. The listener lives in that execution's + state and ends with its run. + workflow_id: Shorthand for ``execution`` naming a workflow. Not + with ``execution``. + run_id: With ``workflow_id``, the run of its chain to address. + rpc_metadata: Headers used on the RPC call. Keys here override + client-level RPC metadata keys. + rpc_timeout: Optional RPC deadline to set for the RPC call. + + Returns: + The listener id the server assigned. + + Raises: + ValueError: Both ``execution`` and ``workflow_id`` were given, or + ``run_id`` without ``workflow_id``. + """ + return await self._impl.register_channel_listener( + RegisterChannelListenerInput( + channel=channel, + callback=callback, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + execution=_channel_execution(execution, workflow_id, run_id), + ) + ) + + async def unregister_channel_listener( + self, + channel: str, + listener_id: str, + *, + execution: temporalio.common.Execution | None = None, + workflow_id: str | None = None, + run_id: str | None = None, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> None: + """Remove a listener from ``channel``. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel. Scoped to the namespace, or to the + execution when one is given. + listener_id: The id :py:meth:`register_channel_listener` returned. + execution: The execution whose linked channel the listener is on, + when it was registered with one. + workflow_id: Shorthand for ``execution`` naming a workflow. Not + with ``execution``. + run_id: With ``workflow_id``, the run of its chain to address. + rpc_metadata: Headers used on the RPC call. Keys here override + client-level RPC metadata keys. + rpc_timeout: Optional RPC deadline to set for the RPC call. + + Raises: + ValueError: Both ``execution`` and ``workflow_id`` were given, or + ``run_id`` without ``workflow_id``. + """ + await self._impl.unregister_channel_listener( + UnregisterChannelListenerInput( + channel=channel, + listener_id=listener_id, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + execution=_channel_execution(execution, workflow_id, run_id), + ) + ) + def create_nexus_client( self, service: type[NexusServiceType] | str, @@ -3032,3 +3306,24 @@ class ClientConfig(TypedDict, total=False): temporalio.common.QueryRejectCondition | None ] header_codec_behavior: Required[HeaderCodecBehavior] + + +def _channel_execution( + execution: temporalio.common.Execution | None, + workflow_id: str | None, + run_id: str | None, +) -> temporalio.common.Execution | None: + """The owner a channel call names, by ``execution`` or by the workflow id shorthand. + + The two spellings do not combine, and a run id needs its workflow id. + ``None`` names the independent channel. + """ + if execution is not None: + if workflow_id is not None or run_id is not None: + raise ValueError("pass execution or workflow_id, not both") + return execution + if workflow_id is None: + if run_id is not None: + raise ValueError("run_id needs workflow_id") + return None + return temporalio.common.Execution.workflow(workflow_id, run_id) diff --git a/temporalio/client/_impl.py b/temporalio/client/_impl.py index abe3f8c09..4be55aa90 100644 --- a/temporalio/client/_impl.py +++ b/temporalio/client/_impl.py @@ -23,6 +23,7 @@ import temporalio.api.enums.v1 import temporalio.api.errordetails.v1 import temporalio.api.failure.v1 +import temporalio.api.notification.v1 import temporalio.api.schedule.v1 import temporalio.api.taskqueue.v1 import temporalio.api.update.v1 @@ -32,6 +33,7 @@ import temporalio.exceptions import temporalio.nexus import temporalio.nexus._operation_context +import temporalio.workflow from temporalio.activity import ActivityCancellationDetails from temporalio.converter import ( ActivitySerializationContext, @@ -54,6 +56,7 @@ ActivityHandle, AsyncActivityIDReference, ) +from ._channel import ChannelDescription, ChannelKind, ChannelListener from ._exceptions import ( AsyncActivityCancelledError, ScheduleAlreadyRunningError, @@ -74,6 +77,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -87,10 +91,13 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseActivityInput, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -104,6 +111,7 @@ TriggerScheduleInput, UnpauseActivityInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateActivityOptionsInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, @@ -1764,6 +1772,137 @@ async def count_nexus_operations( ) ) + ### Notification channel calls + + async def notify_channel(self, input: NotifyChannelInput) -> int: + notification = temporalio.api.notification.v1.Notification( + channel=input.channel, position=input.position, counter=input.counter + ) + for key, value in (input.metadata or {}).items(): + [payload] = await self._client.data_converter.encode([value]) + notification.metadata[key].CopyFrom(payload) + resp = await self._client.workflow_service.notify_channel( + temporalio.api.workflowservice.v1.NotifyChannelRequest( + namespace=self._client.namespace, + notification=notification, + identity=self._client.identity, + request_id=str(uuid.uuid4()), + execution=_channel_owner(input.execution), + ), + retry=True, + metadata=input.rpc_metadata, + timeout=input.rpc_timeout, + ) + return resp.listener_count + + async def poll_channel( + self, input: PollChannelInput + ) -> list[temporalio.workflow.Notification]: + req = temporalio.api.workflowservice.v1.PollChannelRequest( + namespace=self._client.namespace, + channel=input.channel, + after_counter=input.after_counter, + max_notifications=input.max_notifications, + execution=_channel_owner(input.execution), + ) + if input.wait is not None: + req.wait.FromTimedelta(input.wait) + resp = await self._client.workflow_service.poll_channel( + req, retry=True, metadata=input.rpc_metadata, timeout=input.rpc_timeout + ) + return [await self._notification_from_proto(n) for n in resp.notifications] + + async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescription: + resp = await self._client.workflow_service.describe_channel( + temporalio.api.workflowservice.v1.DescribeChannelRequest( + namespace=self._client.namespace, + channel=input.channel, + execution=_channel_owner(input.execution), + ), + retry=True, + metadata=input.rpc_metadata, + timeout=input.rpc_timeout, + ) + return ChannelDescription( + listeners=[ + ChannelListener._from_proto(listener) for listener in resp.listeners + ], + latest=( + await self._notification_from_proto(resp.latest) + if resp.HasField("latest") + else None + ), + retained_count=resp.retained_count, + kind=ChannelKind(resp.kind), + linked_to=( + temporalio.common.Execution.from_proto(resp.linked_to) + if resp.HasField("linked_to") + else None + ), + ) + + async def register_channel_listener( + self, input: RegisterChannelListenerInput + ) -> str: + resp = await self._client.workflow_service.register_channel_listener( + temporalio.api.workflowservice.v1.RegisterChannelListenerRequest( + namespace=self._client.namespace, + channel=input.channel, + callback=temporalio.api.common.v1.Callback( + nexus=temporalio.api.common.v1.Callback.Nexus( + url=input.callback.url, header=input.callback.headers + ) + ), + request_id=str(uuid.uuid4()), + identity=self._client.identity, + execution=_channel_owner(input.execution), + ), + retry=True, + metadata=input.rpc_metadata, + timeout=input.rpc_timeout, + ) + return resp.listener_id + + async def unregister_channel_listener( + self, input: UnregisterChannelListenerInput + ) -> None: + await self._client.workflow_service.unregister_channel_listener( + temporalio.api.workflowservice.v1.UnregisterChannelListenerRequest( + namespace=self._client.namespace, + channel=input.channel, + listener_id=input.listener_id, + identity=self._client.identity, + execution=_channel_owner(input.execution), + ), + retry=True, + metadata=input.rpc_metadata, + timeout=input.rpc_timeout, + ) + + async def _notification_from_proto( + self, proto: temporalio.api.notification.v1.Notification + ) -> temporalio.workflow.Notification: + # The worker runs the codec over a workflow's notifications before + # they reach workflow code. The client does the same here, so both + # sides hand out payloads a converter can read. + metadata = dict(proto.metadata.items()) + codec = self._client.data_converter.payload_codec + if codec and metadata: + keys = list(metadata) + decoded = await codec.decode([metadata[k] for k in keys]) + metadata = dict(zip(keys, decoded)) + return temporalio.workflow.Notification( + channel=proto.channel, + position=proto.position, + counter=proto.counter, + metadata=metadata, + linked_to=( + temporalio.common.Execution.from_proto(proto.linked_to) + if proto.HasField("linked_to") + else None + ), + ) + async def _apply_headers( self, source: Mapping[str, temporalio.api.common.v1.Payload] | None, @@ -1776,3 +1915,10 @@ async def _apply_headers( == HeaderCodecBehavior.CODEC, self._client.data_converter, ) + + +def _channel_owner( + execution: temporalio.common.Execution | None, +) -> temporalio.api.common.v1.Execution | None: + """The owner on the wire, or ``None`` for an independent channel.""" + return None if execution is None else execution.to_proto() diff --git a/temporalio/client/_interceptor.py b/temporalio/client/_interceptor.py index 68077ebc9..c35a1281f 100644 --- a/temporalio/client/_interceptor.py +++ b/temporalio/client/_interceptor.py @@ -21,6 +21,8 @@ from temporalio.converter import DataConverter if TYPE_CHECKING: + import temporalio.workflow + from ._activity import ( ActivityExecutionAsyncIterator, ActivityExecutionCount, @@ -30,6 +32,8 @@ ActivityOptionsUpdate, AsyncActivityIDReference, ) + from ._callback import Callback + from ._channel import ChannelDescription from ._nexus import ( NexusOperationExecutionAsyncIterator, NexusOperationExecutionCount, @@ -698,6 +702,84 @@ class CountNexusOperationsInput: rpc_timeout: timedelta | None +@dataclass +class NotifyChannelInput: + """Input for :py:meth:`OutboundInterceptor.notify_channel`. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + position: bytes + counter: int + metadata: Mapping[str, Any] | None + rpc_metadata: Mapping[str, str | bytes] + rpc_timeout: timedelta | None + execution: temporalio.common.Execution | None = None + + +@dataclass +class PollChannelInput: + """Input for :py:meth:`OutboundInterceptor.poll_channel`. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + after_counter: int + wait: timedelta | None + max_notifications: int + rpc_metadata: Mapping[str, str | bytes] + rpc_timeout: timedelta | None + execution: temporalio.common.Execution | None = None + + +@dataclass +class DescribeChannelInput: + """Input for :py:meth:`OutboundInterceptor.describe_channel`. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + rpc_metadata: Mapping[str, str | bytes] + rpc_timeout: timedelta | None + execution: temporalio.common.Execution | None = None + + +@dataclass +class RegisterChannelListenerInput: + """Input for :py:meth:`OutboundInterceptor.register_channel_listener`. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + callback: Callback + rpc_metadata: Mapping[str, str | bytes] + rpc_timeout: timedelta | None + execution: temporalio.common.Execution | None = None + + +@dataclass +class UnregisterChannelListenerInput: + """Input for :py:meth:`OutboundInterceptor.unregister_channel_listener`. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + listener_id: str + rpc_metadata: Mapping[str, str | bytes] + rpc_timeout: timedelta | None + execution: temporalio.common.Execution | None = None + + @dataclass class Interceptor: """Interceptor for clients. @@ -1001,3 +1083,51 @@ async def count_nexus_operations( This API is experimental and unstable. """ return await self.next.count_nexus_operations(input) + + ### Notification channel calls + + async def notify_channel(self, input: NotifyChannelInput) -> int: + """Called for every :py:meth:`Client.notify_channel` call. + + .. warning:: + This API is experimental and unstable. + """ + return await self.next.notify_channel(input) + + async def poll_channel( + self, input: PollChannelInput + ) -> list[temporalio.workflow.Notification]: + """Called for every :py:meth:`Client.poll_channel` call. + + .. warning:: + This API is experimental and unstable. + """ + return await self.next.poll_channel(input) + + async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescription: + """Called for every :py:meth:`Client.describe_channel` call. + + .. warning:: + This API is experimental and unstable. + """ + return await self.next.describe_channel(input) + + async def register_channel_listener( + self, input: RegisterChannelListenerInput + ) -> str: + """Called for every :py:meth:`Client.register_channel_listener` call. + + .. warning:: + This API is experimental and unstable. + """ + return await self.next.register_channel_listener(input) + + async def unregister_channel_listener( + self, input: UnregisterChannelListenerInput + ) -> None: + """Called for every :py:meth:`Client.unregister_channel_listener` call. + + .. warning:: + This API is experimental and unstable. + """ + await self.next.unregister_channel_listener(input) diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index c631fd88e..6856c1a20 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -57,6 +57,7 @@ as_completed, wait, ) +from ._channels import Notification from ._context import ( Info, ParentInfo, @@ -258,6 +259,7 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "Notification", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py new file mode 100644 index 000000000..94378c4e4 --- /dev/null +++ b/temporalio/workflow/_channels.py @@ -0,0 +1,81 @@ +"""Notification channels from inside workflow code. + +.. warning:: + This module is experimental and may change in future versions. + +A channel carries notifications, not data. A writer tells the channel's +listeners that a source they consume has moved, and each listener reads the +source itself. A workflow that subscribes gets the notifications the server +folded for it with each Workflow Task. They travel in History, so a replay +sees the same ones at the same points. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field + +import temporalio.api.common.v1 +import temporalio.api.notification.v1 +import temporalio.common + +__all__ = [ + "Notification", +] + + +@dataclass(frozen=True) +class Notification: + """One notification from a channel. + + The server folds notifications per listener while one is pending and no + task has been scheduled for it, keeping the one with the highest counter. + A listener therefore sees where a burst of writes ended, not every write. + """ + + channel: str + """The channel the writer notified.""" + + position: bytes + """Where the source stands after the write, in the writer's terms. + + Opaque to the server and handed over as sent. + """ + + counter: int + """Orders notifications from one channel's writers. Higher is later.""" + + metadata: Mapping[str, temporalio.api.common.v1.Payload] = field( + default_factory=dict + ) + """Details for the listener, such as which topic moved, as payloads. + + A codec has been applied. Convert a value with + :py:meth:`temporalio.converter.PayloadConverter.from_payload` on the + converter :py:func:`temporalio.workflow.payload_converter` returns. + """ + + linked_to: temporalio.common.Execution | None = None + """The execution a linked channel belongs to, and the run that received this. + + ``None`` for a notification from an independent channel. A workflow that + holds both kinds of handle on one name gets a notification on the handle + its kind names: :func:`temporalio.workflow.linked_channel` when set, + :func:`temporalio.workflow.subscribe_channel` otherwise. + """ + + @staticmethod + def _from_proto( + proto: temporalio.api.notification.v1.Notification, + ) -> Notification: + return Notification( + channel=proto.channel, + position=proto.position, + counter=proto.counter, + metadata=dict(proto.metadata.items()), + linked_to=( + temporalio.common.Execution.from_proto(proto.linked_to) + if proto.HasField("linked_to") + else None + ), + ) From 7838fa1d5da36ce83d58e4b41d29655e86327e06 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:40:44 -0700 Subject: [PATCH 3/3] Covered the client channel calls. The unit case fakes the service. The live cases skip unless -E names a server that serves channels. --- tests/streams/__init__.py | 0 tests/streams/conftest.py | 34 ++++++ tests/streams/test_channels.py | 184 +++++++++++++++++++++++++++++++++ 3 files changed, 218 insertions(+) create mode 100644 tests/streams/__init__.py create mode 100644 tests/streams/conftest.py create mode 100644 tests/streams/test_channels.py diff --git a/tests/streams/__init__.py b/tests/streams/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py new file mode 100644 index 000000000..a456d5975 --- /dev/null +++ b/tests/streams/conftest.py @@ -0,0 +1,34 @@ +import pytest + +_ENVIRONMENTS_WITHOUT_CHANNELS = ("local", "time-skipping", "envconfig") + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "needs_channel_server: the case needs a server that serves notification " + "channels, named with -E host:port", + ) + + +def pytest_collection_modifyitems( + config: pytest.Config, items: list[pytest.Item] +) -> None: + # The dev server the suite starts for itself does not accept the + # subscribe command, so the live channel cases run only against a server + # the caller points at. + skips: list[tuple[str, pytest.MarkDecorator]] = [] + if config.getoption("--workflow-environment") in _ENVIRONMENTS_WITHOUT_CHANNELS: + skips.append( + ( + "needs_channel_server", + pytest.mark.skip( + reason="needs a server that serves notification channels; " + "name one with -E" + ), + ) + ) + for item in items: + for marker, skip in skips: + if item.get_closest_marker(marker): + item.add_marker(skip) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py new file mode 100644 index 000000000..9a81b31c5 --- /dev/null +++ b/tests/streams/test_channels.py @@ -0,0 +1,184 @@ +"""The notification channel calls on the client. + +The unit case fakes the service. The live cases need a server that serves +channels, named with -E, and skip otherwise. +""" + +from __future__ import annotations + +import uuid +from datetime import timedelta +from typing import Any + +import pytest + +import temporalio.api.enums.v1 +import temporalio.api.notification.v1 +import temporalio.api.workflowservice.v1 +import temporalio.common +from temporalio.client import ( + Callback, + ChannelKind, + Client, +) +from temporalio.service import RPCError, RPCStatusCode + +Notification = temporalio.api.notification.v1.Notification + + +def _linked(channel: str, counter: int, position: bytes = b"") -> Notification: + """A notification the way a linked channel's owner receives it.""" + return Notification( + channel=channel, + counter=counter, + position=position, + linked_to=temporalio.common.Execution.workflow("wf", "run").to_proto(), + ) + + +async def test_the_client_addresses_a_linked_channel_by_execution( + client: Client, monkeypatch: pytest.MonkeyPatch +): + """Every channel call carries the owner it was given, and only then.""" + requests: list[Any] = [] + owner = temporalio.common.Execution.workflow("wf", "run") + describe_response = temporalio.api.workflowservice.v1.DescribeChannelResponse( + kind=temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED, + linked_to=owner.to_proto(), + latest=_linked("orders", 3, b"3-0"), + ) + responses = { + "notify_channel": temporalio.api.workflowservice.v1.NotifyChannelResponse(), + "poll_channel": temporalio.api.workflowservice.v1.PollChannelResponse( + notifications=[_linked("orders", 3, b"3-0")] + ), + "describe_channel": describe_response, + "register_channel_listener": ( + temporalio.api.workflowservice.v1.RegisterChannelListenerResponse( + listener_id="listener" + ) + ), + "unregister_channel_listener": ( + temporalio.api.workflowservice.v1.UnregisterChannelListenerResponse() + ), + } + for name, response in responses.items(): + + async def call(req: Any, *, _response: Any = response, **_: Any) -> Any: + requests.append(req) + return _response + + monkeypatch.setattr(client.workflow_service, name, call) + + callback = Callback(url="http://localhost:1/never-called", headers={}) + await client.notify_channel("orders", counter=1, workflow_id="wf", run_id="run") + [polled] = await client.poll_channel("orders", workflow_id="wf", wait=False) + description = await client.describe_channel("orders", workflow_id="wf") + await client.register_channel_listener("orders", callback, workflow_id="wf") + await client.unregister_channel_listener("orders", "listener", workflow_id="wf") + # The workflow id is shorthand for a workflow execution, run id and all. + by_workflow_id = temporalio.common.Execution.workflow("wf").to_proto() + assert [req.execution for req in requests] == [ + owner.to_proto(), + by_workflow_id, + by_workflow_id, + by_workflow_id, + by_workflow_id, + ] + assert polled.linked_to == owner + assert description.kind == ChannelKind.LINKED + assert description.linked_to == owner + assert description.latest is not None and description.latest.linked_to == owner + + # An execution names any owner, a standalone activity included. + requests.clear() + activity = temporalio.common.Execution.activity("act", "run") + await client.notify_channel("orders", counter=1, execution=activity) + await client.poll_channel("orders", execution=activity, wait=False) + await client.describe_channel("orders", execution=activity) + await client.register_channel_listener("orders", callback, execution=activity) + await client.unregister_channel_listener("orders", "listener", execution=activity) + assert [req.execution for req in requests] == [activity.to_proto()] * 5 + assert activity.to_proto().type == ( + temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_ACTIVITY + ) + + # Without an owner the calls address the independent channel. + requests.clear() + await client.notify_channel("orders", counter=1) + await client.poll_channel("orders", wait=False) + await client.describe_channel("orders") + await client.register_channel_listener("orders", callback) + await client.unregister_channel_listener("orders", "listener") + assert [req.HasField("execution") for req in requests] == [False] * 5 + + with pytest.raises(ValueError, match="run_id needs workflow_id"): + await client.notify_channel("orders", counter=1, run_id="run") + # The shorthand and the execution are two ways to say one thing. + with pytest.raises(ValueError, match="not both"): + await client.notify_channel( + "orders", counter=1, execution=activity, workflow_id="wf" + ) + with pytest.raises(ValueError, match="not both"): + await client.poll_channel("orders", execution=activity, run_id="run") + + +@pytest.mark.needs_channel_server +async def test_a_callback_listener_registers_and_unregisters(client: Client): + channel = f"orders-{uuid.uuid4()}" + callback = Callback(url="http://localhost:1/never-called", headers={}) + listener_id = await client.register_channel_listener(channel, callback) + description = await client.describe_channel(channel) + assert [listener.listener_id for listener in description.listeners] == [listener_id] + assert description.listeners[0].callback == callback + await client.unregister_channel_listener(channel, listener_id) + description = await client.describe_channel(channel) + assert description.listeners == [] + + +@pytest.mark.needs_channel_server +async def test_a_channel_retains_notifications_for_pollers(client: Client): + channel = f"orders-{uuid.uuid4()}" + # A channel nobody has touched does not exist. + with pytest.raises(RPCError) as untouched: + await client.describe_channel(f"untouched-{uuid.uuid4()}") + assert untouched.value.status == RPCStatusCode.NOT_FOUND + # Nobody listens yet: the notification is kept for pollers and the count + # says zero. + assert await client.notify_channel(channel, position=b"2-0", counter=2) == 0 + description = await client.describe_channel(channel) + assert description.listeners == [] + assert description.latest is not None + assert (description.latest.position, description.latest.counter) == (b"2-0", 2) + assert description.retained_count == 1 + polled = await client.poll_channel(channel, wait=False) + assert [(n.position, n.counter) for n in polled] == [(b"2-0", 2)] + # At or below the latest counter a notify changes nothing and is not kept. + assert await client.notify_channel(channel, position=b"1-0", counter=1) == 0 + assert await client.notify_channel(channel, position=b"2-0", counter=2) == 0 + description = await client.describe_channel(channel) + assert description.latest is not None and description.latest.counter == 2 + assert description.retained_count == 1 + # Above it the notification is kept, metadata and all, and a poll after + # the earlier counter sees only the new one. + assert ( + await client.notify_channel( + channel, position=b"3-0", counter=3, metadata={"topic": "inputs"} + ) + == 0 + ) + [newest] = await client.poll_channel(channel, after_counter=2, wait=False) + assert newest.counter == 3 + assert ( + client.data_converter.payload_converter.from_payload( + newest.metadata["topic"], str + ) + == "inputs" + ) + polled = await client.poll_channel(channel, wait=False) + assert [n.counter for n in polled] == [2, 3] + # A poll above the latest waits its bound out and comes back empty. + polled = await client.poll_channel( + channel, after_counter=3, wait=timedelta(seconds=1) + ) + assert polled == []