diff --git a/CHANGELOG.md b/CHANGELOG.md index f1d184c32..fa3038107 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,6 +33,16 @@ to include examples, links to docs, or any other relevant information. through `ExternalOutputStreamClient`. Workflow output is staged outside History and becomes readable only after its compact Workflow Task marker is committed. + +- **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. + - A workflow subscribes with `workflow.subscribe_channel(name)`, reads a channel linked to it with + `workflow.linked_channel(name)` and ends a subscription with `unsubscribe()`. Notifications + arrive with the workflow's tasks. `WorkflowExecutionDescription.channel_subscriptions` lists + the channels a run listens on. ### Changed - Standalone Activities are now generally available (GA). (Standalone Activities as Nexus operations diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index cdb34b860..164c31bfc 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -59,6 +59,13 @@ from ._callback import ( Callback, ) +from ._channel import ( + ChannelAddress, + ChannelDescription, + ChannelKind, + ChannelListener, + ChannelSubscriptionInfo, +) from ._client import ( Client, ClientConfig, @@ -101,6 +108,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -115,9 +123,12 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -130,6 +141,7 @@ TerminateWorkflowInput, TriggerScheduleInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, UpdateWithStartUpdateWorkflowInput, @@ -300,6 +312,11 @@ "TerminateNexusOperationInput", "ListNexusOperationsInput", "CountNexusOperationsInput", + "DescribeChannelInput", + "NotifyChannelInput", + "PollChannelInput", + "RegisterChannelListenerInput", + "UnregisterChannelListenerInput", "StartWorkflowUpdateInput", "UpdateWithStartUpdateWorkflowInput", "UpdateWithStartStartWorkflowInput", @@ -336,6 +353,11 @@ "CloudOperationsClient", "Plugin", "Callback", + "ChannelDescription", + "ChannelKind", + "ChannelListener", + "ChannelAddress", + "ChannelSubscriptionInfo", "_ClientImpl", "_apply_headers", "_decode_user_metadata", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py new file mode 100644 index 000000000..0decc2265 --- /dev/null +++ b/temporalio/client/_channel.py @@ -0,0 +1,228 @@ +"""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.api.workflow.v1 +import temporalio.common +from temporalio.workflow import Notification + +from ._callback import Callback + +__all__ = [ + "ChannelAddress", + "ChannelDescription", + "ChannelKind", + "ChannelListener", + "ChannelSubscriptionInfo", +] + + +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 ChannelSubscriptionInfo: + """A workflow's standing on one channel, as its description reports it. + + An independent channel is listed from the subscribe event until the run + unsubscribes or closes. A linked channel is listed once it holds state. A + closed run keeps listing what it stood on, and a continue-as-new + successor starts with nothing. + + .. warning:: + This API is experimental and unstable. + """ + + channel: str + """Channel name.""" + + kind: ChannelKind + """:attr:`ChannelKind.INDEPENDENT` for a subscription the workflow made by + command, :attr:`ChannelKind.LINKED` for a channel linked to it.""" + + subscribed_event_id: int + """Id of the event that recorded the subscription. Zero for the linked kind.""" + + last_counter: int + """Highest counter the workflow has accepted from the channel. Zero when + none has arrived.""" + + pending_notification: Notification | None + """The notification held for the workflow's next Workflow Task, when one + is pending.""" + + scheduled_counter: int + """Counter carried by the scheduled event of a Workflow Task that has not + started yet. Zero otherwise.""" + + listener_count: int + """Linked kind: callback listeners registered on the channel.""" + + retained_count: int + """Linked kind: notifications retained for pollers.""" + + accepted_count: int + """Linked kind: notifications the channel has accepted over its life.""" + + @staticmethod + def _from_proto( + proto: temporalio.api.workflow.v1.ChannelSubscriptionInfo, + ) -> ChannelSubscriptionInfo: + return ChannelSubscriptionInfo( + channel=proto.channel, + kind=ChannelKind(proto.kind), + subscribed_event_id=proto.subscribed_event_id, + last_counter=proto.last_counter, + pending_notification=( + Notification._from_proto(proto.pending_notification) + if proto.HasField("pending_notification") + else None + ), + scheduled_counter=proto.scheduled_counter, + listener_count=proto.listener_count, + retained_count=proto.retained_count, + accepted_count=proto.accepted_count, + ) + + +@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 5acdfe476..bea1c81cc 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -64,22 +64,28 @@ 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, ) @@ -115,6 +121,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. @@ -2880,6 +2889,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, @@ -3070,3 +3343,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 9595351f2..66fc97504 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, @@ -53,6 +55,7 @@ ActivityHandle, AsyncActivityIDReference, ) +from ._channel import ChannelDescription, ChannelKind, ChannelListener from ._exceptions import ( AsyncActivityCancelledError, ScheduleAlreadyRunningError, @@ -73,6 +76,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -86,9 +90,12 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -101,6 +108,7 @@ TerminateWorkflowInput, TriggerScheduleInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, UpdateWithStartUpdateWorkflowInput, @@ -1697,6 +1705,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, @@ -1709,3 +1848,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 a6daedd45..12a16b8b3 100644 --- a/temporalio/client/_interceptor.py +++ b/temporalio/client/_interceptor.py @@ -25,6 +25,8 @@ from ._callback import Callback if TYPE_CHECKING: + import temporalio.workflow + from ._activity import ( ActivityExecutionAsyncIterator, ActivityExecutionCount, @@ -32,6 +34,8 @@ ActivityHandle, AsyncActivityIDReference, ) + from ._callback import Callback + from ._channel import ChannelDescription from ._nexus import ( NexusOperationExecutionAsyncIterator, NexusOperationExecutionCount, @@ -679,6 +683,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. @@ -979,3 +1061,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/client/_workflow.py b/temporalio/client/_workflow.py index 6b3559c31..e2532e23d 100644 --- a/temporalio/client/_workflow.py +++ b/temporalio/client/_workflow.py @@ -60,6 +60,7 @@ SelfType, ) from ._callback import Callback +from ._channel import ChannelSubscriptionInfo from ._exceptions import ( WorkflowContinuedAsNewError, WorkflowFailureError, @@ -1426,6 +1427,18 @@ class WorkflowExecutionDescription(WorkflowExecution): raw_description: temporalio.api.workflowservice.v1.DescribeWorkflowExecutionResponse """Underlying protobuf description.""" + channel_subscriptions: Sequence[ChannelSubscriptionInfo] = () + """The notification channels this run stands on. + + The independent channels it subscribed to and the channels linked to it + that hold any state, sorted by name with the independent kind first. + Empty when there are none. See + :py:class:`temporalio.client.ChannelSubscriptionInfo`. + + .. warning:: + This API is experimental and unstable. + """ + _static_summary: str | None = None _static_details: str | None = None _metadata_decoded: bool = False @@ -1464,6 +1477,10 @@ async def _from_raw_description( namespace=namespace, converter=converter, raw_description=description, + channel_subscriptions=tuple( + ChannelSubscriptionInfo._from_proto(info) + for info in description.channel_subscriptions + ), ) diff --git a/temporalio/common.py b/temporalio/common.py index ad75b56b9..a4e8e7c5a 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/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index c3db3f3ce..7924d5ca0 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -18,6 +18,7 @@ from collections.abc import ( Awaitable, Callable, + Collection, Coroutine, Generator, Iterable, @@ -329,6 +330,18 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: ) self._patch_activation_callback = det.patch_activation_callback self._default_workflow_logic_flags = det.default_workflow_logic_flags + self._subscribed_channels: set[str] = set() + # Keyed by channel name; one subscription per channel per run + self._channel_subscriptions: dict[ + str, temporalio.workflow.ChannelSubscription + ] = {} + #: The channels linked to this workflow that the run listens on. No + #: command: the owner is the listener by construction, so the set only + #: tells a notification carrying ``linked_to`` from one nobody asked for. + self._linked_channels: set[str] = set() + self._linked_channel_subscriptions: dict[ + str, temporalio.workflow.ChannelSubscription + ] = {} self._primary_task: asyncio.Task[None] | None = None self._cancel_primary_task_pending = False self._time_ns = 0 @@ -533,6 +546,10 @@ def activate( job_sets[0].append(job) elif job.HasField("signal_workflow") or job.HasField("do_update"): job_sets[1].append(job) + elif job.HasField("notifications_received"): + # Ordered with the Signals, where Core puts them: what a + # task was woken for is known before anything it resolves. + job_sets[1].append(job) elif not job.HasField("query_workflow"): if job.HasField("initialize_workflow"): start_job = job.initialize_workflow @@ -701,6 +718,8 @@ def _apply( self._apply_resolve_external_stream_waits(job.resolve_external_stream_waits) elif job.HasField("replay_external_streams"): self._apply_replay_external_streams(job.replay_external_streams) + elif job.HasField("notifications_received"): + self._apply_notifications_received(job.notifications_received) elif job.HasField("resolve_child_workflow_execution"): self._apply_resolve_child_workflow_execution( job.resolve_child_workflow_execution @@ -892,6 +911,80 @@ def _apply_resolve_external_stream_waits( if self._external_stream_runtime is not None: self._external_stream_runtime.resolve_all_pending() + def _apply_notifications_received( + self, + job: temporalio.bridge.proto.workflow_activation.NotificationsReceived, + ) -> None: + """Hands the notifications this Workflow Task was woken with to their takers. + + A notification names the channel it came over and, for a channel + linked to this workflow, the owner in ``linked_to``; a public + subscription of the matching kind gets it in arrival order. + """ + for proto in job.notifications: + if proto.HasField("linked_to"): + subscriptions = self._linked_channel_subscriptions + listened: Collection[str] = self._linked_channels + else: + subscriptions = self._channel_subscriptions + listened = self._subscribed_channels + subscription = subscriptions.get(proto.channel) + if subscription is not None: + subscription._deliver( + temporalio.workflow.Notification._from_proto(proto) + ) + elif proto.channel not in listened: + # The server fans out to whatever listened at the time, so a + # channel this run never asked for is not the workflow's + # concern. + logger.debug( + "Dropping a notification on channel %r, which this run does " + "not listen on", + proto.channel, + ) + + def workflow_subscribe_channel( + self, channel: str + ) -> temporalio.workflow.ChannelSubscription: + existing = self._channel_subscriptions.get(channel) + if existing is not None: + return existing + self._issue_channel_subscription(channel) + subscription = temporalio.workflow.ChannelSubscription(channel) + self._channel_subscriptions[channel] = subscription + return subscription + + def workflow_unsubscribe_channel(self, channel: str) -> None: + command = self._add_command() + command.unsubscribe_notification_channel.channel = channel + # Out of the map before the next activation: the server may still hand + # this run a notification it folded onto a task ahead of the command. + self._channel_subscriptions.pop(channel, None) + # Out of the once-per-run set too, so a later subscribe reaches the + # server again instead of being folded into the one it left. + self._subscribed_channels.discard(channel) + + def workflow_linked_channel( + self, channel: str + ) -> temporalio.workflow.ChannelSubscription: + existing = self._linked_channel_subscriptions.get(channel) + if existing is not None: + return existing + self._linked_channels.add(channel) + subscription = temporalio.workflow.ChannelSubscription(channel, linked=True) + self._linked_channel_subscriptions[channel] = subscription + return subscription + + def _issue_channel_subscription(self, channel: str) -> None: + """Emits the subscribe command once per channel per run. + + A second command for one channel would only record a second event. + """ + if channel in self._subscribed_channels: + return + self._subscribed_channels.add(channel) + self._add_command().subscribe_notification_channel.channel = channel + def _apply_replay_external_streams( self, job: temporalio.bridge.proto.workflow_activation.ReplayExternalStreams, diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index fa2681139..edb054bd6 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -57,6 +57,12 @@ as_completed, wait, ) +from ._channels import ( + ChannelSubscription, + Notification, + linked_channel, + subscribe_channel, +) from ._context import ( Info, ParentInfo, @@ -252,6 +258,10 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "ChannelSubscription", + "Notification", + "linked_channel", + "subscribe_channel", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py new file mode 100644 index 000000000..07c3bde7a --- /dev/null +++ b/temporalio/workflow/_channels.py @@ -0,0 +1,252 @@ +"""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 + +import asyncio +from collections import deque +from collections.abc import Mapping +from dataclasses import dataclass, field + +import temporalio.api.common.v1 +import temporalio.api.notification.v1 +import temporalio.common +from temporalio.workflow._context import _Runtime + +__all__ = [ + "ChannelSubscription", + "Notification", + "linked_channel", + "subscribe_channel", +] + + +@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 + ), + ) + + +class ChannelSubscription: + """A workflow's handle on one channel, of either kind. + + Prefer :func:`temporalio.workflow.subscribe_channel` for an independent + channel and :func:`temporalio.workflow.linked_channel` for one linked to + this workflow. The handle is an async iterator over the notifications as + they arrive, and :meth:`receive` takes them one at a time. Notifications + wait in arrival order until taken. Two loops on one handle share its + buffer and interleave. :meth:`unsubscribe` ends an independent + subscription; a linked channel lasts as long as the run. + """ + + def __init__(self, channel: str, *, linked: bool = False) -> None: + """Prefer the two module functions named above.""" + self._channel = channel + self._linked = linked + self._closed = False + self._pending: deque[Notification] = deque() + self._waiters: deque[asyncio.Future[None]] = deque() + + @property + def channel(self) -> str: + """The channel this handle is on.""" + return self._channel + + @property + def linked(self) -> bool: + """Whether the channel is the one linked to this workflow. + + A linked handle gets the notifications that carry + :attr:`Notification.linked_to`; an independent one gets the rest. + """ + return self._linked + + @property + def closed(self) -> bool: + """Whether :meth:`unsubscribe` has ended this subscription. + + A closed handle still hands out the notifications it had queued, then + :meth:`receive` raises and iteration ends. + """ + return self._closed + + def unsubscribe(self) -> None: + """End this workflow's subscription to the channel. + + Issues the unsubscribe command once; a second call changes nothing. + Notifications already queued on this handle can still be read, and + one the server put on a scheduled Workflow Task before the command + landed is dropped on arrival. A later + :func:`temporalio.workflow.subscribe_channel` for the same name opens + a new subscription with a new command. + + Raises: + ValueError: The handle is from + :func:`temporalio.workflow.linked_channel`. A linked channel is + part of the run and has no subscription to end. + """ + if self._linked: + raise ValueError("a linked channel has no subscription") + if self._closed: + return + _Runtime.current().workflow_unsubscribe_channel(self._channel) + self._closed = True + # The waiters wake to find the handle closed with nothing queued. + self._wake() + + async def receive(self) -> Notification: + """The next notification on this channel, waiting for one to arrive. + + The wait is a future the delivery resolves, so it adds no command and + replays the same way. + + Raises: + RuntimeError: The subscription is closed and nothing is queued. + """ + notification = await self._next() + if notification is None: + raise RuntimeError("channel subscription closed") + return notification + + def __aiter__(self) -> ChannelSubscription: + """The subscription is its own iterator.""" + return self + + async def __anext__(self) -> Notification: + """The next notification. The iteration ends once the handle is closed and drained.""" + notification = await self._next() + if notification is None: + raise StopAsyncIteration + return notification + + async def _next(self) -> Notification | None: + while not self._pending: + if self._closed: + return None + waiter: asyncio.Future[None] = asyncio.Future() + self._waiters.append(waiter) + try: + await waiter + finally: + if waiter in self._waiters: + self._waiters.remove(waiter) + return self._pending.popleft() + + def _deliver(self, notification: Notification) -> None: + self._pending.append(notification) + self._wake() + + def _wake(self) -> None: + # Every waiter wakes; the ones that find the buffer empty again wait + # once more. + while self._waiters: + waiter = self._waiters.popleft() + if not waiter.done(): + waiter.set_result(None) + + +def subscribe_channel(channel: str) -> ChannelSubscription: + """Subscribe this workflow to ``channel``. + + The first call for a channel in a run issues a command, so gate a new + channel with :func:`temporalio.workflow.patched` as you would a timer. A + second call for the same channel returns the subscription already open, + and the two share its buffer. From then on each Workflow Task carries the + notifications the server folded for this workflow on the channel, and + they arrive here. A successor run after continue-as-new starts with no + subscriptions. + + Args: + channel: Name of the channel, scoped to the namespace. + + Raises: + ValueError: ``channel`` is empty. + """ + if not channel: + raise ValueError("channel must not be empty") + return _Runtime.current().workflow_subscribe_channel(channel) + + +def linked_channel(channel: str) -> ChannelSubscription: + """Listen on the channel named ``channel`` that is linked to this workflow. + + A linked channel lives in this workflow's own state, so the workflow is + its listener by construction: no command, no event, and no gate needed + for a new name. A writer reaches it by naming this workflow as the + execution, as in :py:meth:`temporalio.client.Client.notify_channel` with + ``workflow_id``, and a successor run after continue-as-new is reached by + the same calls. + A second call for the same name returns the handle already open, and the + two share its buffer. The name does not collide with an independent + channel's: a notification carrying :attr:`Notification.linked_to` comes + here, one without it goes to :func:`subscribe_channel`. + + Args: + channel: Name of the channel, scoped to this workflow. + + Raises: + ValueError: ``channel`` is empty. + """ + if not channel: + raise ValueError("channel must not be empty") + return _Runtime.current().workflow_linked_channel(channel) diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index b33f83150..f474973eb 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -22,6 +22,7 @@ if TYPE_CHECKING: from ._activities import ActivityCancellationType, ActivityHandle + from ._channels import ChannelSubscription from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle from ._workflow_ops import ( @@ -475,6 +476,21 @@ async def workflow_start_nexus_operation( summary: str | None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... + + @abstractmethod + def workflow_unsubscribe_channel(self, channel: str) -> None: + """Record the unsubscribe command for ``channel`` and forget its handle. + + Called once per subscription, by the handle that owns it, so a late + notification for the channel finds no handle and is dropped. + """ + ... + + @abstractmethod + def workflow_linked_channel(self, channel: str) -> ChannelSubscription: ... + @abstractmethod def workflow_time_ns(self) -> int: ... diff --git a/tests/contrib/external_workflow_streams/conftest.py b/tests/contrib/external_workflow_streams/conftest.py index d6e8483c0..52dec2f67 100644 --- a/tests/contrib/external_workflow_streams/conftest.py +++ b/tests/contrib/external_workflow_streams/conftest.py @@ -12,10 +12,101 @@ import uuid from collections.abc import AsyncGenerator from dataclasses import dataclass +from typing import Any import pytest import pytest_asyncio +#: The environments whose server the suite starts for itself. None of them +#: accepts the subscribe-notification-channel command. +_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", + ) + config.addinivalue_line( + "markers", + "needs_linked_server: the case needs a server that serves channels linked " + "to a workflow, named with -E host:port; the case skips itself on one " + "with only independent channels", + ) + config.addinivalue_line( + "markers", + "needs_unsubscribe_server: the case needs a server that accepts the " + "unsubscribe-notification-channel command, named with -E host:port; an " + "older channel server fails the Workflow Task that carries it", + ) + + +#: The markers naming a server capability the suite's own servers lack. +_CHANNEL_SERVER_MARKERS = ( + "needs_channel_server", + "needs_linked_server", + "needs_unsubscribe_server", +) + + +def pytest_collection_modifyitems( + config: pytest.Config, items: list[pytest.Item] +) -> None: + if config.getoption("--workflow-environment") not in _ENVIRONMENTS_WITHOUT_CHANNELS: + return + skip = pytest.mark.skip( + reason="needs a server that serves notification channels; name one with -E" + ) + for item in items: + if any(item.get_closest_marker(marker) for marker in _CHANNEL_SERVER_MARKERS): + item.add_marker(skip) + + +#: The channel the probe describes. Nobody writes to it. +_PROBE_CHANNEL = "external-stream/probe" + + +async def server_channel_support(client: Any) -> Any: + """Which channel kinds the server serves, asked through the client's own calls. + + ``None`` when it serves no channels, otherwise :attr:`ChannelKind.LINKED` + or :attr:`ChannelKind.INDEPENDENT`. The linked kind shows only on a running + workflow, so one is started on a task queue nobody polls and asked about. + """ + from temporalio.client import ChannelKind + from temporalio.service import RPCError, RPCStatusCode + + try: + await client.describe_channel(_PROBE_CHANNEL) + except RPCError as err: + if err.status == RPCStatusCode.UNIMPLEMENTED: + return None + if err.status != RPCStatusCode.NOT_FOUND: + raise + probe = uuid.uuid4().hex + handle = await client.start_workflow( + "ChannelSupportProbe", + id=f"channel-support-probe-{probe}", + task_queue=f"nobody-polls-{probe}", + ) + try: + description = await client.describe_channel( + _PROBE_CHANNEL, workflow_id=handle.id + ) + except RPCError as err: + # A server with only independent channels ignores the owner and has + # never seen the name. + if err.status == RPCStatusCode.NOT_FOUND: + return ChannelKind.INDEPENDENT + raise + finally: + await handle.terminate() + if description.kind is ChannelKind.LINKED: + return ChannelKind.LINKED + return ChannelKind.INDEPENDENT + + DEFAULT_REDIS_URL = "redis://127.0.0.1:6379" #: Every key this suite creates starts with this, so a leaked key is diff --git a/tests/contrib/external_workflow_streams/test_channels.py b/tests/contrib/external_workflow_streams/test_channels.py new file mode 100644 index 000000000..75252e32f --- /dev/null +++ b/tests/contrib/external_workflow_streams/test_channels.py @@ -0,0 +1,931 @@ +"""The notification channel surface: the command, the delivery and the client calls. + +The workflow instance is driven with activations directly, the way Core +drives it, because the dev server this chain tests against does not accept +the subscribe command. The live cases at the end need a server that does and +skip otherwise; the ones on the linked kind need a server with it. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import uuid +from datetime import datetime, timedelta, timezone +from typing import Any + +import pytest + +import temporalio.api.common.v1 +import temporalio.api.enums.v1 +import temporalio.api.notification.v1 +import temporalio.api.workflow.v1 +import temporalio.api.workflowservice.v1 +import temporalio.bridge.proto.workflow_activation +import temporalio.bridge.proto.workflow_completion +import temporalio.common +import temporalio.converter +from temporalio import workflow +from temporalio.api.enums.v1 import EventType +from temporalio.client import ( + Callback, + ChannelKind, + ChannelSubscriptionInfo, + Client, + WorkflowExecutionDescription, +) +from temporalio.client._client import _channel_execution +from temporalio.client._impl import _channel_owner +from temporalio.common import Execution, ExecutionType +from temporalio.service import RPCError, RPCStatusCode +from temporalio.worker._workflow_instance import ( + UnsandboxedWorkflowRunner, + WorkflowInstance, + WorkflowInstanceDetails, +) +from tests.contrib.external_workflow_streams.conftest import server_channel_support +from tests.helpers import assert_eventually, new_worker + +WorkflowActivation = temporalio.bridge.proto.workflow_activation.WorkflowActivation +WorkflowActivationJob = ( + temporalio.bridge.proto.workflow_activation.WorkflowActivationJob +) +WorkflowActivationCompletion = ( + temporalio.bridge.proto.workflow_completion.WorkflowActivationCompletion +) +Notification = temporalio.api.notification.v1.Notification +ExecutionProto = temporalio.api.common.v1.Execution +WORKFLOW_TYPE = temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_WORKFLOW + + +def _linked_to(notification: workflow.Notification) -> dict[str, Any] | None: + if notification.linked_to is None: + return None + return { + "type": notification.linked_to.type.name, + "business_id": notification.linked_to.business_id, + "run_id": notification.linked_to.run_id, + } + + +@workflow.defn +class ReceiveOne: + """Subscribes to one channel and reports the first notification.""" + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + notification = await workflow.subscribe_channel(channel).receive() + return { + "channel": notification.channel, + "counter": notification.counter, + "position": notification.position.decode(), + "topic": ( + workflow.payload_converter().from_payload( + notification.metadata["topic"], str + ) + if "topic" in notification.metadata + else None + ), + } + + +@workflow.defn +class CountToTwo: + """Subscribes twice to one channel and counts notifications up to counter two.""" + + @workflow.run + async def run(self, channel: str) -> int: + first = workflow.subscribe_channel(channel) + second = workflow.subscribe_channel(channel) + assert first is second + seen = 0 + async for notification in second: + seen += 1 + if notification.counter >= 2: + break + return seen + + +@workflow.defn +class EmptyChannel: + """Asks for a channel with no name.""" + + @workflow.run + async def run(self) -> None: + workflow.subscribe_channel("") + + +@workflow.defn +class EmptyLinkedChannel: + """Asks for a linked channel with no name.""" + + @workflow.run + async def run(self) -> None: + workflow.linked_channel("") + + +@workflow.defn +class ReceiveLinked: + """Reads its own linked channel and reports the first notification.""" + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + subscription = workflow.linked_channel(channel) + assert subscription is workflow.linked_channel(channel) + assert subscription.linked + notification = await subscription.receive() + return { + "channel": notification.channel, + "counter": notification.counter, + "linked_to": _linked_to(notification), + } + + +@workflow.defn +class BothKinds: + """Listens on the independent and the linked channel of one name.""" + + @workflow.run + async def run(self, channel: str) -> list[str]: + independent = workflow.subscribe_channel(channel) + linked = workflow.linked_channel(channel) + seen: list[str] = [] + + async def take(kind: str, subscription: workflow.ChannelSubscription) -> None: + notification = await subscription.receive() + assert (notification.linked_to is not None) == (kind == "linked") + seen.append(kind) + + await asyncio.gather(take("independent", independent), take("linked", linked)) + return seen + + +async def _drain(handle: workflow.ChannelSubscription) -> dict[str, Any]: + """What a closed handle still gives: the queue, then the end, then the refusal.""" + drained = [notification.counter async for notification in handle] + try: + await handle.receive() + except RuntimeError as err: + refused: str | None = str(err) + else: + refused = None + return {"closed": handle.closed, "drained": drained, "refused": refused} + + +@workflow.defn +class ReceiveThenUnsubscribe: + """Takes the first notification, unsubscribes twice, then waits to be finished. + + The wait keeps the run open so a late notification can be aimed at it. + """ + + def __init__(self) -> None: + self._done = False + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + handle = workflow.subscribe_channel(channel) + first = await handle.receive() + assert not handle.closed + handle.unsubscribe() + handle.unsubscribe() + drained = await _drain(handle) + await workflow.wait_condition(lambda: self._done) + return {"first": first.counter, **drained} + + @workflow.signal + def finish(self) -> None: + self._done = True + + +@workflow.defn +class UnsubscribeWithOneQueued: + """Unsubscribes with a notification still queued and reads it afterwards.""" + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + handle = workflow.subscribe_channel(channel) + first = await handle.receive() + handle.unsubscribe() + return {"first": first.counter, **(await _drain(handle))} + + +@workflow.defn +class UnsubscribeLinked: + """Tries to unsubscribe from its linked channel.""" + + @workflow.run + async def run(self, channel: str) -> None: + workflow.linked_channel(channel).unsubscribe() + + +@workflow.defn +class Resubscribe: + """Subscribes, unsubscribes and subscribes again, then receives on the new handle.""" + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + first = workflow.subscribe_channel(channel) + first.unsubscribe() + second = workflow.subscribe_channel(channel) + assert second is not first + assert first.closed and not second.closed + notification = await second.receive() + return {"counter": notification.counter, **(await _drain(first))} + + +def _instance(workflow_class: type) -> WorkflowInstance: + """Build an instance the way the worker does, without a worker. + + Needs a running event loop, since the constructor puts the runtime on it. + """ + defn = workflow._Definition.must_from_class(workflow_class) + now = datetime.now(timezone.utc) + info = workflow.Info( + attempt=1, + continued_run_id=None, + cron_schedule=None, + execution_timeout=None, + first_execution_run_id="run", + headers={}, + namespace="default", + parent=None, + root=None, + priority=temporalio.common.Priority.default, + raw_memo={}, + retry_policy=None, + run_id="run", + run_timeout=None, + search_attributes={}, + start_time=now, + task_queue="tq", + task_timeout=timedelta(seconds=10), + typed_search_attributes=temporalio.common.TypedSearchAttributes.empty, + workflow_id="wf", + workflow_start_time=now, + workflow_type=defn.name or "", + ) + converter = temporalio.converter.DataConverter.default + return UnsandboxedWorkflowRunner().create_instance( + WorkflowInstanceDetails( + payload_converter_factory=converter._new_payload_converter, + failure_converter_class=converter.failure_converter_class, + interceptor_classes=[], + defn=defn, + info=info, + randomness_seed=0, + extern_functions={}, + disable_eager_activity_execution=False, + worker_level_failure_exception_types=[], + patch_activation_callback=None, + last_completion_result=temporalio.api.common.v1.Payloads(), + last_failure=None, + ) + ) + + +def _start(workflow_class: type, *args: Any) -> WorkflowActivation: + job = WorkflowActivationJob() + init = job.initialize_workflow + init.workflow_type = workflow._Definition.must_from_class(workflow_class).name or "" + init.workflow_id = "wf" + init.arguments.extend( + temporalio.converter.PayloadConverter.default.to_payloads(args) + ) + return WorkflowActivation(run_id="run", jobs=[job]) + + +def _notified(*notifications: Notification) -> WorkflowActivation: + job = WorkflowActivationJob() + job.notifications_received.notifications.extend(notifications) + return WorkflowActivation(run_id="run", jobs=[job]) + + +def _signalled(name: str) -> WorkflowActivation: + job = WorkflowActivationJob() + job.signal_workflow.signal_name = name + return WorkflowActivation(run_id="run", jobs=[job]) + + +def _subscribed(completion: WorkflowActivationCompletion) -> list[str]: + assert completion.HasField("successful"), completion.failed.failure.message + return [ + command.subscribe_notification_channel.channel + for command in completion.successful.commands + if command.HasField("subscribe_notification_channel") + ] + + +def _channel_commands( + completion: WorkflowActivationCompletion, +) -> list[tuple[str, str]]: + """The channel commands of a completion in order, as (verb, channel) pairs.""" + assert completion.HasField("successful"), completion.failed.failure.message + commands: list[tuple[str, str]] = [] + for command in completion.successful.commands: + if command.HasField("subscribe_notification_channel"): + commands.append( + ("subscribe", command.subscribe_notification_channel.channel) + ) + elif command.HasField("unsubscribe_notification_channel"): + commands.append( + ("unsubscribe", command.unsubscribe_notification_channel.channel) + ) + return commands + + +def _completed(completion: WorkflowActivationCompletion) -> bool: + assert completion.HasField("successful"), completion.failed.failure.message + return any( + command.HasField("complete_workflow_execution") + for command in completion.successful.commands + ) + + +def _result(completion: WorkflowActivationCompletion) -> Any: + assert completion.HasField("successful"), completion.failed.failure.message + [done] = [ + command + for command in completion.successful.commands + if command.HasField("complete_workflow_execution") + ] + return temporalio.converter.PayloadConverter.default.from_payload( + done.complete_workflow_execution.result + ) + + +async def test_the_first_subscription_is_a_command_and_the_second_shares_it(): + instance = _instance(CountToTwo) + completion = instance.activate(_start(CountToTwo, "orders")) + assert _subscribed(completion) == ["orders"] + assert not _completed(completion) + + +async def test_a_notifications_received_job_wakes_the_receiver(): + instance = _instance(ReceiveOne) + assert _subscribed(instance.activate(_start(ReceiveOne, "orders"))) == ["orders"] + [topic] = temporalio.converter.PayloadConverter.default.to_payloads(["inputs"]) + completion = instance.activate( + _notified( + Notification( + channel="orders", position=b"7-0", counter=7, metadata={"topic": topic} + ) + ) + ) + assert _result(completion) == { + "channel": "orders", + "counter": 7, + "position": "7-0", + "topic": "inputs", + } + + +async def test_notifications_arrive_in_order_and_other_channels_are_dropped(): + instance = _instance(CountToTwo) + instance.activate(_start(CountToTwo, "orders")) + # A channel this run never subscribed to is not the workflow's concern. + completion = instance.activate(_notified(Notification(channel="other", counter=9))) + assert not _completed(completion) + completion = instance.activate( + _notified( + Notification(channel="orders", counter=1), + Notification(channel="orders", counter=2), + ) + ) + assert _result(completion) == 2 + + +async def test_an_empty_channel_name_is_refused(): + completion = _instance(EmptyChannel).activate(_start(EmptyChannel)) + assert completion.HasField("failed") + assert "channel must not be empty" in completion.failed.failure.message + + +async def test_an_empty_linked_channel_name_is_refused(): + completion = _instance(EmptyLinkedChannel).activate(_start(EmptyLinkedChannel)) + assert completion.HasField("failed") + assert "channel must not be empty" in completion.failed.failure.message + + +# --- the linked kind ---------------------------------------------------------- + + +def _linked(channel: str, counter: int, run_id: str = "run") -> Notification: + return Notification( + channel=channel, + counter=counter, + linked_to=ExecutionProto(type=WORKFLOW_TYPE, business_id="wf", run_id=run_id), + ) + + +async def test_a_linked_channel_issues_no_command_and_gets_its_notification(): + instance = _instance(ReceiveLinked) + completion = instance.activate(_start(ReceiveLinked, "orders")) + assert _subscribed(completion) == [] + assert not _completed(completion) + completion = instance.activate(_notified(_linked("orders", 7))) + assert _result(completion) == { + "channel": "orders", + "counter": 7, + "linked_to": {"type": "WORKFLOW", "business_id": "wf", "run_id": "run"}, + } + + +async def test_a_notification_without_an_owner_does_not_reach_the_linked_channel(): + instance = _instance(ReceiveLinked) + instance.activate(_start(ReceiveLinked, "orders")) + # The independent channel of the same name is somebody else's. + completion = instance.activate(_notified(Notification(channel="orders", counter=1))) + assert not _completed(completion) + assert _result(instance.activate(_notified(_linked("orders", 2))))["counter"] == 2 + + +async def test_the_two_kinds_under_one_name_are_told_apart_by_the_owner(): + instance = _instance(BothKinds) + completion = instance.activate(_start(BothKinds, "orders")) + assert _subscribed(completion) == ["orders"] + completion = instance.activate( + _notified(Notification(channel="orders", counter=1), _linked("orders", 1)) + ) + assert sorted(_result(completion)) == ["independent", "linked"] + + +async def test_a_linked_notification_nobody_asked_for_is_dropped(): + instance = _instance(ReceiveLinked) + instance.activate(_start(ReceiveLinked, "orders")) + completion = instance.activate(_notified(_linked("other", 9))) + assert completion.HasField("successful") + assert not _completed(completion) + + +# --- ending a subscription ---------------------------------------------------- + + +_CLOSED = "channel subscription closed" + + +async def test_an_unsubscribe_is_one_command_and_a_late_notification_is_dropped(): + instance = _instance(ReceiveThenUnsubscribe) + assert _channel_commands( + instance.activate(_start(ReceiveThenUnsubscribe, "orders")) + ) == [("subscribe", "orders")] + # The first notification is taken, then the two unsubscribe calls cost one + # command between them. + completion = instance.activate(_notified(Notification(channel="orders", counter=1))) + assert _channel_commands(completion) == [("unsubscribe", "orders")] + assert not _completed(completion) + # The server may still hand the run a notification it folded onto a task + # before the command landed. Nothing listens, so it changes nothing. + completion = instance.activate(_notified(Notification(channel="orders", counter=2))) + assert _channel_commands(completion) == [] + assert not _completed(completion) + assert _result(instance.activate(_signalled("finish"))) == { + "first": 1, + "closed": True, + "drained": [], + "refused": _CLOSED, + } + + +async def test_a_queued_notification_survives_the_unsubscribe_then_the_iteration_ends(): + instance = _instance(UnsubscribeWithOneQueued) + instance.activate(_start(UnsubscribeWithOneQueued, "orders")) + completion = instance.activate( + _notified( + Notification(channel="orders", counter=1), + Notification(channel="orders", counter=2), + ) + ) + assert _channel_commands(completion) == [("unsubscribe", "orders")] + assert _result(completion) == { + "first": 1, + "closed": True, + "drained": [2], + "refused": _CLOSED, + } + + +async def test_a_linked_handle_has_no_subscription_to_end(): + completion = _instance(UnsubscribeLinked).activate( + _start(UnsubscribeLinked, "orders") + ) + assert completion.HasField("failed") + assert "a linked channel has no subscription" in completion.failed.failure.message + + +async def test_a_subscription_after_an_unsubscribe_is_a_new_one(): + instance = _instance(Resubscribe) + completion = instance.activate(_start(Resubscribe, "orders")) + assert _channel_commands(completion) == [ + ("subscribe", "orders"), + ("unsubscribe", "orders"), + ("subscribe", "orders"), + ] + assert not _completed(completion) + # The notification reaches the open handle, and the closed one stays closed. + assert _result( + instance.activate(_notified(Notification(channel="orders", counter=3))) + ) == { + "counter": 3, + "closed": True, + "drained": [], + "refused": _CLOSED, + } + + +# --- the description ---------------------------------------------------------- + + +async def test_the_description_maps_every_channel_subscription_field(): + [topic] = temporalio.converter.PayloadConverter.default.to_payloads(["inputs"]) + pending = Notification(channel="orders", position=b"4-0", counter=4) + pending.metadata["topic"].CopyFrom(topic) + raw = temporalio.api.workflowservice.v1.DescribeWorkflowExecutionResponse( + workflow_execution_info=temporalio.api.workflow.v1.WorkflowExecutionInfo( + execution=temporalio.api.common.v1.WorkflowExecution( + workflow_id="wf", run_id="run" + ), + type=temporalio.api.common.v1.WorkflowType(name="ReceiveOne"), + status=temporalio.api.enums.v1.WorkflowExecutionStatus.WORKFLOW_EXECUTION_STATUS_RUNNING, + task_queue="tq", + ), + channel_subscriptions=[ + temporalio.api.workflow.v1.ChannelSubscriptionInfo( + channel="orders", + kind=temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_INDEPENDENT, + subscribed_event_id=5, + last_counter=3, + pending_notification=pending, + scheduled_counter=4, + ), + temporalio.api.workflow.v1.ChannelSubscriptionInfo( + channel="orders", + kind=temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED, + last_counter=2, + listener_count=1, + retained_count=2, + accepted_count=7, + ), + ], + ) + description = await WorkflowExecutionDescription._from_raw_description( + raw, "default", temporalio.converter.DataConverter.default + ) + assert description.id == "wf" + assert description.channel_subscriptions == ( + ChannelSubscriptionInfo( + channel="orders", + kind=ChannelKind.INDEPENDENT, + subscribed_event_id=5, + last_counter=3, + pending_notification=workflow.Notification( + channel="orders", position=b"4-0", counter=4, metadata={"topic": topic} + ), + scheduled_counter=4, + listener_count=0, + retained_count=0, + accepted_count=0, + ), + ChannelSubscriptionInfo( + channel="orders", + kind=ChannelKind.LINKED, + subscribed_event_id=0, + last_counter=2, + pending_notification=None, + scheduled_counter=0, + listener_count=1, + retained_count=2, + accepted_count=7, + ), + ) + raw.ClearField("channel_subscriptions") + description = await WorkflowExecutionDescription._from_raw_description( + raw, "default", temporalio.converter.DataConverter.default + ) + assert description.channel_subscriptions == () + + +def test_a_channel_call_names_the_execution_it_is_linked_to(): + # The short form names a workflow. The long form names any execution. + assert _channel_execution(None, None, None) is None + assert _channel_execution(None, "wf", None) == Execution.workflow("wf") + assert _channel_execution(None, "wf", "run") == Execution.workflow("wf", "run") + activity = Execution.activity("act") + assert _channel_execution(activity, None, None) is activity + with pytest.raises(ValueError, match="workflow_id"): + _channel_execution(None, None, "run") + with pytest.raises(ValueError, match="not both"): + _channel_execution(activity, "wf", None) + with pytest.raises(ValueError, match="not both"): + _channel_execution(activity, None, "run") + # Unset, the request addresses the independent channel of that name. + request = temporalio.api.workflowservice.v1.DescribeChannelRequest( + channel="c", execution=_channel_owner(None) + ) + assert not request.HasField("execution") + owner = _channel_owner(Execution.workflow("wf")) + assert owner is not None + assert (owner.type, owner.business_id, owner.run_id) == (WORKFLOW_TYPE, "wf", "") + owner = _channel_owner(Execution.workflow("wf", "run")) + assert owner is not None + assert (owner.type, owner.business_id, owner.run_id) == ( + WORKFLOW_TYPE, + "wf", + "run", + ) + assert Execution.from_proto(owner) == Execution.workflow("wf", "run") + assert Execution.from_proto(_channel_owner(activity)) == activity # type: ignore[arg-type] + assert Execution.from_proto(ExecutionProto(business_id="x")) == Execution( + ExecutionType.UNSPECIFIED, "x" + ) + + +async def test_the_client_describes_a_linked_channel_by_its_owner( + client: Client, monkeypatch: pytest.MonkeyPatch +): + described: list[temporalio.api.workflowservice.v1.DescribeChannelRequest] = [] + + async def describe(request, **kwargs): # type: ignore[no-untyped-def] + described.append(request) + return temporalio.api.workflowservice.v1.DescribeChannelResponse( + kind=temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED, + linked_to=ExecutionProto( + type=WORKFLOW_TYPE, business_id="wf", run_id="run" + ), + latest=_linked("orders", 3), + ) + + monkeypatch.setattr(client.workflow_service, "describe_channel", describe) + description = await client.describe_channel( + "orders", workflow_id="wf", run_id="run" + ) + [request] = described + assert request.channel == "orders" + assert request.execution.type == WORKFLOW_TYPE + assert request.execution.business_id == "wf" + assert request.execution.run_id == "run" + assert description.kind is ChannelKind.LINKED + assert description.linked_to == Execution.workflow("wf", "run") + assert description.latest is not None + assert _linked_to(description.latest) == { + "type": "WORKFLOW", + "business_id": "wf", + "run_id": "run", + } + # The long form carries any execution as given. + activity = Execution.activity("act", "run-2") + await client.describe_channel("orders", execution=activity) + assert Execution.from_proto(described[-1].execution) == activity + with pytest.raises(ValueError, match="not both"): + await client.describe_channel("orders", execution=activity, workflow_id="wf") + # A server that predates the kinds reports none. + monkeypatch.setattr( + client.workflow_service, + "describe_channel", + lambda request, **kwargs: _answer( # type: ignore[no-untyped-def] + temporalio.api.workflowservice.v1.DescribeChannelResponse() + ), + ) + description = await client.describe_channel("orders") + assert description.kind is ChannelKind.UNSPECIFIED + assert description.linked_to is None + + +async def _answer(response: Any) -> Any: + return response + + +async def _require_linked(client: Client) -> None: + if await server_channel_support(client) is not ChannelKind.LINKED: + pytest.skip("the server does not serve channels linked to a workflow") + + +@pytest.mark.needs_linked_server +async def test_a_workflow_receives_a_notification_on_its_linked_channel(client: Client): + await _require_linked(client) + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, ReceiveLinked) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + ReceiveLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + # The channel exists with the run: no listener to register, nothing + # retained yet, and the owner named. + description = await client.describe_channel(channel, workflow_id=handle.id) + assert description.kind is ChannelKind.LINKED + assert description.linked_to is not None + assert description.linked_to.type is ExecutionType.WORKFLOW + assert description.linked_to.business_id == handle.id + assert description.listeners == [] + assert description.latest is None + await client.notify_channel( + channel, position=b"1-0", counter=1, workflow_id=handle.id + ) + result = await asyncio.wait_for(handle.result(), 30) + assert (result["channel"], result["counter"]) == (channel, 1) + # The owner and the run that received it, so a listener holding both + # kinds under one name can route it. + assert result["linked_to"]["type"] == "WORKFLOW" + assert result["linked_to"]["business_id"] == handle.id + assert result["linked_to"]["run_id"] + events = [event.event_type async for event in handle.fetch_history_events()] + assert ( + EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_SUBSCRIBED not in events + ) + # The channel's state dies with the run. + with pytest.raises(RPCError) as closed: + await client.notify_channel(channel, counter=2, workflow_id=handle.id) + assert closed.value.status == RPCStatusCode.NOT_FOUND + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + await asyncio.wait_for(worker.shutdown(), 15) + await running + + +@pytest.mark.needs_linked_server +async def test_a_linked_channel_retains_for_pollers_and_takes_callbacks(client: Client): + await _require_linked(client) + channel = f"orders-{uuid.uuid4()}" + # Nobody polls this queue, so the run stays open for the calls below. + handle = await client.start_workflow( + ReceiveLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=f"nobody-polls-{uuid.uuid4()}", + ) + try: + callback = Callback(url="http://localhost:1/never-called", headers={}) + listener_id = await client.register_channel_listener( + channel, callback, workflow_id=handle.id + ) + description = await client.describe_channel(channel, workflow_id=handle.id) + # The owner listens by construction and the server may list it beside + # the callback once the channel holds state; the callback is the one + # listener that was registered. + assert {listener.workflow_id for listener in description.listeners} <= { + None, + handle.id, + } + [registered] = [ + listener for listener in description.listeners if listener.callback + ] + assert (registered.listener_id, registered.callback) == (listener_id, callback) + await client.unregister_channel_listener( + channel, listener_id, workflow_id=handle.id + ) + description = await client.describe_channel(channel, workflow_id=handle.id) + assert [ + listener for listener in description.listeners if listener.callback + ] == [] + # The owner is woken, so a notify counts no registered callback but + # still retains for pollers. + await client.notify_channel( + channel, position=b"2-0", counter=2, workflow_id=handle.id + ) + polled = await client.poll_channel(channel, workflow_id=handle.id, wait=False) + assert [(n.counter, _linked_to(n) is not None) for n in polled] == [(2, True)] + description = await client.describe_channel(channel, workflow_id=handle.id) + assert description.latest is not None and description.latest.counter == 2 + # The independent channel of the same name is untouched. + with pytest.raises(RPCError) as untouched: + await client.describe_channel(channel) + assert untouched.value.status == RPCStatusCode.NOT_FOUND + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + + +@pytest.mark.needs_channel_server +async def test_a_workflow_receives_a_client_notification(client: Client): + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, ReceiveOne) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + ReceiveOne.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + + async def subscribed() -> None: + # The worker's Core has to carry the subscribe command to the + # server. One that refuses it fails every completion, and the task + # times out instead; say so rather than wait on the result forever. + events = [event.event_type async for event in handle.fetch_history_events()] + if EventType.EVENT_TYPE_WORKFLOW_TASK_TIMED_OUT in events: + pytest.fail( + "the worker could not complete the task that subscribes; the Core " + "the bridge pins must carry the subscribe command" + ) + assert ( + EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_SUBSCRIBED in events + ) + + await assert_eventually(subscribed, timeout=timedelta(seconds=30)) + + async def listening() -> None: + # The channel exists once the subscribe lands, so a describe that + # races it is answered with not found. + try: + description = await client.describe_channel(channel) + except RPCError as err: + assert err.status != RPCStatusCode.NOT_FOUND, "channel not created yet" + raise + assert [listener.workflow_id for listener in description.listeners] == [ + handle.id + ] + + await assert_eventually(listening) + listeners = await client.notify_channel( + channel, position=b"1-0", counter=1, metadata={"topic": "inputs"} + ) + assert listeners == 1 + assert await asyncio.wait_for(handle.result(), 30) == { + "channel": channel, + "counter": 1, + "position": "1-0", + "topic": "inputs", + } + polled = await client.poll_channel(channel, wait=False) + assert [n.counter for n in polled] == [1] + description = await client.describe_channel(channel) + assert description.latest is not None and description.latest.counter == 1 + finally: + # A Core that refuses the subscribe command leaves the task in a + # timeout loop and the worker's shutdown waiting on it, so end the run + # first and give the shutdown a bound. + with contextlib.suppress(RPCError): + await handle.terminate() + try: + await asyncio.wait_for(worker.shutdown(), 15) + except asyncio.TimeoutError: + running.cancel() + with contextlib.suppress(asyncio.CancelledError): + await running + + +@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 == []