diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 2eef41a39..7ed7187be 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -64,6 +64,14 @@ from ._callback import ( Callback, ) +from ._channel import ( + ChannelAddress, + ChannelDescription, + ChannelKind, + ChannelListener, + ChannelSubscriptionInfo, + stream_channel, +) from ._client import ( Client, ClientConfig, @@ -106,6 +114,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -120,10 +129,13 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseActivityInput, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -137,6 +149,7 @@ TriggerScheduleInput, UnpauseActivityInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateActivityOptionsInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, @@ -315,6 +328,11 @@ "TerminateNexusOperationInput", "ListNexusOperationsInput", "CountNexusOperationsInput", + "DescribeChannelInput", + "NotifyChannelInput", + "PollChannelInput", + "RegisterChannelListenerInput", + "UnregisterChannelListenerInput", "StartWorkflowUpdateInput", "UpdateWithStartUpdateWorkflowInput", "UpdateWithStartStartWorkflowInput", @@ -351,6 +369,12 @@ "CloudOperationsClient", "Plugin", "Callback", + "ChannelDescription", + "ChannelKind", + "ChannelListener", + "ChannelAddress", + "ChannelSubscriptionInfo", + "stream_channel", "_ClientImpl", "_apply_headers", "_decode_user_metadata", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py new file mode 100644 index 000000000..f3996bd8e --- /dev/null +++ b/temporalio/client/_channel.py @@ -0,0 +1,272 @@ +"""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.streams._ref import StreamRef +from temporalio.workflow import Notification + +from ._callback import Callback + +__all__ = [ + "ChannelAddress", + "ChannelDescription", + "ChannelKind", + "ChannelListener", + "ChannelSubscriptionInfo", + "stream_channel", +] + +STREAM_CHANNEL_PREFIX = "stream/" +"""The first segment of the channel a native stream notifies.""" + + +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 workflow exists by construction, so a + describe with ``workflow_id`` 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 + + +def stream_channel(ref: StreamRef) -> ChannelAddress: + """The channel a native stream notifies on every append and on its close. + + The server derives the name from the stream's identity, and this helper + derives the same one, so a client polls or registers a callback without + asking. A stream a workflow owns notifies ``stream/`` linked to the + owning workflow. A stream an activity owns notifies + ``stream//`` linked to the workflow that scheduled + the activity, or ``stream/`` linked to the activity execution + itself when the activity is a standalone one. A standalone stream notifies + the independent channel ``stream/``, whatever the topic, since + its topics share one stream on the server. The address names the owner + without a run, so it reaches the owner's current run. + + Each change arrives as one notification: the stream's change sequence as + the counter, the head after the change as the position, and ``closed`` + set in the metadata on the close. + """ + if ref.kind == "workflow": + assert ref.workflow_id is not None + return ChannelAddress( + STREAM_CHANNEL_PREFIX + ref.topic, + temporalio.common.Execution.workflow(ref.workflow_id), + ) + if ref.kind == "activity": + assert ref.activity_id is not None + if not ref.workflow_id: + return ChannelAddress( + STREAM_CHANNEL_PREFIX + ref.topic, + temporalio.common.Execution.activity(ref.activity_id), + ) + return ChannelAddress( + f"{STREAM_CHANNEL_PREFIX}{ref.activity_id}/{ref.topic}", + temporalio.common.Execution.workflow(ref.workflow_id), + ) + assert ref.stream_id is not None + return ChannelAddress(STREAM_CHANNEL_PREFIX + ref.stream_id, None) diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 999ad8bba..63a1dace5 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -63,22 +63,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, ) @@ -114,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. @@ -2845,6 +2855,250 @@ 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. + """ + 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. + """ + 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. + """ + 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. + """ + 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. + """ + 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, @@ -3036,3 +3290,24 @@ class ClientConfig(TypedDict, total=False): ] header_codec_behavior: Required[HeaderCodecBehavior] stream_provider: temporalio.streams.StreamProvider | None + + +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/client/_workflow.py b/temporalio/client/_workflow.py index 0607ade0a..9fbc60138 100644 --- a/temporalio/client/_workflow.py +++ b/temporalio/client/_workflow.py @@ -59,6 +59,7 @@ ReturnType, SelfType, ) +from ._channel import ChannelSubscriptionInfo from ._exceptions import ( WorkflowContinuedAsNewError, WorkflowFailureError, @@ -1418,6 +1419,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 @@ -1456,6 +1469,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 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/worker/_replayer.py b/temporalio/worker/_replayer.py index 1291f5066..c7ae8d59d 100644 --- a/temporalio/worker/_replayer.py +++ b/temporalio/worker/_replayer.py @@ -290,6 +290,7 @@ def on_eviction_hook( default_workflow_logic_flags=frozenset( self._default_workflow_logic_flags ), + stream_provider=self._config.get("stream_provider"), ) external_storage = data_converter.external_storage storage_driver_types = ( diff --git a/temporalio/worker/_worker.py b/temporalio/worker/_worker.py index 25dcf697d..f4b7eddad 100644 --- a/temporalio/worker/_worker.py +++ b/temporalio/worker/_worker.py @@ -468,6 +468,11 @@ def _init_from_config(self, client: temporalio.client.Client, config: WorkerConf # Prepend applicable client interceptors to the given ones client_config = config["client"].config(active_config=True) # type: ignore[reportTypedDictNotRequiredAccess] + # A provider registered on the client serves its workers too, so one + # registration covers every context that asks for a stream. + stream_provider = config.get("stream_provider") or client_config.get( + "stream_provider" + ) interceptors_from_client = cast( list[Interceptor], [i for i in client_config["interceptors"] if isinstance(i, Interceptor)], @@ -572,6 +577,7 @@ def check_activity(activity: str): encode_headers=client_config["header_codec_behavior"] != HeaderCodecBehavior.NO_CODEC, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, + stream_provider=stream_provider, ) tuner = config.get("tuner") diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 01304d014..79f5ca6e8 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -13,6 +13,7 @@ from dataclasses import dataclass from datetime import timezone from types import TracebackType +from typing import Any import temporalio.api.common.v1 import temporalio.bridge.proto.common @@ -24,6 +25,7 @@ import temporalio.converter import temporalio.converter._extstore import temporalio.exceptions +import temporalio.streams import temporalio.workflow from temporalio.bridge.worker import PollShutdownError from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo @@ -35,6 +37,7 @@ _relax_sandbox_for_debugger, ) from ._interceptor import ( + ExecuteWorkflowInput, Interceptor, WorkflowInboundInterceptor, WorkflowInterceptorClassInput, @@ -55,6 +58,38 @@ LOG_PROTOS = False +class _StreamHooksInterceptor(WorkflowInboundInterceptor): + """Brackets the workflow function with the stream provider's lifecycle hooks. + + Installed by the worker when it has a stream provider, so no workflow + code has to call anything before it runs or before it returns. The finish + hook runs when the function returns, raises or continues as new, because + a provider that parked a reader against the run has to let go either way. + It does not run when the run is being evicted from the cache or when the + abandoned coroutine is collected: neither is the workflow ending, the + instance's state is not to be touched during eviction, and at collection + time the runtime on the thread belongs to whichever workflow happens to + be running, so the hook would act on that one. + """ + + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + runtime = temporalio.workflow._Runtime.current() + provider = runtime.workflow_streams().provider + provider.on_workflow_start() + try: + result = await self.next.execute_workflow(input) + except GeneratorExit: + raise + except BaseException: + # Eviction cancels the primary task the same way a workflow + # cancellation does, and only the cancellation is a run ending. + if not runtime.workflow_is_evicting(): + await provider.on_workflow_finish() + raise + await provider.on_workflow_finish() + return result + + # Value was chosen abitrarily as a small number that allows some concurrency and prevents # large numbers of concurrent external storage operations causing resource contention. # This default limit is per workflow task activation and does not limit the total number @@ -104,6 +139,7 @@ def __init__( encode_headers: bool, max_workflow_task_external_storage_concurrency: int, default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] | None = None, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: # Debug mode is enabled if specified or if the TEMPORAL_DEBUG env var is truthy debug_mode = debug_mode or bool(os.environ.get("TEMPORAL_DEBUG")) @@ -162,6 +198,11 @@ def __init__( __temporal_assert_local_activity_valid=assert_local_activity_valid, ) ) + self._stream_provider = stream_provider + if stream_provider is not None: + # Innermost, so the lifecycle hooks bracket the workflow function + # itself, after every user interceptor has done its own setup. + self._interceptor_classes.append(_StreamHooksInterceptor) self._workflow_failure_exception_types = workflow_failure_exception_types self._patch_activation_callback = patch_activation_callback @@ -743,6 +784,7 @@ def _create_workflow_instance( last_completion_result=init.last_completion_result, last_failure=last_failure, default_workflow_logic_flags=frozenset(self._default_workflow_logic_flags), + stream_provider=self._stream_provider, ) if defn.sandboxed: return self._workflow_runner.create_instance(det) diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 5dcc4aa96..13f2a1490 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -58,7 +58,9 @@ import temporalio.converter import temporalio.exceptions import temporalio.nexus.system +import temporalio.streams import temporalio.workflow +import temporalio.workflow._streams from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo from temporalio.converter._payload_converter import ( _TemporalTransferTypePayloadConverter, @@ -184,6 +186,7 @@ class WorkflowInstanceDetails: default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] = field( default_factory=lambda: _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS ) + stream_provider: temporalio.streams.StreamProvider | None = None class WorkflowInstance(ABC): @@ -303,6 +306,18 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: det.worker_level_failure_exception_types ) self._patch_activation_callback = det.patch_activation_callback + self._stream_provider = det.stream_provider + self._streams: temporalio.workflow._streams._WorkflowStreams | None = None + # Keyed by channel name; one subscription per channel per run + self._channel_subscriptions: dict[ + str, temporalio.workflow.ChannelSubscription + ] = {} + # The channels linked to this workflow, keyed by name as well. No + # command: the owner is the listener by construction, so the map only + # routes a notification carrying ``linked_to`` to its handle. + self._linked_channel_subscriptions: dict[ + str, temporalio.workflow.ChannelSubscription + ] = {} self._default_workflow_logic_flags = det.default_workflow_logic_flags self._primary_task: asyncio.Task[None] | None = None self._cancel_primary_task_pending = False @@ -624,6 +639,8 @@ def _apply( self._apply_query_workflow(job.query_workflow) elif job.HasField("notify_has_patch"): self._apply_notify_has_patch(job.notify_has_patch) + elif job.HasField("notifications_received"): + self._apply_notifications_received(job.notifications_received) elif job.HasField("remove_from_cache"): self._apply_remove_from_cache(job.remove_from_cache) elif job.HasField("resolve_activity"): @@ -865,6 +882,28 @@ async def run_query() -> None: # Schedule it self.create_task(run_query(), name=f"query: {job.query_type}") + def _apply_notifications_received( + self, job: temporalio.bridge.proto.workflow_activation.NotificationsReceived + ) -> None: + for proto in job.notifications: + # A name may be open as both kinds; the kind the server stamped + # on the notification picks the handle. + if proto.HasField("linked_to"): + subscription = self._linked_channel_subscriptions.get(proto.channel) + else: + subscription = self._channel_subscriptions.get(proto.channel) + if subscription is None: + # 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, + ) + continue + subscription._deliver(temporalio.workflow.Notification._from_proto(proto)) + def _apply_notify_has_patch( self, job: temporalio.bridge.proto.workflow_activation.NotifyHasPatch ) -> None: @@ -1358,6 +1397,9 @@ def workflow_instance(self) -> Any: def workflow_is_continue_as_new_suggested(self) -> bool: return self._continue_as_new_suggested + def workflow_is_evicting(self) -> bool: + return self._deleting + def workflow_is_target_worker_deployment_version_changed(self) -> bool: return self._target_worker_deployment_version_changed @@ -1810,6 +1852,49 @@ async def workflow_start_nexus_operation( ) ) + def workflow_streams(self) -> temporalio.workflow._streams._WorkflowStreams: + if self._streams is None: + if self._stream_provider is None: + raise RuntimeError( + "no stream provider is configured on this worker; pass one with " + "Worker(plugins=[provider]) or stream_provider=" + ) + # The workflow half is made per instance, so whatever it keeps + # dies with the instance the way handlers do. + self._streams = temporalio.workflow._streams._WorkflowStreams( + self._stream_provider.workflow_provider() + ) + return self._streams + + def workflow_subscribe_channel( + self, channel: str + ) -> temporalio.workflow.ChannelSubscription: + existing = self._channel_subscriptions.get(channel) + if existing is not None: + return existing + command = self._add_command() + command.subscribe_notification_channel.channel = 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) + + def workflow_linked_channel( + self, channel: str + ) -> temporalio.workflow.ChannelSubscription: + existing = self._linked_channel_subscriptions.get(channel) + if existing is not None: + return existing + subscription = temporalio.workflow.ChannelSubscription(channel, linked=True) + self._linked_channel_subscriptions[channel] = subscription + return subscription + def workflow_time_ns(self) -> int: return self._time_ns diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index c631fd88e..212625511 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, @@ -151,6 +157,12 @@ logger, unsafe, ) +from ._streams import ( + StreamReader, + StreamWriter, + stream_reader, + stream_writer, +) from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowConfig, @@ -258,6 +270,14 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "ChannelSubscription", + "Notification", + "linked_channel", + "subscribe_channel", + "StreamReader", + "StreamWriter", + "stream_reader", + "stream_writer", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py new file mode 100644 index 000000000..929f99b2d --- /dev/null +++ b/temporalio/workflow/_channels.py @@ -0,0 +1,251 @@ +"""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 with the workflow id, 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 37928afb9..a04791baf 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -22,9 +22,11 @@ if TYPE_CHECKING: from ._activities import ActivityCancellationType, ActivityHandle + from ._channels import ChannelSubscription from ._event_groups import EventGroup from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle + from ._streams import _WorkflowStreams from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowHandle, @@ -353,6 +355,16 @@ def workflow_instance(self) -> Any: ... @abstractmethod def workflow_is_continue_as_new_suggested(self) -> bool: ... + @abstractmethod + def workflow_is_evicting(self) -> bool: + """Whether this instance is being dropped from the cache rather than ending. + + Eviction cancels the primary task the way a workflow cancellation + does, so anything that runs on the way out has to be able to tell the + two apart. Instance state must not be touched while this is true. + """ + ... + @abstractmethod def workflow_is_target_worker_deployment_version_changed(self) -> bool: ... @@ -499,6 +511,24 @@ async def workflow_start_nexus_operation( event_groups: Sequence[EventGroup] | None = None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_streams(self) -> _WorkflowStreams: ... + + @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/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py new file mode 100644 index 000000000..54353508d --- /dev/null +++ b/temporalio/workflow/_streams.py @@ -0,0 +1,336 @@ +"""Streams from inside workflow code. + +.. warning:: + This module is experimental and may change in future versions. + +The reader and writer here are the same on every provider. They convert +values, synthesize supersession and buffer nothing the provider did not hand +them; everything provider-specific sits behind the ``ReadSource`` and +``WriteSink`` that the provider's workflow half opens. The provider itself +comes from the worker, the way the payload converter does. +""" + +from __future__ import annotations + +import asyncio +from collections import deque +from collections.abc import AsyncIterator, Callable +from typing import Any, Generic, TypeVar, cast, overload + +from temporalio.streams._provider import ReadSource, WorkflowStreamProvider, WriteSink +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import RecordDecoder, to_wire +from temporalio.workflow._context import _Runtime, payload_converter +from temporalio.workflow._sandbox import logger + +__all__ = ["StreamReader", "StreamWriter", "stream_reader", "stream_writer"] + +T = TypeVar("T") + + +class _WorkflowStreams: + """The stream state a workflow instance carries: its provider half and open readers.""" + + def __init__(self, provider: WorkflowStreamProvider) -> None: + self.provider = provider + self.readers: dict[str, StreamReader[Any]] = {} + # Finishing is a statement about the topic, not about the writer + # object that made it, and stream_writer() hands out a new object on + # every call. Rebuilt in order on replay, so it stays deterministic. + self.finished: set[str] = set() + + +class StreamReader(Generic[T]): + """Reads one topic of this workflow's stream from inside workflow code. + + The reader is the async iterator: ``async for record in reader`` yields + every kind of record, including the supersession the reader synthesizes + when a producer's newer attempt appears. Check ``record.kind``, or + iterate :meth:`values` when the workflow only wants data. Iteration ends + when the reader is closed or the provider ends the subscription. One loop + per reader: two loops on one reader share its buffer and interleave. + """ + + def __init__( + self, + source: ReadSource, + *, + topic: str, + result_type: type | None, + after: Cursor, + on_close: Callable[[], None], + ) -> None: + """Prefer :func:`temporalio.workflow.stream_reader`.""" + self._source = source + self._topic = topic + self._result_type = result_type + self._decoder = RecordDecoder( + payload_converter(), result_type, after=after, warn=logger.warning + ) + self._pending: deque[StreamRecord[T]] = deque() + self._lock = asyncio.Lock() + self._closed = False + self._ended = False + self._on_close = on_close + + @property + def topic(self) -> str: + """The name of the topic this reader is subscribed to.""" + return self._topic + + def __aiter__(self) -> StreamReader[T]: + """The reader is its own iterator.""" + return self + + async def __anext__(self) -> StreamRecord[T]: + """The next record, waiting for one to arrive.""" + while True: + if self._pending: + return self._pending.popleft() + if self._closed or self._ended: + raise StopAsyncIteration + await self._fill() + + async def _fill(self) -> None: + # Two loops on one reader must not race the source, so one batch + # fetch is in flight at a time and the second loop takes what the + # first one buffered. + async with self._lock: + if self._pending or self._closed or self._ended: + return + try: + batch = await self._source.next_batch() + except StopAsyncIteration: + # The provider ended the subscription. Iteration stops rather + # than raising, so a workflow that reads to the end of a + # finished stream leaves the loop instead of failing its task. + self._ended = True + return + for cursor, wire in batch: + self._pending.extend(self._decoder.decode(cursor, wire)) + + async def values(self) -> AsyncIterator[T]: + """Iterate the data values, dropping control records.""" + async for record in self: + if record.kind is RecordKind.DATA: + yield cast("T", record.value) + + def close(self) -> None: + """End the subscription. Idempotent. + + A later :func:`temporalio.workflow.stream_reader` on the same topic + opens a new subscription, which is a new command. + """ + if self._closed: + return + self._closed = True + self._source.close() + self._on_close() + + +class StreamWriter(Generic[T]): + """Publishes to one topic of this workflow's stream. + + A workflow can only publish transactionally to its own stream, on every + provider. Writing to somebody else's stream is an activity's job, and it + gets the weaker guarantee that goes with doing I/O. The type parameter is + the topic definition's value type; a writer on a string-named topic takes + any value. + """ + + def __init__(self, sink: WriteSink, topic: str, finished: set[str]) -> None: + """Prefer :func:`temporalio.workflow.stream_writer`.""" + self._sink = sink + self._topic = topic + self._finished = finished + + @property + def topic(self) -> str: + """The name of the topic this writer is bound to.""" + return self._topic + + def publish(self, value: T) -> None: + """Append ``value`` to this topic. + + Synchronous, because there is nothing to wait for inside a task: the + record becomes visible when this Workflow Task is accepted, and never + at all if the task fails, so a reader cannot see a decision the + workflow did not commit. A :class:`temporalio.common.RawValue` passes + through pre-encoded. + + Raises: + ValueError: The topic was already finished in this run, by this + writer or by another one on the same topic. + """ + if self._topic in self._finished: + raise ValueError(f"topic {self._topic!r} was already finished") + self._sink.publish( + to_wire( + payload_converter(), + topic=self._topic, + kind=RecordKind.DATA, + value=value, + ) + ) + + def finish(self) -> None: + """Write ``FINISH`` for this workflow on this topic. Idempotent. + + Says this workflow has nothing more to send on the topic. It does not + say the workflow succeeded, and it does not end anyone's read. The + marker belongs to the topic, so a second writer on the same topic in + the same run finds it already written. + """ + if self._topic in self._finished: + return + self._finished.add(self._topic) + self._sink.publish( + to_wire(payload_converter(), topic=self._topic, kind=RecordKind.FINISH) + ) + + +@overload +def stream_reader( + topic: StreamTopic[T], *, after: Cursor = ..., last: int | None = None +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str | None = None, + *, + result_type: type[T], + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str | None = None, + *, + result_type: None = None, + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[Any]: ... + + +def stream_reader( + topic: str | StreamTopic[Any] | None = None, + *, + result_type: type | None = None, + after: Cursor = BEGINNING, + last: int | None = None, +) -> StreamReader[Any]: + """Subscribe this workflow to ``topic`` of its own stream. + + ``topic`` is a :func:`temporalio.streams.topic` definition, which carries + the record type, or a plain string with ``result_type=`` for a name + decided at runtime, and without one the reader is on + :data:`temporalio.streams.DEFAULT_TOPIC`, decoded as a string-named topic + is. One subscription per topic per run. A second call + for the same topic returns the reader already open on it, so records go + to whichever loop pulls first; such a call may pass no ``after``, no + ``last`` and no different type. Adding a reader on a new topic is a new command, so + gate it with :func:`temporalio.workflow.patched` as you would a timer. A + reader in a successor run starts a new subscription: nothing crosses + continue-as-new implicitly. + + Args: + topic: The topic, relative to this workflow's stream. Omit it for + the default topic. + result_type: The value type for a string-named or default topic, used as the + decode hint. :class:`temporalio.common.RawValue` returns the + payload untouched. + after: Resume strictly after this record. ``BEGINNING`` starts at + the oldest record the topic still holds and + :data:`temporalio.streams.END` at whatever is appended after the + subscription is registered. Honoured on the first subscription + of a run, because after that the recorded observations decide. + last: Start at the newest ``last`` records instead, or at all of + them when there are fewer. Records of every kind count. Exclusive + with a cursor in ``after``. Where it lands is resolved once, + outside the workflow, and replay reproduces it. + + Raises: + ValueError: ``topic`` is empty, ``result_type`` was passed with a + definition, ``last`` is not positive or came with a cursor, or a + reader on the topic is already open and this call asked for a + different position or type. + temporalio.streams.StreamCursorError: ``after`` was minted by another + provider. + temporalio.streams.StreamUnsupportedError: The provider cannot start + where ``END`` or ``last`` asks. + """ + check_read_start(after, last) + name, result_type = resolve_topic(topic, result_type) + state: _WorkflowStreams = _Runtime.current().workflow_streams() + existing = state.readers.get(name) + if existing is not None: + if ( + after != BEGINNING + or last is not None + or result_type is not existing._result_type + ): + raise ValueError( + f"topic {name!r} already has a reader in this run; a second " + "stream_reader shares it and takes no after=, last= or other type" + ) + return existing + # Passed only when given, so a provider written before last= existed + # still serves every read it can. + source = ( + state.provider.open_reader(name, after=after) + if last is None + else state.provider.open_reader(name, after=after, last=last) + ) + + def forget() -> None: + state.readers.pop(name, None) + + # The decoder positions a synthesized record at the one before it. Where + # END or last= lands is not known here, so BEGINNING stands in: resuming + # from it may deliver a record twice, where END would skip one. + previous = BEGINNING if after == END or last is not None else after + reader: StreamReader[Any] = StreamReader( + source, topic=name, result_type=result_type, after=previous, on_close=forget + ) + state.readers[name] = reader + return reader + + +@overload +def stream_writer(topic: StreamTopic[T]) -> StreamWriter[T]: ... + + +@overload +def stream_writer(topic: str | None = None) -> StreamWriter[Any]: ... + + +def stream_writer(topic: str | StreamTopic[Any] | None = None) -> StreamWriter[Any]: + """Publish to ``topic`` of this workflow's stream, or to its default topic. + + Every call returns a new writer, and they all share the run's record of + which topics were finished, so ``finish()`` on one is seen by the next. + + Args: + topic: A :func:`temporalio.streams.topic` definition, whose value type + the writer's ``publish`` takes, or a plain string for a name + decided at runtime. Omitted, the writer is on + :data:`temporalio.streams.DEFAULT_TOPIC`. Encoding follows each + published value. + + Raises: + ValueError: ``topic`` is empty. + """ + name, _ = resolve_topic(topic) + state = _Runtime.current().workflow_streams() + return StreamWriter(state.provider.open_writer(name), name, state.finished) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index b8a152f20..049fa2da1 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -1,5 +1,23 @@ import pytest +_ENVIRONMENTS_WITHOUT_CHANNELS = ("local", "time-skipping", "envconfig") + +# The Core the bridge pins decides whether a workflow's subscribe command +# reaches the server. The protos-only pin refuses it; the delivery pin that +# the native layers move to handles it. +PINNED_CORE_HANDLES_CHANNEL_COMMAND = False + +# A channel linked to a workflow needs no command, but a notification reaches +# workflow code as the `NotificationsReceived` job Core builds from the +# scheduled event, and the protos-only pin ignores that job. So a linked case +# in which the workflow receives waits for the delivery pin as well; one that +# only talks to the server from the client runs on every layer. +PINNED_CORE_HANDLES_LINKED_CHANNEL = False + +# The unsubscribe is a command like the subscribe, refused by the protos-only +# pin and matched against its event by the delivery pin. +PINNED_CORE_HANDLES_UNSUBSCRIBE = False + def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( @@ -15,3 +33,127 @@ def pytest_configure(config: pytest.Config) -> None: "markers", "truncates: the case needs a way to drop a topic's oldest records", ) + 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_channel_core: the case needs the pinned Core to handle the " + "subscribe-notification-channel command", + ) + config.addinivalue_line( + "markers", + "needs_linked_server: the case needs a server that serves channels linked " + "to a workflow, named with -E host:port", + ) + config.addinivalue_line( + "markers", + "needs_linked_core: the case needs the pinned Core to hand a linked " + "channel's notifications to workflow code", + ) + config.addinivalue_line( + "markers", + "needs_describe_server: the case needs a server whose workflow description " + "lists the channel subscriptions, named with -E host:port", + ) + config.addinivalue_line( + "markers", + "needs_unsubscribe_server: the case needs a server that accepts the " + "unsubscribe-notification-channel command, named with -E host:port", + ) + config.addinivalue_line( + "markers", + "needs_unsubscribe_core: the case needs the pinned Core to handle the " + "unsubscribe-notification-channel command", + ) + config.addinivalue_line( + "markers", + "needs_stream_channel_server: the case needs a server on which a native " + "stream notifies the channel named by the stream, named with -E host:port", + ) + config.addinivalue_line( + "markers", + "needs_execution_server: the case needs a server that addresses a linked " + "channel by execution, a standalone activity's included, 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" + ), + ) + ) + if not PINNED_CORE_HANDLES_CHANNEL_COMMAND: + skips.append( + ( + "needs_channel_core", + pytest.mark.skip( + reason="the pinned Core refuses the subscribe command; py-05 " + "pins one that handles it" + ), + ) + ) + if config.getoption("--workflow-environment") in _ENVIRONMENTS_WITHOUT_CHANNELS: + skips.append( + ( + "needs_linked_server", + pytest.mark.skip( + reason="needs a server with channels linked to a workflow; " + "name one with -E" + ), + ) + ) + if not PINNED_CORE_HANDLES_LINKED_CHANNEL: + skips.append( + ( + "needs_linked_core", + pytest.mark.skip( + reason="the pinned Core ignores the notifications job; py-05 " + "pins one that delivers it" + ), + ) + ) + if config.getoption("--workflow-environment") in _ENVIRONMENTS_WITHOUT_CHANNELS: + for marker, what in ( + ("needs_describe_server", "lists channel subscriptions on describe"), + ("needs_unsubscribe_server", "accepts the unsubscribe command"), + ("needs_stream_channel_server", "notifies a stream's channel"), + ("needs_execution_server", "addresses a linked channel by execution"), + ): + skips.append( + ( + marker, + pytest.mark.skip( + reason=f"needs a server that {what}; name one with -E" + ), + ) + ) + if not PINNED_CORE_HANDLES_UNSUBSCRIBE: + skips.append( + ( + "needs_unsubscribe_core", + pytest.mark.skip( + reason="the pinned Core refuses the unsubscribe command; py-05 " + "pins one that handles it" + ), + ) + ) + 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..afcdc4a44 --- /dev/null +++ b/tests/streams/test_channels.py @@ -0,0 +1,1272 @@ +"""The notification channel surface: the command, the delivery and the client calls. + +Both kinds of channel are covered: the independent one a workflow subscribes +to by command, and the one linked to the workflow, which needs none. 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, or one +with the linked kind, one whose describe lists the subscriptions, or one that +accepts the unsubscribe, and skip otherwise. +""" + +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, + ChannelAddress, + ChannelKind, + ChannelSubscriptionInfo, + Client, + WorkflowExecutionDescription, + stream_channel, +) +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import DEFAULT_TOPIC, StreamRef +from temporalio.worker._workflow_instance import ( + UnsandboxedWorkflowRunner, + WorkflowInstance, + WorkflowInstanceDetails, +) +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 + + +@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("") + + +def _describe(notification: workflow.Notification) -> dict[str, Any]: + 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 + ), + "owner": ( + notification.linked_to.business_id if notification.linked_to else None + ), + "owner_run": notification.linked_to.run_id if notification.linked_to else None, + } + + +@workflow.defn +class ReceiveLinked: + """Listens on its linked channel and reports the first notification.""" + + @workflow.run + async def run(self, channel: str) -> dict[str, Any]: + handle = workflow.linked_channel(channel) + assert handle is workflow.linked_channel(channel) + assert handle.linked + return _describe(await handle.receive()) + + +@workflow.defn +class BothKinds: + """Holds both kinds of handle on one name and keeps what each receives. + + Ends once the linked handle has seen counter two. + """ + + @workflow.run + async def run(self, channel: str) -> dict[str, list[int]]: + independent = workflow.subscribe_channel(channel) + linked = workflow.linked_channel(channel) + assert not independent.linked and linked.linked + seen: dict[str, list[int]] = {"independent": [], "linked": []} + + async def collect_independent() -> None: + async for notification in independent: + assert notification.linked_to is None + seen["independent"].append(notification.counter) + + collector = asyncio.create_task(collect_independent()) + async for notification in linked: + assert notification.linked_to is not None + seen["linked"].append(notification.counter) + if notification.counter >= 2: + break + collector.cancel() + return seen + + +@workflow.defn +class CountLinked: + """Counts the notifications on its linked channel up to counter two.""" + + @workflow.run + async def run(self, channel: str) -> int: + seen = 0 + async for notification in workflow.linked_channel(channel): + seen += 1 + if notification.counter >= 2: + break + 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", + original_execution_run_id="run", + 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_internal_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 _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(), + ) + + +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 + + +async def test_a_linked_channel_issues_no_command_and_gets_its_own_notifications(): + instance = _instance(ReceiveLinked) + completion = instance.activate(_start(ReceiveLinked, "orders")) + assert completion.HasField("successful"), completion.failed.failure.message + assert list(completion.successful.commands) == [] + # Without an owner the notification is the independent channel's, which + # this run never subscribed to. + completion = instance.activate(_notified(Notification(channel="orders", counter=1))) + assert not _completed(completion) + [topic] = temporalio.converter.PayloadConverter.default.to_payloads(["inputs"]) + notification = _linked("orders", 7, b"7-0") + notification.metadata["topic"].CopyFrom(topic) + assert _result(instance.activate(_notified(notification))) == { + "channel": "orders", + "counter": 7, + "position": "7-0", + "topic": "inputs", + "owner": "wf", + "owner_run": "run", + } + + +async def test_the_owner_on_a_notification_picks_the_handle_of_its_kind(): + instance = _instance(BothKinds) + completion = instance.activate(_start(BothKinds, "orders")) + # Only the independent handle costs a command. + assert _subscribed(completion) == ["orders"] + assert len(completion.successful.commands) == 1 + assert not _completed(instance.activate(_notified(_linked("orders", 1)))) + assert not _completed( + instance.activate(_notified(Notification(channel="orders", counter=5))) + ) + completion = instance.activate(_notified(_linked("orders", 2))) + assert _result(completion) == {"independent": [5], "linked": [1, 2]} + + +_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, + } + + +async def test_a_stream_names_the_channel_it_notifies(): + owner = temporalio.common.Execution.workflow("wf") + assert stream_channel(StreamRef.for_workflow("wf", topic="out")) == ChannelAddress( + "stream/out", owner + ) + # The address names the owner without a run, so it follows the chain. + assert stream_channel(StreamRef.for_workflow("wf", run_id="run")) == ChannelAddress( + "stream/" + DEFAULT_TOPIC, owner + ) + assert stream_channel( + StreamRef.for_activity("act", workflow_id="wf", topic="out") + ) == ChannelAddress("stream/act/out", owner) + # A standalone activity is an execution of its own, so its stream's + # channel is linked to it under the topic's name alone. + assert stream_channel(StreamRef.for_activity("act", topic="out")) == ChannelAddress( + "stream/out", temporalio.common.Execution.activity("act") + ) + # A standalone stream's topics share one stream on the server, so the + # topic is not part of the name. + assert stream_channel( + StreamRef.for_standalone("sid", topic="out") + ) == ChannelAddress("stream/sid", None) + # The workflow id stays readable for a caller that addresses by it, and + # only names a workflow. + assert stream_channel(StreamRef.for_workflow("wf")).workflow_id == "wf" + assert stream_channel(StreamRef.for_activity("act")).workflow_id is None + assert stream_channel(StreamRef.for_standalone("sid")).workflow_id is None + + +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 == () + + +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 +@pytest.mark.needs_channel_core +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 == [] + + +async def _stop(handle: Any, worker: Any, running: asyncio.Task[None]) -> None: + """End the run and the worker, with a bound on the shutdown.""" + 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_linked_server +@pytest.mark.needs_linked_core +async def test_a_workflow_receives_a_notification_on_its_linked_channel( + client: 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 registers, nothing is + # retained yet, and the owner is named. + description = await client.describe_channel(channel, workflow_id=handle.id) + assert description.kind == ChannelKind.LINKED + assert description.listeners == [] + assert description.latest is None + assert description.retained_count == 0 + assert description.linked_to is not None + assert description.linked_to.business_id == handle.id + # The owner is the one listener. + listeners = await client.notify_channel( + channel, + position=b"1-0", + counter=1, + metadata={"topic": "inputs"}, + workflow_id=handle.id, + ) + assert listeners == 1 + assert await asyncio.wait_for(handle.result(), 30) == { + "channel": channel, + "counter": 1, + "position": "1-0", + "topic": "inputs", + "owner": handle.id, + "owner_run": handle.first_execution_run_id, + } + # The owner listens by construction, so History holds no subscribe + # event; the notification rode a scheduled event. + events = [event.event_type async for event in handle.fetch_history_events()] + assert ( + EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_SUBSCRIBED not in events + ) + finally: + await _stop(handle, worker, running) + + +@pytest.mark.needs_linked_server +async def test_a_linked_channel_lives_and_dies_with_its_workflow(client: Client): + """The client side alone: no worker polls, so the run stays open until ended.""" + channel = f"orders-{uuid.uuid4()}" + handle = await client.start_workflow( + CountLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=f"nobody-polls-{uuid.uuid4()}", + ) + run_id = handle.first_execution_run_id + assert run_id is not None + # A name nobody has notified exists all the same, with nothing in it. + description = await client.describe_channel(channel, workflow_id=handle.id) + assert description.kind == ChannelKind.LINKED + assert (description.listeners, description.latest) == ([], None) + assert description.retained_count == 0 + assert description.linked_to is not None + assert (description.linked_to.business_id, description.linked_to.run_id) == ( + handle.id, + run_id, + ) + # The independent channel of that name is a different thing and does not + # exist. + with pytest.raises(RPCError) as independent: + await client.describe_channel(channel) + assert independent.value.status == RPCStatusCode.NOT_FOUND + # A run id names that run; one that is not the chain's is not found, and + # neither is a workflow that never ran. + description = await client.describe_channel( + channel, workflow_id=handle.id, run_id=run_id + ) + assert description.kind == ChannelKind.LINKED + for wrong in ( + client.describe_channel( + channel, workflow_id=handle.id, run_id=str(uuid.uuid4()) + ), + client.notify_channel( + channel, counter=1, workflow_id=handle.id, run_id=str(uuid.uuid4()) + ), + client.notify_channel(channel, counter=1, workflow_id=f"never-{uuid.uuid4()}"), + ): + with pytest.raises(RPCError) as missing: + await wrong + assert missing.value.status == RPCStatusCode.NOT_FOUND + # The channel ends with the run. + await handle.terminate() + with pytest.raises(RPCError) as closed: + await client.notify_channel(channel, counter=1, workflow_id=handle.id) + assert closed.value.status == RPCStatusCode.NOT_FOUND + + +@pytest.mark.needs_linked_server +@pytest.mark.needs_execution_server +async def test_a_linked_channel_names_its_owner_as_an_execution(client: Client): + """The client side alone: the owner comes back typed, by either spelling.""" + channel = f"orders-{uuid.uuid4()}" + handle = await client.start_workflow( + CountLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=f"nobody-polls-{uuid.uuid4()}", + ) + run_id = handle.first_execution_run_id + assert run_id is not None + by_id = temporalio.common.Execution.workflow(handle.id) + by_run = temporalio.common.Execution.workflow(handle.id, run_id) + try: + description = await client.describe_channel(channel, execution=by_id) + assert description.kind == ChannelKind.LINKED + assert description.linked_to is not None + assert description.linked_to == by_run + assert description.linked_to.type is temporalio.common.ExecutionType.WORKFLOW + # The execution and the workflow id shorthand reach one channel. + assert ( + await client.notify_channel( + channel, position=b"1-0", counter=1, execution=by_id + ) + == 1 + ) + [polled] = await client.poll_channel(channel, workflow_id=handle.id, wait=False) + assert (polled.counter, polled.linked_to) == (1, by_run) + [polled] = await client.poll_channel(channel, execution=by_run, wait=False) + assert polled.counter == 1 + # The same id as an activity names an execution that never ran. + with pytest.raises(RPCError) as missing: + await client.describe_channel( + channel, execution=temporalio.common.Execution.activity(handle.id) + ) + assert missing.value.status == RPCStatusCode.NOT_FOUND + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + + +@pytest.mark.needs_linked_server +@pytest.mark.needs_linked_core +async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, CountLinked) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + CountLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + assert ( + await client.notify_channel( + channel, position=b"1-0", counter=1, workflow_id=handle.id + ) + == 1 + ) + polled = await client.poll_channel(channel, workflow_id=handle.id, wait=False) + assert [(n.position, n.counter) for n in polled] == [(b"1-0", 1)] + assert polled[0].linked_to is not None + assert polled[0].linked_to.business_id == handle.id + # The run id reaches the same channel. + assert [ + n.counter + for n in await client.poll_channel( + channel, + workflow_id=handle.id, + run_id=handle.first_execution_run_id, + wait=False, + ) + ] == [1] + description = await client.describe_channel(channel, workflow_id=handle.id) + assert description.kind == ChannelKind.LINKED + assert description.latest is not None and description.latest.counter == 1 + assert description.retained_count == 1 + # A poll above the latest waits its bound out and comes back empty. + polled = await client.poll_channel( + channel, workflow_id=handle.id, after_counter=1, wait=timedelta(seconds=1) + ) + assert polled == [] + assert ( + await client.notify_channel( + channel, position=b"2-0", counter=2, workflow_id=handle.id + ) + == 1 + ) + assert await asyncio.wait_for(handle.result(), 30) == 2 + # The ring went with the run, so a poll after the close finds nothing + # to read. + with pytest.raises(RPCError) as closed: + await client.poll_channel( + channel, workflow_id=handle.id, after_counter=1, wait=False + ) + assert closed.value.status == RPCStatusCode.NOT_FOUND + finally: + await _stop(handle, worker, running) + + +_SUBSCRIBED = EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_SUBSCRIBED +_UNSUBSCRIBED = EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_UNSUBSCRIBED + + +async def _event_ids(handle: Any, event_type: Any) -> list[int]: + """The ids of the events of ``event_type`` in the run's History so far.""" + return [ + event.event_id + async for event in handle.fetch_history_events() + if event.event_type == event_type + ] + + +async def _one_event(handle: Any, event_type: Any) -> int: + """The id of the one event of ``event_type``, failing until it is there.""" + ids = await _event_ids(handle, event_type) + assert len(ids) == 1, ids + return ids[0] + + +@pytest.mark.needs_describe_server +@pytest.mark.needs_channel_core +async def test_a_description_lists_an_independent_subscription(client: Client): + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, CountToTwo) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + CountToTwo.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + event_id = await assert_eventually( + lambda: _one_event(handle, _SUBSCRIBED), timeout=timedelta(seconds=30) + ) + # Listed from the subscribe event on, with nothing accepted yet. The + # counts belong to the channel execution and stay zero here. + description = await handle.describe() + assert description.channel_subscriptions == ( + ChannelSubscriptionInfo( + channel=channel, + kind=ChannelKind.INDEPENDENT, + subscribed_event_id=event_id, + last_counter=0, + pending_notification=None, + scheduled_counter=0, + listener_count=0, + retained_count=0, + accepted_count=0, + ), + ) + assert await client.notify_channel(channel, position=b"1-0", counter=1) == 1 + + async def accepted() -> None: + # Once the task that carried it completes, the counter is the + # run's and nothing is pending or scheduled any more. + [info] = (await handle.describe()).channel_subscriptions + assert info.last_counter == 1 + assert info.pending_notification is None + assert info.scheduled_counter == 0 + + await assert_eventually(accepted) + assert await client.notify_channel(channel, position=b"2-0", counter=2) == 1 + assert await asyncio.wait_for(handle.result(), 30) == 2 + # A closed run keeps listing what it stood on. + [info] = (await handle.describe()).channel_subscriptions + assert (info.kind, info.last_counter) == (ChannelKind.INDEPENDENT, 2) + finally: + await _stop(handle, worker, running) + + +@pytest.mark.needs_describe_server +async def test_a_description_lists_a_linked_channel_once_it_holds_state( + client: Client, +): + """The client side alone: no worker polls, so the run stays open until ended.""" + channel = f"orders-{uuid.uuid4()}" + handle = await client.start_workflow( + CountLinked.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=f"nobody-polls-{uuid.uuid4()}", + ) + try: + # An untouched linked name exists by construction and holds nothing, + # so it is not listed. + assert (await handle.describe()).channel_subscriptions == () + assert ( + await client.notify_channel( + channel, position=b"1-0", counter=1, workflow_id=handle.id + ) + == 1 + ) + [info] = (await handle.describe()).channel_subscriptions + assert (info.channel, info.kind) == (channel, ChannelKind.LINKED) + # The owner's state took the notification in the write that accepted + # it, so the counter is the run's at once. Nobody polls, and the first + # task was scheduled without a counter when the run started, so the + # notification waits behind it as the pending entry. + assert (info.subscribed_event_id, info.last_counter) == (0, 1) + assert info.pending_notification is not None + assert info.pending_notification.counter == 1 + assert info.pending_notification.linked_to is not None + assert info.pending_notification.linked_to.business_id == handle.id + assert info.scheduled_counter == 0 + assert (info.listener_count, info.retained_count, info.accepted_count) == ( + 0, + 1, + 1, + ) + # A callback on the linked channel shows up in the owner's count. + callback = Callback(url="http://localhost:1/never-called", headers={}) + listener_id = await client.register_channel_listener( + channel, callback, workflow_id=handle.id + ) + [info] = (await handle.describe()).channel_subscriptions + assert info.listener_count == 1 + await client.unregister_channel_listener( + channel, listener_id, workflow_id=handle.id + ) + [info] = (await handle.describe()).channel_subscriptions + assert info.listener_count == 0 + # A closed run keeps listing what it stood on. + await handle.terminate() + [info] = (await handle.describe()).channel_subscriptions + assert (info.kind, info.last_counter) == (ChannelKind.LINKED, 1) + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + + +@pytest.mark.needs_unsubscribe_server +@pytest.mark.needs_unsubscribe_core +async def test_a_workflow_unsubscribes_and_a_later_notify_wakes_nothing( + client: Client, +): + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, ReceiveThenUnsubscribe) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + ReceiveThenUnsubscribe.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + subscribed_id = await assert_eventually( + lambda: _one_event(handle, _SUBSCRIBED), timeout=timedelta(seconds=30) + ) + assert await client.notify_channel(channel, position=b"1-0", counter=1) == 1 + unsubscribed_id = await assert_eventually( + lambda: _one_event(handle, _UNSUBSCRIBED), timeout=timedelta(seconds=30) + ) + assert unsubscribed_id > subscribed_id + # The event names the subscription it ended. + [event] = [ + event + async for event in handle.fetch_history_events() + if event.event_id == unsubscribed_id + ] + attrs = event.workflow_notification_channel_unsubscribed_event_attributes + assert (attrs.channel, attrs.subscribed_event_id) == (channel, subscribed_id) + # Gone from both sides: the channel's listeners and the run's standing. + description = await client.describe_channel(channel) + assert [listener.workflow_id for listener in description.listeners] == [] + assert (await handle.describe()).channel_subscriptions == () + # Nothing listens any more, so a notify wakes nobody and is only + # retained for pollers. + assert await client.notify_channel(channel, position=b"2-0", counter=2) == 0 + await handle.signal(ReceiveThenUnsubscribe.finish) + assert await asyncio.wait_for(handle.result(), 30) == { + "first": 1, + "closed": True, + "drained": [], + "refused": _CLOSED, + } + assert await _event_ids(handle, _SUBSCRIBED) == [subscribed_id] + assert await _event_ids(handle, _UNSUBSCRIBED) == [unsubscribed_id] + finally: + await _stop(handle, worker, running) diff --git a/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py new file mode 100644 index 000000000..a547daf3e --- /dev/null +++ b/tests/streams/test_stream_hooks.py @@ -0,0 +1,157 @@ +"""The lifecycle interceptor's rules about when the finish hook runs. + +Driven directly, with a stand-in runtime on the loop, because the two cases +that matter here are the ones a workflow test cannot stage on purpose: a +coroutine collected after its worker went away, and a run evicted from the +cache. Both used to reach the finish hook, and at collection time the hook +acted on whichever workflow was running on the thread. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.worker._interceptor import ( + ExecuteWorkflowInput, + WorkflowInboundInterceptor, +) +from temporalio.worker._workflow import _StreamHooksInterceptor +from temporalio.worker._workflow_instance import _WorkflowInstanceImpl + + +class _Provider: + def __init__(self) -> None: + self.calls: list[str] = [] + + def on_workflow_start(self) -> None: + self.calls.append("start") + + async def on_workflow_finish(self) -> None: + self.calls.append("finish") + + +class _Streams: + def __init__(self, provider: _Provider) -> None: + self.provider = provider + + +class _FakeRuntime: + """Only what the interceptor reads, and only through the runtime interface. + + It deliberately carries no ``_deleting`` attribute. An interceptor that + reads the eviction state by attribute name instead of by method would see + a runtime that is never evicting here, and the eviction case below would + catch it. + """ + + def __init__(self, provider: _Provider, *, evicting: bool = False) -> None: + self._streams = _Streams(provider) + self._evicting = evicting + + def workflow_streams(self) -> _Streams: + return self._streams + + def workflow_is_evicting(self) -> bool: + return self._evicting + + +class _Body(WorkflowInboundInterceptor): + """The workflow function's stand-in: returns, raises or parks forever.""" + + def __init__(self, outcome: Any) -> None: # type: ignore[reportMissingSuperCall] + self._outcome = outcome + + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + del input + if self._outcome is _PARK: + await asyncio.Event().wait() + if isinstance(self._outcome, BaseException): + raise self._outcome + return self._outcome + + +_PARK = object() + + +async def _unused_run_fn() -> None: + pass + + +_INPUT = ExecuteWorkflowInput(type=object, run_fn=_unused_run_fn, args=(), headers={}) + + +@pytest.fixture +async def provider() -> Any: + fake = _Provider() + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake)) # type: ignore[arg-type] + yield fake + workflow._Runtime.set_on_loop(loop, None) + + +def _evicting(fake: _Provider) -> None: + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake, evicting=True)) # type: ignore[arg-type] + + +async def test_the_finish_hook_runs_on_return(provider: _Provider): + assert ( + await _StreamHooksInterceptor(_Body("done")).execute_workflow(_INPUT) == "done" + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_when_the_function_raises(provider: _Provider): + with pytest.raises(RuntimeError, match="boom"): + await _StreamHooksInterceptor(_Body(RuntimeError("boom"))).execute_workflow( + _INPUT + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_on_continue_as_new(provider: _Provider): + error = workflow.ContinueAsNewError.__new__(workflow.ContinueAsNewError) + with pytest.raises(workflow.ContinueAsNewError): + await _StreamHooksInterceptor(_Body(error)).execute_workflow(_INPUT) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_on_a_workflow_cancellation(provider: _Provider): + # A cancelled primary task is the run ending, so the provider lets go. + with pytest.raises(asyncio.CancelledError): + await _StreamHooksInterceptor(_Body(asyncio.CancelledError())).execute_workflow( + _INPUT + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_does_not_run_when_the_coroutine_is_collected( + provider: _Provider, +): + coroutine = _StreamHooksInterceptor(_Body(_PARK)).execute_workflow(_INPUT) + # Run up to the park, the way a worker that shut down without evicting + # leaves the primary task, then close it as garbage collection would. + coroutine.send(None) + coroutine.close() + assert provider.calls == ["start"] + + +async def test_the_finish_hook_does_not_run_during_eviction(provider: _Provider): + _evicting(provider) + with pytest.raises(asyncio.CancelledError): + await _StreamHooksInterceptor(_Body(asyncio.CancelledError())).execute_workflow( + _INPUT + ) + assert provider.calls == ["start"] + + +def test_the_runtime_answers_the_eviction_question_itself(): + # The interceptor asks the runtime rather than reading a private field, + # so the coupling is declared on the base class and the real instance has + # to answer it or fail to construct. + assert "workflow_is_evicting" in workflow._Runtime.__abstractmethods__ + assert "workflow_is_evicting" not in _WorkflowInstanceImpl.__abstractmethods__ diff --git a/tests/streams/test_stream_reader.py b/tests/streams/test_stream_reader.py new file mode 100644 index 000000000..d3918fe40 --- /dev/null +++ b/tests/streams/test_stream_reader.py @@ -0,0 +1,178 @@ +"""The reader's rules that a workflow test cannot pin down. + +``stream_reader`` promises one subscription per topic per run, that a second +loop on one reader shares its buffer rather than racing the source, and that +closing it and opening the topic again is a new subscription. The memory +provider's source hands back everything it has without ever suspending, so a +workflow test cannot tell a reader that serialises its fetches from one that +does not. These drive the reader directly, with a source that suspends where +a real provider would, and a stand-in runtime on the loop. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.converter import DataConverter +from temporalio.streams import BEGINNING, Cursor, RecordKind, topic +from temporalio.streams._wire import WireRecord, to_wire +from temporalio.workflow._streams import _WorkflowStreams + +INPUTS = topic("inputs", dict) + +_CONVERTER = DataConverter.default.payload_converter + + +def _record(n: int) -> WireRecord: + return to_wire( + _CONVERTER, + topic=INPUTS.name, + kind=RecordKind.DATA, + value={"n": n}, + producer_id="model", + attempt=1, + sequence=n, + ) + + +class _SlowSource: + """Hands over one record per batch, suspending first the way a real one does.""" + + def __init__(self, count: int) -> None: + self.opened = True + self.next_calls = 0 + self.in_flight = 0 + self.overlapped = False + self._offset = 0 + self._count = count + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + self.next_calls += 1 + self.in_flight += 1 + if self.in_flight > 1: + self.overlapped = True + try: + # Where a provider waits for the worker to deliver. Two fetches + # that both get here have both left the reader's buffer behind. + await asyncio.sleep(0) + await asyncio.sleep(0) + if self._offset >= self._count: + raise StopAsyncIteration + self._offset += 1 + return [(Cursor(f"fake:{self._offset - 1}"), _record(self._offset - 1))] + finally: + self.in_flight -= 1 + + def close(self) -> None: + self.opened = False + + +class _Half: + """A workflow half that hands out one source per topic and counts opens.""" + + def __init__(self) -> None: + self.sources: list[_SlowSource] = [] + + def open_reader(self, topic: str, *, after: Cursor) -> Any: + del topic, after + source = _SlowSource(4) + self.sources.append(source) + return source + + def open_writer(self, topic: str) -> Any: + del topic + raise NotImplementedError + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +class _FakeRuntime: + """Only what the reader reads: the stream state and the payload converter.""" + + def __init__(self, half: _Half) -> None: + self._streams = _WorkflowStreams(half) # type: ignore[arg-type] + + def workflow_streams(self) -> _WorkflowStreams: + return self._streams + + def workflow_payload_converter(self) -> Any: + return _CONVERTER + + +@pytest.fixture +async def half() -> Any: + fake = _Half() + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake)) # type: ignore[arg-type] + yield fake + workflow._Runtime.set_on_loop(loop, None) + + +async def test_two_loops_on_one_reader_split_the_records(half: _Half): + reader = workflow.stream_reader(INPUTS) + seen: list[int] = [] + + async def pull(count: int) -> None: + for _ in range(count): + record = await reader.__anext__() + assert record.value is not None + seen.append(record.value["n"]) + + await asyncio.wait_for(asyncio.gather(pull(2), pull(2)), 5.0) + # Every record once and none lost, whichever loop got there first, and + # the source was never asked for two batches at the same time. + assert sorted(seen) == [0, 1, 2, 3] + assert half.sources[0].overlapped is False + + +@pytest.mark.usefixtures("half") +async def test_a_cancelled_read_does_not_lose_the_record_it_waited_for(): + reader = workflow.stream_reader(INPUTS) + pending = asyncio.ensure_future(reader.__anext__()) + await asyncio.sleep(0) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + + # The next read picks up where the cancelled one was, rather than finding + # the reader wedged behind a lock the cancelled fetch still holds. + record = await asyncio.wait_for(reader.__anext__(), 5.0) + assert record.value == {"n": 0} + + +async def test_closing_a_reader_lets_the_topic_be_opened_again(half: _Half): + reader = workflow.stream_reader(INPUTS) + assert workflow.stream_reader(INPUTS) is reader + reader.close() + assert half.sources[0].opened is False + + # A new subscription, not the closed one handed back. It is also a new + # command on a real provider, which is why the docstring says to gate it. + again = workflow.stream_reader(INPUTS) + assert again is not reader + assert len(half.sources) == 2 + assert (await asyncio.wait_for(again.__anext__(), 5.0)).value == {"n": 0} + + +@pytest.mark.usefixtures("half") +async def test_a_closed_reader_stops_iterating(): + reader = workflow.stream_reader(INPUTS) + assert (await asyncio.wait_for(reader.__anext__(), 5.0)).value == {"n": 0} + reader.close() + with pytest.raises(StopAsyncIteration): + await asyncio.wait_for(reader.__anext__(), 5.0) + + +@pytest.mark.usefixtures("half") +async def test_a_reader_ends_when_the_source_ends(): + reader = workflow.stream_reader(INPUTS, after=BEGINNING) + values = [record.value async for record in reader] + assert values == [{"n": 0}, {"n": 1}, {"n": 2}, {"n": 3}] diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py new file mode 100644 index 000000000..b666a4adc --- /dev/null +++ b/tests/streams/test_streams_workflow.py @@ -0,0 +1,618 @@ +"""Workflow-side conformance for the stream contract. + +Runs the reader and writer inside a real workflow on the memory provider +with a warm cache, and states the two rules about Workflow Tasks as tests: a +publish commits with its task (rule 1), and reads are recorded observations +that replay re-supplies (rule 2). The memory provider keeps neither and says +so in its docstring, so those two are strict expected failures here. A +storage provider that runs this module turns them into passes; that is the +measurement they exist for. + +The rest is what the portable surface promises on every provider: the +lifecycle hooks the worker calls, one subscription per topic per run, a read +that ends when the chain closes, a handle that follows continue-as-new, and +a topic shared by the workflow and an outside producer. +""" + +from __future__ import annotations + +import asyncio +import uuid +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.client import Client +from temporalio.streams import ( + DEFAULT_TOPIC, + END, + Cursor, + ReadSource, + RecordKind, + StreamCursorError, + WriteSink, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Replayer +from tests.helpers import new_worker +from tests.streams.test_streams_conformance import take + +INPUTS = topic("inputs", dict) +DECISIONS = topic("decisions", dict) + + +@pytest.fixture +def provider(env: WorkflowEnvironment): # pyright: ignore[reportUnusedFunction] + if env.supports_time_skipping: + pytest.skip( + "the memory provider polls on a timer, which time skipping turns into a spin" + ) + streams = MemoryStreams() + yield streams + streams.reset() + + +@workflow.defn +class ContractLoop: + """Reads ``inputs``, publishes a decision per value, reports control records.""" + + @workflow.run + async def run(self) -> list[dict[str, Any]]: + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + try: + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "replaced": record.supersession.previous_attempt, + "attempt": record.supersession.attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decisions.publish({"decided": record.value["n"]}) + trace.append( + { + "kind": "decision", + "n": record.value["n"], + "attempt": record.attempt, + } + ) + finally: + inputs.close() + decisions.finish() + # Twice on purpose: a finished topic stays finished, with one marker. + decisions.finish() + return trace + + +async def _run_the_loop( + client: Client, provider: MemoryStreams +) -> tuple[Any, list[dict[str, Any]]]: + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, ContractLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"n": 1}, {"n": 2}) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"n": 3}) + await second.finish() + trace = await handle.result() + return handle, trace + + +async def test_workflow_reads_decides_and_publishes( + client: Client, provider: MemoryStreams +): + handle, trace = await _run_the_loop(client, provider) + assert trace == [ + {"kind": "decision", "n": 1, "attempt": 1}, + {"kind": "decision", "n": 2, "attempt": 1}, + {"kind": "superseded", "replaced": 1, "attempt": 2}, + {"kind": "decision", "n": 3, "attempt": 2}, + {"kind": "finish", "producer": "model"}, + ] + + # The outside view of what the workflow published, on its own topic. + stream = provider.get_stream_handle(client, handle.id) + records = await take(stream.read(topic=DECISIONS), 5) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"decided": 1}), + (RecordKind.DATA, {"decided": 2}), + (RecordKind.DATA, {"retracting_attempt": 1}), + (RecordKind.DATA, {"decided": 3}), + (RecordKind.FINISH, None), + ] + assert all(r.producer_id == "" and r.topic == DECISIONS.name for r in records) + # The second finish() wrote nothing: the marker is the newest record. + assert await stream.latest(topic=DECISIONS) == records[-1].cursor + + +async def test_read_ends_when_the_workflow_closes_and_the_tail_is_delivered( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + stream = provider.get_stream_handle(client, handle.id) + + async def read_everything() -> list[Any]: + return [r.value async for r in stream.read(topic=DECISIONS)] + + # No count and no early break: the read ends by itself once the workflow + # is closed and everything it retained has been handed over. + values = await asyncio.wait_for(read_everything(), 30) + assert values == [ + {"decided": 1}, + {"decided": 2}, + {"retracting_attempt": 1}, + {"decided": 3}, + None, + ] + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 2: reading is a recorded observation, so replaying the history " + "with the store gone must re-supply the same records; the memory " + "provider reads live process memory instead" + ), +) +async def test_replay_without_the_store_resupplies_the_records( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + history = await handle.fetch_history() + + provider.reset() + replayer = Replayer(workflows=[ContractLoop], plugins=[provider]) + await replayer.replay_workflow(history) + + +# Run ids whose first workflow task already failed, shared with the workflow +# thread so the retry can tell it is the retry. Outside the sandbox on +# purpose: the sandbox re-imports this module per run and would hide the set. +_failed_once: set[str] = set() + + +@workflow.defn(sandboxed=False) +class PublishThenFail: + """Publishes, then fails its first workflow task; the retry publishes again.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + run_id = workflow.info().run_id + committed = run_id in _failed_once + decisions.publish({"committed": committed}) + if not committed: + _failed_once.add(run_id) + raise RuntimeError("the first task fails after publishing") + decisions.finish() + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 1: a publish commits with its workflow task, so no reader sees " + "a record from a task that failed; the memory provider makes it " + "visible at publish time" + ), +) +async def test_a_failed_task_publishes_nothing(client: Client, provider: MemoryStreams): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, PublishThenFail, plugins=[provider]) as worker: + handle = await client.start_workflow( + PublishThenFail.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS), 2, timeout=30) + await handle.result() + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"committed": True}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class Relay: + """Publishes one record per run and continues as new once.""" + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"run": run}) + if run == 0: + workflow.continue_as_new(run + 1) + decisions.finish() + + +class _RecordingHalf: + """A workflow half that logs the hooks the worker calls, then delegates.""" + + def __init__(self, inner: Any, calls: list[tuple[str, str]]) -> None: + self._inner = inner + self._calls = calls + + def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + return self._inner.open_reader(topic, after=after) + + def open_writer(self, topic: str) -> WriteSink: + return self._inner.open_writer(topic) + + def on_workflow_start(self) -> None: + self._calls.append(("start", workflow.info().run_id)) + + async def on_workflow_finish(self) -> None: + self._calls.append(("finish", workflow.info().run_id)) + + +class HookedMemory(MemoryStreams): + """The memory provider with its lifecycle hooks made visible.""" + + def __init__(self) -> None: + super().__init__() + self.calls: list[tuple[str, str]] = [] + + def workflow_provider(self) -> Any: + return _RecordingHalf(super().workflow_provider(), self.calls) + + +@pytest.mark.usefixtures("provider") +async def test_the_worker_calls_the_lifecycle_hooks_around_every_run(client: Client): + hooked = HookedMemory() + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[hooked]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + # Start before the function, finish after it, on both runs: the finish + # hook runs on the continue-as-new exit too, so a provider that parked + # something against the first run can let go before the successor starts. + kinds = [kind for kind, _ in hooked.calls] + assert kinds == ["start", "finish", "start", "finish"] + runs = [run_id for _, run_id in hooked.calls] + assert runs[0] == runs[1] and runs[2] == runs[3] and runs[0] != runs[2] + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[provider]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + await handle.result() + # The chain is followed: the successor's records arrive on the same read, + # and the read ends only when the last run of the chain is closed. + assert records == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class SharedReaders: + """Opens the same topic twice and pulls from both readers in turn.""" + + @workflow.run + async def run(self) -> list[Any]: + first = workflow.stream_reader(INPUTS) + second = workflow.stream_reader(INPUTS) + trace: list[Any] = ["shared" if first is second else "separate"] + try: + workflow.stream_reader(INPUTS, after=Cursor("memory:0")) + except ValueError: + trace.append("after-rejected") + try: + workflow.stream_reader(INPUTS.name, result_type=list) + except ValueError: + trace.append("type-rejected") + trace.append((await first.__anext__()).value) + trace.append((await second.__anext__()).value) + first.close() + return trace + + +async def test_a_second_reader_on_a_topic_shares_the_subscription( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, SharedReaders, plugins=[provider]) as worker: + handle = await client.start_workflow( + SharedReaders.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).append( + {"n": 1}, {"n": 2} + ) + assert await handle.result() == [ + "shared", + "after-rejected", + "type-rejected", + {"n": 1}, + {"n": 2}, + ] + + +@workflow.defn +class FinishThenPublish: + """Finishes a topic on one writer and publishes on a second one.""" + + @workflow.run + async def run(self) -> str: + workflow.stream_writer(DECISIONS).finish() + # A fresh writer object, the same topic. The marker is already there, + # so this publish would land after the end of the topic. + try: + workflow.stream_writer(DECISIONS).publish({"after": "finish"}) + except ValueError as error: + return str(error) + return "published" + + +async def test_a_finished_topic_stays_finished_across_writers( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, FinishThenPublish, plugins=[provider]) as worker: + handle = await client.start_workflow( + FinishThenPublish.run, id=workflow_id, task_queue=worker.task_queue + ) + assert "already finished" in await handle.result() + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + # The marker is all that landed: the publish behind it never reached the + # topic, which is the point of the guard reading like a topic-wide rule. + assert records == [(RecordKind.FINISH, None)] + + +@workflow.defn +class NoStreams: + """Touches no stream at all, on a worker that has a provider.""" + + @workflow.run + async def run(self) -> str: + await asyncio.sleep(0) + return "done" + + +@pytest.mark.usefixtures("provider") +async def test_a_workflow_that_touches_no_stream_runs_unchanged(client: Client): + hooked = HookedMemory() + async with new_worker(client, NoStreams, plugins=[hooked]) as worker: + result = await client.execute_workflow( + NoStreams.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert result == "done" + # The interceptor is installed per worker, not per workflow, so the hooks + # still bracket a run that never opened a reader or a writer. A provider's + # hooks therefore have to be cheap and safe on a workflow that uses none. + assert [kind for kind, _ in hooked.calls] == ["start", "finish"] + + +@workflow.defn +class ForeignCursor: + """Resumes from a cursor another provider minted.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS, after=Cursor("elsewhere:1")) + except StreamCursorError: + return "refused" + return "accepted" + + +async def test_the_workflow_reader_refuses_a_foreign_cursor( + client: Client, provider: MemoryStreams +): + async with new_worker(client, ForeignCursor, plugins=[provider]) as worker: + result = await client.execute_workflow( + ForeignCursor.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert result == "refused" + + +@workflow.defn +class NoProvider: + """Opens a stream on a worker that has no provider.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS) + except RuntimeError as error: + return str(error) + return "opened" + + +async def test_a_worker_without_a_provider_says_so(client: Client): + async with new_worker(client, NoProvider) as worker: + result = await client.execute_workflow( + NoProvider.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert "no stream provider is configured" in result + + +@workflow.defn +class OneLine: + """Publishes one record and finishes the topic.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"from": "workflow"}) + decisions.finish() + + +async def test_an_outside_producer_and_the_workflow_share_a_topic( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, OneLine, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=DECISIONS, producer_id="tool", attempt=1).append( + {"from": "producer"} + ) + handle = await client.start_workflow( + OneLine.run, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + # Both writers land on one topic in one order, each under its own + # identity: the producer's records carry its id, the workflow's carry none. + assert [(r.producer_id, r.kind, r.value) for r in records] == [ + ("tool", RecordKind.DATA, {"from": "producer"}), + ("", RecordKind.DATA, {"from": "workflow"}), + ("", RecordKind.FINISH, None), + ] + + +@workflow.defn +class DefaultTopicEcho: + """Reads one value on its default topic and answers on the same topic.""" + + @workflow.run + async def run(self) -> None: + reader = workflow.stream_reader(result_type=dict) + writer = workflow.stream_writer() + async for value in reader.values(): + writer.publish({"echo": value["n"] * 2}) + reader.close() + writer.finish() + + +async def test_a_workflow_reads_and_writes_its_default_topic( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, DefaultTopicEcho, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(producer_id="client", attempt=1).append({"n": 21}) + handle = await client.start_workflow( + DefaultTopicEcho.run, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DEFAULT_TOPIC)] + + records = await asyncio.wait_for(read_everything(), 30) + assert [(r.topic, r.kind, r.value) for r in records] == [ + (DEFAULT_TOPIC, RecordKind.DATA, {"n": 21}), + (DEFAULT_TOPIC, RecordKind.DATA, {"echo": 42}), + (DEFAULT_TOPIC, RecordKind.FINISH, None), + ] + + +@workflow.defn +class NewestTwo: + """Starts at the newest two records on ``inputs`` and returns their values.""" + + @workflow.run + async def run(self) -> list[Any]: + reader = workflow.stream_reader(INPUTS, last=2) + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == 2: + reader.close() + return values + + +async def test_a_workflow_reader_starts_at_the_last_n_records( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, NewestTwo, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="tool", attempt=1).append( + {"n": 1}, {"n": 2}, {"n": 3}, {"n": 4} + ) + handle = await client.start_workflow( + NewestTwo.run, id=workflow_id, task_queue=worker.task_queue + ) + assert await handle.result() == [3, 4] + + +@workflow.defn +class FromNow: + """Follows ``inputs`` from when it subscribes and returns the first value.""" + + @workflow.run + async def run(self) -> Any: + reader = workflow.stream_reader(INPUTS, after=END) + async for value in reader.values(): + reader.close() + return value["n"] + return None + + +async def test_a_workflow_reader_at_end_skips_what_was_there( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, FromNow, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "old"}) + handle = await client.start_workflow( + FromNow.run, id=workflow_id, task_queue=worker.task_queue + ) + result = asyncio.ensure_future(handle.result()) + # The subscription starts when the workflow runs, which the test does + # not observe, so appends keep coming until the workflow takes one. + for _ in range(150): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({result}, timeout=0.2) + if done: + break + assert await asyncio.wait_for(result, 10) == "new" + + +def test_a_workflow_reader_start_names_one_place(): + # Checked before the reader needs a running workflow, so a mistake says + # what it is rather than that there is no workflow. + with pytest.raises(ValueError, match="either after= or last="): + workflow.stream_reader(INPUTS, after=END, last=1) + with pytest.raises(ValueError, match="positive"): + workflow.stream_reader(INPUTS, last=0)