From fdd81d3ce3f75c5e19bc6c9c2d56cd5d88678f4a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:51:59 -0700 Subject: [PATCH 1/3] Added workflow.stream_reader and workflow.stream_writer. Workflow code reads and writes a stream through the provider's workflow half, made once per instance so its state dies with the instance. --- temporalio/worker/_workflow_instance.py | 22 ++ temporalio/workflow/__init__.py | 10 + temporalio/workflow/_context.py | 14 + temporalio/workflow/_streams.py | 336 ++++++++++++++++++++++++ 4 files changed, 382 insertions(+) create mode 100644 temporalio/workflow/_streams.py diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index e392671a3..13f2a1490 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -58,7 +58,9 @@ import temporalio.converter import temporalio.exceptions import temporalio.nexus.system +import temporalio.streams import temporalio.workflow +import temporalio.workflow._streams from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo from temporalio.converter._payload_converter import ( _TemporalTransferTypePayloadConverter, @@ -184,6 +186,7 @@ class WorkflowInstanceDetails: default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] = field( default_factory=lambda: _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS ) + stream_provider: temporalio.streams.StreamProvider | None = None class WorkflowInstance(ABC): @@ -303,6 +306,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 # Keyed by channel name; one subscription per channel per run self._channel_subscriptions: dict[ str, temporalio.workflow.ChannelSubscription @@ -1392,6 +1397,9 @@ def workflow_instance(self) -> Any: def workflow_is_continue_as_new_suggested(self) -> bool: return self._continue_as_new_suggested + def workflow_is_evicting(self) -> bool: + return self._deleting + def workflow_is_target_worker_deployment_version_changed(self) -> bool: return self._target_worker_deployment_version_changed @@ -1844,6 +1852,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_subscribe_channel( self, channel: str ) -> temporalio.workflow.ChannelSubscription: diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index 8c33a4852..212625511 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -157,6 +157,12 @@ logger, unsafe, ) +from ._streams import ( + StreamReader, + StreamWriter, + stream_reader, + stream_writer, +) from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowConfig, @@ -268,6 +274,10 @@ "Notification", "linked_channel", "subscribe_channel", + "StreamReader", + "StreamWriter", + "stream_reader", + "stream_writer", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 2a9ff6355..a04791baf 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -26,6 +26,7 @@ from ._event_groups import EventGroup from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle + from ._streams import _WorkflowStreams from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowHandle, @@ -354,6 +355,16 @@ def workflow_instance(self) -> Any: ... @abstractmethod def workflow_is_continue_as_new_suggested(self) -> bool: ... + @abstractmethod + def workflow_is_evicting(self) -> bool: + """Whether this instance is being dropped from the cache rather than ending. + + Eviction cancels the primary task the way a workflow cancellation + does, so anything that runs on the way out has to be able to tell the + two apart. Instance state must not be touched while this is true. + """ + ... + @abstractmethod def workflow_is_target_worker_deployment_version_changed(self) -> bool: ... @@ -500,6 +511,9 @@ async def workflow_start_nexus_operation( event_groups: Sequence[EventGroup] | None = None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_streams(self) -> _WorkflowStreams: ... + @abstractmethod def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py new file mode 100644 index 000000000..54353508d --- /dev/null +++ b/temporalio/workflow/_streams.py @@ -0,0 +1,336 @@ +"""Streams from inside workflow code. + +.. warning:: + This module is experimental and may change in future versions. + +The reader and writer here are the same on every provider. They convert +values, synthesize supersession and buffer nothing the provider did not hand +them; everything provider-specific sits behind the ``ReadSource`` and +``WriteSink`` that the provider's workflow half opens. The provider itself +comes from the worker, the way the payload converter does. +""" + +from __future__ import annotations + +import asyncio +from collections import deque +from collections.abc import AsyncIterator, Callable +from typing import Any, Generic, TypeVar, cast, overload + +from temporalio.streams._provider import ReadSource, WorkflowStreamProvider, WriteSink +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import RecordDecoder, to_wire +from temporalio.workflow._context import _Runtime, payload_converter +from temporalio.workflow._sandbox import logger + +__all__ = ["StreamReader", "StreamWriter", "stream_reader", "stream_writer"] + +T = TypeVar("T") + + +class _WorkflowStreams: + """The stream state a workflow instance carries: its provider half and open readers.""" + + def __init__(self, provider: WorkflowStreamProvider) -> None: + self.provider = provider + self.readers: dict[str, StreamReader[Any]] = {} + # Finishing is a statement about the topic, not about the writer + # object that made it, and stream_writer() hands out a new object on + # every call. Rebuilt in order on replay, so it stays deterministic. + self.finished: set[str] = set() + + +class StreamReader(Generic[T]): + """Reads one topic of this workflow's stream from inside workflow code. + + The reader is the async iterator: ``async for record in reader`` yields + every kind of record, including the supersession the reader synthesizes + when a producer's newer attempt appears. Check ``record.kind``, or + iterate :meth:`values` when the workflow only wants data. Iteration ends + when the reader is closed or the provider ends the subscription. One loop + per reader: two loops on one reader share its buffer and interleave. + """ + + def __init__( + self, + source: ReadSource, + *, + topic: str, + result_type: type | None, + after: Cursor, + on_close: Callable[[], None], + ) -> None: + """Prefer :func:`temporalio.workflow.stream_reader`.""" + self._source = source + self._topic = topic + self._result_type = result_type + self._decoder = RecordDecoder( + payload_converter(), result_type, after=after, warn=logger.warning + ) + self._pending: deque[StreamRecord[T]] = deque() + self._lock = asyncio.Lock() + self._closed = False + self._ended = False + self._on_close = on_close + + @property + def topic(self) -> str: + """The name of the topic this reader is subscribed to.""" + return self._topic + + def __aiter__(self) -> StreamReader[T]: + """The reader is its own iterator.""" + return self + + async def __anext__(self) -> StreamRecord[T]: + """The next record, waiting for one to arrive.""" + while True: + if self._pending: + return self._pending.popleft() + if self._closed or self._ended: + raise StopAsyncIteration + await self._fill() + + async def _fill(self) -> None: + # Two loops on one reader must not race the source, so one batch + # fetch is in flight at a time and the second loop takes what the + # first one buffered. + async with self._lock: + if self._pending or self._closed or self._ended: + return + try: + batch = await self._source.next_batch() + except StopAsyncIteration: + # The provider ended the subscription. Iteration stops rather + # than raising, so a workflow that reads to the end of a + # finished stream leaves the loop instead of failing its task. + self._ended = True + return + for cursor, wire in batch: + self._pending.extend(self._decoder.decode(cursor, wire)) + + async def values(self) -> AsyncIterator[T]: + """Iterate the data values, dropping control records.""" + async for record in self: + if record.kind is RecordKind.DATA: + yield cast("T", record.value) + + def close(self) -> None: + """End the subscription. Idempotent. + + A later :func:`temporalio.workflow.stream_reader` on the same topic + opens a new subscription, which is a new command. + """ + if self._closed: + return + self._closed = True + self._source.close() + self._on_close() + + +class StreamWriter(Generic[T]): + """Publishes to one topic of this workflow's stream. + + A workflow can only publish transactionally to its own stream, on every + provider. Writing to somebody else's stream is an activity's job, and it + gets the weaker guarantee that goes with doing I/O. The type parameter is + the topic definition's value type; a writer on a string-named topic takes + any value. + """ + + def __init__(self, sink: WriteSink, topic: str, finished: set[str]) -> None: + """Prefer :func:`temporalio.workflow.stream_writer`.""" + self._sink = sink + self._topic = topic + self._finished = finished + + @property + def topic(self) -> str: + """The name of the topic this writer is bound to.""" + return self._topic + + def publish(self, value: T) -> None: + """Append ``value`` to this topic. + + Synchronous, because there is nothing to wait for inside a task: the + record becomes visible when this Workflow Task is accepted, and never + at all if the task fails, so a reader cannot see a decision the + workflow did not commit. A :class:`temporalio.common.RawValue` passes + through pre-encoded. + + Raises: + ValueError: The topic was already finished in this run, by this + writer or by another one on the same topic. + """ + if self._topic in self._finished: + raise ValueError(f"topic {self._topic!r} was already finished") + self._sink.publish( + to_wire( + payload_converter(), + topic=self._topic, + kind=RecordKind.DATA, + value=value, + ) + ) + + def finish(self) -> None: + """Write ``FINISH`` for this workflow on this topic. Idempotent. + + Says this workflow has nothing more to send on the topic. It does not + say the workflow succeeded, and it does not end anyone's read. The + marker belongs to the topic, so a second writer on the same topic in + the same run finds it already written. + """ + if self._topic in self._finished: + return + self._finished.add(self._topic) + self._sink.publish( + to_wire(payload_converter(), topic=self._topic, kind=RecordKind.FINISH) + ) + + +@overload +def stream_reader( + topic: StreamTopic[T], *, after: Cursor = ..., last: int | None = None +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str | None = None, + *, + result_type: type[T], + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str | None = None, + *, + result_type: None = None, + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[Any]: ... + + +def stream_reader( + topic: str | StreamTopic[Any] | None = None, + *, + result_type: type | None = None, + after: Cursor = BEGINNING, + last: int | None = None, +) -> StreamReader[Any]: + """Subscribe this workflow to ``topic`` of its own stream. + + ``topic`` is a :func:`temporalio.streams.topic` definition, which carries + the record type, or a plain string with ``result_type=`` for a name + decided at runtime, and without one the reader is on + :data:`temporalio.streams.DEFAULT_TOPIC`, decoded as a string-named topic + is. One subscription per topic per run. A second call + for the same topic returns the reader already open on it, so records go + to whichever loop pulls first; such a call may pass no ``after``, no + ``last`` and no different type. Adding a reader on a new topic is a new command, so + gate it with :func:`temporalio.workflow.patched` as you would a timer. A + reader in a successor run starts a new subscription: nothing crosses + continue-as-new implicitly. + + Args: + topic: The topic, relative to this workflow's stream. Omit it for + the default topic. + result_type: The value type for a string-named or default topic, used as the + decode hint. :class:`temporalio.common.RawValue` returns the + payload untouched. + after: Resume strictly after this record. ``BEGINNING`` starts at + the oldest record the topic still holds and + :data:`temporalio.streams.END` at whatever is appended after the + subscription is registered. Honoured on the first subscription + of a run, because after that the recorded observations decide. + last: Start at the newest ``last`` records instead, or at all of + them when there are fewer. Records of every kind count. Exclusive + with a cursor in ``after``. Where it lands is resolved once, + outside the workflow, and replay reproduces it. + + Raises: + ValueError: ``topic`` is empty, ``result_type`` was passed with a + definition, ``last`` is not positive or came with a cursor, or a + reader on the topic is already open and this call asked for a + different position or type. + temporalio.streams.StreamCursorError: ``after`` was minted by another + provider. + temporalio.streams.StreamUnsupportedError: The provider cannot start + where ``END`` or ``last`` asks. + """ + check_read_start(after, last) + name, result_type = resolve_topic(topic, result_type) + state: _WorkflowStreams = _Runtime.current().workflow_streams() + existing = state.readers.get(name) + if existing is not None: + if ( + after != BEGINNING + or last is not None + or result_type is not existing._result_type + ): + raise ValueError( + f"topic {name!r} already has a reader in this run; a second " + "stream_reader shares it and takes no after=, last= or other type" + ) + return existing + # Passed only when given, so a provider written before last= existed + # still serves every read it can. + source = ( + state.provider.open_reader(name, after=after) + if last is None + else state.provider.open_reader(name, after=after, last=last) + ) + + def forget() -> None: + state.readers.pop(name, None) + + # The decoder positions a synthesized record at the one before it. Where + # END or last= lands is not known here, so BEGINNING stands in: resuming + # from it may deliver a record twice, where END would skip one. + previous = BEGINNING if after == END or last is not None else after + reader: StreamReader[Any] = StreamReader( + source, topic=name, result_type=result_type, after=previous, on_close=forget + ) + state.readers[name] = reader + return reader + + +@overload +def stream_writer(topic: StreamTopic[T]) -> StreamWriter[T]: ... + + +@overload +def stream_writer(topic: str | None = None) -> StreamWriter[Any]: ... + + +def stream_writer(topic: str | StreamTopic[Any] | None = None) -> StreamWriter[Any]: + """Publish to ``topic`` of this workflow's stream, or to its default topic. + + Every call returns a new writer, and they all share the run's record of + which topics were finished, so ``finish()`` on one is seen by the next. + + Args: + topic: A :func:`temporalio.streams.topic` definition, whose value type + the writer's ``publish`` takes, or a plain string for a name + decided at runtime. Omitted, the writer is on + :data:`temporalio.streams.DEFAULT_TOPIC`. Encoding follows each + published value. + + Raises: + ValueError: ``topic`` is empty. + """ + name, _ = resolve_topic(topic) + state = _Runtime.current().workflow_streams() + return StreamWriter(state.provider.open_writer(name), name, state.finished) From 732e0d28afe798cd520fed896426dc309e3b4ba1 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:51:59 -0700 Subject: [PATCH 2/3] Ran the stream provider's lifecycle hooks around the workflow. The start hook runs on the workflow's loop before the first task's handlers, so a provider's handlers exist for an Update that arrives with that task. The finish hook skips eviction. --- temporalio/streams/_provider.py | 9 +++-- temporalio/worker/_replayer.py | 1 + temporalio/worker/_worker.py | 6 ++++ temporalio/worker/_workflow.py | 60 +++++++++++++++++++++++++++++++++ 4 files changed, 73 insertions(+), 3 deletions(-) diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py index 662bc7918..98c4851e0 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -322,10 +322,13 @@ def open_writer(self, topic: str) -> WriteSink: ... def on_workflow_start(self) -> None: - """Called before the workflow function runs. + """Called before the workflow function runs, and before the first task's handlers. - A provider that serves outside readers through handlers on the - workflow registers them here, before the first task completes. + After the workflow's own ``__init__`` and before any Signal or Update + of the first task is handled, which the SDK does ahead of the + workflow function. A provider that serves outside readers through + handlers on the workflow registers them here, so an Update that + arrives with the first task finds them. """ ... diff --git a/temporalio/worker/_replayer.py b/temporalio/worker/_replayer.py index 1291f5066..c7ae8d59d 100644 --- a/temporalio/worker/_replayer.py +++ b/temporalio/worker/_replayer.py @@ -290,6 +290,7 @@ def on_eviction_hook( default_workflow_logic_flags=frozenset( self._default_workflow_logic_flags ), + stream_provider=self._config.get("stream_provider"), ) external_storage = data_converter.external_storage storage_driver_types = ( diff --git a/temporalio/worker/_worker.py b/temporalio/worker/_worker.py index 25dcf697d..f4b7eddad 100644 --- a/temporalio/worker/_worker.py +++ b/temporalio/worker/_worker.py @@ -468,6 +468,11 @@ def _init_from_config(self, client: temporalio.client.Client, config: WorkerConf # Prepend applicable client interceptors to the given ones client_config = config["client"].config(active_config=True) # type: ignore[reportTypedDictNotRequiredAccess] + # A provider registered on the client serves its workers too, so one + # registration covers every context that asks for a stream. + stream_provider = config.get("stream_provider") or client_config.get( + "stream_provider" + ) interceptors_from_client = cast( list[Interceptor], [i for i in client_config["interceptors"] if isinstance(i, Interceptor)], @@ -572,6 +577,7 @@ def check_activity(activity: str): encode_headers=client_config["header_codec_behavior"] != HeaderCodecBehavior.NO_CODEC, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, + stream_provider=stream_provider, ) tuner = config.get("tuner") diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 01304d014..f4001b32e 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, cast 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,9 +37,11 @@ _relax_sandbox_for_debugger, ) from ._interceptor import ( + ExecuteWorkflowInput, Interceptor, WorkflowInboundInterceptor, WorkflowInterceptorClassInput, + WorkflowOutboundInterceptor, ) from ._workflow_instance import ( _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS, @@ -55,6 +59,55 @@ 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 start + hook runs when the instance's loop first turns, after the workflow's own + ``__init__`` and before the first task's Signals and Updates are handled, + because the SDK handles those ahead of the workflow function and a handler + registered any later would be missed by an Update that arrives with that + task. 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. + """ + + def init(self, outbound: WorkflowOutboundInterceptor) -> None: + super().init(outbound) + # The hook has to run after the workflow's own __init__, which may + # register handlers the provider adopts, and before the first task's + # Signals and Updates are handled, which the SDK does ahead of the + # workflow function. The instance is its own event loop and nothing + # is queued on it yet, so a callback queued now runs first when that + # loop first turns, which is after every job of the activation has + # been applied and before any task they created takes a step. + runtime = temporalio.workflow._Runtime.current() + loop = cast(asyncio.AbstractEventLoop, cast(object, runtime)) + loop.call_soon(lambda: runtime.workflow_streams().provider.on_workflow_start()) + + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + runtime = temporalio.workflow._Runtime.current() + provider = runtime.workflow_streams().provider + try: + result = await self.next.execute_workflow(input) + except GeneratorExit: + raise + except BaseException: + # Eviction cancels the primary task the same way a workflow + # cancellation does, and only the cancellation is a run ending. + if not runtime.workflow_is_evicting(): + await provider.on_workflow_finish() + raise + await provider.on_workflow_finish() + return result + + # Value was chosen abitrarily as a small number that allows some concurrency and prevents # large numbers of concurrent external storage operations causing resource contention. # This default limit is per workflow task activation and does not limit the total number @@ -104,6 +157,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 +216,11 @@ def __init__( __temporal_assert_local_activity_valid=assert_local_activity_valid, ) ) + self._stream_provider = stream_provider + if stream_provider is not None: + # Innermost, so the lifecycle hooks bracket the workflow function + # itself, after every user interceptor has done its own setup. + self._interceptor_classes.append(_StreamHooksInterceptor) self._workflow_failure_exception_types = workflow_failure_exception_types self._patch_activation_callback = patch_activation_callback @@ -743,6 +802,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) From 7bc1c11d2e8d31b3594614f89a66c425b6eef6ba Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:51:59 -0700 Subject: [PATCH 3/3] Covered the workflow stream runtime. The hooks, the reader's start positions and a workflow reading and writing through the memory provider. --- tests/streams/test_stream_hooks.py | 169 +++++++ tests/streams/test_stream_reader.py | 178 +++++++ tests/streams/test_streams_workflow.py | 618 +++++++++++++++++++++++++ 3 files changed, 965 insertions(+) create mode 100644 tests/streams/test_stream_hooks.py create mode 100644 tests/streams/test_stream_reader.py create mode 100644 tests/streams/test_streams_workflow.py diff --git a/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py new file mode 100644 index 000000000..7908dfc57 --- /dev/null +++ b/tests/streams/test_stream_hooks.py @@ -0,0 +1,169 @@ +"""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, cast + +import pytest + +from temporalio import workflow +from temporalio.worker._interceptor import ( + ExecuteWorkflowInput, + WorkflowInboundInterceptor, +) +from temporalio.worker._workflow import _StreamHooksInterceptor +from temporalio.worker._workflow_instance import _WorkflowInstanceImpl + + +class _Provider: + def __init__(self) -> None: + self.calls: list[str] = [] + + def on_workflow_start(self) -> None: + self.calls.append("start") + + async def on_workflow_finish(self) -> None: + self.calls.append("finish") + + +class _Streams: + def __init__(self, provider: _Provider) -> None: + self.provider = provider + + +class _FakeRuntime: + """Only what the interceptor reads, and only through the runtime interface. + + It deliberately carries no ``_deleting`` attribute. An interceptor that + reads the eviction state by attribute name instead of by method would see + a runtime that is never evicting here, and the eviction case below would + catch it. + """ + + def __init__(self, provider: _Provider, *, evicting: bool = False) -> None: + self._streams = _Streams(provider) + self._evicting = evicting + + def workflow_streams(self) -> _Streams: + return self._streams + + def workflow_is_evicting(self) -> bool: + return self._evicting + + def call_soon(self, callback: Any) -> None: + # The real runtime is the workflow's event loop; here the test's loop + # stands in for it. + asyncio.get_running_loop().call_soon(callback) + + +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 + + def init(self, outbound: Any) -> None: + del outbound + + 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={}) + + +async def _hooked(body: _Body) -> _StreamHooksInterceptor: + """The interceptor over ``body``, initialised the way the instance does it. + + ``init`` queues the start hook for the loop's first turn, so one turn + runs it, as the instance's loop does before any task takes a step. + """ + interceptor = _StreamHooksInterceptor(body) + interceptor.init(cast(Any, None)) + await asyncio.sleep(0) + return interceptor + + +@pytest.fixture +async def provider() -> Any: + fake = _Provider() + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake)) # type: ignore[arg-type] + yield fake + workflow._Runtime.set_on_loop(loop, None) + + +def _evicting(fake: _Provider) -> None: + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake, evicting=True)) # type: ignore[arg-type] + + +async def test_the_finish_hook_runs_on_return(provider: _Provider): + assert await (await _hooked(_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 (await _hooked(_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 (await _hooked(_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 (await _hooked(_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 = (await _hooked(_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 (await _hooked(_Body(asyncio.CancelledError()))).execute_workflow(_INPUT) + assert provider.calls == ["start"] + + +def test_the_runtime_answers_the_eviction_question_itself(): + # The interceptor asks the runtime rather than reading a private field, + # so the coupling is declared on the base class and the real instance has + # to answer it or fail to construct. + assert "workflow_is_evicting" in workflow._Runtime.__abstractmethods__ + assert "workflow_is_evicting" not in _WorkflowInstanceImpl.__abstractmethods__ diff --git a/tests/streams/test_stream_reader.py b/tests/streams/test_stream_reader.py new file mode 100644 index 000000000..d3918fe40 --- /dev/null +++ b/tests/streams/test_stream_reader.py @@ -0,0 +1,178 @@ +"""The reader's rules that a workflow test cannot pin down. + +``stream_reader`` promises one subscription per topic per run, that a second +loop on one reader shares its buffer rather than racing the source, and that +closing it and opening the topic again is a new subscription. The memory +provider's source hands back everything it has without ever suspending, so a +workflow test cannot tell a reader that serialises its fetches from one that +does not. These drive the reader directly, with a source that suspends where +a real provider would, and a stand-in runtime on the loop. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.converter import DataConverter +from temporalio.streams import BEGINNING, Cursor, RecordKind, topic +from temporalio.streams._wire import WireRecord, to_wire +from temporalio.workflow._streams import _WorkflowStreams + +INPUTS = topic("inputs", dict) + +_CONVERTER = DataConverter.default.payload_converter + + +def _record(n: int) -> WireRecord: + return to_wire( + _CONVERTER, + topic=INPUTS.name, + kind=RecordKind.DATA, + value={"n": n}, + producer_id="model", + attempt=1, + sequence=n, + ) + + +class _SlowSource: + """Hands over one record per batch, suspending first the way a real one does.""" + + def __init__(self, count: int) -> None: + self.opened = True + self.next_calls = 0 + self.in_flight = 0 + self.overlapped = False + self._offset = 0 + self._count = count + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + self.next_calls += 1 + self.in_flight += 1 + if self.in_flight > 1: + self.overlapped = True + try: + # Where a provider waits for the worker to deliver. Two fetches + # that both get here have both left the reader's buffer behind. + await asyncio.sleep(0) + await asyncio.sleep(0) + if self._offset >= self._count: + raise StopAsyncIteration + self._offset += 1 + return [(Cursor(f"fake:{self._offset - 1}"), _record(self._offset - 1))] + finally: + self.in_flight -= 1 + + def close(self) -> None: + self.opened = False + + +class _Half: + """A workflow half that hands out one source per topic and counts opens.""" + + def __init__(self) -> None: + self.sources: list[_SlowSource] = [] + + def open_reader(self, topic: str, *, after: Cursor) -> Any: + del topic, after + source = _SlowSource(4) + self.sources.append(source) + return source + + def open_writer(self, topic: str) -> Any: + del topic + raise NotImplementedError + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +class _FakeRuntime: + """Only what the reader reads: the stream state and the payload converter.""" + + def __init__(self, half: _Half) -> None: + self._streams = _WorkflowStreams(half) # type: ignore[arg-type] + + def workflow_streams(self) -> _WorkflowStreams: + return self._streams + + def workflow_payload_converter(self) -> Any: + return _CONVERTER + + +@pytest.fixture +async def half() -> Any: + fake = _Half() + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake)) # type: ignore[arg-type] + yield fake + workflow._Runtime.set_on_loop(loop, None) + + +async def test_two_loops_on_one_reader_split_the_records(half: _Half): + reader = workflow.stream_reader(INPUTS) + seen: list[int] = [] + + async def pull(count: int) -> None: + for _ in range(count): + record = await reader.__anext__() + assert record.value is not None + seen.append(record.value["n"]) + + await asyncio.wait_for(asyncio.gather(pull(2), pull(2)), 5.0) + # Every record once and none lost, whichever loop got there first, and + # the source was never asked for two batches at the same time. + assert sorted(seen) == [0, 1, 2, 3] + assert half.sources[0].overlapped is False + + +@pytest.mark.usefixtures("half") +async def test_a_cancelled_read_does_not_lose_the_record_it_waited_for(): + reader = workflow.stream_reader(INPUTS) + pending = asyncio.ensure_future(reader.__anext__()) + await asyncio.sleep(0) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + + # The next read picks up where the cancelled one was, rather than finding + # the reader wedged behind a lock the cancelled fetch still holds. + record = await asyncio.wait_for(reader.__anext__(), 5.0) + assert record.value == {"n": 0} + + +async def test_closing_a_reader_lets_the_topic_be_opened_again(half: _Half): + reader = workflow.stream_reader(INPUTS) + assert workflow.stream_reader(INPUTS) is reader + reader.close() + assert half.sources[0].opened is False + + # A new subscription, not the closed one handed back. It is also a new + # command on a real provider, which is why the docstring says to gate it. + again = workflow.stream_reader(INPUTS) + assert again is not reader + assert len(half.sources) == 2 + assert (await asyncio.wait_for(again.__anext__(), 5.0)).value == {"n": 0} + + +@pytest.mark.usefixtures("half") +async def test_a_closed_reader_stops_iterating(): + reader = workflow.stream_reader(INPUTS) + assert (await asyncio.wait_for(reader.__anext__(), 5.0)).value == {"n": 0} + reader.close() + with pytest.raises(StopAsyncIteration): + await asyncio.wait_for(reader.__anext__(), 5.0) + + +@pytest.mark.usefixtures("half") +async def test_a_reader_ends_when_the_source_ends(): + reader = workflow.stream_reader(INPUTS, after=BEGINNING) + values = [record.value async for record in reader] + assert values == [{"n": 0}, {"n": 1}, {"n": 2}, {"n": 3}] diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py new file mode 100644 index 000000000..b666a4adc --- /dev/null +++ b/tests/streams/test_streams_workflow.py @@ -0,0 +1,618 @@ +"""Workflow-side conformance for the stream contract. + +Runs the reader and writer inside a real workflow on the memory provider +with a warm cache, and states the two rules about Workflow Tasks as tests: a +publish commits with its task (rule 1), and reads are recorded observations +that replay re-supplies (rule 2). The memory provider keeps neither and says +so in its docstring, so those two are strict expected failures here. A +storage provider that runs this module turns them into passes; that is the +measurement they exist for. + +The rest is what the portable surface promises on every provider: the +lifecycle hooks the worker calls, one subscription per topic per run, a read +that ends when the chain closes, a handle that follows continue-as-new, and +a topic shared by the workflow and an outside producer. +""" + +from __future__ import annotations + +import asyncio +import uuid +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.client import Client +from temporalio.streams import ( + DEFAULT_TOPIC, + END, + Cursor, + ReadSource, + RecordKind, + StreamCursorError, + WriteSink, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Replayer +from tests.helpers import new_worker +from tests.streams.test_streams_conformance import take + +INPUTS = topic("inputs", dict) +DECISIONS = topic("decisions", dict) + + +@pytest.fixture +def provider(env: WorkflowEnvironment): # pyright: ignore[reportUnusedFunction] + if env.supports_time_skipping: + pytest.skip( + "the memory provider polls on a timer, which time skipping turns into a spin" + ) + streams = MemoryStreams() + yield streams + streams.reset() + + +@workflow.defn +class ContractLoop: + """Reads ``inputs``, publishes a decision per value, reports control records.""" + + @workflow.run + async def run(self) -> list[dict[str, Any]]: + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + try: + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "replaced": record.supersession.previous_attempt, + "attempt": record.supersession.attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decisions.publish({"decided": record.value["n"]}) + trace.append( + { + "kind": "decision", + "n": record.value["n"], + "attempt": record.attempt, + } + ) + finally: + inputs.close() + decisions.finish() + # Twice on purpose: a finished topic stays finished, with one marker. + decisions.finish() + return trace + + +async def _run_the_loop( + client: Client, provider: MemoryStreams +) -> tuple[Any, list[dict[str, Any]]]: + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, ContractLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"n": 1}, {"n": 2}) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"n": 3}) + await second.finish() + trace = await handle.result() + return handle, trace + + +async def test_workflow_reads_decides_and_publishes( + client: Client, provider: MemoryStreams +): + handle, trace = await _run_the_loop(client, provider) + assert trace == [ + {"kind": "decision", "n": 1, "attempt": 1}, + {"kind": "decision", "n": 2, "attempt": 1}, + {"kind": "superseded", "replaced": 1, "attempt": 2}, + {"kind": "decision", "n": 3, "attempt": 2}, + {"kind": "finish", "producer": "model"}, + ] + + # The outside view of what the workflow published, on its own topic. + stream = provider.get_stream_handle(client, handle.id) + records = await take(stream.read(topic=DECISIONS), 5) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"decided": 1}), + (RecordKind.DATA, {"decided": 2}), + (RecordKind.DATA, {"retracting_attempt": 1}), + (RecordKind.DATA, {"decided": 3}), + (RecordKind.FINISH, None), + ] + assert all(r.producer_id == "" and r.topic == DECISIONS.name for r in records) + # The second finish() wrote nothing: the marker is the newest record. + assert await stream.latest(topic=DECISIONS) == records[-1].cursor + + +async def test_read_ends_when_the_workflow_closes_and_the_tail_is_delivered( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + stream = provider.get_stream_handle(client, handle.id) + + async def read_everything() -> list[Any]: + return [r.value async for r in stream.read(topic=DECISIONS)] + + # No count and no early break: the read ends by itself once the workflow + # is closed and everything it retained has been handed over. + values = await asyncio.wait_for(read_everything(), 30) + assert values == [ + {"decided": 1}, + {"decided": 2}, + {"retracting_attempt": 1}, + {"decided": 3}, + None, + ] + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 2: reading is a recorded observation, so replaying the history " + "with the store gone must re-supply the same records; the memory " + "provider reads live process memory instead" + ), +) +async def test_replay_without_the_store_resupplies_the_records( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + history = await handle.fetch_history() + + provider.reset() + replayer = Replayer(workflows=[ContractLoop], plugins=[provider]) + await replayer.replay_workflow(history) + + +# Run ids whose first workflow task already failed, shared with the workflow +# thread so the retry can tell it is the retry. Outside the sandbox on +# purpose: the sandbox re-imports this module per run and would hide the set. +_failed_once: set[str] = set() + + +@workflow.defn(sandboxed=False) +class PublishThenFail: + """Publishes, then fails its first workflow task; the retry publishes again.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + run_id = workflow.info().run_id + committed = run_id in _failed_once + decisions.publish({"committed": committed}) + if not committed: + _failed_once.add(run_id) + raise RuntimeError("the first task fails after publishing") + decisions.finish() + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 1: a publish commits with its workflow task, so no reader sees " + "a record from a task that failed; the memory provider makes it " + "visible at publish time" + ), +) +async def test_a_failed_task_publishes_nothing(client: Client, provider: MemoryStreams): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, PublishThenFail, plugins=[provider]) as worker: + handle = await client.start_workflow( + PublishThenFail.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS), 2, timeout=30) + await handle.result() + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"committed": True}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class Relay: + """Publishes one record per run and continues as new once.""" + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"run": run}) + if run == 0: + workflow.continue_as_new(run + 1) + decisions.finish() + + +class _RecordingHalf: + """A workflow half that logs the hooks the worker calls, then delegates.""" + + def __init__(self, inner: Any, calls: list[tuple[str, str]]) -> None: + self._inner = inner + self._calls = calls + + def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + return self._inner.open_reader(topic, after=after) + + def open_writer(self, topic: str) -> WriteSink: + return self._inner.open_writer(topic) + + def on_workflow_start(self) -> None: + self._calls.append(("start", workflow.info().run_id)) + + async def on_workflow_finish(self) -> None: + self._calls.append(("finish", workflow.info().run_id)) + + +class HookedMemory(MemoryStreams): + """The memory provider with its lifecycle hooks made visible.""" + + def __init__(self) -> None: + super().__init__() + self.calls: list[tuple[str, str]] = [] + + def workflow_provider(self) -> Any: + return _RecordingHalf(super().workflow_provider(), self.calls) + + +@pytest.mark.usefixtures("provider") +async def test_the_worker_calls_the_lifecycle_hooks_around_every_run(client: Client): + hooked = HookedMemory() + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[hooked]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + # Start before the function, finish after it, on both runs: the finish + # hook runs on the continue-as-new exit too, so a provider that parked + # something against the first run can let go before the successor starts. + kinds = [kind for kind, _ in hooked.calls] + assert kinds == ["start", "finish", "start", "finish"] + runs = [run_id for _, run_id in hooked.calls] + assert runs[0] == runs[1] and runs[2] == runs[3] and runs[0] != runs[2] + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[provider]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + await handle.result() + # The chain is followed: the successor's records arrive on the same read, + # and the read ends only when the last run of the chain is closed. + assert records == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class SharedReaders: + """Opens the same topic twice and pulls from both readers in turn.""" + + @workflow.run + async def run(self) -> list[Any]: + first = workflow.stream_reader(INPUTS) + second = workflow.stream_reader(INPUTS) + trace: list[Any] = ["shared" if first is second else "separate"] + try: + workflow.stream_reader(INPUTS, after=Cursor("memory:0")) + except ValueError: + trace.append("after-rejected") + try: + workflow.stream_reader(INPUTS.name, result_type=list) + except ValueError: + trace.append("type-rejected") + trace.append((await first.__anext__()).value) + trace.append((await second.__anext__()).value) + first.close() + return trace + + +async def test_a_second_reader_on_a_topic_shares_the_subscription( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, SharedReaders, plugins=[provider]) as worker: + handle = await client.start_workflow( + SharedReaders.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).append( + {"n": 1}, {"n": 2} + ) + assert await handle.result() == [ + "shared", + "after-rejected", + "type-rejected", + {"n": 1}, + {"n": 2}, + ] + + +@workflow.defn +class FinishThenPublish: + """Finishes a topic on one writer and publishes on a second one.""" + + @workflow.run + async def run(self) -> str: + workflow.stream_writer(DECISIONS).finish() + # A fresh writer object, the same topic. The marker is already there, + # so this publish would land after the end of the topic. + try: + workflow.stream_writer(DECISIONS).publish({"after": "finish"}) + except ValueError as error: + return str(error) + return "published" + + +async def test_a_finished_topic_stays_finished_across_writers( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, FinishThenPublish, plugins=[provider]) as worker: + handle = await client.start_workflow( + FinishThenPublish.run, id=workflow_id, task_queue=worker.task_queue + ) + assert "already finished" in await handle.result() + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + # The marker is all that landed: the publish behind it never reached the + # topic, which is the point of the guard reading like a topic-wide rule. + assert records == [(RecordKind.FINISH, None)] + + +@workflow.defn +class NoStreams: + """Touches no stream at all, on a worker that has a provider.""" + + @workflow.run + async def run(self) -> str: + await asyncio.sleep(0) + return "done" + + +@pytest.mark.usefixtures("provider") +async def test_a_workflow_that_touches_no_stream_runs_unchanged(client: Client): + hooked = HookedMemory() + async with new_worker(client, NoStreams, plugins=[hooked]) as worker: + result = await client.execute_workflow( + NoStreams.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert result == "done" + # The interceptor is installed per worker, not per workflow, so the hooks + # still bracket a run that never opened a reader or a writer. A provider's + # hooks therefore have to be cheap and safe on a workflow that uses none. + assert [kind for kind, _ in hooked.calls] == ["start", "finish"] + + +@workflow.defn +class ForeignCursor: + """Resumes from a cursor another provider minted.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS, after=Cursor("elsewhere:1")) + except StreamCursorError: + return "refused" + return "accepted" + + +async def test_the_workflow_reader_refuses_a_foreign_cursor( + client: Client, provider: MemoryStreams +): + async with new_worker(client, ForeignCursor, plugins=[provider]) as worker: + result = await client.execute_workflow( + ForeignCursor.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert result == "refused" + + +@workflow.defn +class NoProvider: + """Opens a stream on a worker that has no provider.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS) + except RuntimeError as error: + return str(error) + return "opened" + + +async def test_a_worker_without_a_provider_says_so(client: Client): + async with new_worker(client, NoProvider) as worker: + result = await client.execute_workflow( + NoProvider.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert "no stream provider is configured" in result + + +@workflow.defn +class OneLine: + """Publishes one record and finishes the topic.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"from": "workflow"}) + decisions.finish() + + +async def test_an_outside_producer_and_the_workflow_share_a_topic( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, OneLine, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=DECISIONS, producer_id="tool", attempt=1).append( + {"from": "producer"} + ) + handle = await client.start_workflow( + OneLine.run, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + # Both writers land on one topic in one order, each under its own + # identity: the producer's records carry its id, the workflow's carry none. + assert [(r.producer_id, r.kind, r.value) for r in records] == [ + ("tool", RecordKind.DATA, {"from": "producer"}), + ("", RecordKind.DATA, {"from": "workflow"}), + ("", RecordKind.FINISH, None), + ] + + +@workflow.defn +class DefaultTopicEcho: + """Reads one value on its default topic and answers on the same topic.""" + + @workflow.run + async def run(self) -> None: + reader = workflow.stream_reader(result_type=dict) + writer = workflow.stream_writer() + async for value in reader.values(): + writer.publish({"echo": value["n"] * 2}) + reader.close() + writer.finish() + + +async def test_a_workflow_reads_and_writes_its_default_topic( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, DefaultTopicEcho, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(producer_id="client", attempt=1).append({"n": 21}) + handle = await client.start_workflow( + DefaultTopicEcho.run, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DEFAULT_TOPIC)] + + records = await asyncio.wait_for(read_everything(), 30) + assert [(r.topic, r.kind, r.value) for r in records] == [ + (DEFAULT_TOPIC, RecordKind.DATA, {"n": 21}), + (DEFAULT_TOPIC, RecordKind.DATA, {"echo": 42}), + (DEFAULT_TOPIC, RecordKind.FINISH, None), + ] + + +@workflow.defn +class NewestTwo: + """Starts at the newest two records on ``inputs`` and returns their values.""" + + @workflow.run + async def run(self) -> list[Any]: + reader = workflow.stream_reader(INPUTS, last=2) + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == 2: + reader.close() + return values + + +async def test_a_workflow_reader_starts_at_the_last_n_records( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, NewestTwo, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="tool", attempt=1).append( + {"n": 1}, {"n": 2}, {"n": 3}, {"n": 4} + ) + handle = await client.start_workflow( + NewestTwo.run, id=workflow_id, task_queue=worker.task_queue + ) + assert await handle.result() == [3, 4] + + +@workflow.defn +class FromNow: + """Follows ``inputs`` from when it subscribes and returns the first value.""" + + @workflow.run + async def run(self) -> Any: + reader = workflow.stream_reader(INPUTS, after=END) + async for value in reader.values(): + reader.close() + return value["n"] + return None + + +async def test_a_workflow_reader_at_end_skips_what_was_there( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, FromNow, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "old"}) + handle = await client.start_workflow( + FromNow.run, id=workflow_id, task_queue=worker.task_queue + ) + result = asyncio.ensure_future(handle.result()) + # The subscription starts when the workflow runs, which the test does + # not observe, so appends keep coming until the workflow takes one. + for _ in range(150): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({result}, timeout=0.2) + if done: + break + assert await asyncio.wait_for(result, 10) == "new" + + +def test_a_workflow_reader_start_names_one_place(): + # Checked before the reader needs a running workflow, so a mistake says + # what it is rather than that there is no workflow. + with pytest.raises(ValueError, match="either after= or last="): + workflow.stream_reader(INPUTS, after=END, last=1) + with pytest.raises(ValueError, match="positive"): + workflow.stream_reader(INPUTS, last=0)