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/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. 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 + ), + ) 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 == []