From 34cf5f4f95e2a1db1deb35136ace112c7ba2bf16 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:44:00 -0700 Subject: [PATCH 1/3] Added channel subscriptions to the workflow runtime. A workflow subscribes by command or listens on its linked channel, and Core hands the notifications over as a job. The handle routes them by kind and ends with unsubscribe(). --- temporalio/worker/_workflow_instance.py | 63 +++++++++ temporalio/workflow/__init__.py | 10 +- temporalio/workflow/_channels.py | 171 ++++++++++++++++++++++++ temporalio/workflow/_context.py | 16 +++ 4 files changed, 259 insertions(+), 1 deletion(-) diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 5dcc4aa96..e392671a3 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -303,6 +303,16 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: det.worker_level_failure_exception_types ) self._patch_activation_callback = det.patch_activation_callback + # 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 +634,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 +877,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: @@ -1810,6 +1844,35 @@ async def workflow_start_nexus_operation( ) ) + 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 6856c1a20..8c33a4852 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -57,7 +57,12 @@ as_completed, wait, ) -from ._channels import Notification +from ._channels import ( + ChannelSubscription, + Notification, + linked_channel, + subscribe_channel, +) from ._context import ( Info, ParentInfo, @@ -259,7 +264,10 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "ChannelSubscription", "Notification", + "linked_channel", + "subscribe_channel", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py index 94378c4e4..07c3bde7a 100644 --- a/temporalio/workflow/_channels.py +++ b/temporalio/workflow/_channels.py @@ -12,15 +12,21 @@ 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", ] @@ -79,3 +85,168 @@ def _from_proto( else None ), ) + + +class ChannelSubscription: + """A workflow's handle on one channel, of either kind. + + Prefer :func:`temporalio.workflow.subscribe_channel` for an independent + channel and :func:`temporalio.workflow.linked_channel` for one linked to + this workflow. The handle is an async iterator over the notifications as + they arrive, and :meth:`receive` takes them one at a time. Notifications + wait in arrival order until taken. Two loops on one handle share its + buffer and interleave. :meth:`unsubscribe` ends an independent + subscription; a linked channel lasts as long as the run. + """ + + def __init__(self, channel: str, *, linked: bool = False) -> None: + """Prefer the two module functions named above.""" + self._channel = channel + self._linked = linked + self._closed = False + self._pending: deque[Notification] = deque() + self._waiters: deque[asyncio.Future[None]] = deque() + + @property + def channel(self) -> str: + """The channel this handle is on.""" + return self._channel + + @property + def linked(self) -> bool: + """Whether the channel is the one linked to this workflow. + + A linked handle gets the notifications that carry + :attr:`Notification.linked_to`; an independent one gets the rest. + """ + return self._linked + + @property + def closed(self) -> bool: + """Whether :meth:`unsubscribe` has ended this subscription. + + A closed handle still hands out the notifications it had queued, then + :meth:`receive` raises and iteration ends. + """ + return self._closed + + def unsubscribe(self) -> None: + """End this workflow's subscription to the channel. + + Issues the unsubscribe command once; a second call changes nothing. + Notifications already queued on this handle can still be read, and + one the server put on a scheduled Workflow Task before the command + landed is dropped on arrival. A later + :func:`temporalio.workflow.subscribe_channel` for the same name opens + a new subscription with a new command. + + Raises: + ValueError: The handle is from + :func:`temporalio.workflow.linked_channel`. A linked channel is + part of the run and has no subscription to end. + """ + if self._linked: + raise ValueError("a linked channel has no subscription") + if self._closed: + return + _Runtime.current().workflow_unsubscribe_channel(self._channel) + self._closed = True + # The waiters wake to find the handle closed with nothing queued. + self._wake() + + async def receive(self) -> Notification: + """The next notification on this channel, waiting for one to arrive. + + The wait is a future the delivery resolves, so it adds no command and + replays the same way. + + Raises: + RuntimeError: The subscription is closed and nothing is queued. + """ + notification = await self._next() + if notification is None: + raise RuntimeError("channel subscription closed") + return notification + + def __aiter__(self) -> ChannelSubscription: + """The subscription is its own iterator.""" + return self + + async def __anext__(self) -> Notification: + """The next notification. The iteration ends once the handle is closed and drained.""" + notification = await self._next() + if notification is None: + raise StopAsyncIteration + return notification + + async def _next(self) -> Notification | None: + while not self._pending: + if self._closed: + return None + waiter: asyncio.Future[None] = asyncio.Future() + self._waiters.append(waiter) + try: + await waiter + finally: + if waiter in self._waiters: + self._waiters.remove(waiter) + return self._pending.popleft() + + def _deliver(self, notification: Notification) -> None: + self._pending.append(notification) + self._wake() + + def _wake(self) -> None: + # Every waiter wakes; the ones that find the buffer empty again wait + # once more. + while self._waiters: + waiter = self._waiters.popleft() + if not waiter.done(): + waiter.set_result(None) + + +def subscribe_channel(channel: str) -> ChannelSubscription: + """Subscribe this workflow to ``channel``. + + The first call for a channel in a run issues a command, so gate a new + channel with :func:`temporalio.workflow.patched` as you would a timer. A + second call for the same channel returns the subscription already open, + and the two share its buffer. From then on each Workflow Task carries the + notifications the server folded for this workflow on the channel, and + they arrive here. A successor run after continue-as-new starts with no + subscriptions. + + Args: + channel: Name of the channel, scoped to the namespace. + + Raises: + ValueError: ``channel`` is empty. + """ + if not channel: + raise ValueError("channel must not be empty") + return _Runtime.current().workflow_subscribe_channel(channel) + + +def linked_channel(channel: str) -> ChannelSubscription: + """Listen on the channel named ``channel`` that is linked to this workflow. + + A linked channel lives in this workflow's own state, so the workflow is + its listener by construction: no command, no event, and no gate needed + for a new name. A writer reaches it by naming this workflow as the + execution, as in :py:meth:`temporalio.client.Client.notify_channel` with + ``workflow_id``, and a successor run after continue-as-new is reached by + the same calls. + A second call for the same name returns the handle already open, and the + two share its buffer. The name does not collide with an independent + channel's: a notification carrying :attr:`Notification.linked_to` comes + here, one without it goes to :func:`subscribe_channel`. + + Args: + channel: Name of the channel, scoped to this workflow. + + Raises: + ValueError: ``channel`` is empty. + """ + if not channel: + raise ValueError("channel must not be empty") + return _Runtime.current().workflow_linked_channel(channel) diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 37928afb9..2a9ff6355 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -22,6 +22,7 @@ if TYPE_CHECKING: from ._activities import ActivityCancellationType, ActivityHandle + from ._channels import ChannelSubscription from ._event_groups import EventGroup from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle @@ -499,6 +500,21 @@ async def workflow_start_nexus_operation( event_groups: Sequence[EventGroup] | None = None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... + + @abstractmethod + def workflow_unsubscribe_channel(self, channel: str) -> None: + """Record the unsubscribe command for ``channel`` and forget its handle. + + Called once per subscription, by the handle that owns it, so a late + notification for the channel finds no handle and is dropped. + """ + ... + + @abstractmethod + def workflow_linked_channel(self, channel: str) -> ChannelSubscription: ... + @abstractmethod def workflow_time_ns(self) -> int: ... From f04fe1143fa7eb25e00ac7bd6554064208f86bf1 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:44:00 -0700 Subject: [PATCH 2/3] Listed a run's channel subscriptions on its description. Describe is how an operator sees what a run listens on and what is pending for it. --- CHANGELOG.md | 4 ++ temporalio/client/__init__.py | 2 + temporalio/client/_channel.py | 67 ++++++++++++++++++++++++++++++++++ temporalio/client/_workflow.py | 17 +++++++++ 4 files changed, 90 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 005a5eb16..6d589a3fb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -42,6 +42,10 @@ to include examples, links to docs, or any other relevant information. `unregister_channel_listener` reach a named channel on the server. The channel is either independent or linked to an execution, which `execution=` (`temporalio.common.Execution`) or the `workflow_id=` shorthand names. + - A workflow subscribes with `workflow.subscribe_channel(name)`, reads a channel linked to it with + `workflow.linked_channel(name)` and ends a subscription with `unsubscribe()`. Notifications + arrive with the workflow's tasks. `WorkflowExecutionDescription.channel_subscriptions` lists + the channels a run listens on. ### Changed diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 87e70537c..d77cf9693 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -69,6 +69,7 @@ ChannelDescription, ChannelKind, ChannelListener, + ChannelSubscriptionInfo, ) from ._client import ( Client, @@ -371,6 +372,7 @@ "ChannelKind", "ChannelListener", "ChannelAddress", + "ChannelSubscriptionInfo", "_ClientImpl", "_apply_headers", "_decode_user_metadata", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py index 91ef53e05..0decc2265 100644 --- a/temporalio/client/_channel.py +++ b/temporalio/client/_channel.py @@ -8,6 +8,7 @@ from enum import IntEnum import temporalio.api.notification.v1 +import temporalio.api.workflow.v1 import temporalio.common from temporalio.workflow import Notification @@ -18,6 +19,7 @@ "ChannelDescription", "ChannelKind", "ChannelListener", + "ChannelSubscriptionInfo", ] @@ -128,6 +130,71 @@ class ChannelDescription: """ +@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. 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 + ), ) From ef6ca30a7f36adce5db2d1a498f83e4eb57639e5 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:44:00 -0700 Subject: [PATCH 3/3] Covered the workflow channel surface. The unit cases drive the instance with hand-built activations. The live cases skip unless -E names a channel server. --- tests/streams/conftest.py | 45 ++ tests/streams/test_channels.py | 1059 +++++++++++++++++++++++++++++++- 2 files changed, 1100 insertions(+), 4 deletions(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index a456d5975..c824515d0 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -9,6 +9,27 @@ def pytest_configure(config: pytest.Config) -> None: "needs_channel_server: the case needs a server that serves notification " "channels, named with -E host:port", ) + config.addinivalue_line( + "markers", + "needs_linked_server: the case needs a server that serves channels linked " + "to a workflow, named with -E host:port", + ) + 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_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( @@ -28,6 +49,30 @@ def pytest_collection_modifyitems( ), ) ) + 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 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_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" + ), + ) + ) for item in items: for marker, skip in skips: if item.get_closest_marker(marker): diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index 9a81b31c5..ab290734e 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -1,31 +1,332 @@ -"""The notification channel calls on the client. +"""The notification channel surface: the command, the delivery and the client calls. -The unit case fakes the service. The live cases need a server that serves -channels, named with -E, and skip otherwise. +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 timedelta +from datetime import datetime, timedelta, timezone from typing import Any import pytest +import temporalio.api.common.v1 import temporalio.api.enums.v1 import temporalio.api.notification.v1 +import temporalio.api.workflow.v1 import temporalio.api.workflowservice.v1 +import temporalio.bridge.proto.workflow_activation +import temporalio.bridge.proto.workflow_completion import temporalio.common +import temporalio.converter +from temporalio import workflow +from temporalio.api.enums.v1 import EventType from temporalio.client import ( Callback, ChannelKind, + ChannelSubscriptionInfo, Client, + WorkflowExecutionDescription, ) from temporalio.service import RPCError, RPCStatusCode +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( @@ -36,6 +337,289 @@ def _linked(channel: str, counter: int, position: bytes = b"") -> Notification: ) +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_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 ): @@ -123,6 +707,76 @@ async def call(req: Any, *, _response: Any = response, **_: Any) -> Any: await client.poll_channel("orders", execution=activity, run_id="run") +@pytest.mark.needs_channel_server +async def test_a_workflow_receives_a_client_notification(client: Client): + channel = f"orders-{uuid.uuid4()}" + worker = new_worker(client, ReceiveOne) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + ReceiveOne.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + + async def subscribed() -> None: + # The worker's Core has to carry the subscribe command to the + # server. One that refuses it fails every completion, and the task + # times out instead; say so rather than wait on the result forever. + events = [event.event_type async for event in handle.fetch_history_events()] + if EventType.EVENT_TYPE_WORKFLOW_TASK_TIMED_OUT in events: + pytest.fail( + "the worker could not complete the task that subscribes; the Core " + "the bridge pins must carry the subscribe command" + ) + assert ( + EventType.EVENT_TYPE_WORKFLOW_NOTIFICATION_CHANNEL_SUBSCRIBED in events + ) + + await assert_eventually(subscribed, timeout=timedelta(seconds=30)) + + async def listening() -> None: + # The channel exists once the subscribe lands, so a describe that + # races it is answered with not found. + try: + description = await client.describe_channel(channel) + except RPCError as err: + assert err.status != RPCStatusCode.NOT_FOUND, "channel not created yet" + raise + assert [listener.workflow_id for listener in description.listeners] == [ + handle.id + ] + + await assert_eventually(listening) + listeners = await client.notify_channel( + channel, position=b"1-0", counter=1, metadata={"topic": "inputs"} + ) + assert listeners == 1 + assert await asyncio.wait_for(handle.result(), 30) == { + "channel": channel, + "counter": 1, + "position": "1-0", + "topic": "inputs", + } + polled = await client.poll_channel(channel, wait=False) + assert [n.counter for n in polled] == [1] + description = await client.describe_channel(channel) + assert description.latest is not None and description.latest.counter == 1 + finally: + # A Core that refuses the subscribe command leaves the task in a + # timeout loop and the worker's shutdown waiting on it, so end the run + # first and give the shutdown a bound. + with contextlib.suppress(RPCError): + await handle.terminate() + try: + await asyncio.wait_for(worker.shutdown(), 15) + except asyncio.TimeoutError: + running.cancel() + with contextlib.suppress(asyncio.CancelledError): + await running + + @pytest.mark.needs_channel_server async def test_a_callback_listener_registers_and_unregisters(client: Client): channel = f"orders-{uuid.uuid4()}" @@ -182,3 +836,400 @@ async def test_a_channel_retains_notifications_for_pollers(client: Client): 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 +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 +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 +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 +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)