From 565783fd25b58963522f6d0430b8a639f395953f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:08:32 -0700 Subject: [PATCH 01/19] Added workflow.stream_reader and workflow.stream_writer. Workflow code reaches its stream the way it reaches the payload converter, through the runtime on the thread. The workflow half of the provider is made per instance, so whatever it keeps dies with the instance the way handlers do. --- temporalio/worker/_workflow_instance.py | 19 ++ temporalio/workflow/__init__.py | 10 + temporalio/workflow/_context.py | 4 + temporalio/workflow/_streams.py | 279 ++++++++++++++++++++++++ 4 files changed, 312 insertions(+) create mode 100644 temporalio/workflow/_streams.py diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 108431021..26f796e57 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, @@ -179,6 +181,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): @@ -298,6 +301,8 @@ 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 self._default_workflow_logic_flags = det.default_workflow_logic_flags self._primary_task: asyncio.Task[None] | None = None self._cancel_primary_task_pending = False @@ -1774,6 +1779,20 @@ 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_time_ns(self) -> int: return self._time_ns diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index fa2681139..f63135bf4 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -147,6 +147,12 @@ logger, unsafe, ) +from ._streams import ( + StreamReader, + StreamWriter, + stream_reader, + stream_writer, +) from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowConfig, @@ -252,6 +258,10 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "StreamReader", + "StreamWriter", + "stream_reader", + "stream_writer", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index b33f83150..25523836a 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -24,6 +24,7 @@ from ._activities import ActivityCancellationType, ActivityHandle from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle + from ._streams import _WorkflowStreams from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowHandle, @@ -475,6 +476,9 @@ async def workflow_start_nexus_operation( summary: str | None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_streams(self) -> _WorkflowStreams: ... + @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..5d319d50a --- /dev/null +++ b/temporalio/workflow/_streams.py @@ -0,0 +1,279 @@ +"""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, Cursor, RecordKind, StreamRecord +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]] = {} + + +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) -> None: + """Prefer :func:`temporalio.workflow.stream_writer`.""" + self._sink = sink + self._topic = topic + self._finished = False + + @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: :meth:`finish` was already called on this writer. + """ + if 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. + """ + if self._finished: + return + self._finished = True + self._sink.publish( + to_wire(payload_converter(), topic=self._topic, kind=RecordKind.FINISH) + ) + + +@overload +def stream_reader(topic: StreamTopic[T], *, after: Cursor = ...) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, *, result_type: type[T], after: Cursor = ... +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, *, result_type: None = None, after: Cursor = ... +) -> StreamReader[Any]: ... + + +def stream_reader( + topic: str | StreamTopic[Any], + *, + result_type: type | None = None, + after: Cursor = BEGINNING, +) -> 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. 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 neither ``after`` nor + a 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. + result_type: The value type for a string-named topic, used as the + decode hint. :class:`temporalio.common.RawValue` returns the + payload untouched. + after: Resume strictly after this record. Honoured on the first + subscription of a run, because after that the recorded + observations decide. + + Raises: + ValueError: ``topic`` is empty, ``result_type`` was passed with a + definition, 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. + """ + 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 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= or other type" + ) + return existing + source = state.provider.open_reader(name, after=after) + + def forget() -> None: + state.readers.pop(name, None) + + reader: StreamReader[Any] = StreamReader( + source, topic=name, result_type=result_type, after=after, on_close=forget + ) + state.readers[name] = reader + return reader + + +@overload +def stream_writer(topic: StreamTopic[T]) -> StreamWriter[T]: ... + + +@overload +def stream_writer(topic: str) -> StreamWriter[Any]: ... + + +def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: + """Publish to ``topic`` of this workflow's stream. + + 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. Encoding follows each published value. + + Raises: + ValueError: ``topic`` is empty. + """ + name, _ = resolve_topic(topic) + provider = _Runtime.current().workflow_streams().provider + return StreamWriter(provider.open_writer(name), name) From a2587273eb4557ef837a218650b6e8ab3b6e5e38 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:08:32 -0700 Subject: [PATCH 02/19] Carried the stream provider into the workflow runtime. The worker and the replayer hand the provider to each workflow instance, and the worker brackets the workflow function with the provider's lifecycle hooks so no workflow code has to call them. The finish hook stays off an evicted run and off a collected coroutine, because neither is the workflow ending. --- temporalio/worker/_replayer.py | 1 + temporalio/worker/_worker.py | 6 + temporalio/worker/_workflow.py | 47 +++ tests/streams/test_stream_hooks.py | 139 ++++++++ tests/streams/test_streams_workflow.py | 443 +++++++++++++++++++++++++ 5 files changed, 636 insertions(+) create mode 100644 tests/streams/test_stream_hooks.py create mode 100644 tests/streams/test_streams_workflow.py 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 1b217b4a5..9b94de95a 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,43 @@ 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: + if not _evicting(runtime): + await provider.on_workflow_finish() + raise + await provider.on_workflow_finish() + return result + + +def _evicting(runtime: temporalio.workflow._Runtime) -> bool: + # Eviction cancels the primary task the same way a workflow cancellation + # does; the flag the instance sets before cancelling is what tells them + # apart, and only the cancellation is a run ending. + return bool(getattr(runtime, "_deleting", False)) + + # 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 +144,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 +203,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 @@ -740,6 +786,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/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py new file mode 100644 index 000000000..9f98f75f6 --- /dev/null +++ b/tests/streams/test_stream_hooks.py @@ -0,0 +1,139 @@ +"""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 + + +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: the stream state and the eviction flag.""" + + def __init__(self, provider: _Provider, *, deleting: bool = False) -> None: + self._streams = _Streams(provider) + self._deleting = deleting + + def workflow_streams(self) -> _Streams: + return self._streams + + +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, deleting=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"] diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py new file mode 100644 index 000000000..71f2bb8bf --- /dev/null +++ b/tests/streams/test_streams_workflow.py @@ -0,0 +1,443 @@ +"""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 ( + 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 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), + ] From 2e556aa84ef44b3eb02eb9f054abf598cba81061 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 13:41:46 -0700 Subject: [PATCH 03/19] Asked the runtime whether it is evicting instead of guessing. Reading the eviction state off a private attribute name failed open: rename it upstream and the finish hook starts running during eviction, with every test still green. --- temporalio/worker/_workflow.py | 11 +++-------- temporalio/worker/_workflow_instance.py | 3 +++ temporalio/workflow/_context.py | 10 ++++++++++ tests/streams/test_stream_hooks.py | 26 +++++++++++++++++++++---- 4 files changed, 38 insertions(+), 12 deletions(-) diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 9b94de95a..9de338a7d 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -81,20 +81,15 @@ async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: except GeneratorExit: raise except BaseException: - if not _evicting(runtime): + # 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 -def _evicting(runtime: temporalio.workflow._Runtime) -> bool: - # Eviction cancels the primary task the same way a workflow cancellation - # does; the flag the instance sets before cancelling is what tells them - # apart, and only the cancellation is a run ending. - return bool(getattr(runtime, "_deleting", False)) - - # 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 diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 26f796e57..7fa4405f0 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -1353,6 +1353,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 diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 25523836a..160da8520 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -345,6 +345,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: ... diff --git a/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py index 9f98f75f6..a547daf3e 100644 --- a/tests/streams/test_stream_hooks.py +++ b/tests/streams/test_stream_hooks.py @@ -20,6 +20,7 @@ WorkflowInboundInterceptor, ) from temporalio.worker._workflow import _StreamHooksInterceptor +from temporalio.worker._workflow_instance import _WorkflowInstanceImpl class _Provider: @@ -39,15 +40,24 @@ def __init__(self, provider: _Provider) -> None: class _FakeRuntime: - """Only what the interceptor reads: the stream state and the eviction flag.""" + """Only what the interceptor reads, and only through the runtime interface. - def __init__(self, provider: _Provider, *, deleting: bool = False) -> None: + 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._deleting = deleting + 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.""" @@ -85,7 +95,7 @@ async def provider() -> Any: def _evicting(fake: _Provider) -> None: loop = asyncio.get_running_loop() - workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake, deleting=True)) # type: ignore[arg-type] + 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): @@ -137,3 +147,11 @@ async def test_the_finish_hook_does_not_run_during_eviction(provider: _Provider) _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__ From 296324666ae00617fb53a2f9e0b044989b2da3c0 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 13:41:53 -0700 Subject: [PATCH 04/19] Held the finished topics on the run, not on the writer object. stream_writer() hands out a new object per call, so a guard kept on one of them let a publish land behind the finish marker it was meant to refuse. --- temporalio/workflow/_streams.py | 28 +++++++++++++------- tests/streams/test_streams_workflow.py | 36 ++++++++++++++++++++++++++ 2 files changed, 55 insertions(+), 9 deletions(-) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index 5d319d50a..49f06902c 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -35,6 +35,10 @@ class _WorkflowStreams: 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]): @@ -135,11 +139,11 @@ class StreamWriter(Generic[T]): any value. """ - def __init__(self, sink: WriteSink, topic: str) -> None: + 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 = False + self._finished = finished @property def topic(self) -> str: @@ -156,9 +160,10 @@ def publish(self, value: T) -> None: through pre-encoded. Raises: - ValueError: :meth:`finish` was already called on this writer. + ValueError: The topic was already finished in this run, by this + writer or by another one on the same topic. """ - if self._finished: + if self._topic in self._finished: raise ValueError(f"topic {self._topic!r} was already finished") self._sink.publish( to_wire( @@ -173,11 +178,13 @@ 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. + 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._finished: + if self._topic in self._finished: return - self._finished = True + self._finished.add(self._topic) self._sink.publish( to_wire(payload_converter(), topic=self._topic, kind=RecordKind.FINISH) ) @@ -266,6 +273,9 @@ def stream_writer(topic: str) -> StreamWriter[Any]: ... def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: """Publish to ``topic`` of this workflow's stream. + 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 @@ -275,5 +285,5 @@ def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: ValueError: ``topic`` is empty. """ name, _ = resolve_topic(topic) - provider = _Runtime.current().workflow_streams().provider - return StreamWriter(provider.open_writer(name), name) + state = _Runtime.current().workflow_streams() + return StreamWriter(state.provider.open_writer(name), name, state.finished) diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py index 71f2bb8bf..50572d419 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -357,6 +357,42 @@ async def test_a_second_reader_on_a_topic_shares_the_subscription( ] +@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 ForeignCursor: """Resumes from a cursor another provider minted.""" From 4695c7df3f7f5b2307d71608262a22efbb61a675 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 13:42:02 -0700 Subject: [PATCH 05/19] Covered the reader promises nothing measured. A second loop on one reader, a cancelled read, a reopened topic and a workflow that touches no stream were all promised in a docstring and pinned by no case. --- tests/streams/test_stream_reader.py | 178 +++++++++++++++++++++++++ tests/streams/test_streams_workflow.py | 26 ++++ 2 files changed, 204 insertions(+) create mode 100644 tests/streams/test_stream_reader.py 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 index 50572d419..afc0b4cff 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -393,6 +393,32 @@ async def read_everything() -> list[Any]: 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.""" From 00f9b2d58a8cac94f750759f8f70e151191159f5 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 14:46:33 -0700 Subject: [PATCH 06/19] Let stream_reader() and stream_writer() default to the default topic. A workflow with one stream should not have to invent a topic name for it; with no topic both resolve to DEFAULT_TOPIC and decode as a string-named topic does. --- temporalio/workflow/_streams.py | 25 ++++++++++------- tests/streams/test_streams_workflow.py | 38 ++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 10 deletions(-) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index 49f06902c..7a545444a 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -196,18 +196,18 @@ def stream_reader(topic: StreamTopic[T], *, after: Cursor = ...) -> StreamReader @overload def stream_reader( - topic: str, *, result_type: type[T], after: Cursor = ... + topic: str | None = None, *, result_type: type[T], after: Cursor = ... ) -> StreamReader[T]: ... @overload def stream_reader( - topic: str, *, result_type: None = None, after: Cursor = ... + topic: str | None = None, *, result_type: None = None, after: Cursor = ... ) -> StreamReader[Any]: ... def stream_reader( - topic: str | StreamTopic[Any], + topic: str | StreamTopic[Any] | None = None, *, result_type: type | None = None, after: Cursor = BEGINNING, @@ -216,7 +216,9 @@ def stream_reader( ``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. One subscription per topic per run. A second call + 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 neither ``after`` nor a different type. Adding a reader on a new topic is a new command, so @@ -225,8 +227,9 @@ def stream_reader( continue-as-new implicitly. Args: - topic: The topic, relative to this workflow's stream. - result_type: The value type for a string-named topic, used as the + 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. Honoured on the first @@ -267,11 +270,11 @@ def stream_writer(topic: StreamTopic[T]) -> StreamWriter[T]: ... @overload -def stream_writer(topic: str) -> StreamWriter[Any]: ... +def stream_writer(topic: str | None = None) -> StreamWriter[Any]: ... -def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: - """Publish to ``topic`` of this workflow's stream. +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. @@ -279,7 +282,9 @@ def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: 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. Encoding follows each published value. + decided at runtime. Omitted, the writer is on + :data:`temporalio.streams.DEFAULT_TOPIC`. Encoding follows each + published value. Raises: ValueError: ``topic`` is empty. diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py index afc0b4cff..f13c13f44 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -25,6 +25,7 @@ from temporalio import workflow from temporalio.client import Client from temporalio.streams import ( + DEFAULT_TOPIC, Cursor, ReadSource, RecordKind, @@ -503,3 +504,40 @@ async def read_everything() -> list[Any]: ("", 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), + ] From a3edfbbfbc80261e7b2b7bfb414c41902a02417b Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:55:12 -0700 Subject: [PATCH 07/19] Let stream_reader() start at END or at the last N records. A workflow could only start from BEGINNING or a cursor, so it had no way to follow from now. The provider resolves the start outside the workflow and records it, so replay reproduces it; last= is passed only when given, which keeps older providers working. --- temporalio/workflow/_streams.py | 72 +++++++++++++++++++------ tests/streams/test_streams_workflow.py | 75 ++++++++++++++++++++++++++ 2 files changed, 132 insertions(+), 15 deletions(-) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index 7a545444a..54353508d 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -18,7 +18,14 @@ from typing import Any, Generic, TypeVar, cast, overload from temporalio.streams._provider import ReadSource, WorkflowStreamProvider, WriteSink -from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +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 @@ -191,18 +198,28 @@ def finish(self) -> None: @overload -def stream_reader(topic: StreamTopic[T], *, after: Cursor = ...) -> StreamReader[T]: ... +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 = ... + 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 = ... + topic: str | None = None, + *, + result_type: None = None, + after: Cursor = ..., + last: int | None = None, ) -> StreamReader[Any]: ... @@ -211,6 +228,7 @@ def stream_reader( *, result_type: type | None = None, after: Cursor = BEGINNING, + last: int | None = None, ) -> StreamReader[Any]: """Subscribe this workflow to ``topic`` of its own stream. @@ -220,8 +238,8 @@ def stream_reader( :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 neither ``after`` nor - a different type. Adding a reader on a new topic is a new command, so + 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. @@ -232,34 +250,58 @@ def stream_reader( 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. Honoured on the first - subscription of a run, because after that the recorded - observations decide. + 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, or a reader on the topic is already open and this - call asked for a different position or type. + 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 result_type is not existing._result_type: + 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= or other type" + "stream_reader shares it and takes no after=, last= or other type" ) return existing - source = state.provider.open_reader(name, after=after) + # 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=after, on_close=forget + source, topic=name, result_type=result_type, after=previous, on_close=forget ) state.readers[name] = reader return reader diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py index f13c13f44..b666a4adc 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -26,6 +26,7 @@ from temporalio.client import Client from temporalio.streams import ( DEFAULT_TOPIC, + END, Cursor, ReadSource, RecordKind, @@ -541,3 +542,77 @@ async def read_everything() -> list[Any]: (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) From aab2644615f3d870ded6c0e983ec3b32bc60017f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 18:24:51 -0700 Subject: [PATCH 08/19] Added channel subscriptions to the workflow runtime. A workflow subscribes to a notification channel once per run and 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. --- temporalio/worker/_workflow_instance.py | 35 ++++ temporalio/workflow/__init__.py | 8 + temporalio/workflow/_channels.py | 144 +++++++++++++++ temporalio/workflow/_context.py | 4 + tests/streams/test_channels.py | 224 ++++++++++++++++++++++++ 5 files changed, 415 insertions(+) create mode 100644 temporalio/workflow/_channels.py create mode 100644 tests/streams/test_channels.py diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index d2b588c88..d37a001c4 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -308,6 +308,10 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: 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 + ] = {} self._default_workflow_logic_flags = det.default_workflow_logic_flags self._primary_task: asyncio.Task[None] | None = None self._cancel_primary_task_pending = False @@ -629,6 +633,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"): @@ -870,6 +876,23 @@ 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: + 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 subscribed to is not the workflow's + # concern. + logger.debug( + "Dropping a notification on channel %r, which this run has not " + "subscribed to", + 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: @@ -1832,6 +1855,18 @@ def workflow_streams(self) -> temporalio.workflow._streams._WorkflowStreams: ) 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_time_ns(self) -> int: return self._time_ns diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index 71354331b..0cf3b058a 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -57,6 +57,11 @@ as_completed, wait, ) +from ._channels import ( + ChannelSubscription, + Notification, + subscribe_channel, +) from ._context import ( Info, ParentInfo, @@ -264,6 +269,9 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "ChannelSubscription", + "Notification", + "subscribe_channel", "StreamReader", "StreamWriter", "stream_reader", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py new file mode 100644 index 000000000..8fd045519 --- /dev/null +++ b/temporalio/workflow/_channels.py @@ -0,0 +1,144 @@ +"""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 +from temporalio.workflow._context import _Runtime + +__all__ = ["ChannelSubscription", "Notification", "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. + """ + + @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()), + ) + + +class ChannelSubscription: + """A workflow's subscription to one channel. + + Prefer :func:`temporalio.workflow.subscribe_channel`. The subscription 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 subscription share its buffer and + interleave. + """ + + def __init__(self, channel: str) -> None: + """Prefer :func:`temporalio.workflow.subscribe_channel`.""" + self._channel = channel + self._pending: deque[Notification] = deque() + self._waiters: deque[asyncio.Future[None]] = deque() + + @property + def channel(self) -> str: + """The channel this subscription is on.""" + return self._channel + + 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. + """ + while not self._pending: + 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 __aiter__(self) -> ChannelSubscription: + """The subscription is its own iterator.""" + return self + + async def __anext__(self) -> Notification: + """The next notification; the iteration never ends on its own.""" + return await self.receive() + + def _deliver(self, notification: Notification) -> None: + self._pending.append(notification) + # 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) diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index fcea4b920..e0cb59a8d 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 @@ -513,6 +514,9 @@ async def workflow_start_nexus_operation( @abstractmethod def workflow_streams(self) -> _WorkflowStreams: ... + @abstractmethod + def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... + @abstractmethod def workflow_time_ns(self) -> int: ... diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py new file mode 100644 index 000000000..a8b6a6d26 --- /dev/null +++ b/tests/streams/test_channels.py @@ -0,0 +1,224 @@ +"""The notification channel surface: the command, the delivery and the client calls. + +The workflow instance is driven with activations directly, the way Core +drives it, because the dev server this chain tests against does not accept +the subscribe command. +""" + +from __future__ import annotations + +from datetime import datetime, timedelta, timezone +from typing import Any + +import temporalio.api.common.v1 +import temporalio.api.notification.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.worker._workflow_instance import ( + UnsandboxedWorkflowRunner, + WorkflowInstance, + WorkflowInstanceDetails, +) + +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("") + + +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 _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 _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 From 5f28a14f7f649d70ab5f8b0d1a9e3e6e9586cd0b Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 18:24:51 -0700 Subject: [PATCH 09/19] Added the notification channel calls to the client. Notify, poll, describe, register and unregister go through the outbound interceptor with one input each. Metadata values are encoded with the data converter on the way out and come back as payloads, the shape workflow code sees. --- temporalio/client/__init__.py | 16 +++ temporalio/client/_channel.py | 78 ++++++++++++ temporalio/client/_client.py | 191 ++++++++++++++++++++++++++++++ temporalio/client/_impl.py | 122 +++++++++++++++++++ temporalio/client/_interceptor.py | 125 +++++++++++++++++++ tests/streams/conftest.py | 23 ++++ tests/streams/test_channels.py | 55 ++++++++- 7 files changed, 609 insertions(+), 1 deletion(-) create mode 100644 temporalio/client/_channel.py diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 2eef41a39..1a879923c 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -64,6 +64,10 @@ from ._callback import ( Callback, ) +from ._channel import ( + ChannelDescription, + ChannelListener, +) from ._client import ( Client, ClientConfig, @@ -106,6 +110,7 @@ CreateScheduleInput, DeleteScheduleInput, DescribeActivityInput, + DescribeChannelInput, DescribeNexusOperationInput, DescribeScheduleInput, DescribeWorkflowInput, @@ -120,10 +125,13 @@ ListNexusOperationsInput, ListSchedulesInput, ListWorkflowsInput, + NotifyChannelInput, OutboundInterceptor, PauseActivityInput, PauseScheduleInput, + PollChannelInput, QueryWorkflowInput, + RegisterChannelListenerInput, ReportCancellationAsyncActivityInput, SignalWorkflowInput, StartActivityInput, @@ -137,6 +145,7 @@ TriggerScheduleInput, UnpauseActivityInput, UnpauseScheduleInput, + UnregisterChannelListenerInput, UpdateActivityOptionsInput, UpdateScheduleInput, UpdateWithStartStartWorkflowInput, @@ -315,6 +324,11 @@ "TerminateNexusOperationInput", "ListNexusOperationsInput", "CountNexusOperationsInput", + "DescribeChannelInput", + "NotifyChannelInput", + "PollChannelInput", + "RegisterChannelListenerInput", + "UnregisterChannelListenerInput", "StartWorkflowUpdateInput", "UpdateWithStartUpdateWorkflowInput", "UpdateWithStartStartWorkflowInput", @@ -351,6 +365,8 @@ "CloudOperationsClient", "Plugin", "Callback", + "ChannelDescription", + "ChannelListener", "_ClientImpl", "_apply_headers", "_decode_user_metadata", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py new file mode 100644 index 000000000..09d5e1d23 --- /dev/null +++ b/temporalio/client/_channel.py @@ -0,0 +1,78 @@ +"""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 + +import temporalio.api.notification.v1 +from temporalio.workflow import Notification + +from ._callback import Callback + +__all__ = ["ChannelDescription", "ChannelListener"] + + +@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.""" diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 999ad8bba..5368e2f5d 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,187 @@ 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, + 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` + 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. + 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. + 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, + ) + ) + + async def poll_channel( + self, + channel: str, + *, + after_counter: int = 0, + wait: bool | timedelta = True, + max_notifications: int = 100, + 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. + 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. + 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, + ) + ) + + async def describe_channel( + self, + channel: str, + *, + rpc_metadata: Mapping[str, str | bytes] = {}, + rpc_timeout: timedelta | None = None, + ) -> ChannelDescription: + """Describe ``channel``: its listeners and what it retains. + + .. warning:: + This API is experimental and unstable. + + Args: + channel: Name of the channel, scoped to the namespace. + 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 + ) + ) + + async def register_channel_listener( + self, + channel: str, + callback: Callback, + *, + 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. + callback: The callback to invoke. + 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, + ) + ) + + async def unregister_channel_listener( + self, + channel: str, + listener_id: str, + *, + 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. + listener_id: The id :py:meth:`register_channel_listener` returned. + 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, + ) + ) + def create_nexus_client( self, service: type[NexusServiceType] | str, diff --git a/temporalio/client/_impl.py b/temporalio/client/_impl.py index abe3f8c09..36ea327e7 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, 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,120 @@ 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()), + ), + 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, + ) + 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 + ), + 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, + ) + + 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, + ), + 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, + ), + 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, + ) + async def _apply_headers( self, source: Mapping[str, temporalio.api.common.v1.Payload] | None, diff --git a/temporalio/client/_interceptor.py b/temporalio/client/_interceptor.py index 68077ebc9..b08246f2b 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,79 @@ 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 + + +@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 + + +@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 + + +@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 + + +@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 + + @dataclass class Interceptor: """Interceptor for clients. @@ -1001,3 +1078,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/tests/streams/conftest.py b/tests/streams/conftest.py index b8a152f20..a476d8f0c 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -1,5 +1,7 @@ import pytest +_ENVIRONMENTS_WITHOUT_CHANNELS = ("local", "time-skipping", "envconfig") + def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( @@ -15,3 +17,24 @@ 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", + ) + + +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. + if config.getoption("--workflow-environment") not in _ENVIRONMENTS_WITHOUT_CHANNELS: + return + skip = pytest.mark.skip( + reason="needs a server that serves notification channels; name one with -E" + ) + for item in items: + if item.get_closest_marker("needs_channel_server"): + item.add_marker(skip) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index a8b6a6d26..64965c789 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -2,14 +2,18 @@ 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 subscribe command. The live cases at the end need a server that does and +skip otherwise. """ from __future__ import annotations +import uuid from datetime import datetime, timedelta, timezone from typing import Any +import pytest + import temporalio.api.common.v1 import temporalio.api.notification.v1 import temporalio.bridge.proto.workflow_activation @@ -17,11 +21,13 @@ import temporalio.common import temporalio.converter from temporalio import workflow +from temporalio.client import Callback, Client 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 = ( @@ -222,3 +228,50 @@ 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 + + +@pytest.mark.needs_channel_server +async def test_a_workflow_receives_a_client_notification(client: Client): + channel = f"orders-{uuid.uuid4()}" + async with new_worker(client, ReceiveOne) as worker: + handle = await client.start_workflow( + ReceiveOne.run, + channel, + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + + async def listening() -> None: + description = await client.describe_channel(channel) + 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 handle.result() == { + "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 + + +@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 == [] From f8a99a9c0603b2db96f60c99c469d198f0830507 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 19:11:31 -0700 Subject: [PATCH 10/19] Made the live channel case fail fast on a Core without the command. A describe that races the subscribe is answered not found, and a worker whose Core refuses the subscribe command leaves the task in a timeout loop that the worker's shutdown waits on. The case now tolerates the first, names the second and bounds the shutdown. --- tests/streams/test_channels.py | 58 ++++++++++++++++++++++++++++------ 1 file changed, 49 insertions(+), 9 deletions(-) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index 64965c789..973a12c45 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -8,6 +8,8 @@ from __future__ import annotations +import asyncio +import contextlib import uuid from datetime import datetime, timedelta, timezone from typing import Any @@ -21,7 +23,9 @@ import temporalio.common import temporalio.converter from temporalio import workflow +from temporalio.api.enums.v1 import EventType from temporalio.client import Callback, Client +from temporalio.service import RPCError, RPCStatusCode from temporalio.worker._workflow_instance import ( UnsandboxedWorkflowRunner, WorkflowInstance, @@ -233,16 +237,40 @@ async def test_an_empty_channel_name_is_refused(): @pytest.mark.needs_channel_server async def test_a_workflow_receives_a_client_notification(client: Client): channel = f"orders-{uuid.uuid4()}" - async with new_worker(client, ReceiveOne) as worker: - handle = await client.start_workflow( - ReceiveOne.run, - channel, - id=f"wf-{uuid.uuid4()}", - task_queue=worker.task_queue, - ) + 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: - description = await client.describe_channel(channel) + # 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 ] @@ -252,7 +280,7 @@ async def listening() -> None: channel, position=b"1-0", counter=1, metadata={"topic": "inputs"} ) assert listeners == 1 - assert await handle.result() == { + assert await asyncio.wait_for(handle.result(), 30) == { "channel": channel, "counter": 1, "position": "1-0", @@ -262,6 +290,18 @@ async def listening() -> None: 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 From 82265749150a289ca9c7d194cb87750a32cd0ff1 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 19:18:06 -0700 Subject: [PATCH 11/19] Skipped the workflow channel case where the pinned Core lacks the command. The protos-only Core refuses the subscribe command, so the case is gated on a conftest constant the native layers flip with their pin. A client-only case covers what the server keeps for pollers: a notify with no listener is retained and answers zero, a counter at or below the latest is not kept, and a bounded poll above the latest comes back empty. --- tests/streams/conftest.py | 41 ++++++++++++++++++++++++++++------ tests/streams/test_channels.py | 40 +++++++++++++++++++++++++++++++++ 2 files changed, 74 insertions(+), 7 deletions(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index a476d8f0c..8515273bf 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -2,6 +2,11 @@ _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 + def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( @@ -22,6 +27,11 @@ 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_channel_core: the case needs the pinned Core to handle the " + "subscribe-notification-channel command", + ) def pytest_collection_modifyitems( @@ -30,11 +40,28 @@ def pytest_collection_modifyitems( # 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. - if config.getoption("--workflow-environment") not in _ENVIRONMENTS_WITHOUT_CHANNELS: - return - skip = pytest.mark.skip( - reason="needs a server that serves notification channels; name one with -E" - ) + 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" + ), + ) + ) for item in items: - if item.get_closest_marker("needs_channel_server"): - item.add_marker(skip) + 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 index 973a12c45..4521bb1ab 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -235,6 +235,7 @@ async def test_an_empty_channel_name_is_refused(): @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) @@ -315,3 +316,42 @@ async def test_a_callback_listener_registers_and_unregisters(client: Client): 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()}" + # 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 + 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 == [] From 4d049cd9c18bb6175ff91651692f6c16f5cb0751 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 19:33:45 -0700 Subject: [PATCH 12/19] Checked the untouched channel and the retained position in the poller case. A channel nobody has touched is not found, and the first notify on a fresh channel shows up in describe with its position and counter before any poll. --- tests/streams/test_channels.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index 4521bb1ab..950ad28d1 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -321,9 +321,18 @@ async def test_a_callback_listener_registers_and_unregisters(client: Client): @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. From 741c97b9999343fd777c0dd020cb1cc15b24bcb3 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 23:22:23 -0700 Subject: [PATCH 13/19] Added the channel linked to a workflow to the runtime and the client. --- temporalio/client/__init__.py | 2 + temporalio/client/_channel.py | 46 +++- temporalio/client/_client.py | 73 ++++- temporalio/client/_impl.py | 26 +- temporalio/client/_interceptor.py | 10 + temporalio/worker/_workflow_instance.py | 29 +- temporalio/workflow/__init__.py | 2 + temporalio/workflow/_channels.py | 72 ++++- temporalio/workflow/_context.py | 3 + tests/streams/conftest.py | 24 ++ tests/streams/test_channels.py | 338 +++++++++++++++++++++++- 11 files changed, 592 insertions(+), 33 deletions(-) diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 1a879923c..2577111d2 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -66,6 +66,7 @@ ) from ._channel import ( ChannelDescription, + ChannelKind, ChannelListener, ) from ._client import ( @@ -366,6 +367,7 @@ "Plugin", "Callback", "ChannelDescription", + "ChannelKind", "ChannelListener", "_ClientImpl", "_apply_headers", diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py index 09d5e1d23..d05d2875b 100644 --- a/temporalio/client/_channel.py +++ b/temporalio/client/_channel.py @@ -5,13 +5,43 @@ from collections.abc import Sequence from dataclasses import dataclass from datetime import datetime, timezone +from enum import IntEnum +import temporalio.api.common.v1 import temporalio.api.notification.v1 from temporalio.workflow import Notification from ._callback import Callback -__all__ = ["ChannelDescription", "ChannelListener"] +__all__ = ["ChannelDescription", "ChannelKind", "ChannelListener"] + + +class ChannelKind(IntEnum): + """Where a channel lives, which decides how a call addresses it. + + .. warning:: + This API is experimental and unstable. + """ + + UNSPECIFIED = int( + temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_UNSPECIFIED + ) + """The server did not say; an older server answers this.""" + + INDEPENDENT = int( + temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_INDEPENDENT + ) + """Its own execution, keyed by namespace and channel name. + + Any number of workflows subscribe to it and callbacks register on it. + """ + + LINKED = int(temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED) + """Kept in one workflow's state, keyed by namespace, workflow id and name. + + The owning workflow is its listener by construction; a call reaches it + with the ``workflow_id`` argument. + """ @dataclass(frozen=True) @@ -76,3 +106,17 @@ class ChannelDescription: 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.api.common.v1.WorkflowExecution | None = None + """The owner of a linked channel and the run that holds it. + + ``None`` for an independent channel. + """ diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 5368e2f5d..f8725afcc 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -2862,27 +2862,36 @@ async def notify_channel( position: bytes = b"", counter: int = 0, metadata: Mapping[str, Any] | 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` - 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. + 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. + channel: Name of the channel. Scoped to the namespace, or to the + workflow when ``workflow_id`` 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. + workflow_id: Address the channel linked to this workflow instead of + the independent channel of that name. + run_id: With ``workflow_id``, a run of its chain; the call reaches + the chain's current run, as a Signal does. Unset means the + current run under the workflow id. 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. @@ -2898,6 +2907,8 @@ async def notify_channel( metadata=metadata, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, + workflow_id=workflow_id, + run_id=run_id, ) ) @@ -2908,6 +2919,8 @@ async def poll_channel( after_counter: int = 0, wait: bool | timedelta = True, max_notifications: int = 100, + 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]: @@ -2917,13 +2930,18 @@ async def poll_channel( This API is experimental and unstable. Args: - channel: Name of the channel, scoped to the namespace. + channel: Name of the channel. Scoped to the namespace, or to the + workflow when ``workflow_id`` 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. + workflow_id: Address the channel linked to this workflow instead of + the independent channel of that name. + run_id: With ``workflow_id``, a run of its chain; the call reaches + the chain's current run. 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. @@ -2945,6 +2963,8 @@ async def poll_channel( max_notifications=max_notifications, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, + workflow_id=workflow_id, + run_id=run_id, ) ) @@ -2952,23 +2972,37 @@ async def describe_channel( self, channel: str, *, + 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 listeners and what it retains. + """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. + channel: Name of the channel. Scoped to the namespace, or to the + workflow when ``workflow_id`` is given. + workflow_id: Describe the channel linked to this workflow instead + of the independent channel of that name. A linked channel of a + running workflow 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. + run_id: With ``workflow_id``, a run of its chain; the call reaches + the chain's current run. 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 + channel=channel, + rpc_metadata=rpc_metadata, + rpc_timeout=rpc_timeout, + workflow_id=workflow_id, + run_id=run_id, ) ) @@ -2977,6 +3011,8 @@ async def register_channel_listener( channel: str, callback: Callback, *, + workflow_id: str | None = None, + run_id: str | None = None, rpc_metadata: Mapping[str, str | bytes] = {}, rpc_timeout: timedelta | None = None, ) -> str: @@ -2989,8 +3025,14 @@ async def register_channel_listener( This API is experimental and unstable. Args: - channel: Name of the channel, scoped to the namespace. + channel: Name of the channel. Scoped to the namespace, or to the + workflow when ``workflow_id`` is given. callback: The callback to invoke. + workflow_id: Listen on the channel linked to this workflow instead + of the independent channel of that name. The listener lives in + that workflow's state and ends with its run. + run_id: With ``workflow_id``, a run of its chain; the call reaches + the chain's current run. 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. @@ -3004,6 +3046,8 @@ async def register_channel_listener( callback=callback, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, + workflow_id=workflow_id, + run_id=run_id, ) ) @@ -3012,6 +3056,8 @@ async def unregister_channel_listener( channel: str, listener_id: str, *, + workflow_id: str | None = None, + run_id: str | None = None, rpc_metadata: Mapping[str, str | bytes] = {}, rpc_timeout: timedelta | None = None, ) -> None: @@ -3021,8 +3067,13 @@ async def unregister_channel_listener( This API is experimental and unstable. Args: - channel: Name of the channel, scoped to the namespace. + channel: Name of the channel. Scoped to the namespace, or to the + workflow when ``workflow_id`` is given. listener_id: The id :py:meth:`register_channel_listener` returned. + workflow_id: The workflow whose linked channel the listener is on, + when it was registered with one. + run_id: With ``workflow_id``, a run of its chain; the call reaches + the chain's current run. 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. @@ -3033,6 +3084,8 @@ async def unregister_channel_listener( listener_id=listener_id, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, + workflow_id=workflow_id, + run_id=run_id, ) ) diff --git a/temporalio/client/_impl.py b/temporalio/client/_impl.py index 36ea327e7..d29838195 100644 --- a/temporalio/client/_impl.py +++ b/temporalio/client/_impl.py @@ -56,7 +56,7 @@ ActivityHandle, AsyncActivityIDReference, ) -from ._channel import ChannelDescription, ChannelListener +from ._channel import ChannelDescription, ChannelKind, ChannelListener from ._exceptions import ( AsyncActivityCancelledError, ScheduleAlreadyRunningError, @@ -1787,6 +1787,7 @@ async def notify_channel(self, input: NotifyChannelInput) -> int: notification=notification, identity=self._client.identity, request_id=str(uuid.uuid4()), + workflow_execution=_channel_owner(input.workflow_id, input.run_id), ), retry=True, metadata=input.rpc_metadata, @@ -1802,6 +1803,7 @@ async def poll_channel( channel=input.channel, after_counter=input.after_counter, max_notifications=input.max_notifications, + workflow_execution=_channel_owner(input.workflow_id, input.run_id), ) if input.wait is not None: req.wait.FromTimedelta(input.wait) @@ -1813,7 +1815,9 @@ async def poll_channel( 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 + namespace=self._client.namespace, + channel=input.channel, + workflow_execution=_channel_owner(input.workflow_id, input.run_id), ), retry=True, metadata=input.rpc_metadata, @@ -1829,6 +1833,8 @@ async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescript else None ), retained_count=resp.retained_count, + kind=ChannelKind(resp.kind), + linked_to=resp.linked_to if resp.HasField("linked_to") else None, ) async def register_channel_listener( @@ -1845,6 +1851,7 @@ async def register_channel_listener( ), request_id=str(uuid.uuid4()), identity=self._client.identity, + workflow_execution=_channel_owner(input.workflow_id, input.run_id), ), retry=True, metadata=input.rpc_metadata, @@ -1861,6 +1868,7 @@ async def unregister_channel_listener( channel=input.channel, listener_id=input.listener_id, identity=self._client.identity, + workflow_execution=_channel_owner(input.workflow_id, input.run_id), ), retry=True, metadata=input.rpc_metadata, @@ -1884,6 +1892,7 @@ async def _notification_from_proto( position=proto.position, counter=proto.counter, metadata=metadata, + linked_to=proto.linked_to if proto.HasField("linked_to") else None, ) async def _apply_headers( @@ -1898,3 +1907,16 @@ async def _apply_headers( == HeaderCodecBehavior.CODEC, self._client.data_converter, ) + + +def _channel_owner( + workflow_id: str | None, run_id: str | None +) -> temporalio.api.common.v1.WorkflowExecution | None: + """The workflow a channel call addresses, or ``None`` for an independent channel.""" + if workflow_id is None: + if run_id is not None: + raise ValueError("run_id needs workflow_id") + return None + return temporalio.api.common.v1.WorkflowExecution( + workflow_id=workflow_id, run_id=run_id or "" + ) diff --git a/temporalio/client/_interceptor.py b/temporalio/client/_interceptor.py index b08246f2b..78fa1b6d9 100644 --- a/temporalio/client/_interceptor.py +++ b/temporalio/client/_interceptor.py @@ -716,6 +716,8 @@ class NotifyChannelInput: metadata: Mapping[str, Any] | None rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None + workflow_id: str | None = None + run_id: str | None = None @dataclass @@ -732,6 +734,8 @@ class PollChannelInput: max_notifications: int rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None + workflow_id: str | None = None + run_id: str | None = None @dataclass @@ -745,6 +749,8 @@ class DescribeChannelInput: channel: str rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None + workflow_id: str | None = None + run_id: str | None = None @dataclass @@ -759,6 +765,8 @@ class RegisterChannelListenerInput: callback: Callback rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None + workflow_id: str | None = None + run_id: str | None = None @dataclass @@ -773,6 +781,8 @@ class UnregisterChannelListenerInput: listener_id: str rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None + workflow_id: str | None = None + run_id: str | None = None @dataclass diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index d37a001c4..0ffebe5c1 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -312,6 +312,12 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: 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 @@ -880,14 +886,19 @@ def _apply_notifications_received( self, job: temporalio.bridge.proto.workflow_activation.NotificationsReceived ) -> None: for proto in job.notifications: - subscription = self._channel_subscriptions.get(proto.channel) + # 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 subscribed to is not the workflow's + # channel this run never asked for is not the workflow's # concern. logger.debug( - "Dropping a notification on channel %r, which this run has not " - "subscribed to", + "Dropping a notification on channel %r, which this run does " + "not listen on", proto.channel, ) continue @@ -1867,6 +1878,16 @@ def workflow_subscribe_channel( self._channel_subscriptions[channel] = subscription return subscription + 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 0cf3b058a..212625511 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -60,6 +60,7 @@ from ._channels import ( ChannelSubscription, Notification, + linked_channel, subscribe_channel, ) from ._context import ( @@ -271,6 +272,7 @@ "unsafe", "ChannelSubscription", "Notification", + "linked_channel", "subscribe_channel", "StreamReader", "StreamWriter", diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py index 8fd045519..fbbf6843d 100644 --- a/temporalio/workflow/_channels.py +++ b/temporalio/workflow/_channels.py @@ -21,7 +21,12 @@ import temporalio.api.notification.v1 from temporalio.workflow._context import _Runtime -__all__ = ["ChannelSubscription", "Notification", "subscribe_channel"] +__all__ = [ + "ChannelSubscription", + "Notification", + "linked_channel", + "subscribe_channel", +] @dataclass(frozen=True) @@ -55,6 +60,15 @@ class Notification: converter :py:func:`temporalio.workflow.payload_converter` returns. """ + linked_to: temporalio.api.common.v1.WorkflowExecution | None = None + """The workflow 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, @@ -64,30 +78,42 @@ def _from_proto( position=proto.position, counter=proto.counter, metadata=dict(proto.metadata.items()), + linked_to=proto.linked_to if proto.HasField("linked_to") else None, ) class ChannelSubscription: - """A workflow's subscription to one channel. - - Prefer :func:`temporalio.workflow.subscribe_channel`. The subscription 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 subscription share its buffer and - interleave. + """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. """ - def __init__(self, channel: str) -> None: - """Prefer :func:`temporalio.workflow.subscribe_channel`.""" + def __init__(self, channel: str, *, linked: bool = False) -> None: + """Prefer the two module functions named above.""" self._channel = channel + self._linked = linked self._pending: deque[Notification] = deque() self._waiters: deque[asyncio.Future[None]] = deque() @property def channel(self) -> str: - """The channel this subscription is on.""" + """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 + async def receive(self) -> Notification: """The next notification on this channel, waiting for one to arrive. @@ -142,3 +168,27 @@ def subscribe_channel(channel: str) -> ChannelSubscription: 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 e0cb59a8d..1371afc29 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -517,6 +517,9 @@ def workflow_streams(self) -> _WorkflowStreams: ... @abstractmethod def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... + @abstractmethod + def workflow_linked_channel(self, channel: str) -> ChannelSubscription: ... + @abstractmethod def workflow_time_ns(self) -> int: ... diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 8515273bf..22e3fa18e 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -7,6 +7,11 @@ # the native layers move to handles it. PINNED_CORE_HANDLES_CHANNEL_COMMAND = False +# A channel linked to a workflow needs no command or machine in Core, only +# the protos that carry `linked_to` on the notification, so the protos-only +# pin already handles it and the live linked cases run from this layer on. +PINNED_CORE_HANDLES_LINKED_CHANNEL = True + def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( @@ -32,6 +37,12 @@ def pytest_configure(config: pytest.Config) -> None: "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, and a Core whose protos carry " + "the linked contract", + ) def pytest_collection_modifyitems( @@ -61,6 +72,19 @@ def pytest_collection_modifyitems( ), ) ) + if ( + config.getoption("--workflow-environment") in _ENVIRONMENTS_WITHOUT_CHANNELS + or not PINNED_CORE_HANDLES_LINKED_CHANNEL + ): + skips.append( + ( + "needs_linked_server", + pytest.mark.skip( + reason="needs a server with channels linked to a workflow, " + "named with -E, and a Core pin with the linked contract" + ), + ) + ) 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 950ad28d1..ac22e2a4c 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -1,9 +1,11 @@ """The notification channel surface: the command, the delivery and the client calls. -The workflow instance is driven with activations directly, the way Core -drives it, because the dev server this chain tests against does not accept -the subscribe command. The live cases at the end need a server that does and -skip otherwise. +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, and skip otherwise. """ from __future__ import annotations @@ -18,13 +20,14 @@ import temporalio.api.common.v1 import temporalio.api.notification.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, Client +from temporalio.client import Callback, ChannelKind, Client from temporalio.service import RPCError, RPCStatusCode from temporalio.worker._workflow_instance import ( UnsandboxedWorkflowRunner, @@ -90,6 +93,89 @@ 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.workflow_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 + + def _instance(workflow_class: type) -> WorkflowInstance: """Build an instance the way the worker does, without a worker. @@ -158,6 +244,18 @@ def _notified(*notifications: Notification) -> WorkflowActivation: 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.api.common.v1.WorkflowExecution( + workflow_id="wf", run_id="run" + ), + ) + + def _subscribed(completion: WorkflowActivationCompletion) -> list[str]: assert completion.HasField("successful"), completion.failed.failure.message return [ @@ -234,6 +332,113 @@ async def test_an_empty_channel_name_is_refused(): 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]} + + +async def test_the_client_addresses_a_linked_channel_by_workflow( + client: Client, monkeypatch: pytest.MonkeyPatch +): + """Every channel call carries the owner it was given, and only then.""" + requests: list[Any] = [] + owner = temporalio.api.common.v1.WorkflowExecution(workflow_id="wf", run_id="run") + describe_response = temporalio.api.workflowservice.v1.DescribeChannelResponse( + kind=temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED, + linked_to=owner, + 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") + assert [req.workflow_execution for req in requests] == [ + owner, + temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), + temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), + temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), + temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), + ] + 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 + + # 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("workflow_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") + + @pytest.mark.needs_channel_server @pytest.mark.needs_channel_core async def test_a_workflow_receives_a_client_notification(client: Client): @@ -364,3 +569,126 @@ 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.workflow_id == handle.id + await client.notify_channel( + channel, + position=b"1-0", + counter=1, + metadata={"topic": "inputs"}, + workflow_id=handle.id, + ) + 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_of_a_running_workflow_exists_untouched(client: Client): + worker = new_worker(client, CountLinked) + running = asyncio.create_task(worker.run()) + handle = await client.start_workflow( + CountLinked.run, + f"orders-{uuid.uuid4()}", + id=f"wf-{uuid.uuid4()}", + task_queue=worker.task_queue, + ) + try: + untouched = f"untouched-{uuid.uuid4()}" + description = await client.describe_channel(untouched, 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.workflow_id == handle.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(untouched) + assert independent.value.status == RPCStatusCode.NOT_FOUND + finally: + await _stop(handle, worker, running) + + +@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: + await client.notify_channel( + channel, position=b"1-0", counter=1, workflow_id=handle.id + ) + 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.workflow_id == handle.id + 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 == [] + await client.notify_channel( + channel, position=b"2-0", counter=2, workflow_id=handle.id + ) + assert await asyncio.wait_for(handle.result(), 30) == 2 + finally: + await _stop(handle, worker, running) From 52f950862988506d5432815073f9b1658f572b4f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 23:47:11 -0700 Subject: [PATCH 14/19] Gated the linked receive cases on the Core that delivers the job. --- tests/streams/conftest.py | 37 ++++++++----- tests/streams/test_channels.py | 99 +++++++++++++++++++++++++--------- 2 files changed, 99 insertions(+), 37 deletions(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 22e3fa18e..13b2a999c 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -7,10 +7,12 @@ # the native layers move to handles it. PINNED_CORE_HANDLES_CHANNEL_COMMAND = False -# A channel linked to a workflow needs no command or machine in Core, only -# the protos that carry `linked_to` on the notification, so the protos-only -# pin already handles it and the live linked cases run from this layer on. -PINNED_CORE_HANDLES_LINKED_CHANNEL = True +# 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 def pytest_configure(config: pytest.Config) -> None: @@ -40,8 +42,12 @@ def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( "markers", "needs_linked_server: the case needs a server that serves channels linked " - "to a workflow, named with -E host:port, and a Core whose protos carry " - "the linked contract", + "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", ) @@ -72,16 +78,23 @@ def pytest_collection_modifyitems( ), ) ) - if ( - config.getoption("--workflow-environment") in _ENVIRONMENTS_WITHOUT_CHANNELS - or not PINNED_CORE_HANDLES_LINKED_CHANNEL - ): + 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, " - "named with -E, and a Core pin with the linked contract" + 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" ), ) ) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index ac22e2a4c..e7f3c87f8 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -584,6 +584,7 @@ async def _stop(handle: Any, worker: Any, running: asyncio.Task[None]) -> None: @pytest.mark.needs_linked_server +@pytest.mark.needs_linked_core async def test_a_workflow_receives_a_notification_on_its_linked_channel( client: Client, ): @@ -606,13 +607,15 @@ async def test_a_workflow_receives_a_notification_on_its_linked_channel( assert description.retained_count == 0 assert description.linked_to is not None assert description.linked_to.workflow_id == handle.id - await client.notify_channel( + # 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, @@ -632,33 +635,59 @@ async def test_a_workflow_receives_a_notification_on_its_linked_channel( @pytest.mark.needs_linked_server -async def test_a_linked_channel_of_a_running_workflow_exists_untouched(client: Client): - worker = new_worker(client, CountLinked) - running = asyncio.create_task(worker.run()) +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, - f"orders-{uuid.uuid4()}", + channel, id=f"wf-{uuid.uuid4()}", - task_queue=worker.task_queue, + task_queue=f"nobody-polls-{uuid.uuid4()}", ) - try: - untouched = f"untouched-{uuid.uuid4()}" - description = await client.describe_channel(untouched, 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.workflow_id == handle.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(untouched) - assert independent.value.status == RPCStatusCode.NOT_FOUND - finally: - await _stop(handle, worker, running) + 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.workflow_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_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) @@ -670,13 +699,26 @@ async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): task_queue=worker.task_queue, ) try: - await client.notify_channel( - channel, position=b"1-0", counter=1, workflow_id=handle.id + 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.workflow_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 @@ -686,9 +728,16 @@ async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): channel, workflow_id=handle.id, after_counter=1, wait=timedelta(seconds=1) ) assert polled == [] - await client.notify_channel( - channel, position=b"2-0", counter=2, workflow_id=handle.id + 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 + polled = await client.poll_channel( + channel, workflow_id=handle.id, after_counter=1, wait=False + ) + assert [n.counter for n in polled] == [2] finally: await _stop(handle, worker, running) From 77c4289298ca2c3de6964d01fb474c2f3f95d2ec Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 23:51:26 -0700 Subject: [PATCH 15/19] Expected not found from a poll on a closed workflow's linked channel. --- tests/streams/test_channels.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index e7f3c87f8..687b5c63a 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -735,9 +735,12 @@ async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): == 1 ) assert await asyncio.wait_for(handle.result(), 30) == 2 - polled = await client.poll_channel( - channel, workflow_id=handle.id, after_counter=1, wait=False - ) - assert [n.counter for n in polled] == [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) From 5a85fb2a025484959a893b66c150dac6c3fa04fe Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:26:53 -0700 Subject: [PATCH 16/19] Added channel unsubscribe, subscriptions on describe and stream_channel. A subscription used to end only with the run. The handle now records the unsubscribe command once and closes, so a workflow can rotate channels under the per-run cap. The description lists what a run stands on, and stream_channel derives the channel a native stream notifies so a client can follow a stream without asking the server. --- temporalio/client/__init__.py | 6 + temporalio/client/_channel.py | 127 ++++++- temporalio/client/_workflow.py | 17 + temporalio/worker/_workflow_instance.py | 7 + temporalio/workflow/_channels.py | 70 +++- temporalio/workflow/_context.py | 9 + tests/streams/conftest.py | 48 +++ tests/streams/test_channels.py | 453 +++++++++++++++++++++++- 8 files changed, 725 insertions(+), 12 deletions(-) diff --git a/temporalio/client/__init__.py b/temporalio/client/__init__.py index 2577111d2..7ed7187be 100644 --- a/temporalio/client/__init__.py +++ b/temporalio/client/__init__.py @@ -65,9 +65,12 @@ Callback, ) from ._channel import ( + ChannelAddress, ChannelDescription, ChannelKind, ChannelListener, + ChannelSubscriptionInfo, + stream_channel, ) from ._client import ( Client, @@ -369,6 +372,9 @@ "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 index d05d2875b..5ae2d3371 100644 --- a/temporalio/client/_channel.py +++ b/temporalio/client/_channel.py @@ -9,11 +9,23 @@ import temporalio.api.common.v1 import temporalio.api.notification.v1 +import temporalio.api.workflow.v1 +from temporalio.streams._ref import StreamRef from temporalio.workflow import Notification from ._callback import Callback -__all__ = ["ChannelDescription", "ChannelKind", "ChannelListener"] +__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): @@ -120,3 +132,116 @@ class ChannelDescription: ``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.""" + + workflow_id: str | None + """The workflow 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 ``workflow_id``. + """ + + +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 independent for a standalone activity, which has no + linked channels of its own. A standalone stream notifies the independent + channel ``stream/``, whatever the topic, since its topics share + one stream on the server. + + 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, ref.workflow_id) + if ref.kind == "activity": + assert ref.activity_id is not None + return ChannelAddress( + f"{STREAM_CHANNEL_PREFIX}{ref.activity_id}/{ref.topic}", ref.workflow_id + ) + assert ref.stream_id is not None + return ChannelAddress(STREAM_CHANNEL_PREFIX + ref.stream_id, None) 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/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 0ffebe5c1..13f2a1490 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -1878,6 +1878,13 @@ def workflow_subscribe_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: diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py index fbbf6843d..04d6fb836 100644 --- a/temporalio/workflow/_channels.py +++ b/temporalio/workflow/_channels.py @@ -90,13 +90,15 @@ class ChannelSubscription: 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. + 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() @@ -114,13 +116,68 @@ def linked(self) -> bool: """ 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: @@ -130,16 +187,11 @@ async def receive(self) -> Notification: self._waiters.remove(waiter) return self._pending.popleft() - def __aiter__(self) -> ChannelSubscription: - """The subscription is its own iterator.""" - return self - - async def __anext__(self) -> Notification: - """The next notification; the iteration never ends on its own.""" - return await self.receive() - 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: diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 1371afc29..a04791baf 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -517,6 +517,15 @@ 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: ... diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 13b2a999c..8af40d565 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -14,6 +14,10 @@ # 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( @@ -49,6 +53,26 @@ def pytest_configure(config: pytest.Config) -> None: "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", + ) def pytest_collection_modifyitems( @@ -98,6 +122,30 @@ def pytest_collection_modifyitems( ), ) ) + 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"), + ): + 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): diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index 687b5c63a..3a3ac88ee 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -5,7 +5,8 @@ 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, and skip otherwise. +with the linked kind, one whose describe lists the subscriptions, or one that +accepts the unsubscribe, and skip otherwise. """ from __future__ import annotations @@ -19,7 +20,9 @@ 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 @@ -27,8 +30,17 @@ import temporalio.converter from temporalio import workflow from temporalio.api.enums.v1 import EventType -from temporalio.client import Callback, ChannelKind, Client +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, @@ -176,6 +188,80 @@ async def run(self, channel: str) -> int: 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. @@ -256,6 +342,12 @@ 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 [ @@ -265,6 +357,24 @@ def _subscribed(completion: WorkflowActivationCompletion) -> list[str]: ] +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( @@ -374,6 +484,168 @@ async def test_the_owner_on_a_notification_picks_the_handle_of_its_kind(): 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(): + assert stream_channel(StreamRef.for_workflow("wf", topic="out")) == ChannelAddress( + "stream/out", "wf" + ) + assert stream_channel(StreamRef.for_workflow("wf", run_id="run")) == ChannelAddress( + "stream/" + DEFAULT_TOPIC, "wf" + ) + assert stream_channel( + StreamRef.for_activity("act", workflow_id="wf", topic="out") + ) == ChannelAddress("stream/act/out", "wf") + # A standalone activity is an execution of its own with no linked channels. + assert stream_channel(StreamRef.for_activity("act", topic="out")) == ChannelAddress( + "stream/act/out", None + ) + # 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) + + +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_workflow( client: Client, monkeypatch: pytest.MonkeyPatch ): @@ -744,3 +1016,180 @@ async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): 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) + assert (info.subscribed_event_id, info.last_counter) == (0, 0) + assert (info.listener_count, info.retained_count, info.accepted_count) == ( + 0, + 1, + 1, + ) + # Nobody polls. The notification either rode the task the notify + # scheduled or waits behind the first task, which was scheduled + # without a counter when the run started. + pending = ( + info.pending_notification.counter if info.pending_notification else None + ) + assert (info.scheduled_counter, pending) in {(1, None), (0, 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 + 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) From b4e41f23183a22859f663571eedf8a4b93de039a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:44:08 -0700 Subject: [PATCH 17/19] Pinned the linked channel's describe fields to the server's timing. The owner's state takes a notification in the write that accepts it, so the counter is the run's at once, and with the first task already scheduled the notification waits behind it as the pending entry. A closed run keeps its listing. --- tests/streams/test_channels.py | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index 3a3ac88ee..c41ac75dc 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -1114,19 +1114,21 @@ async def test_a_description_lists_a_linked_channel_once_it_holds_state( ) [info] = (await handle.describe()).channel_subscriptions assert (info.channel, info.kind) == (channel, ChannelKind.LINKED) - assert (info.subscribed_event_id, info.last_counter) == (0, 0) + # 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.workflow_id == handle.id + assert info.scheduled_counter == 0 assert (info.listener_count, info.retained_count, info.accepted_count) == ( 0, 1, 1, ) - # Nobody polls. The notification either rode the task the notify - # scheduled or waits behind the first task, which was scheduled - # without a counter when the run started. - pending = ( - info.pending_notification.counter if info.pending_notification else None - ) - assert (info.scheduled_counter, pending) in {(1, None), (0, 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( @@ -1139,6 +1141,10 @@ async def test_a_description_lists_a_linked_channel_once_it_holds_state( ) [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() From d6b39adaa6ac35d58eef8b003dab3f0108800e06 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 17:56:13 -0700 Subject: [PATCH 18/19] Addressed a linked channel by execution on the client. A standalone activity can own a channel too, so `temporalio.common.Execution` names the owner of a linked channel, the five channel calls take it, and `workflow_id` with `run_id` stays as the shorthand for a workflow owner. `stream_channel` names a standalone activity's channel as `stream/` linked to the activity execution. --- temporalio/client/_channel.py | 55 ++++++++++---- temporalio/client/_client.py | 78 ++++++++++++-------- temporalio/client/_impl.py | 50 +++++++++---- temporalio/client/_interceptor.py | 5 ++ temporalio/common.py | 58 +++++++++++++++ temporalio/workflow/_channels.py | 11 ++- tests/streams/conftest.py | 7 ++ tests/streams/test_channels.py | 117 ++++++++++++++++++++++++------ 8 files changed, 298 insertions(+), 83 deletions(-) diff --git a/temporalio/client/_channel.py b/temporalio/client/_channel.py index 5ae2d3371..f3996bd8e 100644 --- a/temporalio/client/_channel.py +++ b/temporalio/client/_channel.py @@ -7,9 +7,9 @@ from datetime import datetime, timezone from enum import IntEnum -import temporalio.api.common.v1 import temporalio.api.notification.v1 import temporalio.api.workflow.v1 +import temporalio.common from temporalio.streams._ref import StreamRef from temporalio.workflow import Notification @@ -49,10 +49,11 @@ class ChannelKind(IntEnum): """ LINKED = int(temporalio.api.notification.v1.ChannelKind.CHANNEL_KIND_LINKED) - """Kept in one workflow's state, keyed by namespace, workflow id and name. + """Kept in one execution's state, keyed by namespace, execution and name. - The owning workflow is its listener by construction; a call reaches it - with the ``workflow_id`` argument. + 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. """ @@ -127,7 +128,7 @@ class ChannelDescription: listeners and nothing retained for a name nobody has notified yet. """ - linked_to: temporalio.api.common.v1.WorkflowExecution | None = None + linked_to: temporalio.common.Execution | None = None """The owner of a linked channel and the run that holds it. ``None`` for an independent channel. @@ -210,13 +211,27 @@ class ChannelAddress: channel: str """The channel name.""" - workflow_id: str | None - """The workflow the channel is linked to, or ``None`` for an independent one. + 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 ``workflow_id``. + 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. @@ -225,11 +240,12 @@ def stream_channel(ref: StreamRef) -> ChannelAddress: 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 independent for a standalone activity, which has no - linked channels of its own. A standalone stream notifies the independent - channel ``stream/``, whatever the topic, since its topics share - one stream on the server. + ``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`` @@ -237,11 +253,20 @@ def stream_channel(ref: StreamRef) -> ChannelAddress: """ if ref.kind == "workflow": assert ref.workflow_id is not None - return ChannelAddress(STREAM_CHANNEL_PREFIX + ref.topic, ref.workflow_id) + 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}", ref.workflow_id + 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 f8725afcc..29de96098 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -2862,6 +2862,7 @@ async def notify_channel( 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] = {}, @@ -2880,18 +2881,20 @@ async def notify_channel( Args: channel: Name of the channel. Scoped to the namespace, or to the - workflow when ``workflow_id`` is given. + 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. - workflow_id: Address the channel linked to this workflow instead of - the independent channel of that name. - run_id: With ``workflow_id``, a run of its chain; the call reaches - the chain's current run, as a Signal does. Unset means the - current run under the workflow id. + 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. @@ -2909,6 +2912,7 @@ async def notify_channel( rpc_timeout=rpc_timeout, workflow_id=workflow_id, run_id=run_id, + execution=execution, ) ) @@ -2919,6 +2923,7 @@ async def poll_channel( 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] = {}, @@ -2931,17 +2936,20 @@ async def poll_channel( Args: channel: Name of the channel. Scoped to the namespace, or to the - workflow when ``workflow_id`` is given. + 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. - workflow_id: Address the channel linked to this workflow instead of - the independent channel of that name. - run_id: With ``workflow_id``, a run of its chain; the call reaches - the chain's current run. + 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. @@ -2965,6 +2973,7 @@ async def poll_channel( rpc_timeout=rpc_timeout, workflow_id=workflow_id, run_id=run_id, + execution=execution, ) ) @@ -2972,6 +2981,7 @@ 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] = {}, @@ -2984,14 +2994,16 @@ async def describe_channel( Args: channel: Name of the channel. Scoped to the namespace, or to the - workflow when ``workflow_id`` is given. - workflow_id: Describe the channel linked to this workflow instead - of the independent channel of that name. A linked channel of a - running workflow 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. - run_id: With ``workflow_id``, a run of its chain; the call reaches - the chain's current run. + 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. @@ -3003,6 +3015,7 @@ async def describe_channel( rpc_timeout=rpc_timeout, workflow_id=workflow_id, run_id=run_id, + execution=execution, ) ) @@ -3011,6 +3024,7 @@ async def register_channel_listener( 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] = {}, @@ -3026,13 +3040,15 @@ async def register_channel_listener( Args: channel: Name of the channel. Scoped to the namespace, or to the - workflow when ``workflow_id`` is given. + execution when one is given. callback: The callback to invoke. - workflow_id: Listen on the channel linked to this workflow instead - of the independent channel of that name. The listener lives in - that workflow's state and ends with its run. - run_id: With ``workflow_id``, a run of its chain; the call reaches - the chain's current run. + 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. @@ -3048,6 +3064,7 @@ async def register_channel_listener( rpc_timeout=rpc_timeout, workflow_id=workflow_id, run_id=run_id, + execution=execution, ) ) @@ -3056,6 +3073,7 @@ async def unregister_channel_listener( 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] = {}, @@ -3068,12 +3086,13 @@ async def unregister_channel_listener( Args: channel: Name of the channel. Scoped to the namespace, or to the - workflow when ``workflow_id`` is given. + execution when one is given. listener_id: The id :py:meth:`register_channel_listener` returned. - workflow_id: The workflow whose linked channel the listener is on, + execution: The execution whose linked channel the listener is on, when it was registered with one. - run_id: With ``workflow_id``, a run of its chain; the call reaches - the chain's current 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. @@ -3086,6 +3105,7 @@ async def unregister_channel_listener( rpc_timeout=rpc_timeout, workflow_id=workflow_id, run_id=run_id, + execution=execution, ) ) diff --git a/temporalio/client/_impl.py b/temporalio/client/_impl.py index d29838195..f21b7df54 100644 --- a/temporalio/client/_impl.py +++ b/temporalio/client/_impl.py @@ -1787,7 +1787,9 @@ async def notify_channel(self, input: NotifyChannelInput) -> int: notification=notification, identity=self._client.identity, request_id=str(uuid.uuid4()), - workflow_execution=_channel_owner(input.workflow_id, input.run_id), + execution=_channel_owner( + input.execution, input.workflow_id, input.run_id + ), ), retry=True, metadata=input.rpc_metadata, @@ -1803,7 +1805,7 @@ async def poll_channel( channel=input.channel, after_counter=input.after_counter, max_notifications=input.max_notifications, - workflow_execution=_channel_owner(input.workflow_id, input.run_id), + execution=_channel_owner(input.execution, input.workflow_id, input.run_id), ) if input.wait is not None: req.wait.FromTimedelta(input.wait) @@ -1817,7 +1819,9 @@ async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescript temporalio.api.workflowservice.v1.DescribeChannelRequest( namespace=self._client.namespace, channel=input.channel, - workflow_execution=_channel_owner(input.workflow_id, input.run_id), + execution=_channel_owner( + input.execution, input.workflow_id, input.run_id + ), ), retry=True, metadata=input.rpc_metadata, @@ -1834,7 +1838,11 @@ async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescript ), retained_count=resp.retained_count, kind=ChannelKind(resp.kind), - linked_to=resp.linked_to if resp.HasField("linked_to") else None, + linked_to=( + temporalio.common.Execution.from_proto(resp.linked_to) + if resp.HasField("linked_to") + else None + ), ) async def register_channel_listener( @@ -1851,7 +1859,9 @@ async def register_channel_listener( ), request_id=str(uuid.uuid4()), identity=self._client.identity, - workflow_execution=_channel_owner(input.workflow_id, input.run_id), + execution=_channel_owner( + input.execution, input.workflow_id, input.run_id + ), ), retry=True, metadata=input.rpc_metadata, @@ -1868,7 +1878,9 @@ async def unregister_channel_listener( channel=input.channel, listener_id=input.listener_id, identity=self._client.identity, - workflow_execution=_channel_owner(input.workflow_id, input.run_id), + execution=_channel_owner( + input.execution, input.workflow_id, input.run_id + ), ), retry=True, metadata=input.rpc_metadata, @@ -1892,7 +1904,11 @@ async def _notification_from_proto( position=proto.position, counter=proto.counter, metadata=metadata, - linked_to=proto.linked_to if proto.HasField("linked_to") else None, + linked_to=( + temporalio.common.Execution.from_proto(proto.linked_to) + if proto.HasField("linked_to") + else None + ), ) async def _apply_headers( @@ -1910,13 +1926,21 @@ async def _apply_headers( def _channel_owner( - workflow_id: str | None, run_id: str | None -) -> temporalio.api.common.v1.WorkflowExecution | None: - """The workflow a channel call addresses, or ``None`` for an independent channel.""" + execution: temporalio.common.Execution | None, + workflow_id: str | None, + run_id: str | None, +) -> temporalio.api.common.v1.Execution | None: + """The execution a channel call addresses, or ``None`` for an independent channel. + + ``workflow_id`` and ``run_id`` are the shorthand for a workflow owner, so + they do not combine with ``execution``. + """ + 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.to_proto() if workflow_id is None: if run_id is not None: raise ValueError("run_id needs workflow_id") return None - return temporalio.api.common.v1.WorkflowExecution( - workflow_id=workflow_id, run_id=run_id or "" - ) + return temporalio.common.Execution.workflow(workflow_id, run_id).to_proto() diff --git a/temporalio/client/_interceptor.py b/temporalio/client/_interceptor.py index 78fa1b6d9..b0282fe38 100644 --- a/temporalio/client/_interceptor.py +++ b/temporalio/client/_interceptor.py @@ -718,6 +718,7 @@ class NotifyChannelInput: rpc_timeout: timedelta | None workflow_id: str | None = None run_id: str | None = None + execution: temporalio.common.Execution | None = None @dataclass @@ -736,6 +737,7 @@ class PollChannelInput: rpc_timeout: timedelta | None workflow_id: str | None = None run_id: str | None = None + execution: temporalio.common.Execution | None = None @dataclass @@ -751,6 +753,7 @@ class DescribeChannelInput: rpc_timeout: timedelta | None workflow_id: str | None = None run_id: str | None = None + execution: temporalio.common.Execution | None = None @dataclass @@ -767,6 +770,7 @@ class RegisterChannelListenerInput: rpc_timeout: timedelta | None workflow_id: str | None = None run_id: str | None = None + execution: temporalio.common.Execution | None = None @dataclass @@ -783,6 +787,7 @@ class UnregisterChannelListenerInput: rpc_timeout: timedelta | None workflow_id: str | None = None run_id: str | None = None + execution: temporalio.common.Execution | None = None @dataclass diff --git a/temporalio/common.py b/temporalio/common.py index 04081bf19..cf126ff95 100644 --- a/temporalio/common.py +++ b/temporalio/common.py @@ -18,6 +18,7 @@ Generic, TypeAlias, TypeVar, + cast, get_origin, get_type_hints, overload, @@ -108,6 +109,63 @@ def _validate(self) -> None: raise ValueError("Maximum attempts cannot be negative") +class ExecutionType(IntEnum): + """What kind of execution an :class:`Execution` names. + + .. warning:: + This API is experimental and unstable. + """ + + UNSPECIFIED = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_UNSPECIFIED) + WORKFLOW = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_WORKFLOW) + ACTIVITY = int(temporalio.api.enums.v1.ExecutionType.EXECUTION_TYPE_ACTIVITY) + """A standalone activity, one started by a client rather than a workflow.""" + + +@dataclass(frozen=True) +class Execution: + """One execution in a namespace: a workflow or a standalone activity. + + ``business_id`` is the id the caller chose, the workflow id or the + activity id. ``run_id`` pins one run of it. Unset, a call reaches the + current run of a workflow chain, as a Signal does. + + .. warning:: + This API is experimental and unstable. + """ + + type: ExecutionType + business_id: str + run_id: str | None = None + + @classmethod + def workflow(cls, workflow_id: str, run_id: str | None = None) -> Execution: + """A workflow execution.""" + return cls(ExecutionType.WORKFLOW, workflow_id, run_id) + + @classmethod + def activity(cls, activity_id: str, run_id: str | None = None) -> Execution: + """A standalone activity execution.""" + return cls(ExecutionType.ACTIVITY, activity_id, run_id) + + def to_proto(self) -> temporalio.api.common.v1.Execution: + """This execution as the API names it.""" + return temporalio.api.common.v1.Execution( + type=cast( + "temporalio.api.enums.v1.ExecutionType.ValueType", int(self.type) + ), + business_id=self.business_id, + run_id=self.run_id or "", + ) + + @staticmethod + def from_proto(proto: temporalio.api.common.v1.Execution) -> Execution: + """From the API's form. An empty run id reads as unset.""" + return Execution( + ExecutionType(proto.type), proto.business_id, proto.run_id or None + ) + + class WorkflowIDReusePolicy(IntEnum): """How already-in-use workflow IDs are handled on start. diff --git a/temporalio/workflow/_channels.py b/temporalio/workflow/_channels.py index 04d6fb836..929f99b2d 100644 --- a/temporalio/workflow/_channels.py +++ b/temporalio/workflow/_channels.py @@ -19,6 +19,7 @@ import temporalio.api.common.v1 import temporalio.api.notification.v1 +import temporalio.common from temporalio.workflow._context import _Runtime __all__ = [ @@ -60,8 +61,8 @@ class Notification: converter :py:func:`temporalio.workflow.payload_converter` returns. """ - linked_to: temporalio.api.common.v1.WorkflowExecution | None = None - """The workflow a linked channel belongs to, and the run that received this. + 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 @@ -78,7 +79,11 @@ def _from_proto( position=proto.position, counter=proto.counter, metadata=dict(proto.metadata.items()), - linked_to=proto.linked_to if proto.HasField("linked_to") else None, + linked_to=( + temporalio.common.Execution.from_proto(proto.linked_to) + if proto.HasField("linked_to") + else None + ), ) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 8af40d565..049fa2da1 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -73,6 +73,12 @@ def pytest_configure(config: pytest.Config) -> None: "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( @@ -127,6 +133,7 @@ def pytest_collection_modifyitems( ("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( ( diff --git a/tests/streams/test_channels.py b/tests/streams/test_channels.py index c41ac75dc..afcdc4a44 100644 --- a/tests/streams/test_channels.py +++ b/tests/streams/test_channels.py @@ -127,7 +127,7 @@ def _describe(notification: workflow.Notification) -> dict[str, Any]: else None ), "owner": ( - notification.linked_to.workflow_id if notification.linked_to else None + notification.linked_to.business_id if notification.linked_to else None ), "owner_run": notification.linked_to.run_id if notification.linked_to else None, } @@ -336,9 +336,7 @@ def _linked(channel: str, counter: int, position: bytes = b"") -> Notification: channel=channel, counter=counter, position=position, - linked_to=temporalio.api.common.v1.WorkflowExecution( - workflow_id="wf", run_id="run" - ), + linked_to=temporalio.common.Execution.workflow("wf", "run").to_proto(), ) @@ -557,24 +555,32 @@ async def test_a_subscription_after_an_unsubscribe_is_a_new_one(): 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", "wf" + "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, "wf" + "stream/" + DEFAULT_TOPIC, owner ) assert stream_channel( StreamRef.for_activity("act", workflow_id="wf", topic="out") - ) == ChannelAddress("stream/act/out", "wf") - # A standalone activity is an execution of its own with no linked channels. + ) == 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/act/out", None + "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(): @@ -646,15 +652,15 @@ async def test_the_description_maps_every_channel_subscription_field(): assert description.channel_subscriptions == () -async def test_the_client_addresses_a_linked_channel_by_workflow( +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.api.common.v1.WorkflowExecution(workflow_id="wf", run_id="run") + 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, + linked_to=owner.to_proto(), latest=_linked("orders", 3, b"3-0"), ) responses = { @@ -686,18 +692,33 @@ async def call(req: Any, *, _response: Any = response, **_: Any) -> Any: 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") - assert [req.workflow_execution for req in requests] == [ - owner, - temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), - temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), - temporalio.api.common.v1.WorkflowExecution(workflow_id="wf"), - temporalio.api.common.v1.WorkflowExecution(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) @@ -705,10 +726,17 @@ async def call(req: Any, *, _response: Any = response, **_: Any) -> Any: await client.describe_channel("orders") await client.register_channel_listener("orders", callback) await client.unregister_channel_listener("orders", "listener") - assert [req.HasField("workflow_execution") for req in requests] == [False] * 5 + 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 @@ -878,7 +906,7 @@ async def test_a_workflow_receives_a_notification_on_its_linked_channel( assert description.latest is None assert description.retained_count == 0 assert description.linked_to is not None - assert description.linked_to.workflow_id == handle.id + assert description.linked_to.business_id == handle.id # The owner is the one listener. listeners = await client.notify_channel( channel, @@ -924,7 +952,7 @@ async def test_a_linked_channel_lives_and_dies_with_its_workflow(client: Client) assert (description.listeners, description.latest) == ([], None) assert description.retained_count == 0 assert description.linked_to is not None - assert (description.linked_to.workflow_id, description.linked_to.run_id) == ( + assert (description.linked_to.business_id, description.linked_to.run_id) == ( handle.id, run_id, ) @@ -958,6 +986,49 @@ async def test_a_linked_channel_lives_and_dies_with_its_workflow(client: Client) 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): @@ -980,7 +1051,7 @@ async def test_a_linked_channel_is_polled_by_workflow_id(client: Client): 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.workflow_id == handle.id + assert polled[0].linked_to.business_id == handle.id # The run id reaches the same channel. assert [ n.counter @@ -1122,7 +1193,7 @@ async def test_a_description_lists_a_linked_channel_once_it_holds_state( 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.workflow_id == handle.id + 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, From 7cd9b3a938cc7156a266966505103ad34534d93a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 18:09:57 -0700 Subject: [PATCH 19/19] Carried one execution field on the channel inputs. The interceptor inputs name the owner once, as `execution`, and the client resolves the `workflow_id` and `run_id` shorthand before building them, so an interceptor sees one spelling. --- temporalio/client/_client.py | 41 ++++++++++++++++++++----------- temporalio/client/_impl.py | 36 ++++++--------------------- temporalio/client/_interceptor.py | 10 -------- 3 files changed, 33 insertions(+), 54 deletions(-) diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 29de96098..63a1dace5 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -2910,9 +2910,7 @@ async def notify_channel( metadata=metadata, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, - workflow_id=workflow_id, - run_id=run_id, - execution=execution, + execution=_channel_execution(execution, workflow_id, run_id), ) ) @@ -2971,9 +2969,7 @@ async def poll_channel( max_notifications=max_notifications, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, - workflow_id=workflow_id, - run_id=run_id, - execution=execution, + execution=_channel_execution(execution, workflow_id, run_id), ) ) @@ -3013,9 +3009,7 @@ async def describe_channel( channel=channel, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, - workflow_id=workflow_id, - run_id=run_id, - execution=execution, + execution=_channel_execution(execution, workflow_id, run_id), ) ) @@ -3062,9 +3056,7 @@ async def register_channel_listener( callback=callback, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, - workflow_id=workflow_id, - run_id=run_id, - execution=execution, + execution=_channel_execution(execution, workflow_id, run_id), ) ) @@ -3103,9 +3095,7 @@ async def unregister_channel_listener( listener_id=listener_id, rpc_metadata=rpc_metadata, rpc_timeout=rpc_timeout, - workflow_id=workflow_id, - run_id=run_id, - execution=execution, + execution=_channel_execution(execution, workflow_id, run_id), ) ) @@ -3300,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 f21b7df54..4be55aa90 100644 --- a/temporalio/client/_impl.py +++ b/temporalio/client/_impl.py @@ -1787,9 +1787,7 @@ async def notify_channel(self, input: NotifyChannelInput) -> int: notification=notification, identity=self._client.identity, request_id=str(uuid.uuid4()), - execution=_channel_owner( - input.execution, input.workflow_id, input.run_id - ), + execution=_channel_owner(input.execution), ), retry=True, metadata=input.rpc_metadata, @@ -1805,7 +1803,7 @@ async def poll_channel( channel=input.channel, after_counter=input.after_counter, max_notifications=input.max_notifications, - execution=_channel_owner(input.execution, input.workflow_id, input.run_id), + execution=_channel_owner(input.execution), ) if input.wait is not None: req.wait.FromTimedelta(input.wait) @@ -1819,9 +1817,7 @@ async def describe_channel(self, input: DescribeChannelInput) -> ChannelDescript temporalio.api.workflowservice.v1.DescribeChannelRequest( namespace=self._client.namespace, channel=input.channel, - execution=_channel_owner( - input.execution, input.workflow_id, input.run_id - ), + execution=_channel_owner(input.execution), ), retry=True, metadata=input.rpc_metadata, @@ -1859,9 +1855,7 @@ async def register_channel_listener( ), request_id=str(uuid.uuid4()), identity=self._client.identity, - execution=_channel_owner( - input.execution, input.workflow_id, input.run_id - ), + execution=_channel_owner(input.execution), ), retry=True, metadata=input.rpc_metadata, @@ -1878,9 +1872,7 @@ async def unregister_channel_listener( channel=input.channel, listener_id=input.listener_id, identity=self._client.identity, - execution=_channel_owner( - input.execution, input.workflow_id, input.run_id - ), + execution=_channel_owner(input.execution), ), retry=True, metadata=input.rpc_metadata, @@ -1927,20 +1919,6 @@ async def _apply_headers( def _channel_owner( execution: temporalio.common.Execution | None, - workflow_id: str | None, - run_id: str | None, ) -> temporalio.api.common.v1.Execution | None: - """The execution a channel call addresses, or ``None`` for an independent channel. - - ``workflow_id`` and ``run_id`` are the shorthand for a workflow owner, so - they do not combine with ``execution``. - """ - 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.to_proto() - 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).to_proto() + """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 b0282fe38..c35a1281f 100644 --- a/temporalio/client/_interceptor.py +++ b/temporalio/client/_interceptor.py @@ -716,8 +716,6 @@ class NotifyChannelInput: metadata: Mapping[str, Any] | None rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None - workflow_id: str | None = None - run_id: str | None = None execution: temporalio.common.Execution | None = None @@ -735,8 +733,6 @@ class PollChannelInput: max_notifications: int rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None - workflow_id: str | None = None - run_id: str | None = None execution: temporalio.common.Execution | None = None @@ -751,8 +747,6 @@ class DescribeChannelInput: channel: str rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None - workflow_id: str | None = None - run_id: str | None = None execution: temporalio.common.Execution | None = None @@ -768,8 +762,6 @@ class RegisterChannelListenerInput: callback: Callback rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None - workflow_id: str | None = None - run_id: str | None = None execution: temporalio.common.Execution | None = None @@ -785,8 +777,6 @@ class UnregisterChannelListenerInput: listener_id: str rpc_metadata: Mapping[str, str | bytes] rpc_timeout: timedelta | None - workflow_id: str | None = None - run_id: str | None = None execution: temporalio.common.Execution | None = None