From 09dccf42da8b7758d4583a793f30a777abfb754f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 15 Sep 2026 15:40:22 -0700 Subject: [PATCH 01/32] Added the Workflow Streams provider over the shipped wire format. --- .../streams/providers/workflow_streams.py | 341 ++++++++++++++++++ .../streams/test_workflow_streams_provider.py | 144 ++++++++ 2 files changed, 485 insertions(+) create mode 100644 temporalio/streams/providers/workflow_streams.py create mode 100644 tests/streams/test_workflow_streams_provider.py diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py new file mode 100644 index 000000000..6cc1faa4b --- /dev/null +++ b/temporalio/streams/providers/workflow_streams.py @@ -0,0 +1,341 @@ +"""The provider over today's Workflow Streams (Option 0). + +Speaks the shipped contrib feature's wire format, the +``__temporal_workflow_stream_*`` Signal, Update and Query, so interface code +and existing Workflow Streams code interoperate on one stream, and old +histories replay. Records live in the owning workflow's History, which is +also this provider's limit: Option 0's caps (payloads in History, the Signal +cap, bounded subscribers, no reads after the workflow closes) are transport +properties and remain. + +The mapping, in one place: + +- An interface record's frame rides as the item's ``Payload`` data. +- Inbound stream ``s`` is shipped topic ``in:s``; a writer topic ``t`` is + shipped topic ``out:t``, so the two namespaces cannot collide. +- Producer identity dedupes through the shipped publisher state: the + publisher id is ``producer#attempt`` and every publish Signal carries a + monotonic sequence, so a retried batch drops and a new attempt passes. +- ``append`` returns an empty cursor. The Signal transport learns positions + at read time; that is this provider's stated deviation. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator +from datetime import timedelta +from typing import Any + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.client import Client +from temporalio.common import RawValue +from temporalio.contrib.workflow_streams import ( + PublishEntry, + PublishInput, + WorkflowStream, + WorkflowStreamClient, +) +from temporalio.contrib.workflow_streams._stream import _PUBLISH_SIGNAL +from temporalio.contrib.workflow_streams._types import _encode_payload +from temporalio.streams import _frame, _provider +from temporalio.streams._handles import ReadSource, WriteSink +from temporalio.streams._policy import AttemptTracker +from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord + +_RUN_ATTR = "__temporal_streams_ws_runtime" +_IN = "in:" +_OUT = "out:" + + +class _Runtime: + """Per-run holder for the shipped stream object. + + A separate class because ``WorkflowStream`` insists on being constructed + from a method named ``__init__``, and the provider builds it lazily on + the first read or write of a run. + """ + + def __init__(self) -> None: + self.stream = WorkflowStream() + + +def _runtime() -> _Runtime: + instance = workflow.instance() + runtime = getattr(instance, _RUN_ATTR, None) + if runtime is None: + runtime = _Runtime() + setattr(instance, _RUN_ATTR, runtime) + return runtime + + +class _WSReadSource: + """Reads the signal-fed log the shipped feature keeps in workflow state. + + Reaches into the stream's private log rather than ``get_state()``, + because the snapshot copies the whole log per call and drops offsets, + and this runs inside ``workflow.wait_condition``. + """ + + def __init__(self, stream: WorkflowStream, shipped_topic: str, start: int) -> None: + self._stream = stream + self._shipped_topic = shipped_topic + self._cursor = start + self._closed = False + + def _end(self) -> int: + return self._stream._base_offset + len(self._stream._log) + + async def next_batch(self) -> list[tuple[Cursor, bytes]]: + while True: + if self._closed: + raise StopAsyncIteration + base = self._stream._base_offset + if self._cursor < base: + # Truncated below the cursor; resume at what remains. + self._cursor = base + await workflow.wait_condition( + lambda: self._closed or self._end() > self._cursor + ) + if self._closed: + raise StopAsyncIteration + batch: list[tuple[Cursor, bytes]] = [] + end = self._end() + for offset in range(self._cursor, end): + item = self._stream._log[offset - self._stream._base_offset] + if item.topic == self._shipped_topic: + batch.append((Cursor(str(offset)), item.data.data)) + self._cursor = end + if batch: + return batch + + def close(self) -> None: + self._closed = True + + +class _WSWriteSink: + def __init__(self, stream: WorkflowStream, shipped_topic: str) -> None: + self._handle = stream.topic(shipped_topic) + + async def publish(self, frame: bytes) -> None: + # Appending to workflow state commits with the task, and a poll + # Update's result rides the same task completion, so a failed task + # leaks nothing: rule 1 through the shipped mechanics. + payload = workflow.payload_converter().to_payloads([frame])[0] + self._handle.publish(payload) + + +class WorkflowStreamsProducer: + """Appends by sending the shipped publish Signal directly. + + Direct rather than through ``WorkflowStreamClient`` because the interface + owns the publisher identity: it must be ``producer#attempt`` for the + shipped dedupe to drop a retry and pass a new generation, and the client + would use its own random id. + """ + + def __init__( + self, handle: Any, converter: Any, stream: str, producer_id: str, attempt: int + ) -> None: + self._handle = handle + self._converter = converter + self._stream = stream + self._producer_id = producer_id + self._attempt = attempt + self._sequence = 0 + self._signal_sequence = 0 + + @property + def attempt(self) -> int: + return self._attempt + + @property + def _provider_id(self) -> str: + return ( + f"{self._producer_id}#{self._attempt}" if self._attempt else self._producer_id + ) + + async def append(self, *values: Any) -> Cursor: + entries = [] + for value in values: + frame = _frame.encode( + topic=self._stream, + kind=RecordKind.DATA, + producer=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + body=self._encode(value), + ) + self._sequence += 1 + entries.append(self._entry(frame)) + await self._send(entries) + return Cursor("") + + async def finish(self) -> None: + frame = _frame.encode( + topic=self._stream, + kind=RecordKind.FINISH, + producer=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + body=b"", + ) + self._sequence += 1 + await self._send([self._entry(frame)]) + + def _entry(self, frame: bytes) -> PublishEntry: + payload = Payload(data=frame) + return PublishEntry(topic=f"{_IN}{self._stream}", data=_encode_payload(payload)) + + async def _send(self, entries: list[PublishEntry]) -> None: + self._signal_sequence += 1 + await self._handle.signal( + _PUBLISH_SIGNAL, + PublishInput( + items=entries, + publisher_id=self._provider_id, + sequence=self._signal_sequence, + ), + ) + + def _encode(self, value: Any) -> bytes: + payload = ( + value + if isinstance(value, Payload) + else self._converter.to_payloads([value])[0] + ) + return payload.SerializeToString() + + +class WorkflowStreamsConsumer: + """Reads through the shipped long-poll Update, with shared supersession.""" + + def __init__( + self, + stream_client: WorkflowStreamClient, + shipped_topic: str | None, + poll_cooldown: timedelta, + ) -> None: + self._client = stream_client + self._shipped_topic = shipped_topic + self._poll_cooldown = poll_cooldown + + async def read( + self, + *, + start: Cursor = BEGINNING, + topic: str | None = None, + type: type | None = None, + ) -> AsyncIterator[StreamRecord[Any]]: + attempts = AttemptTracker() + subscription = self._client.subscribe( + self._shipped_topic, + from_offset=int(start.token) if start.token else 0, + result_type=RawValue, + poll_cooldown=self._poll_cooldown, + ) + async for item in subscription: + if self._shipped_topic is None and not item.topic.startswith(_OUT): + continue + cursor = Cursor(str(item.offset)) + try: + kind, frame_topic, source, attempt, sequence, body = _frame.decode( + item.data.payload.data + ) + except ValueError: + continue + if topic is not None and frame_topic != topic: + continue + superseded = attempts.note(source, attempt, cursor) + if superseded is not None: + yield superseded + yield StreamRecord( + value=self._decode(body, type) if kind is RecordKind.DATA else None, + cursor=cursor, + kind=kind, + topic=frame_topic, + producer=source, + attempt=attempt, + sequence=sequence, + ) + + def _decode(self, body: bytes, as_type: type | None) -> Any: + payload = Payload() + payload.ParseFromString(body) + converter = self._client._client.data_converter.payload_converter + if as_type is None: + return converter.from_payloads([payload])[0] + return converter.from_payloads([payload], [as_type])[0] + + +class _WorkflowStreamsProvider: + name = "workflow_streams" + + def __init__(self) -> None: + self._poll_cooldown = timedelta(milliseconds=100) + + def configure(self, **options: Any) -> None: + cooldown = options.pop("poll_cooldown", None) + if cooldown is not None: + self._poll_cooldown = cooldown + if options: + raise TypeError( + "the workflow_streams provider takes only poll_cooldown, " + f"got {sorted(options)}" + ) + + def worker_options(self) -> dict[str, Any]: + return {} + + def open_read( + self, + stream: str, + *, + start: Cursor = BEGINNING, + idle_timeout: timedelta | None = None, + ) -> ReadSource: + # Ignored: a publish Signal is a workflow event, so the wait below + # wakes on delivery and nothing is held between records. + del idle_timeout + return _WSReadSource( + _runtime().stream, + f"{_IN}{stream}", + int(start.token) if start.token else 0, + ) + + def open_write(self, topic: str) -> WriteSink: + return _WSWriteSink(_runtime().stream, f"{_OUT}{topic}") + + async def producer( + self, + client: Client, + *, + workflow_id: str, + stream: str, + producer_id: str = "", + attempt: int = 0, + ) -> WorkflowStreamsProducer: + if not producer_id: + from temporalio import activity + + producer_id = activity.info().activity_id + attempt = attempt or activity.info().attempt + return WorkflowStreamsProducer( + client.get_workflow_handle(workflow_id), + client.data_converter.payload_converter, + stream, + producer_id, + attempt, + ) + + async def consumer( + self, client: Client, *, workflow_id: str, stream: str = "" + ) -> WorkflowStreamsConsumer: + return WorkflowStreamsConsumer( + WorkflowStreamClient.create(client, workflow_id), + f"{_IN}{stream}" if stream else None, + self._poll_cooldown, + ) + + +_provider.register("workflow_streams", _WorkflowStreamsProvider) diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py new file mode 100644 index 000000000..2268234ee --- /dev/null +++ b/tests/streams/test_workflow_streams_provider.py @@ -0,0 +1,144 @@ +"""Live conformance for the workflow_streams provider. + +Runs the interface loop over the shipped Option 0 transport against a real +server: an outside producer appends through the publish Signal, the workflow +reads and republishes through its own state, and an outside consumer follows +the poll Update. Gated behind ``STREAMS_LIVE=workflow_streams`` because it +needs a running server (``TEMPORAL_ADDRESS``, default ``localhost:7233``). +""" + +from __future__ import annotations + +import asyncio +import os +import uuid + +import pytest + +from temporalio import streams, workflow +from temporalio.client import Client +from temporalio.streams import RecordKind +from temporalio.worker import Worker + +pytestmark = pytest.mark.skipif( + os.environ.get("STREAMS_LIVE") != "workflow_streams", + reason="needs a live server; run with STREAMS_LIVE=workflow_streams", +) + + +@workflow.defn +class EchoLoop: + """Reads ``inputs``, echoes each value onto ``decisions``, ends on FINISH.""" + + @workflow.run + async def run(self) -> int: + inputs = streams.reader("inputs", type=dict) + decisions = streams.writer("decisions") + seen = 0 + async for record in inputs: + if record.kind is RecordKind.FINISH: + break + if record.kind is not RecordKind.DATA: + continue + seen += 1 + await decisions.publish({"echo": record.value["n"]}) + await decisions.finish() + return seen + + +async def take(records, count: int, timeout: float = 30.0) -> list: + out: list = [] + + async def _collect() -> None: + async for record in records: + out.append(record) + if len(out) >= count: + return + + await asyncio.wait_for(_collect(), timeout) + return out + + +async def test_interface_loop_over_workflow_streams(): + streams.configure(provider="workflow_streams") + client = await Client.connect( + os.environ.get("TEMPORAL_ADDRESS", "localhost:7233") + ) + workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" + + async with Worker( + client, + task_queue=f"tq-{workflow_id}", + workflows=[EchoLoop], + **streams.worker_options(), + ): + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + + producer = await streams.producer( + client, + workflow_id=workflow_id, + stream="inputs", + producer_id="model", + attempt=1, + ) + await producer.append({"n": 1}, {"n": 2}) + await producer.append({"n": 3}) + await producer.finish() + + consumer = await streams.consumer(client, workflow_id=workflow_id) + records = await take(consumer.read(type=dict), 4) + assert [r.kind for r in records] == [ + RecordKind.DATA, + RecordKind.DATA, + RecordKind.DATA, + RecordKind.FINISH, + ] + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + + assert await handle.result() == 3 + + +async def test_retried_producer_dedupes_and_new_attempt_supersedes(): + streams.configure(provider="workflow_streams") + client = await Client.connect( + os.environ.get("TEMPORAL_ADDRESS", "localhost:7233") + ) + workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" + + async with Worker( + client, + task_queue=f"tq-{workflow_id}", + workflows=[EchoLoop], + **streams.worker_options(), + ): + await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + + first = await streams.producer( + client, workflow_id=workflow_id, stream="inputs", + producer_id="model", attempt=1, + ) + await first.append({"n": 1}) + # The retry of the same attempt re-sends its first batch. + retry = await streams.producer( + client, workflow_id=workflow_id, stream="inputs", + producer_id="model", attempt=1, + ) + await retry.append({"n": 1}) + second = await streams.producer( + client, workflow_id=workflow_id, stream="inputs", + producer_id="model", attempt=2, + ) + await second.append({"n": 2}) + + consumer = await streams.consumer( + client, workflow_id=workflow_id, stream="inputs" + ) + records = await take(consumer.read(type=dict), 3) + assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 + assert records[1].kind is RecordKind.SUPERSEDED + assert records[1].value.previous_attempt == 1 + assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 From f088496082159fbe95e368b532115ff27a6fbe0f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 15 Sep 2026 15:55:22 -0700 Subject: [PATCH 02/32] Let a workflow drain parked pollers before returning. --- .../streams/providers/workflow_streams.py | 15 +++++++++++++++ .../streams/test_workflow_streams_provider.py | 19 ++++++++++++++++++- 2 files changed, 33 insertions(+), 1 deletion(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 6cc1faa4b..a12608692 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -69,6 +69,21 @@ def _runtime() -> _Runtime: return runtime +def drain() -> None: + """Release parked pollers so the workflow can return. + + An Option 0 stream dies with its workflow, and a parked long-poll Update + would otherwise hold completion open. Call it right before the workflow + returns, the same obligation the shipped feature's ``detach_pollers`` + documents. A storage provider has no such step, which is one of the + differences the comparison table charges this transport with. + """ + instance = workflow.instance() + runtime = getattr(instance, _RUN_ATTR, None) + if runtime is not None: + runtime.stream.detach_pollers() + + class _WSReadSource: """Reads the signal-fed log the shipped feature keeps in workflow state. diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 2268234ee..2b491cbf1 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -18,6 +18,7 @@ from temporalio import streams, workflow from temporalio.client import Client from temporalio.streams import RecordKind +from temporalio.streams.providers.workflow_streams import drain from temporalio.worker import Worker pytestmark = pytest.mark.skipif( @@ -28,7 +29,19 @@ @workflow.defn class EchoLoop: - """Reads ``inputs``, echoes each value onto ``decisions``, ends on FINISH.""" + """Reads ``inputs``, echoes each value onto ``decisions``, ends on FINISH. + + Lingers until released, because an Option 0 stream dies with its + workflow: a reader that arrives after close finds nothing, which is the + transport limit the doc states rather than a defect to fix here. + """ + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True @workflow.run async def run(self) -> int: @@ -43,6 +56,9 @@ async def run(self) -> int: seen += 1 await decisions.publish({"echo": record.value["n"]}) await decisions.finish() + await workflow.wait_condition(lambda: self._released) + drain() + await workflow.wait_condition(workflow.all_handlers_finished) return seen @@ -97,6 +113,7 @@ async def test_interface_loop_over_workflow_streams(): ] assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + await handle.signal(EchoLoop.release) assert await handle.result() == 3 From 3393901ea9d802577bcfdd55d77f205d3b1d59c8 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 15 Sep 2026 17:58:03 -0700 Subject: [PATCH 03/32] Added the prepare and drain hooks the Workflow Streams transport needs. A provider that answers outside readers through handlers on the workflow has to register them before the first task completes, and one that parks a long-poll update against the run has to let go before the workflow returns. Both are no-ops on the storage providers, so workflow code calls them unconditionally and stays portable. --- .../streams/providers/workflow_streams.py | 28 +++++++++++++++---- 1 file changed, 23 insertions(+), 5 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index a12608692..70ec01293 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -60,12 +60,18 @@ def __init__(self) -> None: self.stream = WorkflowStream() +# Held per run rather than on the workflow instance, because the provider is +# asked to install its handlers from the workflow's constructor, and the +# instance is not registered with the runtime yet at that point. +_runtimes: dict[str, _Runtime] = {} + + def _runtime() -> _Runtime: - instance = workflow.instance() - runtime = getattr(instance, _RUN_ATTR, None) + key = workflow.info().run_id + runtime = _runtimes.get(key) if runtime is None: runtime = _Runtime() - setattr(instance, _RUN_ATTR, runtime) + _runtimes[key] = runtime return runtime @@ -78,8 +84,7 @@ def drain() -> None: documents. A storage provider has no such step, which is one of the differences the comparison table charges this transport with. """ - instance = workflow.instance() - runtime = getattr(instance, _RUN_ATTR, None) + runtime = _runtimes.pop(workflow.info().run_id, None) if runtime is not None: runtime.stream.detach_pollers() @@ -302,6 +307,19 @@ def configure(self, **options: Any) -> None: def worker_options(self) -> dict[str, Any]: return {} + def prepare(self) -> None: + """Register the shipped publish signal and poll update handlers. + + Done here rather than on the first read, because an outside reader + can poll before workflow code has opened anything, and an update with + no handler yet is rejected rather than held. + """ + _runtime() + + def drain(self) -> None: + """Release parked pollers so the workflow can return.""" + drain() + def open_read( self, stream: str, From ee1f64d7d1ae77bbaca149dec857b47e5ca971a0 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 15 Sep 2026 18:01:04 -0700 Subject: [PATCH 04/32] Made cursors exclusive, added latest(), and allowed topic appends. Porting the agent harness onto the interface surfaced three things the small example never needed. A reader resumes after the record a cursor names, so it can store the last cursor it handled without advancing an opaque token. A consumer reports its latest position, so a client can follow a turn it is about to start without the workflow reporting a position. And a producer can append onto a topic of the stream the workflow publishes, so an activity's live output lands next to the workflow's own records for one outside reader. --- .../streams/providers/workflow_streams.py | 44 ++++++++++++++----- 1 file changed, 34 insertions(+), 10 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 70ec01293..d5bc5bc70 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -12,7 +12,9 @@ - An interface record's frame rides as the item's ``Payload`` data. - Inbound stream ``s`` is shipped topic ``in:s``; a writer topic ``t`` is - shipped topic ``out:t``, so the two namespaces cannot collide. + shipped topic ``out:t``, so the two namespaces cannot collide. A producer + appending onto the workflow's own topic ``t`` writes ``out:t`` as well, + which is what today's activities already do through the publish Signal. - Producer identity dedupes through the shipped publisher state: the publisher id is ``producer#attempt`` and every publish Signal carries a monotonic sequence, so a retried batch drops and a new attempt passes. @@ -155,16 +157,31 @@ class WorkflowStreamsProducer: """ def __init__( - self, handle: Any, converter: Any, stream: str, producer_id: str, attempt: int + self, + handle: Any, + converter: Any, + stream: str, + topic: str, + producer_id: str, + attempt: int, ) -> None: self._handle = handle self._converter = converter self._stream = stream + self._topic = topic self._producer_id = producer_id self._attempt = attempt self._sequence = 0 self._signal_sequence = 0 + @property + def _frame_topic(self) -> str: + return self._stream or self._topic + + @property + def _shipped_topic(self) -> str: + return f"{_IN}{self._stream}" if self._stream else f"{_OUT}{self._topic}" + @property def attempt(self) -> int: return self._attempt @@ -179,7 +196,7 @@ async def append(self, *values: Any) -> Cursor: entries = [] for value in values: frame = _frame.encode( - topic=self._stream, + topic=self._frame_topic, kind=RecordKind.DATA, producer=self._producer_id, attempt=self._attempt, @@ -193,7 +210,7 @@ async def append(self, *values: Any) -> Cursor: async def finish(self) -> None: frame = _frame.encode( - topic=self._stream, + topic=self._frame_topic, kind=RecordKind.FINISH, producer=self._producer_id, attempt=self._attempt, @@ -205,7 +222,7 @@ async def finish(self) -> None: def _entry(self, frame: bytes) -> PublishEntry: payload = Payload(data=frame) - return PublishEntry(topic=f"{_IN}{self._stream}", data=_encode_payload(payload)) + return PublishEntry(topic=self._shipped_topic, data=_encode_payload(payload)) async def _send(self, entries: list[PublishEntry]) -> None: self._signal_sequence += 1 @@ -243,14 +260,14 @@ def __init__( async def read( self, *, - start: Cursor = BEGINNING, + after: Cursor = BEGINNING, topic: str | None = None, type: type | None = None, ) -> AsyncIterator[StreamRecord[Any]]: attempts = AttemptTracker() subscription = self._client.subscribe( self._shipped_topic, - from_offset=int(start.token) if start.token else 0, + from_offset=int(after.token) + 1 if after.token else 0, result_type=RawValue, poll_cooldown=self._poll_cooldown, ) @@ -279,6 +296,11 @@ async def read( sequence=sequence, ) + async def latest(self, *, topic: str | None = None) -> Cursor: + del topic # one log per workflow, whatever the topic + head = await self._client.get_offset() + return Cursor(str(head - 1)) if head > 0 else BEGINNING + def _decode(self, body: bytes, as_type: type | None) -> Any: payload = Payload() payload.ParseFromString(body) @@ -324,7 +346,7 @@ def open_read( self, stream: str, *, - start: Cursor = BEGINNING, + after: Cursor = BEGINNING, idle_timeout: timedelta | None = None, ) -> ReadSource: # Ignored: a publish Signal is a workflow event, so the wait below @@ -333,7 +355,7 @@ def open_read( return _WSReadSource( _runtime().stream, f"{_IN}{stream}", - int(start.token) if start.token else 0, + int(after.token) + 1 if after.token else 0, ) def open_write(self, topic: str) -> WriteSink: @@ -344,7 +366,8 @@ async def producer( client: Client, *, workflow_id: str, - stream: str, + stream: str = "", + topic: str = "", producer_id: str = "", attempt: int = 0, ) -> WorkflowStreamsProducer: @@ -357,6 +380,7 @@ async def producer( client.get_workflow_handle(workflow_id), client.data_converter.payload_converter, stream, + topic, producer_id, attempt, ) From 91d11f6559ae1391d3b4324fe5161eecaf9f84a2 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 15 Sep 2026 18:28:01 -0700 Subject: [PATCH 05/32] Served a closed workflow's stream tail by Query on the shipped transport. The poll Update stops answering once the workflow is closing, so a reader between polls at that moment lost whatever the final task published, and nothing could read the stream after completion at all. The log is workflow state, so one Query serves it for as long as the History is retained; the consumer asks for the tail when its subscription ends. --- .../streams/providers/workflow_streams.py | 101 ++++++++++++++---- 1 file changed, 83 insertions(+), 18 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index d5bc5bc70..d849f1244 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -5,8 +5,10 @@ and existing Workflow Streams code interoperate on one stream, and old histories replay. Records live in the owning workflow's History, which is also this provider's limit: Option 0's caps (payloads in History, the Signal -cap, bounded subscribers, no reads after the workflow closes) are transport -properties and remain. +cap, bounded subscribers) are transport properties and remain. Reads after +the workflow closes go through one Query this provider adds, which serves +the log from workflow state for as long as the History is retained, so a +reader between polls when the workflow completed still gets the tail. The mapping, in one place: @@ -24,6 +26,7 @@ from __future__ import annotations +import base64 from collections.abc import AsyncIterator from datetime import timedelta from typing import Any @@ -46,6 +49,7 @@ from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord _RUN_ATTR = "__temporal_streams_ws_runtime" +_TAIL_QUERY = "__temporal_streams_tail" _IN = "in:" _OUT = "out:" @@ -60,6 +64,23 @@ class _Runtime: def __init__(self) -> None: self.stream = WorkflowStream() + # The poll Update stops answering once the workflow is closing, and a + # reader between polls at that moment would lose what the final task + # published. The log is workflow state, so a Query still serves it + # after completion. + workflow.set_query_handler(_TAIL_QUERY, self._tail) + + def _tail(self, from_offset: int) -> list[dict[str, Any]]: + base = self.stream._base_offset + return [ + { + "offset": base + index, + "topic": item.topic, + "data": base64.b64encode(item.data.data).decode("ascii"), + } + for index, item in enumerate(self.stream._log) + if base + index >= from_offset + ] # Held per run rather than on the workflow instance, because the provider is @@ -265,28 +286,70 @@ async def read( type: type | None = None, ) -> AsyncIterator[StreamRecord[Any]]: attempts = AttemptTracker() + next_offset = int(after.token) + 1 if after.token else 0 subscription = self._client.subscribe( self._shipped_topic, - from_offset=int(after.token) + 1 if after.token else 0, + from_offset=next_offset, result_type=RawValue, poll_cooldown=self._poll_cooldown, ) async for item in subscription: - if self._shipped_topic is None and not item.topic.startswith(_OUT): - continue - cursor = Cursor(str(item.offset)) - try: - kind, frame_topic, source, attempt, sequence, body = _frame.decode( - item.data.payload.data - ) - except ValueError: - continue - if topic is not None and frame_topic != topic: - continue - superseded = attempts.note(source, attempt, cursor) - if superseded is not None: - yield superseded - yield StreamRecord( + next_offset = item.offset + 1 + record = self._record( + attempts, item.offset, item.topic, item.data.payload.data, topic, type + ) + for out in record: + yield out + # The subscription ends when the workflow is closing or closed. What + # landed after the last poll is still in workflow state, so the tail + # comes back by Query rather than being lost with the run. + for wire in await self._tail(next_offset): + for out in self._record( + attempts, + wire["offset"], + wire["topic"], + base64.b64decode(wire["data"]), + topic, + type, + ): + yield out + + async def _tail(self, from_offset: int) -> list[dict[str, Any]]: + try: + return await self._client._handle.query( + _TAIL_QUERY, from_offset, result_type=list + ) + except Exception: + # A workflow that never opened a stream has no handler to ask, and + # one whose History is gone has nothing left to serve. + return [] + + def _record( + self, + attempts: AttemptTracker, + offset: int, + shipped_topic: str, + frame: bytes, + topic: str | None, + type: type | None, + ) -> list[StreamRecord[Any]]: + if self._shipped_topic is None and not shipped_topic.startswith(_OUT): + return [] + if self._shipped_topic is not None and shipped_topic != self._shipped_topic: + return [] + cursor = Cursor(str(offset)) + try: + kind, frame_topic, source, attempt, sequence, body = _frame.decode(frame) + except ValueError: + return [] + if topic is not None and frame_topic != topic: + return [] + out: list[StreamRecord[Any]] = [] + superseded = attempts.note(source, attempt, cursor) + if superseded is not None: + out.append(superseded) + out.append( + StreamRecord( value=self._decode(body, type) if kind is RecordKind.DATA else None, cursor=cursor, kind=kind, @@ -295,6 +358,8 @@ async def read( attempt=attempt, sequence=sequence, ) + ) + return out async def latest(self, *, topic: str | None = None) -> Cursor: del topic # one log per workflow, whatever the topic From 87ae2c87ebc852b4b3227a92b1a4ce18b5bc08bd Mon Sep 17 00:00:00 2001 From: Moe Dashti Date: Wed, 16 Sep 2026 15:06:58 -0700 Subject: [PATCH 06/32] Made the Workflow Streams provider pass the linters. The transport already resolves its payload converter and falls back to the default when it holds no client, so the provider asks it rather than reaching through a client that may not be there. --- .../streams/providers/workflow_streams.py | 13 ++++++-- .../streams/test_workflow_streams_provider.py | 32 +++++++++++-------- 2 files changed, 30 insertions(+), 15 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index d849f1244..656056381 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -186,6 +186,7 @@ def __init__( producer_id: str, attempt: int, ) -> None: + """Bind this producer to one stream or topic on ``handle``.""" self._handle = handle self._converter = converter self._stream = stream @@ -205,15 +206,19 @@ def _shipped_topic(self) -> str: @property def attempt(self) -> int: + """The generation this producer is writing.""" return self._attempt @property def _provider_id(self) -> str: return ( - f"{self._producer_id}#{self._attempt}" if self._attempt else self._producer_id + f"{self._producer_id}#{self._attempt}" + if self._attempt + else self._producer_id ) async def append(self, *values: Any) -> Cursor: + """Append ``values`` through the shipped publish Signal.""" entries = [] for value in values: frame = _frame.encode( @@ -230,6 +235,7 @@ async def append(self, *values: Any) -> Cursor: return Cursor("") async def finish(self) -> None: + """Mark this producer done, so a reader stops waiting on it.""" frame = _frame.encode( topic=self._frame_topic, kind=RecordKind.FINISH, @@ -274,6 +280,7 @@ def __init__( shipped_topic: str | None, poll_cooldown: timedelta, ) -> None: + """Read what ``stream_client`` reaches, one long poll at a time.""" self._client = stream_client self._shipped_topic = shipped_topic self._poll_cooldown = poll_cooldown @@ -285,6 +292,7 @@ async def read( topic: str | None = None, type: type | None = None, ) -> AsyncIterator[StreamRecord[Any]]: + """Yield records after ``after``, waiting for ones not written yet.""" attempts = AttemptTracker() next_offset = int(after.token) + 1 if after.token else 0 subscription = self._client.subscribe( @@ -362,6 +370,7 @@ def _record( return out async def latest(self, *, topic: str | None = None) -> Cursor: + """The cursor of the last record written, for following from now.""" del topic # one log per workflow, whatever the topic head = await self._client.get_offset() return Cursor(str(head - 1)) if head > 0 else BEGINNING @@ -369,7 +378,7 @@ async def latest(self, *, topic: str | None = None) -> Cursor: def _decode(self, body: bytes, as_type: type | None) -> Any: payload = Payload() payload.ParseFromString(body) - converter = self._client._client.data_converter.payload_converter + converter = self._client._payload_converter() if as_type is None: return converter.from_payloads([payload])[0] return converter.from_payloads([payload], [as_type])[0] diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 2b491cbf1..612aaa33d 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -12,6 +12,7 @@ import asyncio import os import uuid +from typing import Any import pytest @@ -62,7 +63,7 @@ async def run(self) -> int: return seen -async def take(records, count: int, timeout: float = 30.0) -> list: +async def take(records: Any, count: int, timeout: float = 30.0) -> list: out: list = [] async def _collect() -> None: @@ -77,9 +78,7 @@ async def _collect() -> None: async def test_interface_loop_over_workflow_streams(): streams.configure(provider="workflow_streams") - client = await Client.connect( - os.environ.get("TEMPORAL_ADDRESS", "localhost:7233") - ) + client = await Client.connect(os.environ.get("TEMPORAL_ADDRESS", "localhost:7233")) workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" async with Worker( @@ -119,9 +118,7 @@ async def test_interface_loop_over_workflow_streams(): async def test_retried_producer_dedupes_and_new_attempt_supersedes(): streams.configure(provider="workflow_streams") - client = await Client.connect( - os.environ.get("TEMPORAL_ADDRESS", "localhost:7233") - ) + client = await Client.connect(os.environ.get("TEMPORAL_ADDRESS", "localhost:7233")) workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" async with Worker( @@ -135,19 +132,28 @@ async def test_retried_producer_dedupes_and_new_attempt_supersedes(): ) first = await streams.producer( - client, workflow_id=workflow_id, stream="inputs", - producer_id="model", attempt=1, + client, + workflow_id=workflow_id, + stream="inputs", + producer_id="model", + attempt=1, ) await first.append({"n": 1}) # The retry of the same attempt re-sends its first batch. retry = await streams.producer( - client, workflow_id=workflow_id, stream="inputs", - producer_id="model", attempt=1, + client, + workflow_id=workflow_id, + stream="inputs", + producer_id="model", + attempt=1, ) await retry.append({"n": 1}) second = await streams.producer( - client, workflow_id=workflow_id, stream="inputs", - producer_id="model", attempt=2, + client, + workflow_id=workflow_id, + stream="inputs", + producer_id="model", + attempt=2, ) await second.append({"n": 2}) From 305a0d216fad65315c78e3f96f72ffacb693768b Mon Sep 17 00:00:00 2001 From: Moe Dashti Date: Wed, 16 Sep 2026 15:06:58 -0700 Subject: [PATCH 07/32] Added the changelog entry for the Workflow Streams provider. The changelog checkpoint requires an entry for any user-facing change. --- CHANGELOG.md | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index c74bdf72a..811ebef5d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,6 +33,10 @@ to include examples, links to docs, or any other relevant information. is `temporal.api.stream.v1.StreamRecord` on every provider. `temporalio.streams.providers.memory.MemoryStreams` is the in-memory reference provider the conformance tests run against. +- **Experimental**: `temporalio.streams.providers.workflow_streams` serves the + stream interface over the shipped Workflow Streams transport, so a workflow + reads and publishes through `temporalio.contrib.workflow_streams` without + naming it. ### Changed From f33de01d914be2a63733ea8cf662647a5614ab76 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:44:59 -0700 Subject: [PATCH 08/32] Added public accessors to Workflow Streams for code that rides its wire. A provider that reuses the shipped transport needs the publish signal's name, the log past an offset, and the client's handle and converter. Reading them off private attributes breaks silently on the next contrib change. --- .../contrib/workflow_streams/__init__.py | 6 +++- .../contrib/workflow_streams/_client.py | 14 ++++++++++ .../contrib/workflow_streams/_stream.py | 28 ++++++++++++++++++- 3 files changed, 46 insertions(+), 2 deletions(-) diff --git a/temporalio/contrib/workflow_streams/__init__.py b/temporalio/contrib/workflow_streams/__init__.py index 41f670f0c..3f225bb2f 100644 --- a/temporalio/contrib/workflow_streams/__init__.py +++ b/temporalio/contrib/workflow_streams/__init__.py @@ -13,7 +13,10 @@ """ from temporalio.contrib.workflow_streams._client import WorkflowStreamClient -from temporalio.contrib.workflow_streams._stream import WorkflowStream +from temporalio.contrib.workflow_streams._stream import ( + PUBLISH_SIGNAL_NAME, + WorkflowStream, +) from temporalio.contrib.workflow_streams._topic_handle import ( TopicHandle, WorkflowTopicHandle, @@ -29,6 +32,7 @@ ) __all__ = [ + "PUBLISH_SIGNAL_NAME", "PollInput", "PollResult", "PublishEntry", diff --git a/temporalio/contrib/workflow_streams/_client.py b/temporalio/contrib/workflow_streams/_client.py index 605bf3f03..5c458bdab 100644 --- a/temporalio/contrib/workflow_streams/_client.py +++ b/temporalio/contrib/workflow_streams/_client.py @@ -325,6 +325,20 @@ def topic( self._topic_types[name] = bound return TopicHandle(self, name, bound) + @property + def handle(self) -> WorkflowHandle[Any, Any]: + """The workflow handle this client publishes to and polls. + + Re-targeted when :py:meth:`subscribe` follows a continue-as-new, so + read it when needed rather than caching it. + """ + return self._handle + + @property + def payload_converter(self) -> PayloadConverter: + """The sync payload converter used for per-item encode and decode.""" + return self._payload_converter() + async def flush(self) -> None: """Flush buffered (and pending) items and wait for server confirmation. diff --git a/temporalio/contrib/workflow_streams/_stream.py b/temporalio/contrib/workflow_streams/_stream.py index ae8608c3b..ff0865c29 100644 --- a/temporalio/contrib/workflow_streams/_stream.py +++ b/temporalio/contrib/workflow_streams/_stream.py @@ -50,7 +50,14 @@ _WorkflowStreamWireItem, ) -_PUBLISH_SIGNAL = "__temporal_workflow_stream_publish" +PUBLISH_SIGNAL_NAME = "__temporal_workflow_stream_publish" +"""The signal :class:`WorkflowStream` registers for external publishes. + +Public so code that sends the signal itself, with its own publisher identity, +does not have to copy the name. +""" + +_PUBLISH_SIGNAL = PUBLISH_SIGNAL_NAME _POLL_UPDATE = "__temporal_workflow_stream_poll" _OFFSET_QUERY = "__temporal_workflow_stream_offset" @@ -234,6 +241,25 @@ def topic( self._topic_types[name] = bound return WorkflowTopicHandle(self, name, bound) + @property + def next_offset(self) -> int: + """The global offset the next published item will receive.""" + return self._base_offset + len(self._log) + + def items_from(self, offset: int) -> list[tuple[int, str, Payload]]: + """Return ``(offset, topic, payload)`` for every item at or past ``offset``. + + Reads the log in place, so it is safe to call from a + :func:`temporalio.workflow.wait_condition` predicate. An ``offset`` + below the truncation base starts at the base instead; the offsets in + the result say where the items actually sit. + """ + start = max(offset, self._base_offset) - self._base_offset + return [ + (self._base_offset + index, item.topic, item.data) + for index, item in enumerate(self._log[start:], start) + ] + def get_state( self, *, publisher_ttl: timedelta = timedelta(seconds=900) ) -> WorkflowStreamState: From 643bcdd78f6f80908487448619f8a31ebd4a9973 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:51:11 -0700 Subject: [PATCH 09/32] Keyed the Workflow Streams runtime by workflow instance, not run id. The SDK rebuilds an evicted workflow from history as a new object. A process map keyed by run id handed that object the previous instance's log, with no handlers registered on the new one and the records of a failed task still inside. The stream is now found through the handler the shipped class registers on the instance, and the provider reads the contrib module only through its public surface. --- .../streams/providers/workflow_streams.py | 121 ++++++++---------- 1 file changed, 54 insertions(+), 67 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 656056381..8155f7e9d 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -36,6 +36,7 @@ from temporalio.client import Client from temporalio.common import RawValue from temporalio.contrib.workflow_streams import ( + PUBLISH_SIGNAL_NAME, PublishEntry, PublishInput, WorkflowStream, @@ -48,77 +49,63 @@ from temporalio.streams._policy import AttemptTracker from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord -_RUN_ATTR = "__temporal_streams_ws_runtime" _TAIL_QUERY = "__temporal_streams_tail" _IN = "in:" _OUT = "out:" class _Runtime: - """Per-run holder for the shipped stream object. + """A view over the shipped stream object of the running workflow instance. A separate class because ``WorkflowStream`` insists on being constructed - from a method named ``__init__``, and the provider builds it lazily on - the first read or write of a run. + from a method named ``__init__``. """ - def __init__(self) -> None: - self.stream = WorkflowStream() - # The poll Update stops answering once the workflow is closing, and a - # reader between polls at that moment would lose what the final task - # published. The log is workflow state, so a Query still serves it - # after completion. - workflow.set_query_handler(_TAIL_QUERY, self._tail) + def __init__(self, stream: WorkflowStream | None = None) -> None: + self.stream = WorkflowStream() if stream is None else stream + if workflow.get_query_handler(_TAIL_QUERY) is None: + # The poll Update stops answering once the workflow is closing, + # and a reader between polls at that moment would lose what the + # final task published. The log is workflow state, so a Query + # still serves it after completion. + workflow.set_query_handler(_TAIL_QUERY, self._tail) def _tail(self, from_offset: int) -> list[dict[str, Any]]: - base = self.stream._base_offset return [ { - "offset": base + index, - "topic": item.topic, - "data": base64.b64encode(item.data.data).decode("ascii"), + "offset": offset, + "topic": topic, + "data": base64.b64encode(payload.data).decode("ascii"), } - for index, item in enumerate(self.stream._log) - if base + index >= from_offset + for offset, topic, payload in self.stream.items_from(from_offset) ] -# Held per run rather than on the workflow instance, because the provider is -# asked to install its handlers from the workflow's constructor, and the -# instance is not registered with the runtime yet at that point. -_runtimes: dict[str, _Runtime] = {} +def _registered_stream() -> WorkflowStream | None: + handler = workflow.get_signal_handler(PUBLISH_SIGNAL_NAME) + if handler is None: + return None + stream = getattr(handler, "__self__", None) + if not isinstance(stream, WorkflowStream): + raise RuntimeError( + f"the {PUBLISH_SIGNAL_NAME!r} signal on this workflow is handled by " + "something other than a WorkflowStream, so the workflow_streams " + "provider cannot share its log" + ) + return stream def _runtime() -> _Runtime: - key = workflow.info().run_id - runtime = _runtimes.get(key) - if runtime is None: - runtime = _Runtime() - _runtimes[key] = runtime - return runtime - - -def drain() -> None: - """Release parked pollers so the workflow can return. - - An Option 0 stream dies with its workflow, and a parked long-poll Update - would otherwise hold completion open. Call it right before the workflow - returns, the same obligation the shipped feature's ``detach_pollers`` - documents. A storage provider has no such step, which is one of the - differences the comparison table charges this transport with. - """ - runtime = _runtimes.pop(workflow.info().run_id, None) - if runtime is not None: - runtime.stream.detach_pollers() + # Found on the instance rather than in a process-level map keyed by run + # id: the SDK rebuilds an evicted workflow from history as a new object, + # and a map would hand that object the stale log with its unregistered + # handlers and the records of a task that failed. + stream = _registered_stream() + return _Runtime() if stream is None else _Runtime(stream) class _WSReadSource: - """Reads the signal-fed log the shipped feature keeps in workflow state. - - Reaches into the stream's private log rather than ``get_state()``, - because the snapshot copies the whole log per call and drops offsets, - and this runs inside ``workflow.wait_condition``. - """ + """Reads the signal-fed log the shipped feature keeps in workflow state.""" def __init__(self, stream: WorkflowStream, shipped_topic: str, start: int) -> None: self._stream = stream @@ -126,29 +113,21 @@ def __init__(self, stream: WorkflowStream, shipped_topic: str, start: int) -> No self._cursor = start self._closed = False - def _end(self) -> int: - return self._stream._base_offset + len(self._stream._log) - async def next_batch(self) -> list[tuple[Cursor, bytes]]: while True: if self._closed: raise StopAsyncIteration - base = self._stream._base_offset - if self._cursor < base: - # Truncated below the cursor; resume at what remains. - self._cursor = base await workflow.wait_condition( - lambda: self._closed or self._end() > self._cursor + lambda: self._closed or self._stream.next_offset > self._cursor ) if self._closed: raise StopAsyncIteration - batch: list[tuple[Cursor, bytes]] = [] - end = self._end() - for offset in range(self._cursor, end): - item = self._stream._log[offset - self._stream._base_offset] - if item.topic == self._shipped_topic: - batch.append((Cursor(str(offset)), item.data.data)) - self._cursor = end + batch = [ + (Cursor(str(offset)), payload.data) + for offset, topic, payload in self._stream.items_from(self._cursor) + if topic == self._shipped_topic + ] + self._cursor = self._stream.next_offset if batch: return batch @@ -254,7 +233,7 @@ def _entry(self, frame: bytes) -> PublishEntry: async def _send(self, entries: list[PublishEntry]) -> None: self._signal_sequence += 1 await self._handle.signal( - _PUBLISH_SIGNAL, + PUBLISH_SIGNAL_NAME, PublishInput( items=entries, publisher_id=self._provider_id, @@ -324,7 +303,7 @@ async def read( async def _tail(self, from_offset: int) -> list[dict[str, Any]]: try: - return await self._client._handle.query( + wire = await self._client.handle.query( _TAIL_QUERY, from_offset, result_type=list ) except Exception: @@ -378,7 +357,7 @@ async def latest(self, *, topic: str | None = None) -> Cursor: def _decode(self, body: bytes, as_type: type | None) -> Any: payload = Payload() payload.ParseFromString(body) - converter = self._client._payload_converter() + converter = self._client.payload_converter if as_type is None: return converter.from_payloads([payload])[0] return converter.from_payloads([payload], [as_type])[0] @@ -413,8 +392,16 @@ def prepare(self) -> None: _runtime() def drain(self) -> None: - """Release parked pollers so the workflow can return.""" - drain() + """Release parked pollers so the workflow can return. + + An Option 0 stream dies with its workflow, and a parked long-poll + Update would otherwise hold completion open. Call it right before + the workflow returns, the same obligation the shipped feature's + ``detach_pollers`` documents. + """ + stream = _registered_stream() + if stream is not None: + stream.detach_pollers() def open_read( self, From d015a6ff68484eca813d59b9a048b2cf3f0e4796 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:51:11 -0700 Subject: [PATCH 10/32] Committed producer sequences only after the publish signal was accepted. A retry after an ambiguous failure used to carry fresh sequences, so the shipped dedupe let the same batch land twice. The batch now stays pending under its signal sequence until the server accepts it, and goes out first if the caller moves on. append returns None because this transport learns positions at read time. --- .../streams/providers/workflow_streams.py | 100 +++++++++++------- 1 file changed, 62 insertions(+), 38 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 8155f7e9d..fb685fc2c 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -42,8 +42,7 @@ WorkflowStream, WorkflowStreamClient, ) -from temporalio.contrib.workflow_streams._stream import _PUBLISH_SIGNAL -from temporalio.contrib.workflow_streams._types import _encode_payload +from temporalio.service import RPCError, RPCStatusCode from temporalio.streams import _frame, _provider from temporalio.streams._handles import ReadSource, WriteSink from temporalio.streams._policy import AttemptTracker @@ -147,6 +146,11 @@ async def publish(self, frame: bytes) -> None: self._handle.publish(payload) +def _entry_data(payload: Payload) -> str: + # The documented wire form of PublishEntry.data. + return base64.b64encode(payload.SerializeToString()).decode("ascii") + + class WorkflowStreamsProducer: """Appends by sending the shipped publish Signal directly. @@ -154,6 +158,12 @@ class WorkflowStreamsProducer: owns the publisher identity: it must be ``producer#attempt`` for the shipped dedupe to drop a retry and pass a new generation, and the client would use its own random id. + + Sequences are committed only after the server accepted the Signal. A + batch whose Signal raised stays pending and goes out again under the + same signal sequence, either when the caller retries the same values or + ahead of whatever the caller sends next, so an ambiguous failure writes + the batch once and loses nothing. """ def __init__( @@ -174,10 +184,7 @@ def __init__( self._attempt = attempt self._sequence = 0 self._signal_sequence = 0 - - @property - def _frame_topic(self) -> str: - return self._stream or self._topic + self._pending: tuple[list[PublishEntry], int] | None = None @property def _shipped_topic(self) -> str: @@ -196,50 +203,67 @@ def _provider_id(self) -> str: else self._producer_id ) - async def append(self, *values: Any) -> Cursor: - """Append ``values`` through the shipped publish Signal.""" + async def append(self, *values: Any) -> None: + """Append ``values`` through the shipped publish Signal. + + Always ``None``: this transport learns positions at read time. + """ + if not values: + return None + await self._send([(RecordKind.DATA, self._encode(value)) for value in values]) + return None + + async def finish(self) -> None: + """Mark this producer done, so a reader stops waiting on it.""" + await self._send([(RecordKind.FINISH, b"")]) + + def _frames( + self, bodies: list[tuple[RecordKind, bytes]] + ) -> tuple[list[PublishEntry], int]: + sequence = self._sequence entries = [] - for value in values: + for kind, body in bodies: frame = _frame.encode( - topic=self._frame_topic, - kind=RecordKind.DATA, + topic=self._topic, + kind=kind, producer=self._producer_id, attempt=self._attempt, - sequence=self._sequence, - body=self._encode(value), + sequence=sequence, + body=body, ) - self._sequence += 1 - entries.append(self._entry(frame)) - await self._send(entries) - return Cursor("") - - async def finish(self) -> None: - """Mark this producer done, so a reader stops waiting on it.""" - frame = _frame.encode( - topic=self._frame_topic, - kind=RecordKind.FINISH, - producer=self._producer_id, - attempt=self._attempt, - sequence=self._sequence, - body=b"", - ) - self._sequence += 1 - await self._send([self._entry(frame)]) - - def _entry(self, frame: bytes) -> PublishEntry: - payload = Payload(data=frame) - return PublishEntry(topic=self._shipped_topic, data=_encode_payload(payload)) - - async def _send(self, entries: list[PublishEntry]) -> None: - self._signal_sequence += 1 + sequence += 1 + entries.append( + PublishEntry( + topic=self._shipped_topic, data=_entry_data(Payload(data=frame)) + ) + ) + return entries, sequence + + async def _send(self, bodies: list[tuple[RecordKind, bytes]]) -> None: + entries, next_sequence = self._frames(bodies) + if self._pending is not None and self._pending[0] != entries: + # The caller moved on from a batch whose Signal raised. It goes + # first, under the signal sequence it already had, so a copy the + # server did accept is dropped and one it never saw lands. The + # new batch is then renumbered behind it. + await self._signal(*self._pending) + entries, next_sequence = self._frames(bodies) + await self._signal(entries, next_sequence) + + async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None: + signal_sequence = self._signal_sequence + 1 + self._pending = (entries, next_sequence) await self._handle.signal( PUBLISH_SIGNAL_NAME, PublishInput( items=entries, publisher_id=self._provider_id, - sequence=self._signal_sequence, + sequence=signal_sequence, ), ) + self._signal_sequence = signal_sequence + self._sequence = next_sequence + self._pending = None def _encode(self, value: Any) -> bytes: payload = ( From 0e8236ef00d658edf83e0fd684c0b50c666038fd Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:51:11 -0700 Subject: [PATCH 11/32] Adapted the Workflow Streams provider to the interface changes. Inbound frames carry no topic and a topic filter on an inbound stream is rejected through the shared check; producer identity comes resolved from the package; an unreadable frame is skipped with a warning; read is declared as the generator it is. --- .../streams/providers/workflow_streams.py | 24 ++++++++++--------- 1 file changed, 13 insertions(+), 11 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index fb685fc2c..ef2a88312 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -27,7 +27,8 @@ from __future__ import annotations import base64 -from collections.abc import AsyncIterator +import logging +from collections.abc import AsyncGenerator from datetime import timedelta from typing import Any @@ -52,6 +53,8 @@ _IN = "in:" _OUT = "out:" +logger = logging.getLogger(__name__) + class _Runtime: """A view over the shipped stream object of the running workflow instance. @@ -280,12 +283,13 @@ class WorkflowStreamsConsumer: def __init__( self, stream_client: WorkflowStreamClient, - shipped_topic: str | None, + stream: str, poll_cooldown: timedelta, ) -> None: """Read what ``stream_client`` reaches, one long poll at a time.""" self._client = stream_client - self._shipped_topic = shipped_topic + self._stream = stream + self._shipped_topic = f"{_IN}{stream}" if stream else None self._poll_cooldown = poll_cooldown async def read( @@ -294,8 +298,9 @@ async def read( after: Cursor = BEGINNING, topic: str | None = None, type: type | None = None, - ) -> AsyncIterator[StreamRecord[Any]]: + ) -> AsyncGenerator[StreamRecord[Any], None]: """Yield records after ``after``, waiting for ones not written yet.""" + _provider.check_topic(self._stream, topic) attempts = AttemptTracker() next_offset = int(after.token) + 1 if after.token else 0 subscription = self._client.subscribe( @@ -351,7 +356,9 @@ def _record( cursor = Cursor(str(offset)) try: kind, frame_topic, source, attempt, sequence, body = _frame.decode(frame) - except ValueError: + except ValueError as error: + # Same answer as the workflow-side reader: skip and say so. + logger.warning("skipping stream record at %s: %s", cursor, error) return [] if topic is not None and frame_topic != topic: return [] @@ -456,11 +463,6 @@ async def producer( producer_id: str = "", attempt: int = 0, ) -> WorkflowStreamsProducer: - if not producer_id: - from temporalio import activity - - producer_id = activity.info().activity_id - attempt = attempt or activity.info().attempt return WorkflowStreamsProducer( client.get_workflow_handle(workflow_id), client.data_converter.payload_converter, @@ -475,7 +477,7 @@ async def consumer( ) -> WorkflowStreamsConsumer: return WorkflowStreamsConsumer( WorkflowStreamClient.create(client, workflow_id), - f"{_IN}{stream}" if stream else None, + stream, self._poll_cooldown, ) From dff43413786aa20c3ac55650bba92aee74736258 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:51:11 -0700 Subject: [PATCH 12/32] Narrowed the tail query's error handling and filtered owner reads by topic. Only a missing handler or a vanished history means there is no tail to serve; anything else, such as a query timeout with no worker polling, is now an error rather than a silent empty stream. An owner-stream read with a topic subscribes to that shipped topic alone instead of dropping inbound items client-side. --- .../streams/providers/workflow_streams.py | 73 +++++++++++++------ 1 file changed, 51 insertions(+), 22 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index ef2a88312..049fd60e9 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -20,8 +20,16 @@ - Producer identity dedupes through the shipped publisher state: the publisher id is ``producer#attempt`` and every publish Signal carries a monotonic sequence, so a retried batch drops and a new attempt passes. -- ``append`` returns an empty cursor. The Signal transport learns positions - at read time; that is this provider's stated deviation. +- ``append`` returns ``None``. The Signal transport learns positions at read + time, so a caller that wants to follow from now asks ``Consumer.latest``. +- The workflow-side stream object belongs to the workflow instance, found + through the handler the shipped class registers on it. An evicted and + rebuilt workflow gets its own, so a task that failed leaks nothing into + the next attempt's log and a replayed run does not see records twice. +- An owner-stream read with a topic subscribes to that shipped topic alone. + Without one it has to take every shipped topic and drop the inbound + items client-side, which spends the poll response cap on records the + reader never sees; name a topic on a busy workflow. """ from __future__ import annotations @@ -34,7 +42,7 @@ from temporalio import workflow from temporalio.api.common.v1 import Payload -from temporalio.client import Client +from temporalio.client import Client, WorkflowQueryFailedError from temporalio.common import RawValue from temporalio.contrib.workflow_streams import ( PUBLISH_SIGNAL_NAME, @@ -303,42 +311,63 @@ async def read( _provider.check_topic(self._stream, topic) attempts = AttemptTracker() next_offset = int(after.token) + 1 if after.token else 0 + if self._shipped_topic is not None: + shipped = [self._shipped_topic] + elif topic is not None: + shipped = [f"{_OUT}{topic}"] + else: + shipped = None subscription = self._client.subscribe( - self._shipped_topic, + shipped, from_offset=next_offset, result_type=RawValue, poll_cooldown=self._poll_cooldown, ) - async for item in subscription: - next_offset = item.offset + 1 - record = self._record( - attempts, item.offset, item.topic, item.data.payload.data, topic, type - ) - for out in record: - yield out + try: + async for item in subscription: + next_offset = item.offset + 1 + for out in self._record( + attempts, + item.offset, + item.topic, + item.data.payload.data, + topic, + type, + ): + yield out + finally: + # Lets go of the parked poll when the caller stops early. + if isinstance(subscription, AsyncGenerator): + await subscription.aclose() # The subscription ends when the workflow is closing or closed. What # landed after the last poll is still in workflow state, so the tail # comes back by Query rather than being lost with the run. - for wire in await self._tail(next_offset): + for offset, shipped_topic, frame in await self._tail(next_offset): for out in self._record( - attempts, - wire["offset"], - wire["topic"], - base64.b64decode(wire["data"]), - topic, - type, + attempts, offset, shipped_topic, frame, topic, type ): yield out - async def _tail(self, from_offset: int) -> list[dict[str, Any]]: + async def _tail(self, from_offset: int) -> list[tuple[int, str, bytes]]: try: wire = await self._client.handle.query( _TAIL_QUERY, from_offset, result_type=list ) - except Exception: - # A workflow that never opened a stream has no handler to ask, and - # one whose History is gone has nothing left to serve. + except WorkflowQueryFailedError as error: + if "expected but not found" not in str(error): + raise + # The workflow never opened a stream through this provider, so + # there is no tail to serve. return [] + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + # The History is gone; nothing is left to serve. + return [] + return [ + (item["offset"], item["topic"], base64.b64decode(item["data"])) + for item in wire + ] def _record( self, From 50b69e42797946cd31c99aa7bf28257ff0562f19 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 18 Sep 2026 16:51:11 -0700 Subject: [PATCH 13/32] Ran the Workflow Streams provider tests on the test environment's server. The tests take the shared client fixture instead of an opt-in live gate, the loop workflow calls the portable lifecycle hooks, and new cases cover a cold workflow cache, the tail of a closed run, and the resend of a batch whose signal failed. --- .../streams/test_workflow_streams_provider.py | 227 +++++++++++++----- 1 file changed, 166 insertions(+), 61 deletions(-) diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 612aaa33d..612e99e02 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -1,44 +1,46 @@ -"""Live conformance for the workflow_streams provider. +"""Conformance for the workflow_streams provider on the test environment's server. -Runs the interface loop over the shipped Option 0 transport against a real -server: an outside producer appends through the publish Signal, the workflow -reads and republishes through its own state, and an outside consumer follows -the poll Update. Gated behind ``STREAMS_LIVE=workflow_streams`` because it -needs a running server (``TEMPORAL_ADDRESS``, default ``localhost:7233``). +Runs the interface loop over the shipped Option 0 transport: an outside +producer appends through the publish Signal, the workflow reads and +republishes through its own state, and an outside consumer follows the poll +Update while the run is open and the tail Query once it has closed. """ from __future__ import annotations import asyncio -import os +import base64 import uuid from typing import Any import pytest from temporalio import streams, workflow +from temporalio.api.common.v1 import Payload from temporalio.client import Client -from temporalio.streams import RecordKind -from temporalio.streams.providers.workflow_streams import drain -from temporalio.worker import Worker +from temporalio.contrib.workflow_streams import PublishInput +from temporalio.converter import DataConverter +from temporalio.streams import RecordKind, _frame +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProducer +from tests.helpers import new_worker -pytestmark = pytest.mark.skipif( - os.environ.get("STREAMS_LIVE") != "workflow_streams", - reason="needs a live server; run with STREAMS_LIVE=workflow_streams", -) + +@pytest.fixture(autouse=True) +def _workflow_streams_provider(): # pyright: ignore[reportUnusedFunction] + streams.configure(provider="workflow_streams") @workflow.defn class EchoLoop: """Reads ``inputs``, echoes each value onto ``decisions``, ends on FINISH. - Lingers until released, because an Option 0 stream dies with its - workflow: a reader that arrives after close finds nothing, which is the - transport limit the doc states rather than a defect to fix here. + Lingers until released so a reader can follow it while it runs; whoever + arrives after it closed is served by the tail Query instead. """ def __init__(self) -> None: self._released = False + streams.prepare() @workflow.signal def release(self) -> None: @@ -54,11 +56,12 @@ async def run(self) -> int: break if record.kind is not RecordKind.DATA: continue + assert isinstance(record.value, dict) seen += 1 await decisions.publish({"echo": record.value["n"]}) await decisions.finish() await workflow.wait_condition(lambda: self._released) - drain() + streams.drain() await workflow.wait_condition(workflow.all_handlers_finished) return seen @@ -76,59 +79,40 @@ async def _collect() -> None: return out -async def test_interface_loop_over_workflow_streams(): - streams.configure(provider="workflow_streams") - client = await Client.connect(os.environ.get("TEMPORAL_ADDRESS", "localhost:7233")) - workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" - - async with Worker( - client, - task_queue=f"tq-{workflow_id}", - workflows=[EchoLoop], - **streams.worker_options(), - ): - handle = await client.start_workflow( - EchoLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" - ) +async def _feed(client: Client, workflow_id: str) -> None: + producer = await streams.producer( + client, workflow_id=workflow_id, stream="inputs", producer_id="model", attempt=1 + ) + await producer.append({"n": 1}, {"n": 2}) + await producer.append({"n": 3}) + await producer.finish() - producer = await streams.producer( - client, - workflow_id=workflow_id, - stream="inputs", - producer_id="model", - attempt=1, + +ECHOED = [RecordKind.DATA, RecordKind.DATA, RecordKind.DATA, RecordKind.FINISH] + + +async def test_interface_loop_over_workflow_streams(client: Client): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) - await producer.append({"n": 1}, {"n": 2}) - await producer.append({"n": 3}) - await producer.finish() + await _feed(client, workflow_id) consumer = await streams.consumer(client, workflow_id=workflow_id) records = await take(consumer.read(type=dict), 4) - assert [r.kind for r in records] == [ - RecordKind.DATA, - RecordKind.DATA, - RecordKind.DATA, - RecordKind.FINISH, - ] + assert [r.kind for r in records] == ECHOED assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] await handle.signal(EchoLoop.release) assert await handle.result() == 3 -async def test_retried_producer_dedupes_and_new_attempt_supersedes(): - streams.configure(provider="workflow_streams") - client = await Client.connect(os.environ.get("TEMPORAL_ADDRESS", "localhost:7233")) - workflow_id = f"streams-ws-live-{uuid.uuid4().hex}" - - async with Worker( - client, - task_queue=f"tq-{workflow_id}", - workflows=[EchoLoop], - **streams.worker_options(), - ): - await client.start_workflow( - EchoLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" +async def test_retried_producer_dedupes_and_new_attempt_supersedes(client: Client): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) first = await streams.producer( @@ -147,7 +131,7 @@ async def test_retried_producer_dedupes_and_new_attempt_supersedes(): producer_id="model", attempt=1, ) - await retry.append({"n": 1}) + assert await retry.append({"n": 1}) is None second = await streams.producer( client, workflow_id=workflow_id, @@ -163,5 +147,126 @@ async def test_retried_producer_dedupes_and_new_attempt_supersedes(): records = await take(consumer.read(type=dict), 3) assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 assert records[1].kind is RecordKind.SUPERSEDED + assert isinstance(records[1].value, streams.Supersession) assert records[1].value.previous_attempt == 1 assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 + # Inbound records carry no topic; the stream's name is the address. + assert all(r.topic == "" for r in records) + + await second.finish() + await handle.signal(EchoLoop.release) + await handle.result() + + +async def test_cold_cache_serves_each_record_once(client: Client): + # Every task rebuilds the workflow from history, so the stream object + # has to belong to the instance that is running: a stale one would carry + # the previous instance's log and hand out every record twice. + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, max_cached_workflows=0) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await _feed(client, workflow_id) + + consumer = await streams.consumer(client, workflow_id=workflow_id) + records = await take(consumer.read(type=dict), 4) + assert [r.kind for r in records] == ECHOED + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + inbound = await streams.consumer( + client, workflow_id=workflow_id, stream="inputs" + ) + arrived = await take(inbound.read(type=dict), 4) + assert [r.kind for r in arrived] == ECHOED + assert [r.value["n"] for r in arrived[:3]] == [1, 2, 3] + + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + +async def test_a_closed_run_serves_its_tail_by_query(client: Client): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await _feed(client, workflow_id) + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + # Nothing polled while the run was open. The poll Update is gone with + # the run, so everything below arrives through the tail Query. + consumer = await streams.consumer(client, workflow_id=workflow_id) + records = await take(consumer.read(type=dict), 4) + assert [r.kind for r in records] == ECHOED + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + + checkpoint = records[1].cursor + resumed = await streams.consumer(client, workflow_id=workflow_id) + again = await take(resumed.read(type=dict, after=checkpoint), 2) + assert [r.value["echo"] for r in again[:1]] == [3] + assert again[1].kind is RecordKind.FINISH + + +class _FlakyHandle: + """A workflow handle whose first Signal is accepted and then reported failed.""" + + def __init__(self) -> None: + self.sent: list[PublishInput] = [] + self._fail_next = True + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.sent.append(arg) + if self._fail_next: + self._fail_next = False + raise ConnectionResetError( + "the server accepted the signal, the reply was lost" + ) + + +def _frames( + publish: PublishInput, +) -> list[tuple[RecordKind, str, str, int, int, bytes]]: + out = [] + for entry in publish.items: + payload = Payload() + payload.ParseFromString(base64.b64decode(entry.data)) + out.append(_frame.decode(payload.data)) + return out + + +def _sequences(sent: list[PublishInput]) -> list[tuple[int, list[int]]]: + return [ + (publish.sequence, [frame[4] for frame in _frames(publish)]) for publish in sent + ] + + +def _producer(handle: _FlakyHandle) -> WorkflowStreamsProducer: + return WorkflowStreamsProducer( + handle, DataConverter.default.payload_converter, "inputs", "", "model", 1 + ) + + +async def test_a_retried_append_after_an_ambiguous_failure_writes_once(): + handle = _FlakyHandle() + producer = _producer(handle) + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + await producer.append({"n": 1}) + await producer.append({"n": 2}) + # The retry carries the same signal sequence and the same frame sequence + # as the failed send, so the shipped dedupe drops the copy; the batch + # after it continues the numbering. + assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1])] + + +async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): + handle = _FlakyHandle() + producer = _producer(handle) + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + await producer.append({"n": 2}, {"n": 3}) + await producer.finish() + assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1, 2]), (3, [3])] + assert _frames(handle.sent[-1])[0][0] is RecordKind.FINISH From 06cbd7885272982637b9fab918c88db9738e8d12 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 02:06:16 -0700 Subject: [PATCH 14/32] Rewrote the Workflow Streams provider on the plugin surface. Topics are the shipped topics of the same name, the record is the StreamRecord proto inside the item payload, and cursors name the run because a log is not carried across continue-as-new. --- .../streams/providers/workflow_streams.py | 684 +++++++++++------- 1 file changed, 409 insertions(+), 275 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 049fd60e9..6f18dce46 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -1,35 +1,33 @@ -"""The provider over today's Workflow Streams (Option 0). +"""The provider over the shipped Workflow Streams transport (Option 0). Speaks the shipped contrib feature's wire format, the ``__temporal_workflow_stream_*`` Signal, Update and Query, so interface code -and existing Workflow Streams code interoperate on one stream, and old -histories replay. Records live in the owning workflow's History, which is -also this provider's limit: Option 0's caps (payloads in History, the Signal -cap, bounded subscribers) are transport properties and remain. Reads after -the workflow closes go through one Query this provider adds, which serves -the log from workflow state for as long as the History is retained, so a -reader between polls when the workflow completed still gets the tail. +and existing Workflow Streams code share one log, and old histories replay. +Records live in the owning workflow's History, which is also this provider's +limit: the shipped caps (payloads in History, the Signal cap, bounded +subscribers) are transport properties and remain. The mapping, in one place: -- An interface record's frame rides as the item's ``Payload`` data. -- Inbound stream ``s`` is shipped topic ``in:s``; a writer topic ``t`` is - shipped topic ``out:t``, so the two namespaces cannot collide. A producer - appending onto the workflow's own topic ``t`` writes ``out:t`` as well, - which is what today's activities already do through the publish Signal. +- A topic is the shipped topic of the same name in the workflow's one log. + A record rides as the item's ``Payload``: its data is the serialized + ``StreamRecord`` proto and its encoding is ``binary/plain``. The shipped + code stores and returns that ``Payload`` untouched, so the body's own + encoding never meets the transport. - Producer identity dedupes through the shipped publisher state: the publisher id is ``producer#attempt`` and every publish Signal carries a monotonic sequence, so a retried batch drops and a new attempt passes. -- ``append`` returns ``None``. The Signal transport learns positions at read - time, so a caller that wants to follow from now asks ``Consumer.latest``. +- ``append()`` returns ``None``. The Signal transport learns positions at + read time, so a caller that wants to follow from now asks ``latest()``. +- A log belongs to one run and is not carried across continue-as-new, so a + cursor names the run as well as the offset. A handle without a run id + reads run after run: each log through the poll Update while its run is + open and through the tail Query once it has closed, then the successor's + from its first record. - The workflow-side stream object belongs to the workflow instance, found through the handler the shipped class registers on it. An evicted and rebuilt workflow gets its own, so a task that failed leaks nothing into the next attempt's log and a replayed run does not see records twice. -- An owner-stream read with a topic subscribes to that shipped topic alone. - Without one it has to take every shipped topic and drop the inbound - items client-side, which spends the poll response cap on records the - reader never sees; name a topic on a busy workflow. """ from __future__ import annotations @@ -40,9 +38,17 @@ from datetime import timedelta from typing import Any +from google.protobuf.message import DecodeError + from temporalio import workflow from temporalio.api.common.v1 import Payload -from temporalio.client import Client, WorkflowQueryFailedError +from temporalio.client import ( + Client, + WorkflowExecutionStatus, + WorkflowHandle, + WorkflowHistoryEventFilterType, + WorkflowQueryFailedError, +) from temporalio.common import RawValue from temporalio.contrib.workflow_streams import ( PUBLISH_SIGNAL_NAME, @@ -51,20 +57,80 @@ WorkflowStream, WorkflowStreamClient, ) +from temporalio.converter import PayloadConverter from temporalio.service import RPCError, RPCStatusCode -from temporalio.streams import _frame, _provider -from temporalio.streams._handles import ReadSource, WriteSink -from temporalio.streams._policy import AttemptTracker +from temporalio.streams._errors import StreamCursorError, StreamNotFoundError +from temporalio.streams._provider import ReadSource, WriteSink from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +from temporalio.streams._wire import ( + RecordDecoder, + WireRecord, + cursor_position, + mint_cursor, + producer_identity, + to_wire, +) +from temporalio.streams.providers import ProviderPlugin + +__all__ = [ + "WorkflowStreamsHandle", + "WorkflowStreamsProducer", + "WorkflowStreamsProvider", +] +_PROVIDER = "workflow_streams" _TAIL_QUERY = "__temporal_streams_tail" -_IN = "in:" -_OUT = "out:" +_ENCODING = b"binary/plain" logger = logging.getLogger(__name__) -class _Runtime: +def _require_topic(topic: str) -> None: + if not topic: + raise ValueError("topic must not be empty") + + +def _cursor(run_id: str, offset: int) -> Cursor: + return mint_cursor(_PROVIDER, f"{run_id}:{offset}") + + +def _position(after: Cursor) -> tuple[str, int] | None: + """The ``(run id, offset)`` a cursor of this provider names, or ``None`` for BEGINNING.""" + token = cursor_position(after, provider=_PROVIDER) + if token is None: + return None + run_id, _, offset = token.rpartition(":") + try: + if not run_id: + raise ValueError + return run_id, int(offset) + except ValueError: + raise StreamCursorError( + f"cursor {after.token!r} does not name a run and an offset on the " + "workflow_streams provider" + ) from None + + +def _wrap(record: WireRecord) -> Payload: + return Payload(metadata={"encoding": _ENCODING}, data=record.SerializeToString()) + + +def _unwrap(cursor: Cursor, payload: Payload, warn: Any) -> WireRecord | None: + try: + return WireRecord.FromString(payload.data) + except DecodeError as error: + # An item another publisher put on this topic, or a corrupt one: + # skip and say so, so one bad record cannot pin a reader. + warn("skipping stream record at %s: %s", cursor, error) + return None + + +def _entry_data(payload: Payload) -> str: + # The documented wire form of PublishEntry.data. + return base64.b64encode(payload.SerializeToString()).decode("ascii") + + +class _InstanceStream: """A view over the shipped stream object of the running workflow instance. A separate class because ``WorkflowStream`` insists on being constructed @@ -85,7 +151,7 @@ def _tail(self, from_offset: int) -> list[dict[str, Any]]: { "offset": offset, "topic": topic, - "data": base64.b64encode(payload.data).decode("ascii"), + "data": base64.b64encode(payload.SerializeToString()).decode("ascii"), } for offset, topic, payload in self.stream.items_from(from_offset) ] @@ -105,39 +171,45 @@ def _registered_stream() -> WorkflowStream | None: return stream -def _runtime() -> _Runtime: +def _instance() -> _InstanceStream: # Found on the instance rather than in a process-level map keyed by run # id: the SDK rebuilds an evicted workflow from history as a new object, # and a map would hand that object the stale log with its unregistered # handlers and the records of a task that failed. stream = _registered_stream() - return _Runtime() if stream is None else _Runtime(stream) + return _InstanceStream() if stream is None else _InstanceStream(stream) class _WSReadSource: """Reads the signal-fed log the shipped feature keeps in workflow state.""" - def __init__(self, stream: WorkflowStream, shipped_topic: str, start: int) -> None: + def __init__( + self, stream: WorkflowStream, topic: str, start: int, run_id: str + ) -> None: self._stream = stream - self._shipped_topic = shipped_topic - self._cursor = start + self._topic = topic + self._offset = start + self._run_id = run_id self._closed = False - async def next_batch(self) -> list[tuple[Cursor, bytes]]: + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: while True: if self._closed: raise StopAsyncIteration await workflow.wait_condition( - lambda: self._closed or self._stream.next_offset > self._cursor + lambda: self._closed or self._stream.next_offset > self._offset ) if self._closed: raise StopAsyncIteration - batch = [ - (Cursor(str(offset)), payload.data) - for offset, topic, payload in self._stream.items_from(self._cursor) - if topic == self._shipped_topic - ] - self._cursor = self._stream.next_offset + batch: list[tuple[Cursor, WireRecord]] = [] + for offset, topic, payload in self._stream.items_from(self._offset): + if topic != self._topic: + continue + cursor = _cursor(self._run_id, offset) + wire = _unwrap(cursor, payload, workflow.logger.warning) + if wire is not None: + batch.append((cursor, wire)) + self._offset = self._stream.next_offset if batch: return batch @@ -146,20 +218,51 @@ def close(self) -> None: class _WSWriteSink: - def __init__(self, stream: WorkflowStream, shipped_topic: str) -> None: - self._handle = stream.topic(shipped_topic) + def __init__(self, stream: WorkflowStream, topic: str) -> None: + self._handle = stream.topic(topic) - async def publish(self, frame: bytes) -> None: + def publish(self, record: WireRecord) -> None: # Appending to workflow state commits with the task, and a poll # Update's result rides the same task completion, so a failed task # leaks nothing: rule 1 through the shipped mechanics. - payload = workflow.payload_converter().to_payloads([frame])[0] - self._handle.publish(payload) + self._handle.publish(_wrap(record)) -def _entry_data(payload: Payload) -> str: - # The documented wire form of PublishEntry.data. - return base64.b64encode(payload.SerializeToString()).decode("ascii") +class _WSWorkflowProvider: + """The workflow half: the shipped stream object of the running instance.""" + + def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + _require_topic(topic) + run_id = workflow.info().run_id + start = 0 + named = _position(after) + if named is not None: + if named[0] != run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run; a run's log is its own" + ) + start = named[1] + 1 + return _WSReadSource(_instance().stream, topic, start, run_id) + + def open_writer(self, topic: str) -> WriteSink: + _require_topic(topic) + return _WSWriteSink(_instance().stream, topic) + + def on_workflow_start(self) -> None: + # Registered before the first task completes, because an outside + # reader can poll before workflow code has opened anything, and an + # Update with no handler yet is rejected rather than held. + _instance() + + async def on_workflow_finish(self) -> None: + # An Option 0 stream dies with its run, and a parked long-poll Update + # would otherwise hold completion open. Same recipe the shipped + # feature documents before a return or a continue-as-new. + stream = _registered_stream() + if stream is None: + return + stream.detach_pollers() + await workflow.wait_condition(workflow.all_handlers_finished) class WorkflowStreamsProducer: @@ -180,16 +283,14 @@ class WorkflowStreamsProducer: def __init__( self, handle: Any, - converter: Any, - stream: str, + converter: PayloadConverter, topic: str, producer_id: str, attempt: int, ) -> None: - """Bind this producer to one stream or topic on ``handle``.""" + """Bind this producer to ``topic`` on the workflow behind ``handle``.""" self._handle = handle self._converter = converter - self._stream = stream self._topic = topic self._producer_id = producer_id self._attempt = attempt @@ -198,8 +299,9 @@ def __init__( self._pending: tuple[list[PublishEntry], int] | None = None @property - def _shipped_topic(self) -> str: - return f"{_IN}{self._stream}" if self._stream else f"{_OUT}{self._topic}" + def producer_id(self) -> str: + """Who this producer writes as.""" + return self._producer_id @property def attempt(self) -> int: @@ -207,152 +309,258 @@ def attempt(self) -> int: return self._attempt @property - def _provider_id(self) -> str: + def _publisher_id(self) -> str: return ( f"{self._producer_id}#{self._attempt}" if self._attempt else self._producer_id ) - async def append(self, *values: Any) -> None: + async def append(self, *values: Any) -> Cursor | None: """Append ``values`` through the shipped publish Signal. - Always ``None``: this transport learns positions at read time. + Always ``None``: this transport learns positions at read time, so a + caller that wants to follow from now asks + :meth:`WorkflowStreamsHandle.latest`. """ if not values: return None - await self._send([(RecordKind.DATA, self._encode(value)) for value in values]) + await self._send([(RecordKind.DATA, value) for value in values]) return None async def finish(self) -> None: - """Mark this producer done, so a reader stops waiting on it.""" - await self._send([(RecordKind.FINISH, b"")]) + """Write ``FINISH`` for this producer on this topic.""" + await self._send([(RecordKind.FINISH, None)]) - def _frames( - self, bodies: list[tuple[RecordKind, bytes]] + def _entries( + self, batch: list[tuple[RecordKind, Any]] ) -> tuple[list[PublishEntry], int]: sequence = self._sequence entries = [] - for kind, body in bodies: - frame = _frame.encode( + for kind, value in batch: + wire = to_wire( + self._converter, topic=self._topic, kind=kind, - producer=self._producer_id, + value=value, + producer_id=self._producer_id, attempt=self._attempt, sequence=sequence, - body=body, ) sequence += 1 entries.append( - PublishEntry( - topic=self._shipped_topic, data=_entry_data(Payload(data=frame)) - ) + PublishEntry(topic=self._topic, data=_entry_data(_wrap(wire))) ) return entries, sequence - async def _send(self, bodies: list[tuple[RecordKind, bytes]]) -> None: - entries, next_sequence = self._frames(bodies) + async def _send(self, batch: list[tuple[RecordKind, Any]]) -> None: + entries, next_sequence = self._entries(batch) if self._pending is not None and self._pending[0] != entries: # The caller moved on from a batch whose Signal raised. It goes # first, under the signal sequence it already had, so a copy the # server did accept is dropped and one it never saw lands. The # new batch is then renumbered behind it. await self._signal(*self._pending) - entries, next_sequence = self._frames(bodies) + entries, next_sequence = self._entries(batch) await self._signal(entries, next_sequence) async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None: signal_sequence = self._signal_sequence + 1 self._pending = (entries, next_sequence) - await self._handle.signal( - PUBLISH_SIGNAL_NAME, - PublishInput( - items=entries, - publisher_id=self._provider_id, - sequence=signal_sequence, - ), - ) + try: + await self._handle.signal( + PUBLISH_SIGNAL_NAME, + PublishInput( + items=entries, + publisher_id=self._publisher_id, + sequence=signal_sequence, + ), + ) + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._handle.id!r} was not found, so its stream " + "cannot be appended to" + ) from error + raise self._signal_sequence = signal_sequence self._sequence = next_sequence self._pending = None - def _encode(self, value: Any) -> bytes: - payload = ( - value - if isinstance(value, Payload) - else self._converter.to_payloads([value])[0] - ) - return payload.SerializeToString() - -class WorkflowStreamsConsumer: - """Reads through the shipped long-poll Update, with shared supersession.""" +class WorkflowStreamsHandle: + """One workflow's log from outside, through the shipped poll Update and a tail Query.""" def __init__( self, - stream_client: WorkflowStreamClient, - stream: str, + client: Client, + workflow_id: str, + run_id: str | None, poll_cooldown: timedelta, ) -> None: - """Read what ``stream_client`` reaches, one long poll at a time.""" - self._client = stream_client - self._stream = stream - self._shipped_topic = f"{_IN}{stream}" if stream else None + """Address ``workflow_id``'s log, pinned to ``run_id`` when one is given.""" + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id self._poll_cooldown = poll_cooldown + self._converter = client.data_converter.payload_converter - async def read( + def read( self, *, + topic: str, after: Cursor = BEGINNING, - topic: str | None = None, - type: type | None = None, + result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: - """Yield records after ``after``, waiting for ones not written yet.""" - _provider.check_topic(self._stream, topic) - attempts = AttemptTracker() - next_offset = int(after.token) + 1 if after.token else 0 - if self._shipped_topic is not None: - shipped = [self._shipped_topic] - elif topic is not None: - shipped = [f"{_OUT}{topic}"] - else: - shipped = None - subscription = self._client.subscribe( - shipped, - from_offset=next_offset, - result_type=RawValue, - poll_cooldown=self._poll_cooldown, + """Yield records on ``topic`` after ``after`` until the chain, or the pinned run, closes.""" + _require_topic(topic) + # Parsed here so a foreign cursor fails this call, not the first + # iteration of the generator. + named = _position(after) + if named is not None and self._run_id is not None and named[0] != self._run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run than this handle is pinned to" + ) + return self._read(topic, named, after, result_type) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning ) - try: - async for item in subscription: - next_offset = item.offset + 1 - for out in self._record( - attempts, - item.offset, - item.topic, - item.data.payload.data, - topic, - type, - ): - yield out - finally: - # Lets go of the parked poll when the caller stops early. - if isinstance(subscription, AsyncGenerator): - await subscription.aclose() - # The subscription ends when the workflow is closing or closed. What - # landed after the last poll is still in workflow state, so the tail - # comes back by Query rather than being lost with the run. - for offset, shipped_topic, frame in await self._tail(next_offset): - for out in self._record( - attempts, offset, shipped_topic, frame, topic, type + if named is not None: + run_id, offset = named[0], named[1] + 1 + else: + run_id, offset = self._run_id or await self._first_run(), 0 + while True: + handle = self._handle(run_id) + next_offset = offset + # Pinned to one run, with no client for the shipped chain + # following: a log is not carried across continue-as-new, so + # the successor's offsets start over and this loop is what + # moves from one run to the next. + subscription = WorkflowStreamClient(handle).subscribe( + topic, + from_offset=offset, + result_type=RawValue, + poll_cooldown=self._poll_cooldown, + ) + try: + async for item in subscription: + next_offset = item.offset + 1 + if item.topic != topic: + continue + for record in self._records( + decoder, run_id, item.offset, item.data.payload + ): + yield record + except RPCError as error: + # The run closed and its poll Update went with it, or the + # workflow does not exist; the describe below tells which. + if error.status != RPCStatusCode.NOT_FOUND: + raise + finally: + if isinstance(subscription, AsyncGenerator): + await subscription.aclose() + status = await self._status(handle) + if status is None: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {run_id!r} was not found" + ) + if status == WorkflowExecutionStatus.RUNNING: + # The subscription ended early, on an RPC timeout for + # instance; the run is still open, so pick up where it left. + offset = next_offset + continue + # What landed after the last poll is still in workflow state, so + # the tail comes back by Query rather than being lost with the run. + for offset_, shipped_topic, payload in await self._tail( + handle, next_offset + ): + if shipped_topic != topic: + continue + for record in self._records(decoder, run_id, offset_, payload): + yield record + if ( + self._run_id is not None + or status != WorkflowExecutionStatus.CONTINUED_AS_NEW ): - yield out + return + successor = await self._successor(handle) + if successor is None: + return + run_id, offset = successor, 0 + + def _records( + self, decoder: RecordDecoder, run_id: str, offset: int, payload: Payload + ) -> list[StreamRecord[Any]]: + cursor = _cursor(run_id, offset) + wire = _unwrap(cursor, payload, logger.warning) + if wire is None: + return [] + return decoder.decode(cursor, wire) + + def _handle(self, run_id: str | None) -> WorkflowHandle[Any, Any]: + return self._client.get_workflow_handle(self._workflow_id, run_id=run_id) - async def _tail(self, from_offset: int) -> list[tuple[int, str, bytes]]: + async def _status( + self, handle: WorkflowHandle[Any, Any] + ) -> WorkflowExecutionStatus | None: try: - wire = await self._client.handle.query( - _TAIL_QUERY, from_offset, result_type=list - ) + return (await handle.describe()).status + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + return None + raise + + async def _first_run(self) -> str: + """The oldest retained run of the chain, walking back from the latest.""" + try: + run_id = (await self._handle(None).describe()).run_id + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + assert run_id is not None + while True: + previous = await self._predecessor(run_id) + if previous is None: + return run_id + run_id = previous + + async def _predecessor(self, run_id: str) -> str | None: + try: + async for event in self._handle(run_id).fetch_history_events(page_size=1): + attributes = event.workflow_execution_started_event_attributes + return attributes.continued_execution_run_id or None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + # The run's History is gone: the chain's retained part starts here. + return None + + async def _successor(self, handle: WorkflowHandle[Any, Any]) -> str | None: + events = handle.fetch_history_events( + event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT + ) + async for event in events: + if event.HasField("workflow_execution_continued_as_new_event_attributes"): + attributes = event.workflow_execution_continued_as_new_event_attributes + return attributes.new_execution_run_id or None + return None + + async def _tail( + self, handle: WorkflowHandle[Any, Any], from_offset: int + ) -> list[tuple[int, str, Payload]]: + try: + wire = await handle.query(_TAIL_QUERY, from_offset, result_type=list) except WorkflowQueryFailedError as error: if "expected but not found" not in str(error): raise @@ -365,150 +573,76 @@ async def _tail(self, from_offset: int) -> list[tuple[int, str, bytes]]: # The History is gone; nothing is left to serve. return [] return [ - (item["offset"], item["topic"], base64.b64decode(item["data"])) + ( + item["offset"], + item["topic"], + Payload.FromString(base64.b64decode(item["data"])), + ) for item in wire ] - def _record( - self, - attempts: AttemptTracker, - offset: int, - shipped_topic: str, - frame: bytes, - topic: str | None, - type: type | None, - ) -> list[StreamRecord[Any]]: - if self._shipped_topic is None and not shipped_topic.startswith(_OUT): - return [] - if self._shipped_topic is not None and shipped_topic != self._shipped_topic: - return [] - cursor = Cursor(str(offset)) + async def latest(self, *, topic: str) -> Cursor: + """The newest position in the log, which orders every topic of this workflow. + + The log is one per run, so the cursor names the run it was read from: + the pinned run, or the latest run of the chain. + """ + _require_topic(topic) + handle = self._handle(self._run_id) try: - kind, frame_topic, source, attempt, sequence, body = _frame.decode(frame) - except ValueError as error: - # Same answer as the workflow-side reader: skip and say so. - logger.warning("skipping stream record at %s: %s", cursor, error) - return [] - if topic is not None and frame_topic != topic: - return [] - out: list[StreamRecord[Any]] = [] - superseded = attempts.note(source, attempt, cursor) - if superseded is not None: - out.append(superseded) - out.append( - StreamRecord( - value=self._decode(body, type) if kind is RecordKind.DATA else None, - cursor=cursor, - kind=kind, - topic=frame_topic, - producer=source, - attempt=attempt, - sequence=sequence, - ) + description = await handle.describe() + head = await WorkflowStreamClient(handle).get_offset() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + run_id = description.run_id + assert run_id is not None + if head > 0: + return _cursor(run_id, head - 1) + # An empty log on the chain's first run is the beginning of the + # stream; on a successor it is a position of its own, because + # BEGINNING would send a chain-following read back to the first run. + if await self._predecessor(run_id) is None: + return BEGINNING + return _cursor(run_id, -1) + + def producer( + self, *, topic: str, producer_id: str = "", attempt: int = 0 + ) -> WorkflowStreamsProducer: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + _require_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return WorkflowStreamsProducer( + self._handle(self._run_id), self._converter, topic, producer_id, attempt ) - return out - - async def latest(self, *, topic: str | None = None) -> Cursor: - """The cursor of the last record written, for following from now.""" - del topic # one log per workflow, whatever the topic - head = await self._client.get_offset() - return Cursor(str(head - 1)) if head > 0 else BEGINNING - - def _decode(self, body: bytes, as_type: type | None) -> Any: - payload = Payload() - payload.ParseFromString(body) - converter = self._client.payload_converter - if as_type is None: - return converter.from_payloads([payload])[0] - return converter.from_payloads([payload], [as_type])[0] - - -class _WorkflowStreamsProvider: - name = "workflow_streams" - - def __init__(self) -> None: - self._poll_cooldown = timedelta(milliseconds=100) - - def configure(self, **options: Any) -> None: - cooldown = options.pop("poll_cooldown", None) - if cooldown is not None: - self._poll_cooldown = cooldown - if options: - raise TypeError( - "the workflow_streams provider takes only poll_cooldown, " - f"got {sorted(options)}" - ) - - def worker_options(self) -> dict[str, Any]: - return {} - def prepare(self) -> None: - """Register the shipped publish signal and poll update handlers. - Done here rather than on the first read, because an outside reader - can poll before workflow code has opened anything, and an update with - no handler yet is rejected rather than held. - """ - _runtime() +class WorkflowStreamsProvider(ProviderPlugin): + """The provider over the shipped Workflow Streams transport.""" - def drain(self) -> None: - """Release parked pollers so the workflow can return. + def __init__( + self, *, poll_cooldown: timedelta = timedelta(milliseconds=100) + ) -> None: + """Create the provider. - An Option 0 stream dies with its workflow, and a parked long-poll - Update would otherwise hold completion open. Call it right before - the workflow returns, the same obligation the shipped feature's - ``detach_pollers`` documents. + Args: + poll_cooldown: How long an outside reader that is caught up waits + between polls. Backlogs drain at full speed regardless. """ - stream = _registered_stream() - if stream is not None: - stream.detach_pollers() - - def open_read( - self, - stream: str, - *, - after: Cursor = BEGINNING, - idle_timeout: timedelta | None = None, - ) -> ReadSource: - # Ignored: a publish Signal is a workflow event, so the wait below - # wakes on delivery and nothing is held between records. - del idle_timeout - return _WSReadSource( - _runtime().stream, - f"{_IN}{stream}", - int(after.token) + 1 if after.token else 0, - ) - - def open_write(self, topic: str) -> WriteSink: - return _WSWriteSink(_runtime().stream, f"{_OUT}{topic}") - - async def producer( - self, - client: Client, - *, - workflow_id: str, - stream: str = "", - topic: str = "", - producer_id: str = "", - attempt: int = 0, - ) -> WorkflowStreamsProducer: - return WorkflowStreamsProducer( - client.get_workflow_handle(workflow_id), - client.data_converter.payload_converter, - stream, - topic, - producer_id, - attempt, - ) + self._poll_cooldown = poll_cooldown - async def consumer( - self, client: Client, *, workflow_id: str, stream: str = "" - ) -> WorkflowStreamsConsumer: - return WorkflowStreamsConsumer( - WorkflowStreamClient.create(client, workflow_id), - stream, - self._poll_cooldown, - ) + def workflow_provider(self) -> _WSWorkflowProvider: + """The workflow half, over the running instance's shipped stream object.""" + return _WSWorkflowProvider() + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> WorkflowStreamsHandle: + """A handle on ``workflow_id``'s log; without ``run_id`` it follows the chain.""" + return WorkflowStreamsHandle(client, workflow_id, run_id, self._poll_cooldown) -_provider.register("workflow_streams", _WorkflowStreamsProvider) + async def close(self) -> None: + """Nothing to release: the provider holds no connection of its own.""" From 0de91c292cc4a4b0ec1748b46c9ff1a9f15e0e14 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 02:06:16 -0700 Subject: [PATCH 15/32] Ran the conformance suite over Workflow Streams and covered chain reads. The provider registers in SETUPS with a host workflow that owns the stream, and its own tests cover the tail Query, a read that ends, and a handle that follows continue-as-new run by run. --- CHANGELOG.md | 10 +- tests/streams/test_streams_conformance.py | 49 +++- .../streams/test_workflow_streams_provider.py | 239 +++++++++++------- 3 files changed, 205 insertions(+), 93 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 811ebef5d..8f9e500d6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -33,10 +33,12 @@ to include examples, links to docs, or any other relevant information. is `temporal.api.stream.v1.StreamRecord` on every provider. `temporalio.streams.providers.memory.MemoryStreams` is the in-memory reference provider the conformance tests run against. -- **Experimental**: `temporalio.streams.providers.workflow_streams` serves the - stream interface over the shipped Workflow Streams transport, so a workflow - reads and publishes through `temporalio.contrib.workflow_streams` without - naming it. +- **Experimental**: `temporalio.streams.providers.workflow_streams.WorkflowStreamsProvider` + serves the stream interface over the shipped Workflow Streams transport as a + worker plugin, so a workflow reads and publishes through + `temporalio.contrib.workflow_streams` without naming it. Records are the + `StreamRecord` proto inside the shipped item payload, and a handle without a + run id follows continue-as-new run by run. ### Changed diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 9002ea16a..630814c7f 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -24,12 +24,14 @@ import uuid from collections.abc import AsyncIterator, Awaitable, Callable from dataclasses import dataclass +from datetime import timedelta from typing import Any import pytest +from temporalio import workflow from temporalio.api.common.v1 import Payload -from temporalio.client import Client +from temporalio.client import Client, WorkflowHandle from temporalio.common import RawValue from temporalio.converter import DataConverter from temporalio.streams import ( @@ -46,6 +48,8 @@ ) from temporalio.streams._policy import AttemptTracker from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider +from tests.helpers import new_worker # Defined once and shared by every case, the way an application shares them # between its workflow, its activities and its backend. @@ -91,8 +95,49 @@ async def _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: provider.reset() +@workflow.defn +class StreamHost: + """Owns a stream and lingers, so outside code has a running workflow to address.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: + # No STREAMS_LIVE gate: the store is the workflow's own History, which the + # test environment's server provides. + provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + hosts: dict[str, WorkflowHandle[Any, Any]] = {} + async with new_worker(client, StreamHost, plugins=[provider]) as worker: + + async def host(workflow_id: str) -> None: + if workflow_id not in hosts: + hosts[workflow_id] = await client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + + yield ProviderCase( + "workflow_streams", + provider, + client, + reports_positions=False, + host=host, + ) + for handle in hosts.values(): + await handle.terminate() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { - "memory": _memory_case + "memory": _memory_case, + "workflow_streams": _workflow_streams_case, } _CAPABILITIES = { diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 612e99e02..28f2f631d 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -2,8 +2,10 @@ Runs the interface loop over the shipped Option 0 transport: an outside producer appends through the publish Signal, the workflow reads and -republishes through its own state, and an outside consumer follows the poll -Update while the run is open and the tail Query once it has closed. +republishes through its own state, and an outside reader follows the poll +Update while the run is open and the tail Query once it has closed. The +outside-surface cases shared by every provider run from +``test_streams_conformance``; this file covers what the transport adds. """ from __future__ import annotations @@ -11,23 +13,31 @@ import asyncio import base64 import uuid +from datetime import timedelta from typing import Any import pytest -from temporalio import streams, workflow +from temporalio import workflow from temporalio.api.common.v1 import Payload from temporalio.client import Client from temporalio.contrib.workflow_streams import PublishInput from temporalio.converter import DataConverter -from temporalio.streams import RecordKind, _frame -from temporalio.streams.providers.workflow_streams import WorkflowStreamsProducer +from temporalio.streams import RecordKind, Supersession +from temporalio.streams._wire import WireRecord +from temporalio.streams.providers.workflow_streams import ( + WorkflowStreamsProducer, + WorkflowStreamsProvider, +) from tests.helpers import new_worker +INPUTS = "inputs" +DECISIONS = "decisions" -@pytest.fixture(autouse=True) -def _workflow_streams_provider(): # pyright: ignore[reportUnusedFunction] - streams.configure(provider="workflow_streams") + +@pytest.fixture +def provider() -> WorkflowStreamsProvider: + return WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) @workflow.defn @@ -40,7 +50,6 @@ class EchoLoop: def __init__(self) -> None: self._released = False - streams.prepare() @workflow.signal def release(self) -> None: @@ -48,21 +57,19 @@ def release(self) -> None: @workflow.run async def run(self) -> int: - inputs = streams.reader("inputs", type=dict) - decisions = streams.writer("decisions") + inputs = workflow.stream_reader(INPUTS, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) seen = 0 async for record in inputs: if record.kind is RecordKind.FINISH: break if record.kind is not RecordKind.DATA: continue - assert isinstance(record.value, dict) + assert record.value is not None seen += 1 - await decisions.publish({"echo": record.value["n"]}) - await decisions.finish() + decisions.publish({"echo": record.value["n"]}) + decisions.finish() await workflow.wait_condition(lambda: self._released) - streams.drain() - await workflow.wait_condition(workflow.all_handlers_finished) return seen @@ -79,10 +86,11 @@ async def _collect() -> None: return out -async def _feed(client: Client, workflow_id: str) -> None: - producer = await streams.producer( - client, workflow_id=workflow_id, stream="inputs", producer_id="model", attempt=1 - ) +async def _feed( + provider: WorkflowStreamsProvider, client: Client, workflow_id: str +) -> None: + stream = provider.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) await producer.append({"n": 1}, {"n": 2}) await producer.append({"n": 3}) await producer.finish() @@ -91,92 +99,77 @@ async def _feed(client: Client, workflow_id: str) -> None: ECHOED = [RecordKind.DATA, RecordKind.DATA, RecordKind.DATA, RecordKind.FINISH] -async def test_interface_loop_over_workflow_streams(client: Client): +async def test_interface_loop_over_workflow_streams( + client: Client, provider: WorkflowStreamsProvider +): workflow_id = f"streams-ws-{uuid.uuid4().hex}" - async with new_worker(client, EchoLoop) as worker: + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: handle = await client.start_workflow( EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) - await _feed(client, workflow_id) + await _feed(provider, client, workflow_id) - consumer = await streams.consumer(client, workflow_id=workflow_id) - records = await take(consumer.read(type=dict), 4) + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS, result_type=dict), 4) assert [r.kind for r in records] == ECHOED assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + assert all(r.topic == DECISIONS and r.producer_id == "" for r in records) await handle.signal(EchoLoop.release) assert await handle.result() == 3 -async def test_retried_producer_dedupes_and_new_attempt_supersedes(client: Client): +async def test_retried_producer_dedupes_and_new_attempt_supersedes( + client: Client, provider: WorkflowStreamsProvider +): workflow_id = f"streams-ws-{uuid.uuid4().hex}" - async with new_worker(client, EchoLoop) as worker: + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: handle = await client.start_workflow( EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) + stream = provider.get_stream_handle(client, workflow_id) - first = await streams.producer( - client, - workflow_id=workflow_id, - stream="inputs", - producer_id="model", - attempt=1, - ) - await first.append({"n": 1}) + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + # Positions are learnt at read time on this transport. + assert await first.append({"n": 1}) is None # The retry of the same attempt re-sends its first batch. - retry = await streams.producer( - client, - workflow_id=workflow_id, - stream="inputs", - producer_id="model", - attempt=1, - ) - assert await retry.append({"n": 1}) is None - second = await streams.producer( - client, - workflow_id=workflow_id, - stream="inputs", - producer_id="model", - attempt=2, - ) + retry = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await retry.append({"n": 1}) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) await second.append({"n": 2}) - consumer = await streams.consumer( - client, workflow_id=workflow_id, stream="inputs" - ) - records = await take(consumer.read(type=dict), 3) + records = await take(stream.read(topic=INPUTS, result_type=dict), 3) assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 assert records[1].kind is RecordKind.SUPERSEDED - assert isinstance(records[1].value, streams.Supersession) - assert records[1].value.previous_attempt == 1 + assert records[1].supersession == Supersession("model", 1, 2) assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 - # Inbound records carry no topic; the stream's name is the address. - assert all(r.topic == "" for r in records) + assert all(r.topic == INPUTS for r in records) await second.finish() await handle.signal(EchoLoop.release) await handle.result() -async def test_cold_cache_serves_each_record_once(client: Client): +async def test_cold_cache_serves_each_record_once( + client: Client, provider: WorkflowStreamsProvider +): # Every task rebuilds the workflow from history, so the stream object # has to belong to the instance that is running: a stale one would carry # the previous instance's log and hand out every record twice. workflow_id = f"streams-ws-{uuid.uuid4().hex}" - async with new_worker(client, EchoLoop, max_cached_workflows=0) as worker: + async with new_worker( + client, EchoLoop, plugins=[provider], max_cached_workflows=0 + ) as worker: handle = await client.start_workflow( EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) - await _feed(client, workflow_id) + await _feed(provider, client, workflow_id) - consumer = await streams.consumer(client, workflow_id=workflow_id) - records = await take(consumer.read(type=dict), 4) + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS, result_type=dict), 4) assert [r.kind for r in records] == ECHOED assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] - inbound = await streams.consumer( - client, workflow_id=workflow_id, stream="inputs" - ) - arrived = await take(inbound.read(type=dict), 4) + arrived = await take(stream.read(topic=INPUTS, result_type=dict), 4) assert [r.kind for r in arrived] == ECHOED assert [r.value["n"] for r in arrived[:3]] == [1, 2, 3] @@ -184,33 +177,106 @@ async def test_cold_cache_serves_each_record_once(client: Client): assert await handle.result() == 3 -async def test_a_closed_run_serves_its_tail_by_query(client: Client): +async def test_a_closed_run_serves_its_tail_by_query_and_the_read_ends( + client: Client, provider: WorkflowStreamsProvider +): workflow_id = f"streams-ws-{uuid.uuid4().hex}" - async with new_worker(client, EchoLoop) as worker: + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: handle = await client.start_workflow( EchoLoop.run, id=workflow_id, task_queue=worker.task_queue ) - await _feed(client, workflow_id) + await _feed(provider, client, workflow_id) await handle.signal(EchoLoop.release) assert await handle.result() == 3 # Nothing polled while the run was open. The poll Update is gone with - # the run, so everything below arrives through the tail Query. - consumer = await streams.consumer(client, workflow_id=workflow_id) - records = await take(consumer.read(type=dict), 4) + # the run, so everything below arrives through the tail Query, and + # the read ends by itself once the tail is delivered. + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS, result_type=dict)] + + records = await asyncio.wait_for(read_everything(), 30) assert [r.kind for r in records] == ECHOED assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] checkpoint = records[1].cursor - resumed = await streams.consumer(client, workflow_id=workflow_id) - again = await take(resumed.read(type=dict, after=checkpoint), 2) + again = await take( + stream.read(topic=DECISIONS, result_type=dict, after=checkpoint), 2 + ) assert [r.value["echo"] for r in again[:1]] == [3] assert again[1].kind is RecordKind.FINISH +@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() + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{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 + ) + await handle.result() + first_run = handle.first_execution_run_id + assert first_run is not None + + chain = provider.get_stream_handle(client, workflow_id) + + async def read_everything(stream: Any) -> list[Any]: + return [ + (r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + # Each run keeps its own log, so following the chain means reading + # the first run to its close and then the successor from its start. + assert await asyncio.wait_for(read_everything(chain), 30) == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + pinned = provider.get_stream_handle(client, workflow_id, run_id=first_run) + assert await asyncio.wait_for(read_everything(pinned), 30) == [ + (RecordKind.DATA, {"run": 0}), + ] + # A cursor from the first run resumes into the successor. + records = await take(chain.read(topic=DECISIONS, result_type=dict), 1) + resumed = await asyncio.wait_for( + asyncio.ensure_future( + _values( + chain.read( + topic=DECISIONS, result_type=dict, after=records[0].cursor + ) + ) + ), + 30, + ) + assert resumed == [{"run": 1}, None] + + +async def _values(records: Any) -> list[Any]: + return [r.value async for r in records] + + class _FlakyHandle: """A workflow handle whose first Signal is accepted and then reported failed.""" + id = "flaky" + def __init__(self) -> None: self.sent: list[PublishInput] = [] self._fail_next = True @@ -225,26 +291,24 @@ async def signal(self, name: str, arg: PublishInput) -> None: ) -def _frames( - publish: PublishInput, -) -> list[tuple[RecordKind, str, str, int, int, bytes]]: +def _wires(publish: PublishInput) -> list[WireRecord]: out = [] for entry in publish.items: - payload = Payload() - payload.ParseFromString(base64.b64decode(entry.data)) - out.append(_frame.decode(payload.data)) + payload = Payload.FromString(base64.b64decode(entry.data)) + out.append(WireRecord.FromString(payload.data)) return out def _sequences(sent: list[PublishInput]) -> list[tuple[int, list[int]]]: return [ - (publish.sequence, [frame[4] for frame in _frames(publish)]) for publish in sent + (publish.sequence, [wire.sequence for wire in _wires(publish)]) + for publish in sent ] def _producer(handle: _FlakyHandle) -> WorkflowStreamsProducer: return WorkflowStreamsProducer( - handle, DataConverter.default.payload_converter, "inputs", "", "model", 1 + handle, DataConverter.default.payload_converter, INPUTS, "model", 1 ) @@ -255,10 +319,11 @@ async def test_a_retried_append_after_an_ambiguous_failure_writes_once(): await producer.append({"n": 1}) await producer.append({"n": 1}) await producer.append({"n": 2}) - # The retry carries the same signal sequence and the same frame sequence - # as the failed send, so the shipped dedupe drops the copy; the batch - # after it continues the numbering. + # The retry carries the same signal sequence and the same record + # sequence as the failed send, so the shipped dedupe drops the copy; the + # batch after it continues the numbering. assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1])] + assert all(publish.publisher_id == "model#1" for publish in handle.sent) async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): @@ -269,4 +334,4 @@ async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): await producer.append({"n": 2}, {"n": 3}) await producer.finish() assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1, 2]), (3, [3])] - assert _frames(handle.sent[-1])[0][0] is RecordKind.FINISH + assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) From c44ba22617f43c762955d8b6e414344b9cbda9fa Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 03:46:46 -0700 Subject: [PATCH 16/32] Detached the stream this instance captured, not the thread's current one. The workflow half now keeps the WorkflowStream it found at start, so the finish hook lets go of this run's pollers rather than of whatever the signal handler lookup answers on the thread at the time. --- .../streams/providers/workflow_streams.py | 26 ++++++++++++++----- 1 file changed, 19 insertions(+), 7 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 6f18dce46..cef1cbe3e 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -229,7 +229,20 @@ def publish(self, record: WireRecord) -> None: class _WSWorkflowProvider: - """The workflow half: the shipped stream object of the running instance.""" + """The workflow half: the shipped stream object of one workflow instance. + + Made per instance by the worker, so the stream object it captures on + first use is this instance's, and the finish hook lets go of that one + rather than of whatever the thread's handler lookup answers at the time. + """ + + def __init__(self) -> None: + self._stream: WorkflowStream | None = None + + def _own_stream(self) -> WorkflowStream: + if self._stream is None: + self._stream = _instance().stream + return self._stream def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: _require_topic(topic) @@ -242,26 +255,25 @@ def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: f"cursor {after.token!r} names another run; a run's log is its own" ) start = named[1] + 1 - return _WSReadSource(_instance().stream, topic, start, run_id) + return _WSReadSource(self._own_stream(), topic, start, run_id) def open_writer(self, topic: str) -> WriteSink: _require_topic(topic) - return _WSWriteSink(_instance().stream, topic) + return _WSWriteSink(self._own_stream(), topic) def on_workflow_start(self) -> None: # Registered before the first task completes, because an outside # reader can poll before workflow code has opened anything, and an # Update with no handler yet is rejected rather than held. - _instance() + self._own_stream() async def on_workflow_finish(self) -> None: # An Option 0 stream dies with its run, and a parked long-poll Update # would otherwise hold completion open. Same recipe the shipped # feature documents before a return or a continue-as-new. - stream = _registered_stream() - if stream is None: + if self._stream is None: return - stream.detach_pollers() + self._stream.detach_pollers() await workflow.wait_condition(workflow.all_handlers_finished) From 309dcecf974d948436d6815e736269fdeaadd757 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 04:26:40 -0700 Subject: [PATCH 17/32] Drove the poll Update directly so a cancelled read stays cancelled. The shipped subscribe loop swallows a cancellation of the caller's task, so a bounded read resubscribed forever instead of timing out. The poll name and its failure types are public on the contrib package for this. --- .../contrib/workflow_streams/__init__.py | 6 ++ .../contrib/workflow_streams/_stream.py | 9 +- .../streams/providers/workflow_streams.py | 100 +++++++++++++----- .../streams/test_workflow_streams_provider.py | 33 ++++++ 4 files changed, 122 insertions(+), 26 deletions(-) diff --git a/temporalio/contrib/workflow_streams/__init__.py b/temporalio/contrib/workflow_streams/__init__.py index 3f225bb2f..c44bd25ea 100644 --- a/temporalio/contrib/workflow_streams/__init__.py +++ b/temporalio/contrib/workflow_streams/__init__.py @@ -14,6 +14,7 @@ from temporalio.contrib.workflow_streams._client import WorkflowStreamClient from temporalio.contrib.workflow_streams._stream import ( + POLL_UPDATE_NAME, PUBLISH_SIGNAL_NAME, WorkflowStream, ) @@ -22,6 +23,8 @@ WorkflowTopicHandle, ) from temporalio.contrib.workflow_streams._types import ( + STREAM_DRAINING_ERROR_TYPE, + TRUNCATED_OFFSET_ERROR_TYPE, PollInput, PollResult, PublishEntry, @@ -32,7 +35,10 @@ ) __all__ = [ + "POLL_UPDATE_NAME", "PUBLISH_SIGNAL_NAME", + "STREAM_DRAINING_ERROR_TYPE", + "TRUNCATED_OFFSET_ERROR_TYPE", "PollInput", "PollResult", "PublishEntry", diff --git a/temporalio/contrib/workflow_streams/_stream.py b/temporalio/contrib/workflow_streams/_stream.py index ff0865c29..e1ebda15b 100644 --- a/temporalio/contrib/workflow_streams/_stream.py +++ b/temporalio/contrib/workflow_streams/_stream.py @@ -57,8 +57,15 @@ does not have to copy the name. """ +POLL_UPDATE_NAME = "__temporal_workflow_stream_poll" +"""The update :class:`WorkflowStream` registers for long polls. + +Public so code that drives the poll itself, with its own retry and +cancellation rules, does not have to copy the name. +""" + _PUBLISH_SIGNAL = PUBLISH_SIGNAL_NAME -_POLL_UPDATE = "__temporal_workflow_stream_poll" +_POLL_UPDATE = POLL_UPDATE_NAME _OFFSET_QUERY = "__temporal_workflow_stream_offset" _MAX_POLL_RESPONSE_BYTES = 1_000_000 diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index cef1cbe3e..af297899b 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -32,6 +32,7 @@ from __future__ import annotations +import asyncio import base64 import logging from collections.abc import AsyncGenerator @@ -48,10 +49,17 @@ WorkflowHandle, WorkflowHistoryEventFilterType, WorkflowQueryFailedError, + WorkflowUpdateFailedError, + WorkflowUpdateRPCTimeoutOrCancelledError, + WorkflowUpdateStage, ) -from temporalio.common import RawValue from temporalio.contrib.workflow_streams import ( + POLL_UPDATE_NAME, PUBLISH_SIGNAL_NAME, + STREAM_DRAINING_ERROR_TYPE, + TRUNCATED_OFFSET_ERROR_TYPE, + PollInput, + PollResult, PublishEntry, PublishInput, WorkflowStream, @@ -81,6 +89,9 @@ _PROVIDER = "workflow_streams" _TAIL_QUERY = "__temporal_streams_tail" _ENCODING = b"binary/plain" +# The server's failure type for an accepted Update whose run closed before +# answering it: the poll's way of saying the run is over. +_UPDATE_OUTLIVED_RUN = "AcceptedUpdateCompletedWorkflow" logger = logging.getLogger(__name__) @@ -452,33 +463,14 @@ async def _read( while True: handle = self._handle(run_id) next_offset = offset - # Pinned to one run, with no client for the shipped chain - # following: a log is not carried across continue-as-new, so - # the successor's offsets start over and this loop is what - # moves from one run to the next. - subscription = WorkflowStreamClient(handle).subscribe( - topic, - from_offset=offset, - result_type=RawValue, - poll_cooldown=self._poll_cooldown, - ) + polls = self._poll(handle, topic, offset) try: - async for item in subscription: - next_offset = item.offset + 1 - if item.topic != topic: - continue - for record in self._records( - decoder, run_id, item.offset, item.data.payload - ): + async for item_offset, payload in polls: + next_offset = item_offset + 1 + for record in self._records(decoder, run_id, item_offset, payload): yield record - except RPCError as error: - # The run closed and its poll Update went with it, or the - # workflow does not exist; the describe below tells which. - if error.status != RPCStatusCode.NOT_FOUND: - raise finally: - if isinstance(subscription, AsyncGenerator): - await subscription.aclose() + await polls.aclose() status = await self._status(handle) if status is None: raise StreamNotFoundError( @@ -508,6 +500,64 @@ async def _read( return run_id, offset = successor, 0 + async def _poll( + self, handle: WorkflowHandle[Any, Any], topic: str, offset: int + ) -> AsyncGenerator[tuple[int, Payload], None]: + """Drive the shipped poll Update on one run until that run closes. + + Written here rather than through ``WorkflowStreamClient.subscribe`` + for two reasons. That loop swallows a cancellation of the caller's + task, so a consumer's ``asyncio.timeout`` or task cancel around a read + would end the subscription and let the read resubscribe forever; here + the cancellation leaves ``read()`` as what it was. And it follows + continue-as-new with offsets this provider does not carry across + runs, which is the outer loop's job. + """ + cooldown = self._poll_cooldown.total_seconds() + while True: + try: + update = await handle.start_update( + POLL_UPDATE_NAME, + PollInput(topics=[topic], from_offset=offset), + wait_for_stage=WorkflowUpdateStage.ACCEPTED, + result_type=PollResult, + ) + result = await update.result() + except WorkflowUpdateRPCTimeoutOrCancelledError as error: + if isinstance(error.__cause__, asyncio.CancelledError): + # The SDK wraps a cancelled await in this error; the + # consumer cancelled us, so that is what comes out. + raise asyncio.CancelledError() from error + # The RPC itself timed out while the run is still open. + continue + except WorkflowUpdateFailedError as error: + cause = getattr(error.cause, "type", None) + if cause == TRUNCATED_OFFSET_ERROR_TYPE: + # The log was truncated past this position; zero means + # from whatever the run still retains. + offset = 0 + continue + if cause == STREAM_DRAINING_ERROR_TYPE: + # Pollers are detached because the run is closing; the + # next attempt learns how it closed. + await asyncio.sleep(cooldown) + continue + if cause == _UPDATE_OUTLIVED_RUN: + return + raise + except RPCError as error: + # The run closed and its poll Update went with it, or the + # workflow does not exist; the caller describes to tell which. + if error.status != RPCStatusCode.NOT_FOUND: + raise + return + for item in result.items: + if item.topic == topic: + yield item.offset, Payload.FromString(base64.b64decode(item.data)) + offset = result.next_offset + if not result.more_ready and cooldown > 0: + await asyncio.sleep(cooldown) + def _records( self, decoder: RecordDecoder, run_id: str, offset: int, payload: Payload ) -> list[StreamRecord[Any]]: diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 28f2f631d..186de9094 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -209,6 +209,39 @@ async def read_everything() -> list[Any]: assert again[1].kind is RecordKind.FINISH +async def test_a_bounded_read_is_cancelled_within_its_timeout( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + async def read_forever() -> None: + async for _ in stream.read(topic=DECISIONS, result_type=dict): + pass + + # The run is open and nothing is published, so the read parks on the + # poll Update. The bound has to come out as a timeout rather than + # vanish into a resubscribe. + started = asyncio.get_running_loop().time() + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(read_forever(), timeout=1) + assert asyncio.get_running_loop().time() - started < 10 + # A reader task cancelled outright ends the same way. + task = asyncio.create_task(read_forever()) + await asyncio.sleep(0.2) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).finish() + await handle.signal(EchoLoop.release) + assert await handle.result() == 0 + + @workflow.defn class Relay: """Publishes one record per run and continues as new once.""" From 1842f6ea0d7a278fce706df5133532492276af14 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 10:27:51 -0700 Subject: [PATCH 18/32] Registered the Workflow Streams conformance provider on the client. --- tests/streams/test_streams_conformance.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 630814c7f..86d14373d 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -115,8 +115,13 @@ async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: # No STREAMS_LIVE gate: the store is the workflow's own History, which the # test environment's server provides. provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + # Registered once, on the client: the host's worker inherits it and the + # cases open handles through client.get_stream_handle. + config = client.config() + config["plugins"] = [provider] + client = Client(**config) hosts: dict[str, WorkflowHandle[Any, Any]] = {} - async with new_worker(client, StreamHost, plugins=[provider]) as worker: + async with new_worker(client, StreamHost) as worker: async def host(workflow_id: str) -> None: if workflow_id not in hosts: From 1802e8cec49c11735761e7e72a91a46525e0332f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 21 Sep 2026 10:53:33 -0700 Subject: [PATCH 19/32] Took topic definitions on the Workflow Streams handle. --- .../streams/providers/workflow_streams.py | 27 ++++++++++++------- 1 file changed, 17 insertions(+), 10 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index af297899b..88de82e29 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -37,7 +37,7 @@ import logging from collections.abc import AsyncGenerator from datetime import timedelta -from typing import Any +from typing import Any, Generic, TypeVar from google.protobuf.message import DecodeError @@ -70,6 +70,7 @@ from temporalio.streams._errors import StreamCursorError, StreamNotFoundError from temporalio.streams._provider import ReadSource, WriteSink from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( RecordDecoder, WireRecord, @@ -86,6 +87,8 @@ "WorkflowStreamsProvider", ] +T = TypeVar("T") + _PROVIDER = "workflow_streams" _TAIL_QUERY = "__temporal_streams_tail" _ENCODING = b"binary/plain" @@ -288,7 +291,7 @@ async def on_workflow_finish(self) -> None: await workflow.wait_condition(workflow.all_handlers_finished) -class WorkflowStreamsProducer: +class WorkflowStreamsProducer(Generic[T]): """Appends by sending the shipped publish Signal directly. Direct rather than through ``WorkflowStreamClient`` because the interface @@ -339,7 +342,7 @@ def _publisher_id(self) -> str: else self._producer_id ) - async def append(self, *values: Any) -> Cursor | None: + async def append(self, *values: T) -> Cursor | None: """Append ``values`` through the shipped publish Signal. Always ``None``: this transport learns positions at read time, so a @@ -431,12 +434,12 @@ def __init__( def read( self, *, - topic: str, + topic: str | StreamTopic[Any], after: Cursor = BEGINNING, result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: """Yield records on ``topic`` after ``after`` until the chain, or the pinned run, closes.""" - _require_topic(topic) + topic, result_type = resolve_topic(topic, result_type) # Parsed here so a foreign cursor fails this call, not the first # iteration of the generator. named = _position(after) @@ -643,13 +646,13 @@ async def _tail( for item in wire ] - async def latest(self, *, topic: str) -> Cursor: + async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: """The newest position in the log, which orders every topic of this workflow. The log is one per run, so the cursor names the run it was read from: the pinned run, or the latest run of the chain. """ - _require_topic(topic) + topic, _ = resolve_topic(topic) handle = self._handle(self._run_id) try: description = await handle.describe() @@ -672,10 +675,14 @@ async def latest(self, *, topic: str) -> Cursor: return _cursor(run_id, -1) def producer( - self, *, topic: str, producer_id: str = "", attempt: int = 0 - ) -> WorkflowStreamsProducer: + self, + *, + topic: str | StreamTopic[Any], + producer_id: str = "", + attempt: int = 0, + ) -> WorkflowStreamsProducer[Any]: """A producer on ``topic``; inside an activity its identity is the activity's.""" - _require_topic(topic) + topic, _ = resolve_topic(topic) producer_id, attempt = producer_identity(producer_id, attempt) return WorkflowStreamsProducer( self._handle(self._run_id), self._converter, topic, producer_id, attempt From f171df21f503353f541f1e78fe4bb5d6f8ee9132 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:06:18 -0700 Subject: [PATCH 20/32] Keyed the publish dedupe on where a producer's records end. A count of signals sent is not where the records reach, so a retry that batched its records differently either lost a batch or wrote records the log already held. What this transport still cannot do is refuse a divergent retry, which the provider now says plainly. --- .../streams/providers/workflow_streams.py | 26 ++++++++++--- tests/streams/test_streams_conformance.py | 7 ++++ .../streams/test_workflow_streams_provider.py | 37 ++++++++++++++++++- 3 files changed, 63 insertions(+), 7 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 88de82e29..392ae28cd 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -15,8 +15,20 @@ code stores and returns that ``Payload`` untouched, so the body's own encoding never meets the transport. - Producer identity dedupes through the shipped publisher state: the - publisher id is ``producer#attempt`` and every publish Signal carries a - monotonic sequence, so a retried batch drops and a new attempt passes. + publisher id is ``producer#attempt`` and every publish Signal carries the + sequence its records end at, so a retried batch drops and a new attempt + passes. + + What this transport cannot keep is the rest of that rule. A publish is a + Signal, which has no response, so the dedupe decision is taken in the + workflow and there is nowhere to report it. A repeat that carries + *different* content at a sequence the log already holds is therefore + dropped rather than refused with + :class:`temporalio.streams.StreamProducerError`, the way the memory and + native providers refuse it. Raising in the Signal handler is not an + alternative: it would fail the Workflow Task on every replay and the + caller would still learn nothing. A caller that needs a divergent retry to + be caught wants a provider whose append is a request and a response. - ``append()`` returns ``None``. The Signal transport learns positions at read time, so a caller that wants to follow from now asks ``latest()``. - A log belongs to one run and is not carried across continue-as-new, so a @@ -321,7 +333,6 @@ def __init__( self._producer_id = producer_id self._attempt = attempt self._sequence = 0 - self._signal_sequence = 0 self._pending: tuple[list[PublishEntry], int] | None = None @property @@ -391,7 +402,11 @@ async def _send(self, batch: list[tuple[RecordKind, Any]]) -> None: await self._signal(entries, next_sequence) async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None: - signal_sequence = self._signal_sequence + 1 + # The dedupe sequence is where this producer's records end, not how + # many signals it has sent. The two differ once a retry batches its + # records differently from the send it is repeating, and a counter of + # signals then either drops a batch of new records or lets records + # that are already there through a second time. self._pending = (entries, next_sequence) try: await self._handle.signal( @@ -399,7 +414,7 @@ async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None PublishInput( items=entries, publisher_id=self._publisher_id, - sequence=signal_sequence, + sequence=next_sequence, ), ) except RPCError as error: @@ -409,7 +424,6 @@ async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None "cannot be appended to" ) from error raise - self._signal_sequence = signal_sequence self._sequence = next_sequence self._pending = None diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 86d14373d..6bea3aead 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -69,6 +69,8 @@ class ProviderCase: client: Client | None = None reports_positions: bool = True """``append()`` returns where the records landed.""" + detects_divergent_retries: bool = True + """``append()`` compares a repeat's content with what it already holds.""" host: Callable[[str], Awaitable[None]] | None = None """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" @@ -134,6 +136,10 @@ async def host(workflow_id: str) -> None: provider, client, reports_positions=False, + # A publish is a Signal, so the dedupe decision is taken in the + # workflow with nowhere to report it. See the module docstring of + # the provider. + detects_divergent_retries=False, host=host, ) for handle in hosts.values(): @@ -147,6 +153,7 @@ async def host(workflow_id: str) -> None: _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, + "detects_divergent_retries": lambda case: case.detects_divergent_retries, } diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 186de9094..f345bffa4 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -366,5 +366,40 @@ async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): await producer.append({"n": 1}) await producer.append({"n": 2}, {"n": 3}) await producer.finish() - assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1, 2]), (3, [3])] + assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (3, [1, 2]), (4, [3])] assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) + + +class _PlainHandle: + """A workflow handle that accepts every Signal and remembers it.""" + + id = "plain" + + def __init__(self) -> None: + self.sent: list[PublishInput] = [] + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.sent.append(arg) + + +async def test_the_dedupe_sequence_names_where_the_records_end(): + # The shipped handler drops a batch whose sequence it has already passed, + # so the sequence has to say how far this producer's records reach. A + # count of signals does not: a retry that batches its records differently + # from the send it repeats then carries a sequence the workflow has not + # seen, and the records it already holds go in a second time. + first = _PlainHandle() + original = _producer(first) # type: ignore[arg-type] + await original.append({"n": 1}, {"n": 2}) + + second = _PlainHandle() + retry = _producer(second) # type: ignore[arg-type] + await retry.append({"n": 1}) + await retry.append({"n": 2}) + await retry.append({"n": 3}) + + # The original ended at record 1, so its sequence is 2. Neither half of + # the retry's re-split reaches past it, and only the new record does. + assert _sequences(first.sent) == [(2, [0, 1])] + assert _sequences(second.sent) == [(1, [0]), (2, [1]), (3, [2])] From ffd14d1c9255f9374a5f6aec70d57d80104197fb Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:06:53 -0700 Subject: [PATCH 21/32] Made the read path page its tail and name its failures. The tail Query answered a whole log in one blob, and latest() answered about the log rather than the topic it was given. A truncated position restarted the read from zero instead of saying the cursor is gone, and four failure paths reached the caller untranslated. --- .../streams/providers/workflow_streams.py | 181 ++++++++++---- temporalio/worker/_workflow_instance.py | 10 +- .../streams/test_workflow_streams_provider.py | 225 +++++++++++++++++- 3 files changed, 360 insertions(+), 56 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 392ae28cd..c0a0dfb6e 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -75,11 +75,14 @@ PublishEntry, PublishInput, WorkflowStream, - WorkflowStreamClient, ) from temporalio.converter import PayloadConverter from temporalio.service import RPCError, RPCStatusCode -from temporalio.streams._errors import StreamCursorError, StreamNotFoundError +from temporalio.streams._errors import ( + StreamCursorError, + StreamError, + StreamNotFoundError, +) from temporalio.streams._provider import ReadSource, WriteSink from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord from temporalio.streams._topic import StreamTopic, resolve_topic @@ -92,6 +95,7 @@ to_wire, ) from temporalio.streams.providers import ProviderPlugin +from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND __all__ = [ "WorkflowStreamsHandle", @@ -103,7 +107,11 @@ _PROVIDER = "workflow_streams" _TAIL_QUERY = "__temporal_streams_tail" +_LATEST_QUERY = "__temporal_streams_latest" _ENCODING = b"binary/plain" +# The same cap the shipped poll path answers under, because both are one +# response through the same server. +_MAX_TAIL_RESPONSE_BYTES = 1_000_000 # The server's failure type for an accepted Update whose run closed before # answering it: the poll's way of saying the run is over. _UPDATE_OUTLIVED_RUN = "AcceptedUpdateCompletedWorkflow" @@ -171,16 +179,46 @@ def __init__(self, stream: WorkflowStream | None = None) -> None: # final task published. The log is workflow state, so a Query # still serves it after completion. workflow.set_query_handler(_TAIL_QUERY, self._tail) + if workflow.get_query_handler(_LATEST_QUERY) is None: + workflow.set_query_handler(_LATEST_QUERY, self._latest) - def _tail(self, from_offset: int) -> list[dict[str, Any]]: - return [ - { - "offset": offset, - "topic": topic, - "data": base64.b64encode(payload.SerializeToString()).decode("ascii"), - } - for offset, topic, payload in self.stream.items_from(from_offset) - ] + def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: + """One page of ``topic``'s items at or past ``from_offset``. + + Filtered and capped here rather than at the caller, because a Query + response has to fit the server's blob limit and a log the reader only + wants one topic of can be much larger than that. + """ + items: list[dict[str, Any]] = [] + size = 0 + next_offset = self.stream.next_offset + more_ready = False + for offset, item_topic, payload in self.stream.items_from(from_offset): + if item_topic != topic: + continue + data = base64.b64encode(payload.SerializeToString()).decode("ascii") + if items and size + len(data) > _MAX_TAIL_RESPONSE_BYTES: + next_offset, more_ready = offset, True + break + size += len(data) + items.append({"offset": offset, "topic": item_topic, "data": data}) + return { + "items": items, + "next_offset": next_offset, + "more_ready": more_ready, + } + + def _latest(self, topic: str) -> int: + """The newest offset holding a record on ``topic``, or -1 when it has none. + + The log orders every topic together, so the head of the log is not an + answer about one topic. Scanning here costs one Query rather than + shipping the log to the caller to find the same thing. + """ + for offset, item_topic, _ in reversed(self.stream.items_from(0)): + if item_topic == topic: + return offset + return -1 def _registered_stream() -> WorkflowStream | None: @@ -500,13 +538,14 @@ async def _read( continue # What landed after the last poll is still in workflow state, so # the tail comes back by Query rather than being lost with the run. - for offset_, shipped_topic, payload in await self._tail( - handle, next_offset - ): - if shipped_topic != topic: - continue - for record in self._records(decoder, run_id, offset_, payload): - yield record + tail_offset = next_offset + while True: + page, tail_offset, more = await self._tail(handle, topic, tail_offset) + for offset_, payload in page: + for record in self._records(decoder, run_id, offset_, payload): + yield record + if not more: + break if ( self._run_id is not None or status != WorkflowExecutionStatus.CONTINUED_AS_NEW @@ -550,10 +589,13 @@ async def _poll( except WorkflowUpdateFailedError as error: cause = getattr(error.cause, "type", None) if cause == TRUNCATED_OFFSET_ERROR_TYPE: - # The log was truncated past this position; zero means - # from whatever the run still retains. - offset = 0 - continue + # Restarting from the beginning would hand the caller + # records it already handled, and only the caller can + # decide to do that. + raise StreamCursorError( + f"offset {offset} of workflow {self._workflow_id!r} run " + f"{handle.run_id!r} is no longer retained" + ) from error if cause == STREAM_DRAINING_ERROR_TYPE: # Pollers are detached because the run is closing; the # next attempt learns how it closed. @@ -561,7 +603,9 @@ async def _poll( continue if cause == _UPDATE_OUTLIVED_RUN: return - raise + raise StreamError( + f"the poll update on workflow {self._workflow_id!r} failed: {error}" + ) from error except RPCError as error: # The run closed and its poll Update went with it, or the # workflow does not exist; the caller describes to tell which. @@ -587,16 +631,21 @@ def _records( def _handle(self, run_id: str | None) -> WorkflowHandle[Any, Any]: return self._client.get_workflow_handle(self._workflow_id, run_id=run_id) - async def _status( - self, handle: WorkflowHandle[Any, Any] - ) -> WorkflowExecutionStatus | None: + async def _describe(self, handle: WorkflowHandle[Any, Any]) -> Any | None: + """The description of ``handle``'s run, or ``None`` when it is gone.""" try: - return (await handle.describe()).status + return await handle.describe() except RPCError as error: if error.status == RPCStatusCode.NOT_FOUND: return None raise + async def _status( + self, handle: WorkflowHandle[Any, Any] + ) -> WorkflowExecutionStatus | None: + description = await self._describe(handle) + return None if description is None else description.status + async def _first_run(self) -> str: """The oldest retained run of the chain, walking back from the latest.""" try: @@ -629,48 +678,76 @@ async def _successor(self, handle: WorkflowHandle[Any, Any]) -> str | None: events = handle.fetch_history_events( event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT ) - async for event in events: - if event.HasField("workflow_execution_continued_as_new_event_attributes"): - attributes = event.workflow_execution_continued_as_new_event_attributes - return attributes.new_execution_run_id or None + try: + async for event in events: + if event.HasField( + "workflow_execution_continued_as_new_event_attributes" + ): + attributes = ( + event.workflow_execution_continued_as_new_event_attributes + ) + return attributes.new_execution_run_id or None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {handle.run_id!r} was not found, " + "so its successor cannot be followed" + ) from error return None async def _tail( - self, handle: WorkflowHandle[Any, Any], from_offset: int - ) -> list[tuple[int, str, Payload]]: + self, handle: WorkflowHandle[Any, Any], topic: str, from_offset: int + ) -> tuple[list[tuple[int, Payload]], int, bool]: + """One page of ``topic``'s tail as ``(items, next offset, more to come)``.""" try: - wire = await handle.query(_TAIL_QUERY, from_offset, result_type=list) + wire = await handle.query( + _TAIL_QUERY, args=[from_offset, topic], result_type=dict + ) except WorkflowQueryFailedError as error: - if "expected but not found" not in str(error): - raise + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the tail query on workflow {self._workflow_id!r} failed: {error}" + ) from error # The workflow never opened a stream through this provider, so # there is no tail to serve. - return [] + return [], from_offset, False except RPCError as error: if error.status != RPCStatusCode.NOT_FOUND: raise # The History is gone; nothing is left to serve. - return [] - return [ - ( - item["offset"], - item["topic"], - Payload.FromString(base64.b64decode(item["data"])), - ) - for item in wire + return [], from_offset, False + items = [ + (item["offset"], Payload.FromString(base64.b64decode(item["data"]))) + for item in wire["items"] ] + return items, wire["next_offset"], bool(wire["more_ready"]) async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: - """The newest position in the log, which orders every topic of this workflow. + """The newest position holding a record on ``topic``. The log is one per run, so the cursor names the run it was read from: - the pinned run, or the latest run of the chain. + the pinned run, or the latest run of the chain. One log orders every + topic, so the answer comes from a Query that scans it for this topic + rather than from the head of the log, which usually names some other + topic's record. """ topic, _ = resolve_topic(topic) handle = self._handle(self._run_id) + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") try: - description = await handle.describe() - head = await WorkflowStreamClient(handle).get_offset() + head = await handle.query(_LATEST_QUERY, topic, result_type=int) + except WorkflowQueryFailedError as error: + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the latest query on workflow {self._workflow_id!r} failed: " + f"{error}" + ) from error + # The workflow never opened a stream through this provider, so + # the topic holds nothing. + head = -1 except RPCError as error: if error.status == RPCStatusCode.NOT_FOUND: raise StreamNotFoundError( @@ -679,9 +756,9 @@ async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: raise run_id = description.run_id assert run_id is not None - if head > 0: - return _cursor(run_id, head - 1) - # An empty log on the chain's first run is the beginning of the + if head >= 0: + return _cursor(run_id, head) + # An empty topic on the chain's first run is the beginning of the # stream; on a successor it is a position of its own, because # BEGINNING would send a chain-following read back to the first run. if await self._predecessor(run_id) is None: diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 26f796e57..f1d6d7bf7 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -91,6 +91,13 @@ logger = logging.getLogger(__name__) +QUERY_HANDLER_NOT_FOUND = "expected but not found" +"""The phrase a query for an unregistered handler comes back with. + +A caller that has to recognise the condition has only the failure message to +go on, so it matches this constant rather than a copy of the sentence. +""" + # Set to true to log all cases where we're ignoring things during delete LOG_IGNORE_DURING_DELETE = False @@ -818,7 +825,8 @@ async def run_query() -> None: if not defn: known_queries = sorted([k for k in self._queries.keys() if k]) raise RuntimeError( - f"Query handler for '{job.query_type}' expected but not found, " + f"Query handler for '{job.query_type}' " + f"{QUERY_HANDLER_NOT_FOUND}, " f"known queries: [{' '.join(known_queries)}]" ) diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index f345bffa4..d99b1d1f6 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -13,6 +13,7 @@ import asyncio import base64 import uuid +from collections.abc import AsyncIterator from datetime import timedelta from typing import Any @@ -20,15 +21,30 @@ from temporalio import workflow from temporalio.api.common.v1 import Payload -from temporalio.client import Client -from temporalio.contrib.workflow_streams import PublishInput +from temporalio.client import ( + Client, + WorkflowExecutionStatus, + WorkflowQueryFailedError, +) +from temporalio.contrib.workflow_streams import PublishInput, WorkflowStream from temporalio.converter import DataConverter -from temporalio.streams import RecordKind, Supersession +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import ( + BEGINNING, + RecordKind, + StreamCursorError, + StreamError, + StreamNotFoundError, + Supersession, +) from temporalio.streams._wire import WireRecord from temporalio.streams.providers.workflow_streams import ( + WorkflowStreamsHandle, WorkflowStreamsProducer, WorkflowStreamsProvider, + _InstanceStream, ) +from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND from tests.helpers import new_worker INPUTS = "inputs" @@ -242,6 +258,106 @@ async def read_forever() -> None: assert await handle.result() == 0 +@workflow.defn +class Truncating: + """Publishes on two topics and truncates its log when told to.""" + + def __init__(self) -> None: + # Constructed here so the shipped signal handler is registered before + # the provider looks for it, the way a migrating application holds it. + self._stream = WorkflowStream() + self._released = False + + @workflow.signal + def truncate_to(self, offset: int) -> None: + self._stream.truncate(offset) + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + other = workflow.stream_writer(INPUTS) + decisions.publish({"n": 0}) + other.publish({"side": "inputs"}) + decisions.publish({"n": 1}) + await workflow.wait_condition(lambda: self._released) + + +async def test_a_truncated_position_is_refused_rather_than_restarted( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Truncating, plugins=[provider]) as worker: + handle = await client.start_workflow( + Truncating.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + first = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert first[0].value == {"n": 0} + + # Everything the reader's cursor names is dropped from the log. + await handle.signal(Truncating.truncate_to, 3) + resumed = stream.read(topic=DECISIONS, result_type=dict, after=first[0].cursor) + # Starting over would hand back records the caller already handled, + # and only the caller can decide to do that. + with pytest.raises(StreamCursorError): + await take(resumed, 1, timeout=30) + + await handle.signal(Truncating.release) + await handle.result() + + +async def test_latest_names_the_newest_record_on_the_topic_asked_for( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Truncating, plugins=[provider]) as worker: + handle = await client.start_workflow( + Truncating.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + decisions = await take(stream.read(topic=DECISIONS, result_type=dict), 2) + + # One log orders both topics and the newest item on it belongs to + # `inputs`, so a log-global answer would be the wrong cursor here. + assert await stream.latest(topic=DECISIONS) == decisions[-1].cursor + inputs = await take(stream.read(topic=INPUTS, result_type=dict), 1) + assert await stream.latest(topic=INPUTS) == inputs[0].cursor + assert await stream.latest(topic="never-written") == BEGINNING + + await handle.signal(Truncating.release) + await handle.result() + + +async def test_the_tail_query_pages_instead_of_answering_in_one_blob(): + # The workflow-side half of the tail, driven directly: a Query response + # has to fit the server's blob limit, so a log larger than the cap comes + # back a page at a time with a position to resume from. + big = Payload(metadata={"encoding": b"binary/plain"}, data=b"x" * 400_000) + items = [(offset, DECISIONS, big) for offset in range(6)] + + class _Log: + next_offset = len(items) + + def items_from(self, offset: int) -> list[Any]: + return [item for item in items if item[0] >= offset] + + instance = _InstanceStream.__new__(_InstanceStream) + instance.stream = _Log() # type: ignore[assignment] + + seen: list[int] = [] + offset, more = 0, True + while more: + page = instance._tail(offset, DECISIONS) + assert page["items"], "a page that fits nothing would never finish" + seen.extend(item["offset"] for item in page["items"]) + offset, more = page["next_offset"], page["more_ready"] + assert seen == list(range(6)) + + @workflow.defn class Relay: """Publishes one record per run and continues as new once.""" @@ -370,6 +486,109 @@ async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) +class _Description: + run_id = "the-only-run" + status = WorkflowExecutionStatus.COMPLETED + + +class _StubHandle: + """A workflow handle that answers describe and fails whatever the test names.""" + + id = "stub" + run_id = "the-only-run" + + def __init__( + self, + *, + query_error: BaseException | None = None, + events_error: BaseException | None = None, + ) -> None: + self._query_error = query_error + self._events_error = events_error + + async def describe(self) -> Any: + return _Description() + + async def start_update(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + # The run is closed, so its poll Update is gone with it and the read + # goes on to the tail Query, which is what these cases are about. + raise RPCError("no poll update", RPCStatusCode.NOT_FOUND, b"") + + async def query(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + assert self._query_error is not None + raise self._query_error + + def fetch_history_events(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + error = self._events_error + + async def _events() -> AsyncIterator[Any]: + for event in (): + yield event + if error is not None: + raise error + + return _events() + + +class _OneHandleClient: + """A client that answers every handle request with the same handle.""" + + data_converter = DataConverter.default + + def __init__(self, handle: Any) -> None: + self._handle = handle + + def get_workflow_handle( + self, workflow_id: str, *, run_id: str | None = None + ) -> Any: + del workflow_id, run_id + return self._handle + + +def _handle_over(stub: _StubHandle) -> WorkflowStreamsHandle: + return WorkflowStreamsHandle( + _OneHandleClient(stub), # type: ignore[arg-type] + "wf", + "the-only-run", + timedelta(0), + ) + + +async def test_a_failing_tail_query_comes_back_as_a_stream_error(): + # The handler being absent is the one benign case; anything else is a + # real failure, and the interface says a caller catches stream conditions + # by meaning rather than by the client's own exception types. + stub = _StubHandle(query_error=WorkflowQueryFailedError("the workflow rejected it")) + with pytest.raises(StreamError): + await take(_handle_over(stub).read(topic=DECISIONS, result_type=dict), 1) + + +async def test_a_missing_tail_handler_ends_the_read_instead_of_failing_it(): + stub = _StubHandle( + query_error=WorkflowQueryFailedError( + f"Query handler for 'x' {QUERY_HANDLER_NOT_FOUND}, known queries: []" + ) + ) + # The workflow never opened a stream through this provider, so the run + # holds no tail and the read is simply over. + assert [r async for r in _handle_over(stub).read(topic=DECISIONS)] == [] + + +async def test_a_failing_latest_query_comes_back_as_a_stream_error(): + stub = _StubHandle(query_error=WorkflowQueryFailedError("the workflow rejected it")) + with pytest.raises(StreamError): + await _handle_over(stub).latest(topic=DECISIONS) + + +async def test_a_missing_run_leaves_the_successor_lookup_as_a_stream_error(): + stub = _StubHandle(events_error=RPCError("gone", RPCStatusCode.NOT_FOUND, b"")) + with pytest.raises(StreamNotFoundError): + await _handle_over(stub)._successor(stub) # type: ignore[arg-type] + + class _PlainHandle: """A workflow handle that accepts every Signal and remembers it.""" From bae3a83695693ac94b44a4fe79b07c562a2d0f20 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 20:02:05 -0700 Subject: [PATCH 22/32] Numbered the Workflow Streams producer's records from one. Zero on the wire now says a producer does not number its records, so the WFS producer follows the memory producer and starts at one. --- temporalio/streams/providers/workflow_streams.py | 4 +++- tests/streams/test_workflow_streams_provider.py | 10 +++++----- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index c0a0dfb6e..20e6cc526 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -370,7 +370,9 @@ def __init__( self._topic = topic self._producer_id = producer_id self._attempt = attempt - self._sequence = 0 + # One-based, because zero on the wire says the producer does not + # number its records and this one does. + self._sequence = 1 self._pending: tuple[list[PublishEntry], int] | None = None @property diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index d99b1d1f6..c46d0f791 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -471,7 +471,7 @@ async def test_a_retried_append_after_an_ambiguous_failure_writes_once(): # The retry carries the same signal sequence and the same record # sequence as the failed send, so the shipped dedupe drops the copy; the # batch after it continues the numbering. - assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (2, [1])] + assert _sequences(handle.sent) == [(2, [1]), (2, [1]), (3, [2])] assert all(publish.publisher_id == "model#1" for publish in handle.sent) @@ -482,7 +482,7 @@ async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): await producer.append({"n": 1}) await producer.append({"n": 2}, {"n": 3}) await producer.finish() - assert _sequences(handle.sent) == [(1, [0]), (1, [0]), (3, [1, 2]), (4, [3])] + assert _sequences(handle.sent) == [(2, [1]), (2, [1]), (4, [2, 3]), (5, [4])] assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) @@ -618,7 +618,7 @@ async def test_the_dedupe_sequence_names_where_the_records_end(): await retry.append({"n": 2}) await retry.append({"n": 3}) - # The original ended at record 1, so its sequence is 2. Neither half of + # The original ended at record 2, so its sequence is 3. Neither half of # the retry's re-split reaches past it, and only the new record does. - assert _sequences(first.sent) == [(2, [0, 1])] - assert _sequences(second.sent) == [(1, [0]), (2, [1]), (3, [2])] + assert _sequences(first.sent) == [(3, [1, 2])] + assert _sequences(second.sent) == [(2, [1]), (3, [2]), (4, [3])] From cc6dbbd7aed6bcb355a759182b2778902931922a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 20:02:05 -0700 Subject: [PATCH 23/32] Typed the empty history iterator in the provider test double. --- tests/streams/test_workflow_streams_provider.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index c46d0f791..e520936d0 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -525,7 +525,8 @@ def fetch_history_events(self, *args: Any, **kwargs: Any) -> Any: error = self._events_error async def _events() -> AsyncIterator[Any]: - for event in (): + events: tuple[Any, ...] = () + for event in events: yield event if error is not None: raise error From 3dc0d665323b6fa3138a2c8fca37de6e919a5e34 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 29 Sep 2026 08:30:03 -0700 Subject: [PATCH 24/32] Retried a poll that reaches a run before its first task registers. An Update in a run's first task runs ahead of the start hook, so a live read hit it on every successor after continue-as-new. A second rejection on a running run now names the missing provider instead. --- .../streams/providers/workflow_streams.py | 31 +++- .../streams/test_workflow_streams_provider.py | 163 ++++++++++++++++++ 2 files changed, 191 insertions(+), 3 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 20e6cc526..2b5f1b625 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -115,6 +115,9 @@ # The server's failure type for an accepted Update whose run closed before # answering it: the poll's way of saying the run is over. _UPDATE_OUTLIVED_RUN = "AcceptedUpdateCompletedWorkflow" +# The SDK's rejection of an Update no handler is registered for. It shares +# its wording with the Query one. +_HANDLER_NOT_FOUND = QUERY_HANDLER_NOT_FOUND logger = logging.getLogger(__name__) @@ -326,9 +329,10 @@ def open_writer(self, topic: str) -> WriteSink: return _WSWriteSink(self._own_stream(), topic) def on_workflow_start(self) -> None: - # Registered before the first task completes, because an outside - # reader can poll before workflow code has opened anything, and an - # Update with no handler yet is rejected rather than held. + # Registered on the first task, because an outside reader can poll + # before workflow code has opened anything, and an Update with no + # handler yet is rejected rather than held. A poll that arrives in + # that first task still runs ahead of this hook; the reader retries it. self._own_stream() async def on_workflow_finish(self) -> None: @@ -572,6 +576,7 @@ async def _poll( runs, which is the outer loop's job. """ cooldown = self._poll_cooldown.total_seconds() + unhandled = False while True: try: update = await handle.start_update( @@ -605,6 +610,26 @@ async def _poll( continue if cause == _UPDATE_OUTLIVED_RUN: return + if _HANDLER_NOT_FOUND in str(error.cause): + if await self._status(handle) != WorkflowExecutionStatus.RUNNING: + # The run closed with the rejecting task, as a run + # that continues as new on its first task does; the + # caller describes it and follows the chain. + return + if unhandled: + raise StreamError( + f"workflow {self._workflow_id!r} run {handle.run_id!r} " + "does not serve the poll update: the workflow_streams " + "provider is not installed on its worker, and no " + "WorkflowStream was constructed" + ) from error + # A rejection comes back with the completion of the task + # that made it, so a retry reaches a later task, and the + # start hook registers the handler on the first one. Only + # a second rejection means there is no handler to wait for. + unhandled = True + await asyncio.sleep(cooldown) + continue raise StreamError( f"the poll update on workflow {self._workflow_id!r} failed: {error}" ) from error diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index e520936d0..4b88e16c1 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -417,6 +417,169 @@ async def read_everything(stream: Any) -> list[Any]: assert resumed == [{"run": 1}, None] +@workflow.defn +class Rolling: + """Publishes what it is sent, continues as new onto ``successor_queue`` once.""" + + def __init__(self) -> None: + self._sent: list[int] = [] + self._roll_to = "" + self._released = False + + @workflow.signal + def emit(self, n: int) -> None: + self._sent.append(n) + + @workflow.signal + def roll(self, successor_queue: str) -> None: + self._roll_to = successor_queue + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + while True: + await workflow.wait_condition( + lambda: bool(self._sent or self._roll_to or self._released) + ) + while self._sent: + decisions.publish({"run": run, "n": self._sent.pop(0)}) + if self._roll_to: + workflow.continue_as_new(run + 1, task_queue=self._roll_to) + if self._released: + decisions.finish() + return + + +@workflow.defn +class Idle: + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def _collect_all(records: Any, into: list[Any]) -> None: + async for record in records: + into.append((record.kind, record.value)) + + +async def test_a_live_read_without_a_run_id_follows_continue_as_new( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + successor_queue = f"streams-ws-successor-{uuid.uuid4().hex}" + async with new_worker(client, Rolling, plugins=[provider]) as worker: + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_all( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + await handle.signal(Rolling.emit, 1) + await handle.signal(Rolling.emit, 2) + await assert_eventually_len(seen, 2, reader) + + # The successor runs on a queue nobody polls yet, so the reader's + # first poll on it waits for the successor's first task and lands in + # it, ahead of the hook that registers the handler. + await handle.signal(Rolling.roll, successor_queue) + successor = client.get_workflow_handle(workflow_id) + while (await successor.describe()).run_id == handle.result_run_id: + await asyncio.sleep(0.05) + await asyncio.sleep(1) + assert not reader.done() + async with new_worker( + client, Rolling, plugins=[provider], task_queue=successor_queue + ): + await successor.signal(Rolling.emit, 3) + await successor.signal(Rolling.emit, 4) + await assert_eventually_len(seen, 4, reader) + await successor.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await successor.result() + + assert seen == [ + (RecordKind.DATA, {"run": 0, "n": 1}), + (RecordKind.DATA, {"run": 0, "n": 2}), + (RecordKind.DATA, {"run": 1, "n": 3}), + (RecordKind.DATA, {"run": 1, "n": 4}), + (RecordKind.FINISH, None), + ] + + +async def test_a_poll_that_arrives_before_the_first_task_is_retried( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + task_queue = f"streams-ws-{uuid.uuid4().hex}" + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=task_queue + ) + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_all( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + # No worker yet, so the poll is delivered in the run's first task. + await asyncio.sleep(1) + async with new_worker(client, Rolling, plugins=[provider], task_queue=task_queue): + await handle.signal(Rolling.emit, 1) + await assert_eventually_len(seen, 1, reader) + await handle.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await handle.result() + assert seen == [(RecordKind.DATA, {"run": 0, "n": 1}), (RecordKind.FINISH, None)] + + +async def test_a_running_workflow_without_the_provider_fails_the_read_clearly( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Idle) as worker: + handle = await client.start_workflow( + Idle.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + with pytest.raises(StreamError, match="provider is not installed"): + await take(stream.read(topic=DECISIONS), 1) + await handle.signal(Idle.release) + await handle.result() + + +async def assert_eventually_len( + items: list[Any], count: int, reader: asyncio.Task[None] +) -> None: + async def _wait() -> None: + while len(items) < count: + if reader.done(): + # A read that failed surfaces its error instead of a timeout. + reader.result() + raise AssertionError(f"the read ended early with {items}") + await asyncio.sleep(0.05) + + await asyncio.wait_for(_wait(), 30) + + async def _values(records: Any) -> list[Any]: return [r.value async for r in records] From 2332342d6c28d1dd92179c2b8d08388bca3e29ce Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 29 Sep 2026 10:10:31 -0700 Subject: [PATCH 25/32] Absorbed the current stream contract on Workflow Streams. The provider serves a read that starts at END or at the last N records, defaults the topic when a call names none, and refuses a stream an activity owns, since its log lives inside a running workflow. The conformance setup gained the truncate hook and the worker the shared cases expect. --- .../streams/providers/workflow_streams.py | 147 ++++++++++++++++-- tests/streams/test_streams_conformance.py | 64 +++++++- .../streams/test_workflow_streams_provider.py | 71 +++++++++ 3 files changed, 265 insertions(+), 17 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 2b5f1b625..b03fbbbe4 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -49,7 +49,7 @@ import logging from collections.abc import AsyncGenerator from datetime import timedelta -from typing import Any, Generic, TypeVar +from typing import Any, Generic, NoReturn, TypeVar from google.protobuf.message import DecodeError @@ -82,9 +82,17 @@ StreamCursorError, StreamError, StreamNotFoundError, + StreamUnsupportedError, ) from temporalio.streams._provider import ReadSource, 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, @@ -108,6 +116,7 @@ _PROVIDER = "workflow_streams" _TAIL_QUERY = "__temporal_streams_tail" _LATEST_QUERY = "__temporal_streams_latest" +_START_QUERY = "__temporal_streams_start" _ENCODING = b"binary/plain" # The same cap the shipped poll path answers under, because both are one # response through the same server. @@ -184,6 +193,8 @@ def __init__(self, stream: WorkflowStream | None = None) -> None: workflow.set_query_handler(_TAIL_QUERY, self._tail) if workflow.get_query_handler(_LATEST_QUERY) is None: workflow.set_query_handler(_LATEST_QUERY, self._latest) + if workflow.get_query_handler(_START_QUERY) is None: + workflow.set_query_handler(_START_QUERY, self._start) def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: """One page of ``topic``'s items at or past ``from_offset``. @@ -223,6 +234,33 @@ def _latest(self, topic: str) -> int: return offset return -1 + def _start(self, topic: str, last_n: int) -> int: + """Where a read starts: the log's head, or ``topic``'s newest ``last_n`` records. + + Zero ``last_n`` asks for the head, which is where ``END`` starts. + """ + return start_offset(self.stream, topic, last_n) + + +def start_offset(stream: WorkflowStream, topic: str, last_n: int) -> int: + """The log offset a read of ``topic`` starts at. + + ``last_n`` of zero is the head of the log, so only what is published next + is read. Otherwise it is the offset of ``topic``'s ``last_n``-th newest + item, or the oldest one the log holds when there are fewer. The log is + workflow state, so the workflow half answers the same on every replay. + """ + if last_n <= 0: + return stream.next_offset + items = stream.items_from(0) + seen = 0 + for offset, item_topic, _ in reversed(items): + if item_topic == topic: + seen += 1 + if seen == last_n: + return offset + return items[0][0] if items else stream.next_offset + def _registered_stream() -> WorkflowStream | None: handler = workflow.get_signal_handler(PUBLISH_SIGNAL_NAME) @@ -311,10 +349,25 @@ def _own_stream(self) -> WorkflowStream: self._stream = _instance().stream return self._stream - def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + check_read_start(after, last) _require_topic(topic) run_id = workflow.info().run_id start = 0 + # The log is workflow state, so END and last= resolve against it here + # and land on the same offset on every replay. + if last is not None: + return _WSReadSource( + self._own_stream(), + topic, + start_offset(self._own_stream(), topic, last), + run_id, + ) + if after == END: + stream = self._own_stream() + return _WSReadSource(stream, topic, stream.next_offset, run_id) named = _position(after) if named is not None: if named[0] != run_id: @@ -492,35 +545,48 @@ def __init__( def read( self, *, - topic: str | StreamTopic[Any], + topic: str | StreamTopic[Any] | None = None, after: Cursor = BEGINNING, + last: int | None = None, result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: - """Yield records on ``topic`` after ``after`` until the chain, or the pinned run, closes.""" + """Yield records on ``topic`` from where the read starts until the chain, or the pinned run, closes. + + ``BEGINNING`` is the oldest item the first retained run's log still + holds: the poll Update reads offset zero as the log's base. ``END`` + and ``last=`` start on the current run, or the pinned one, at an + offset its workflow answers by Query when the read starts. + """ + check_read_start(after, last) topic, result_type = resolve_topic(topic, result_type) # Parsed here so a foreign cursor fails this call, not the first # iteration of the generator. - named = _position(after) + named = None if after == END else _position(after) if named is not None and self._run_id is not None and named[0] != self._run_id: raise StreamCursorError( f"cursor {after.token!r} names another run than this handle is pinned to" ) - return self._read(topic, named, after, result_type) + return self._read(topic, named, after, last, result_type) async def _read( self, topic: str, named: tuple[str, int] | None, after: Cursor, + last: int | None, result_type: type | None, ) -> AsyncGenerator[StreamRecord[Any], None]: - decoder = RecordDecoder( - self._converter, result_type, after=after, warn=logger.warning - ) if named is not None: run_id, offset = named[0], named[1] + 1 + elif last is not None or after == END: + run_id, offset = await self._start_on_current_run(topic, last or 0) + # A synthesized record is positioned before the first one read. + after = _cursor(run_id, offset - 1) else: run_id, offset = self._run_id or await self._first_run(), 0 + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) while True: handle = self._handle(run_id) next_offset = offset @@ -750,7 +816,41 @@ async def _tail( ] return items, wire["next_offset"], bool(wire["more_ready"]) - async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: + async def _start_on_current_run(self, topic: str, last_n: int) -> tuple[str, int]: + """The run a read at ``END`` or of the newest records starts on, and the offset. + + Raises: + StreamUnsupportedError: The workflow's worker runs a provider that + predates these reads and cannot answer where they start. + """ + handle = self._handle(self._run_id) + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") + run_id = description.run_id + assert run_id is not None + try: + offset = await self._handle(run_id).query( + _START_QUERY, args=[topic, last_n], result_type=int + ) + except WorkflowQueryFailedError as error: + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the start query on workflow {self._workflow_id!r} failed: {error}" + ) from error + raise StreamUnsupportedError( + f"workflow {self._workflow_id!r} does not answer where a read at END " + "or of the last records starts; its worker's provider predates them" + ) from error + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + return run_id, offset + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: """The newest position holding a record on ``topic``. The log is one per run, so the cursor names the run it was read from: @@ -795,7 +895,7 @@ async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: def producer( self, *, - topic: str | StreamTopic[Any], + topic: str | StreamTopic[Any] | None = None, producer_id: str = "", attempt: int = 0, ) -> WorkflowStreamsProducer[Any]: @@ -831,5 +931,28 @@ def get_stream_handle( """A handle on ``workflow_id``'s log; without ``run_id`` it follows the chain.""" return WorkflowStreamsHandle(client, workflow_id, run_id, self._poll_cooldown) + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> NoReturn: + """Refused: this provider cannot hold a stream an activity owns. + + The log lives inside a running workflow and is served by its handlers, + and a standalone activity has no workflow to host one. An activity + writes to its workflow's topics instead, through + ``activity.stream_handle()`` without a scope. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the workflow streams provider cannot hold a stream an activity owns: its " + "log lives inside a running workflow" + ) + async def close(self) -> None: """Nothing to release: the provider holds no connection of its own.""" diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 8b1560f1d..942595170 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -48,6 +48,7 @@ Supersession, topic, ) +from temporalio.streams.providers import workflow_streams from temporalio.streams.providers.memory import MemoryStreams from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider from tests.helpers import new_worker @@ -74,6 +75,8 @@ class ProviderCase: """``append()`` compares a repeat's content with what it already holds.""" host: Callable[[str], Awaitable[None]] | None = None """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" + task_queue: str | None = None + """Where the setup's worker runs the workflows below, when it has one.""" truncate: Callable[[str, str, int], Awaitable[None]] | None = None """Drops all but the newest records of a workflow's topic, standing in for retention, or ``None`` when the provider offers no way to.""" @@ -106,8 +109,8 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: @workflow.defn -class StreamHost: - """Owns a stream and lingers, so outside code has a running workflow to address.""" +class TruncatingStreamHost: + """A stream host whose log an update can truncate, the way a workflow's retention would.""" def __init__(self) -> None: self._released = False @@ -116,11 +119,28 @@ def __init__(self) -> None: def release(self) -> None: self._released = True + @workflow.update + def truncate(self, topic: str, keep: int) -> None: + stream = workflow_streams._instance().stream # pyright: ignore[reportPrivateUsage] + stream.truncate(workflow_streams.start_offset(stream, topic, keep)) + @workflow.run async def run(self) -> None: await workflow.wait_condition(lambda: self._released) +@workflow.defn +class DefaultTopicAnswer: + """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) + async for value in reader.values(): + workflow.stream_writer().publish({"answer": value["n"] * 2}) + reader.close() + + async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: # No STREAMS_LIVE gate: the store is the workflow's own History, which the # test environment's server provides. @@ -131,14 +151,21 @@ async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: config["plugins"] = [provider] client = Client(**config) hosts: dict[str, WorkflowHandle[Any, Any]] = {} - async with new_worker(client, StreamHost) as worker: + async with new_worker(client, TruncatingStreamHost, DefaultTopicAnswer) as worker: async def host(workflow_id: str) -> None: if workflow_id not in hosts: hosts[workflow_id] = await client.start_workflow( - StreamHost.run, id=workflow_id, task_queue=worker.task_queue + TruncatingStreamHost.run, + id=workflow_id, + task_queue=worker.task_queue, ) + async def truncate(workflow_id: str, topic: str, keep: int) -> None: + await hosts[workflow_id].execute_update( + TruncatingStreamHost.truncate, args=[topic, keep] + ) + yield ProviderCase( "workflow_streams", provider, @@ -149,6 +176,8 @@ async def host(workflow_id: str) -> None: # the provider. detects_divergent_retries=False, host=host, + task_queue=worker.task_queue, + truncate=truncate, ) for handle in hosts.values(): await handle.terminate() @@ -377,6 +406,29 @@ async def test_naming_no_topic_addresses_the_default_topic(case: ProviderCase): assert await stream.latest(topic=DEFAULT_TOPIC) == records[0].cursor +async def test_a_workflow_answers_on_its_default_topic(case: ProviderCase): + if case.client is None or case.task_queue is None: + pytest.skip( + f"the {case.name} setup runs no worker; test_streams_workflow covers " + "its workflow half" + ) + workflow_id = new_workflow_id() + handle = await case.client.start_workflow( + DefaultTopicAnswer.run, id=workflow_id, task_queue=case.task_queue + ) + stream = case.client.get_stream_handle(workflow_id) + await stream.producer(producer_id="client", attempt=1).append({"n": 21}) + await handle.result() + + # The outside producer and the workflow, each naming no topic, meet on + # one topic that a reader naming none sees in order. + records = await take(stream.read(), 2, timeout=30.0) + assert [(r.topic, r.producer_id, r.value) for r in records] == [ + (DEFAULT_TOPIC, "client", {"n": 21}), + (DEFAULT_TOPIC, "", {"answer": 42}), + ] + + async def test_cursor_resumes_where_it_points(case: ProviderCase): workflow_id = new_workflow_id() stream = await case.open(workflow_id) @@ -538,8 +590,10 @@ async def test_beginning_starts_at_the_oldest_record_still_held(case: ProviderCa assert [r.value for r in records] == [{"n": 3}, {"n": 4}] newest = await take(stream.read(topic=OUT, last=3), 2) assert [r.value for r in newest] == [{"n": 3}, {"n": 4}] + # A cursor below the floor is refused. A provider that needs a round trip + # to know says so on the first step rather than on the call. with pytest.raises(StreamCursorError): - stream.read(topic=OUT, after=before[0].cursor) + await take(stream.read(topic=OUT, after=before[0].cursor), 1) async def test_a_read_start_names_one_place(case: ProviderCase): diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 4b88e16c1..6415f07ad 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -31,10 +31,12 @@ from temporalio.service import RPCError, RPCStatusCode from temporalio.streams import ( BEGINNING, + END, RecordKind, StreamCursorError, StreamError, StreamNotFoundError, + StreamUnsupportedError, Supersession, ) from temporalio.streams._wire import WireRecord @@ -786,3 +788,72 @@ async def test_the_dedupe_sequence_names_where_the_records_end(): # the retry's re-split reaches past it, and only the new record does. assert _sequences(first.sent) == [(3, [1, 2])] assert _sequences(second.sent) == [(2, [1]), (3, [2]), (4, [3])] + + +@workflow.defn +class StartsWhenTold: + """Opens a reader on ``inputs`` at the start a signal names and returns what it read.""" + + def __init__(self) -> None: + self._start: str | None = None + + @workflow.signal + def begin(self, start: str) -> None: + self._start = start + + @workflow.run + async def run(self) -> list[Any]: + await workflow.wait_condition(lambda: self._start is not None) + if self._start == "end": + reader = workflow.stream_reader(INPUTS, result_type=dict, after=END) + want = 1 + else: + reader = workflow.stream_reader(INPUTS, result_type=dict, last=2) + want = 2 + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == want: + break + return values + + +async def test_a_workflow_reader_starts_at_end_or_the_newest_records( + client: Client, provider: WorkflowStreamsProvider +): + # The log is workflow state, so both starts resolve against it on the + # workflow thread; a cold cache replays every task and has to land the + # reader on the same offset each time. + async with new_worker( + client, StartsWhenTold, plugins=[provider], max_cached_workflows=0 + ) as worker: + for start, expected in (("last", [3, 4]), ("end", ["new"])): + handle = await client.start_workflow( + StartsWhenTold.run, + id=f"ws-start-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + stream = provider.get_stream_handle(client, handle.id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + await handle.signal(StartsWhenTold.begin, start) + result = asyncio.ensure_future(handle.result()) + if start == "end": + 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, 30) == expected + + +async def test_a_stream_an_activity_owns_is_refused( + client: Client, provider: WorkflowStreamsProvider +): + # The log lives inside a running workflow, so an activity has no place to + # put a stream of its own here. The refusal is the documented error, not + # an AttributeError or a stream silently put somewhere else. + with pytest.raises(StreamUnsupportedError, match="activity"): + provider.get_activity_stream_handle(client, "act") + with pytest.raises(StreamUnsupportedError, match="activity"): + provider.get_activity_stream_handle(client, "act", workflow_id="wf") From 0e567a739ce2f59609c6d0ad6840aefd85416971 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:16:31 -0700 Subject: [PATCH 26/32] Hosted a workflow activity's own streams as reserved topics in the log. An activity a workflow scheduled keeps its streams under activity// in the workflow's log, the way native reserves that prefix, so activity.stream_handle(scope="activity") and the client's activity handle work here. A standalone activity stays refused because no workflow hosts its log, and the read ends with the activity as the workflow describes it. --- .../streams/providers/workflow_streams.py | 190 ++++++++++++++++-- tests/streams/conftest.py | 5 + tests/streams/test_activity_streams.py | 37 +++- .../streams/test_workflow_streams_provider.py | 105 +++++++++- 4 files changed, 307 insertions(+), 30 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index b03fbbbe4..538d9d472 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -40,6 +40,13 @@ through the handler the shipped class registers on it. An evicted and rebuilt workflow gets its own, so a task that failed leaks nothing into the next attempt's log and a replayed run does not see records twice. +- An activity a workflow scheduled keeps its own streams inside that + workflow's log, under the reserved topic ``activity//`` + with ``%`` and ``/`` in the id percent-encoded, the way the native + provider reserves ``activity/`` in the owner's map. The record itself + carries the plain name. A standalone activity has no workflow to host a + log, so its streams are refused, and a workflow's own topic may not start + with the reserved prefix. """ from __future__ import annotations @@ -49,7 +56,7 @@ import logging from collections.abc import AsyncGenerator from datetime import timedelta -from typing import Any, Generic, NoReturn, TypeVar +from typing import Any, Generic, TypeVar from google.protobuf.message import DecodeError @@ -106,6 +113,7 @@ from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND __all__ = [ + "WorkflowStreamsActivityHandle", "WorkflowStreamsHandle", "WorkflowStreamsProducer", "WorkflowStreamsProvider", @@ -118,6 +126,9 @@ _LATEST_QUERY = "__temporal_streams_latest" _START_QUERY = "__temporal_streams_start" _ENCODING = b"binary/plain" +# The topics an activity owns live in its workflow's log under this prefix, +# so no workflow topic may start with it. +_ACTIVITY_PREFIX = "activity/" # The same cap the shipped poll path answers under, because both are one # response through the same server. _MAX_TAIL_RESPONSE_BYTES = 1_000_000 @@ -134,6 +145,21 @@ def _require_topic(topic: str) -> None: if not topic: raise ValueError("topic must not be empty") + if topic.startswith(_ACTIVITY_PREFIX): + raise ValueError( + f"topic {topic!r} is reserved: names under {_ACTIVITY_PREFIX!r} hold the " + "streams of the workflow's activities on the workflow_streams provider" + ) + + +def _activity_topic(activity_id: str, topic: str) -> str: + """The reserved name ``topic`` of ``activity_id``'s streams takes in the log. + + The id is percent-encoded so an id holding ``/`` cannot be read as two + components; the name comes last, so it may hold anything. + """ + escaped = activity_id.replace("%", "%25").replace("/", "%2F") + return f"{_ACTIVITY_PREFIX}{escaped}/{topic}" def _cursor(run_id: str, offset: int) -> Cursor: @@ -207,7 +233,8 @@ def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: size = 0 next_offset = self.stream.next_offset more_ready = False - for offset, item_topic, payload in self.stream.items_from(from_offset): + held = self.stream.items_from(from_offset) + for offset, item_topic, payload in held: if item_topic != topic: continue data = base64.b64encode(payload.SerializeToString()).decode("ascii") @@ -220,6 +247,9 @@ def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: "items": items, "next_offset": next_offset, "more_ready": more_ready, + # Where the retained log starts at or past ``from_offset``: a + # reader whose position is below it has fallen behind truncation. + "base_offset": held[0][0] if held else next_offset, } def _latest(self, topic: str) -> int: @@ -420,11 +450,19 @@ def __init__( topic: str, producer_id: str, attempt: int, + *, + item_topic: str | None = None, ) -> None: - """Bind this producer to ``topic`` on the workflow behind ``handle``.""" + """Bind this producer to ``topic`` on the workflow behind ``handle``. + + ``item_topic`` is the name the log files the records under when it + differs from the name the records carry, as an activity's reserved + topics do. + """ self._handle = handle self._converter = converter self._topic = topic + self._item_topic = topic if item_topic is None else item_topic self._producer_id = producer_id self._attempt = attempt # One-based, because zero on the wire says the producer does not @@ -483,7 +521,7 @@ def _entries( ) sequence += 1 entries.append( - PublishEntry(topic=self._topic, data=_entry_data(_wrap(wire))) + PublishEntry(topic=self._item_topic, data=_entry_data(_wrap(wire))) ) return entries, sequence @@ -559,6 +597,7 @@ def read( """ check_read_start(after, last) topic, result_type = resolve_topic(topic, result_type) + wire_topic = self._wire_topic(topic) # Parsed here so a foreign cursor fails this call, not the first # iteration of the generator. named = None if after == END else _position(after) @@ -566,7 +605,12 @@ def read( raise StreamCursorError( f"cursor {after.token!r} names another run than this handle is pinned to" ) - return self._read(topic, named, after, last, result_type) + return self._read(wire_topic, named, after, last, result_type) + + def _wire_topic(self, topic: str) -> str: + """The name the log files ``topic`` under: the topic itself for a workflow's.""" + _require_topic(topic) + return topic async def _read( self, @@ -810,6 +854,13 @@ async def _tail( raise # The History is gone; nothing is left to serve. return [], from_offset, False + if from_offset and wire.get("base_offset", from_offset) > from_offset: + # Restarting from the base would hand the caller records it + # already handled, and only the caller can decide to do that. + raise StreamCursorError( + f"offset {from_offset} of workflow {self._workflow_id!r} run " + f"{handle.run_id!r} is no longer retained" + ) items = [ (item["offset"], Payload.FromString(base64.b64decode(item["data"]))) for item in wire["items"] @@ -860,12 +911,13 @@ async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor topic's record. """ topic, _ = resolve_topic(topic) + wire_topic = self._wire_topic(topic) handle = self._handle(self._run_id) description = await self._describe(handle) if description is None: raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") try: - head = await handle.query(_LATEST_QUERY, topic, result_type=int) + head = await handle.query(_LATEST_QUERY, wire_topic, result_type=int) except WorkflowQueryFailedError as error: if QUERY_HANDLER_NOT_FOUND not in str(error): raise StreamError( @@ -901,10 +953,108 @@ def producer( ) -> WorkflowStreamsProducer[Any]: """A producer on ``topic``; inside an activity its identity is the activity's.""" topic, _ = resolve_topic(topic) + wire_topic = self._wire_topic(topic) producer_id, attempt = producer_identity(producer_id, attempt) return WorkflowStreamsProducer( - self._handle(self._run_id), self._converter, topic, producer_id, attempt + self._handle(self._run_id), + self._converter, + topic, + producer_id, + attempt, + item_topic=wire_topic, + ) + + +class WorkflowStreamsActivityHandle(WorkflowStreamsHandle): + """The streams one activity of a workflow owns, kept in that workflow's log. + + Each topic is the reserved topic ``activity//`` of the + workflow's log, so an activity's ``tokens`` and its workflow's ``tokens`` + are two streams. The activity belongs to one run, so the handle pins the + run on first use and never follows a successor. A read ends when the + workflow closes, or once the activity has been seen pending, is pending + no longer and the read delivered a record: a stream the activity never + wrote has nothing to close, so a read on it waits for the workflow. + + A read here polls the tail Query and describes the workflow between + polls, every ``poll_cooldown``, rather than parking on the poll Update: + an Update parked in the workflow returns only when the log grows, so it + could not notice the activity ending, and each abandoned one would stay + parked against the run's Update caps. + """ + + def __init__( + self, + client: Client, + workflow_id: str, + run_id: str | None, + activity_id: str, + poll_cooldown: timedelta, + ) -> None: + """Address ``activity_id``'s streams inside ``workflow_id``'s log.""" + super().__init__(client, workflow_id, run_id, poll_cooldown) + self._activity_id = activity_id + + def _wire_topic(self, topic: str) -> str: + return _activity_topic(self._activity_id, topic) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + after: Cursor, + last: int | None, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + if named is not None: + run_id, offset = named[0], named[1] + 1 + elif last is not None or after == END: + run_id, offset = await self._start_on_current_run(topic, last or 0) + after = _cursor(run_id, offset - 1) + else: + run_id, offset = await self._current_run(), 0 + handle = self._handle(run_id) + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning ) + cooldown = self._poll_cooldown.total_seconds() + seen_pending = False + delivered = False + ended = False + while True: + more = True + while more: + page, offset, more = await self._tail(handle, topic, offset) + for offset_, payload in page: + for record in self._records(decoder, run_id, offset_, payload): + delivered = True + yield record + if ended: + return + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {run_id!r} was not found" + ) + pending = any( + info.activity_id == self._activity_id + for info in description.raw_description.pending_activities + ) + seen_pending = seen_pending or pending + # One more pass after learning the stream ended, so a record that + # landed between the tail and the describe is not lost. + ended = description.status != WorkflowExecutionStatus.RUNNING or ( + seen_pending and not pending and delivered + ) + if not ended: + await asyncio.sleep(cooldown) + + async def _current_run(self) -> str: + description = await self._describe(self._handle(self._run_id)) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") + assert description.run_id is not None + return description.run_id class WorkflowStreamsProvider(ProviderPlugin): @@ -938,20 +1088,26 @@ def get_activity_stream_handle( *, workflow_id: str | None = None, run_id: str | None = None, - ) -> NoReturn: - """Refused: this provider cannot hold a stream an activity owns. + ) -> WorkflowStreamsActivityHandle: + """A handle on the streams an activity of ``workflow_id`` owns. - The log lives inside a running workflow and is served by its handlers, - and a standalone activity has no workflow to host one. An activity - writes to its workflow's topics instead, through - ``activity.stream_handle()`` without a scope. + They live in the workflow's log under the reserved topics + ``activity//``, and ``run_id`` pins the workflow's + run. Without ``workflow_id`` the activity is a standalone one, which + has no workflow to host a log, so this provider refuses it; such an + activity addresses a workflow's stream by ``workflow_id`` instead. Raises: - StreamUnsupportedError: Always. + StreamUnsupportedError: ``workflow_id`` was not given. """ - raise StreamUnsupportedError( - "the workflow streams provider cannot hold a stream an activity owns: its " - "log lives inside a running workflow" + if workflow_id is None: + raise StreamUnsupportedError( + "the workflow_streams provider cannot hold a stream a standalone " + "activity owns: its log lives inside a running workflow, and no " + "workflow hosts this activity's" + ) + return WorkflowStreamsActivityHandle( + client, workflow_id, run_id, activity_id, self._poll_cooldown ) async def close(self) -> None: diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index b8a152f20..45e09bfd5 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -15,3 +15,8 @@ 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", + "standalone_activities: the case needs the streams of an activity outside " + "any workflow", + ) diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index 8aa4a653a..af77646aa 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -7,10 +7,13 @@ activity execution, so a retry writes to the same stream and a reader sees the attempt change as ``SUPERSEDED``. -The memory provider always runs. A storage provider adds itself to -``SETUPS`` behind its own ``STREAMS_LIVE`` gate: its setup receives the -environment's client and hands back the provider and a client with it -registered, which the workers and the reads in these cases share. +The memory provider always runs, and so does Workflow Streams, whose store is +the workflow's own History. A storage provider adds itself to ``SETUPS`` +behind its own ``STREAMS_LIVE`` gate: its setup receives the environment's +client and hands back the provider and a client with it registered, which the +workers and the reads in these cases share, and says whether it holds the +streams of a standalone activity, so the cases marked +``standalone_activities`` are skipped with a reason where it does not. """ from __future__ import annotations @@ -34,6 +37,7 @@ topic, ) from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider from temporalio.testing import WorkflowEnvironment from tests.helpers import new_worker @@ -47,6 +51,8 @@ class ActivitySetup: name: str provider: StreamProvider client: Client + standalone_activities: bool = True + """The provider holds the streams of an activity outside any workflow.""" async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: @@ -57,8 +63,20 @@ async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider.reset() +async def _workflow_streams_setup(client: Client) -> AsyncIterator[ActivitySetup]: + provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + config = client.config() + config["plugins"] = [provider] + # An activity's streams live in its workflow's log, so a standalone + # activity has nowhere to put them. See the provider's module docstring. + yield ActivitySetup( + "workflow_streams", provider, Client(**config), standalone_activities=False + ) + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ActivitySetup]]] = { - "memory": _memory_setup + "memory": _memory_setup, + "workflow_streams": _workflow_streams_setup, } @@ -69,6 +87,13 @@ async def setup( if env.supports_time_skipping: pytest.skip("the time-skipping test server has no standalone activities") async for found in SETUPS[request.param](client): + if ( + request.node.get_closest_marker("standalone_activities") + and not found.standalone_activities + ): + pytest.skip( + f"the {found.name} provider does not hold a standalone activity's streams" + ) yield found @@ -167,6 +192,7 @@ async def test_scope_activity_gives_a_workflow_activity_its_own_streams( assert await workflow_stream.latest(topic=TOKENS) == BEGINNING +@pytest.mark.standalone_activities async def test_standalone_activity_defaults_to_its_own_stream(setup: ActivitySetup): client = setup.client activity_id = f"streams-saa-{uuid.uuid4().hex}" @@ -199,6 +225,7 @@ async def fail_once_after_writing() -> None: await producer.finish() +@pytest.mark.standalone_activities async def test_a_retry_inherits_the_stream_and_supersedes(setup: ActivitySetup): client = setup.client activity_id = f"streams-saa-retry-{uuid.uuid4().hex}" diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 6415f07ad..ed5eb106a 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -19,7 +19,7 @@ import pytest -from temporalio import workflow +from temporalio import activity, workflow from temporalio.api.common.v1 import Payload from temporalio.client import ( Client, @@ -41,6 +41,7 @@ ) from temporalio.streams._wire import WireRecord from temporalio.streams.providers.workflow_streams import ( + WorkflowStreamsActivityHandle, WorkflowStreamsHandle, WorkflowStreamsProducer, WorkflowStreamsProvider, @@ -847,13 +848,101 @@ async def test_a_workflow_reader_starts_at_end_or_the_newest_records( assert await asyncio.wait_for(result, 30) == expected -async def test_a_stream_an_activity_owns_is_refused( +async def test_a_standalone_activity_stream_is_refused( client: Client, provider: WorkflowStreamsProvider ): - # The log lives inside a running workflow, so an activity has no place to - # put a stream of its own here. The refusal is the documented error, not - # an AttributeError or a stream silently put somewhere else. - with pytest.raises(StreamUnsupportedError, match="activity"): + # The log lives inside a running workflow, so an activity outside any + # workflow has no place to put a stream of its own here. The refusal is + # the documented error, not an AttributeError or a stream silently put + # somewhere else. An activity a workflow scheduled is served. + with pytest.raises(StreamUnsupportedError, match="standalone"): provider.get_activity_stream_handle(client, "act") - with pytest.raises(StreamUnsupportedError, match="activity"): - provider.get_activity_stream_handle(client, "act", workflow_id="wf") + assert isinstance( + provider.get_activity_stream_handle(client, "act", workflow_id="wf"), + WorkflowStreamsActivityHandle, + ) + + +TOKENS = "tokens" + + +@activity.defn +async def stream_then_finish(count: int) -> None: + producer = activity.stream_handle(scope="activity").producer(topic=TOKENS) + for n in range(count): + await producer.append({"n": n}) + # Long enough for a reader polling every few milliseconds to see this + # activity pending before it finishes. + await asyncio.sleep(1) + await producer.finish() + + +@workflow.defn +class RunsAnActivityThenLingers: + """Runs the streaming activity by a fixed id, then waits to be released.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self, count: int) -> None: + await workflow.execute_activity( + stream_then_finish, + count, + activity_id="streamer", + start_to_close_timeout=timedelta(seconds=30), + ) + await workflow.wait_condition(lambda: self._released) + + +async def test_an_activity_read_ends_with_the_activity_while_the_workflow_runs( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker( + client, + RunsAnActivityThenLingers, + activities=[stream_then_finish], + plugins=[provider], + ) as worker: + handle = await client.start_workflow( + RunsAnActivityThenLingers.run, + 2, + id=workflow_id, + task_queue=worker.task_queue, + ) + own = provider.get_activity_stream_handle( + client, "streamer", workflow_id=workflow_id + ) + + async def read_everything() -> list[Any]: + return [r async for r in own.read(topic=TOKENS, result_type=dict)] + + records = await asyncio.wait_for(read_everything(), 30) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"n": 0}), + (RecordKind.DATA, {"n": 1}), + (RecordKind.FINISH, None), + ] + # The producer is the activity, and the record carries the plain + # topic name; the reserved name is the log's business. + assert all(r.producer_id == "streamer" and r.attempt == 1 for r in records) + assert all(r.topic == TOKENS for r in records) + # The activity's end ended the read: the workflow is still running. + assert (await handle.describe()).status == WorkflowExecutionStatus.RUNNING + + # The workflow's topic of the same name is another stream, and the + # reserved name cannot be reached as a workflow topic. + workflow_stream = provider.get_stream_handle(client, workflow_id) + assert await workflow_stream.latest(topic=TOKENS) == BEGINNING + with pytest.raises(ValueError, match="reserved"): + workflow_stream.read(topic="activity/streamer/tokens") + with pytest.raises(ValueError, match="reserved"): + workflow_stream.producer(topic="activity/x", producer_id="p", attempt=1) + + await handle.signal(RunsAnActivityThenLingers.release) + await handle.result() From 93a4e4554e586da79f04b112f2dca945d2c415d8 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:21:02 -0700 Subject: [PATCH 27/32] Followed a reset run from the position the read had reached. The base run closes with nothing in its own History about the reset, so a chain-following read asks describe for the run reset into it. That run rebuilt the base run's log up to the reset point by replay, so the items before it sit at the same offsets and the read carries on where it was. --- .../streams/providers/workflow_streams.py | 62 ++++++++-- .../streams/test_workflow_streams_provider.py | 116 +++++++++++++++++- 2 files changed, 168 insertions(+), 10 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 538d9d472..e700ba959 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -35,7 +35,11 @@ cursor names the run as well as the offset. A handle without a run id reads run after run: each log through the poll Update while its run is open and through the tail Query once it has closed, then the successor's - from its first record. + from its first record. A run that was reset is followed too, into the run + describe names, at the position the read had reached: the reset run + rebuilt the base run's log up to the reset point by replay, so the items + before it sit at the same offsets. The reset itself is not reported to + the reader as a record; that is a wire change for a later round. - The workflow-side stream object belongs to the workflow instance, found through the handler the shipped class registers on it. An evicted and rebuilt workflow gets its own, so a task that failed leaks nothing into @@ -594,6 +598,11 @@ def read( holds: the poll Update reads offset zero as the log's base. ``END`` and ``last=`` start on the current run, or the pinned one, at an offset its workflow answers by Query when the read starts. + + The chain is followed across continue-as-new, into the successor's + log from its first record, and across a reset, into the run describe + names at the position the read had reached. A handle pinned to a run + ends with that run either way. """ check_read_start(after, last) topic, result_type = resolve_topic(topic, result_type) @@ -662,15 +671,44 @@ async def _read( yield record if not more: break - if ( - self._run_id is not None - or status != WorkflowExecutionStatus.CONTINUED_AS_NEW - ): + if self._run_id is not None: return - successor = await self._successor(handle) - if successor is None: + following = await self._following(handle, status, tail_offset) + if following is None: return - run_id, offset = successor, 0 + run_id, offset = following + + async def _following( + self, + handle: WorkflowHandle[Any, Any], + status: WorkflowExecutionStatus, + position: int, + ) -> tuple[str, int] | None: + """The run that carries on after ``handle``'s, and where to read it from. + + A continue-as-new names its successor in the close event, and the + successor's log starts over at zero. A reset does not: the base run is + closed with no word of it in its own History, and only describe names + the run reset from it. That run rebuilt its log by replaying the base + run's History up to the reset point, so the items before that point + sit at the same offsets in both logs, and the read carries on at the + position it reached. Past the reset point the two logs differ, and a + reader already there is not told: reporting the reset to consumers as + a record is a wire change for a later round. + """ + if status == WorkflowExecutionStatus.CONTINUED_AS_NEW: + successor = await self._successor(handle) + return None if successor is None else (successor, 0) + reset_run = await self._reset_run(handle) + return None if reset_run is None else (reset_run, position) + + async def _reset_run(self, handle: WorkflowHandle[Any, Any]) -> str | None: + """The run ``handle``'s run was reset into, which only describe reports.""" + description = await self._describe(handle) + if description is None: + return None + extended = description.raw_description.workflow_extended_info + return extended.reset_run_id or None async def _poll( self, handle: WorkflowHandle[Any, Any], topic: str, offset: int @@ -804,7 +842,13 @@ async def _predecessor(self, run_id: str) -> str | None: try: async for event in self._handle(run_id).fetch_history_events(page_size=1): attributes = event.workflow_execution_started_event_attributes - return attributes.continued_execution_run_id or None + if attributes.continued_execution_run_id: + return attributes.continued_execution_run_id + # A reset run's start event is the base run's, copied, and the + # original run id it carries is kept across resets, so it names + # the run the chain of resets began from. + original = attributes.original_execution_run_id + return original if original and original != run_id else None except RPCError as error: if error.status != RPCStatusCode.NOT_FOUND: raise diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index ed5eb106a..f37fccd8a 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -20,7 +20,8 @@ import pytest from temporalio import activity, workflow -from temporalio.api.common.v1 import Payload +from temporalio.api.common.v1 import Payload, WorkflowExecution +from temporalio.api.workflowservice.v1 import ResetWorkflowExecutionRequest from temporalio.client import ( Client, WorkflowExecutionStatus, @@ -554,6 +555,119 @@ async def test_a_poll_that_arrives_before_the_first_task_is_retried( assert seen == [(RecordKind.DATA, {"run": 0, "n": 1}), (RecordKind.FINISH, None)] +async def _reset_at_last_completed_task( + client: Client, workflow_id: str, run_id: str +) -> str: + """Reset ``run_id`` at its last completed task; the id of the run reset into.""" + completion_id = 0 + events = client.get_workflow_handle( + workflow_id, run_id=run_id + ).fetch_history_events() + async for event in events: + if event.HasField("workflow_task_completed_event_attributes"): + completion_id = event.event_id + assert completion_id + answer = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=run_id + ), + reason="re-run from the last completed task", + workflow_task_finish_event_id=completion_id, + request_id=uuid.uuid4().hex, + ) + ) + return answer.run_id + + +async def _collect_records(records: Any, into: list[Any]) -> None: + async for record in records: + into.append(record) + + +def _run_and_offset(record: Any) -> tuple[str, int]: + _, run_id, offset = record.cursor.token.split(":") + return run_id, int(offset) + + +async def test_a_live_read_without_a_run_id_follows_a_reset( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Rolling, plugins=[provider]) as worker: + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + base_run = handle.result_run_id + assert base_run is not None + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_records( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + await handle.signal(Rolling.emit, 1) + await handle.signal(Rolling.emit, 2) + await assert_eventually_len(seen, 2, reader) + + # The base run is closed by the reset with nothing in its own History + # to say so; the reader learns where it went from describe. The reset + # run replays the base run's History up to the last completed task, + # so its log holds the same two records at the same offsets, and the + # read carries on from the position it had reached. + reset_run = await _reset_at_last_completed_task(client, workflow_id, base_run) + assert reset_run != base_run + current = client.get_workflow_handle(workflow_id) + assert (await current.describe()).run_id == reset_run + await current.signal(Rolling.emit, 3) + await current.signal(Rolling.emit, 4) + await assert_eventually_len(seen, 4, reader) + await current.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await current.result() + + assert [(r.kind, r.value) for r in seen] == [ + (RecordKind.DATA, {"run": 0, "n": 1}), + (RecordKind.DATA, {"run": 0, "n": 2}), + (RecordKind.DATA, {"run": 0, "n": 3}), + (RecordKind.DATA, {"run": 0, "n": 4}), + (RecordKind.FINISH, None), + ] + # Two records from the base run, then the reset run's, whose offsets + # continue where the base run's log stood at the reset point. + assert [_run_and_offset(r) for r in seen] == [ + (base_run, 0), + (base_run, 1), + (reset_run, 2), + (reset_run, 3), + (reset_run, 4), + ] + # A handle pinned to the base run ends with it. Both closed runs are + # served by the tail Query, which needs the worker still up. + pinned = provider.get_stream_handle(client, workflow_id, run_id=base_run) + pinned_records: list[Any] = [] + await asyncio.wait_for( + _collect_records(pinned.read(topic=DECISIONS), pinned_records), 30 + ) + assert [_run_and_offset(r) for r in pinned_records] == [ + (base_run, 0), + (base_run, 1), + ] + # BEGINNING on the chain starts at the base run, whose start event the + # reset run copied, and a resume from a base run cursor crosses over. + chain = provider.get_stream_handle(client, workflow_id) + resumed = await take(chain.read(topic=DECISIONS, after=seen[1].cursor), 3) + assert [_run_and_offset(r) for r in resumed] == [ + (reset_run, 2), + (reset_run, 3), + (reset_run, 4), + ] + + async def test_a_running_workflow_without_the_provider_fails_the_read_clearly( client: Client, provider: WorkflowStreamsProvider ): From b5321258a4556848b7dc2a64e8923aa273fa087c Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:38:25 -0700 Subject: [PATCH 28/32] Carried outside publishes by an Update that answers and refuses. A Signal has no response, so the workflow could only drop a divergent retry. The publish Update answers with the batch's position, so append() returns a cursor, and refuses a conflicting repeat with a typed error before accepting it. A producer falls back to the Signal on a workflow whose worker predates the Update; the Action cost is one per batch either way. --- CHANGELOG.md | 6 +- .../streams/providers/workflow_streams.py | 377 +++++++++++++++--- tests/streams/test_streams_conformance.py | 7 +- .../streams/test_workflow_streams_provider.py | 192 ++++++++- 4 files changed, 506 insertions(+), 76 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index a7913ec4e..f43009cbb 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -56,7 +56,11 @@ to include examples, links to docs, or any other relevant information. worker plugin, so a workflow reads and publishes through `temporalio.contrib.workflow_streams` without naming it. Records are the `StreamRecord` proto inside the shipped item payload, and a handle without a - run id follows continue-as-new run by run. + run id follows continue-as-new run by run and a reset into the run reset to. + An outside publish is an Update that answers with the batch's position and + refuses a conflicting repeat, falling back to the shipped Signal on a + workflow whose worker predates it. A workflow's activity keeps its own + streams in the workflow's log under `activity//`. ### Changed diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index e700ba959..0ec67b870 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -14,23 +14,31 @@ ``StreamRecord`` proto and its encoding is ``binary/plain``. The shipped code stores and returns that ``Payload`` untouched, so the body's own encoding never meets the transport. -- Producer identity dedupes through the shipped publisher state: the - publisher id is ``producer#attempt`` and every publish Signal carries the - sequence its records end at, so a retried batch drops and a new attempt - passes. - - What this transport cannot keep is the rest of that rule. A publish is a - Signal, which has no response, so the dedupe decision is taken in the - workflow and there is nowhere to report it. A repeat that carries - *different* content at a sequence the log already holds is therefore - dropped rather than refused with - :class:`temporalio.streams.StreamProducerError`, the way the memory and - native providers refuse it. Raising in the Signal handler is not an - alternative: it would fail the Workflow Task on every replay and the - caller would still learn nothing. A caller that needs a divergent retry to - be caught wants a provider whose append is a request and a response. -- ``append()`` returns ``None``. The Signal transport learns positions at - read time, so a caller that wants to follow from now asks ``latest()``. +- An outside publish is an Update, ``__temporal_streams_publish``, carrying + the shipped publish Signal's input. Producer identity dedupes through the + shipped publisher state, the publisher id being ``producer#attempt`` and + the sequence where the batch's records end, and the Update adds what a + Signal has no room for: a response. The workflow answers with the run and + offset the batch landed at, so ``append()`` returns a cursor, and it + refuses a repeat that carries *different* content at a sequence it holds, + or one behind its most recent, before accepting the Update, so the caller + gets :class:`temporalio.streams.StreamProducerError` and the log takes + nothing. Content is compared by a hash the workflow keeps per producer. + The Update's id is derived from producer, sequence and content, so a + retry after a lost reply is answered by the server from the first + outcome. The Action cost is unchanged: one Update per append batch in + place of one Signal. + + The Signal stays as the transport a worker that predates the Update + serves. A producer that is told twice, across a task boundary, that the + workflow has no publish Update falls back to it for the rest of its life + and returns ``None`` from ``append()``, since a Signal learns positions at + read time; ``publish_transport="signal"`` on the provider picks it from + the start. Records land in the same log either way, so a log written by + Signals reads the same. The shipped Signal handler's dedupe is one table + per publisher across topics, so on that transport one identity writing + two topics at the same sequence has its second batch dropped; the Update + keeps one per producer and topic, as the other providers do. - A log belongs to one run and is not carried across continue-as-new, so a cursor names the run as well as the offset. A handle without a run id reads run after run: each log through the poll Update while its run is @@ -57,10 +65,12 @@ import asyncio import base64 +import hashlib import logging from collections.abc import AsyncGenerator +from dataclasses import dataclass from datetime import timedelta -from typing import Any, Generic, TypeVar +from typing import Any, Generic, Literal, TypeVar from google.protobuf.message import DecodeError @@ -88,11 +98,13 @@ WorkflowStream, ) from temporalio.converter import PayloadConverter +from temporalio.exceptions import ApplicationError from temporalio.service import RPCError, RPCStatusCode from temporalio.streams._errors import ( StreamCursorError, StreamError, StreamNotFoundError, + StreamProducerError, StreamUnsupportedError, ) from temporalio.streams._provider import ReadSource, WriteSink @@ -117,6 +129,8 @@ from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND __all__ = [ + "PRODUCER_CONFLICT_ERROR_TYPE", + "PublishTransport", "WorkflowStreamsActivityHandle", "WorkflowStreamsHandle", "WorkflowStreamsProducer", @@ -125,7 +139,18 @@ T = TypeVar("T") +PublishTransport = Literal["update", "signal"] +"""How an outside producer's batches reach the workflow.""" + +PRODUCER_CONFLICT_ERROR_TYPE = "StreamProducerConflict" +"""The ``ApplicationError.type`` the publish Update refuses a producer conflict with. + +Public so a caller driving the Update itself can tell the refusal from any +other failure. +""" + _PROVIDER = "workflow_streams" +_PUBLISH_UPDATE = "__temporal_streams_publish" _TAIL_QUERY = "__temporal_streams_tail" _LATEST_QUERY = "__temporal_streams_latest" _START_QUERY = "__temporal_streams_start" @@ -146,6 +171,17 @@ logger = logging.getLogger(__name__) +def _run_is_closing(error: RPCError) -> bool: + """Whether the server refused an Update because the run is completing. + + The window between a run deciding to close, or continue as new, and its + close being recorded; the next attempt learns how it closed. + """ + return ( + error.status == RPCStatusCode.FAILED_PRECONDITION and "closing" in error.message + ) + + def _require_topic(topic: str) -> None: if not topic: raise ValueError("topic must not be empty") @@ -206,6 +242,51 @@ def _entry_data(payload: Payload) -> str: return base64.b64encode(payload.SerializeToString()).decode("ascii") +def _content(publish: PublishInput) -> str: + """A digest of one batch's content, length-delimited so a resplit cannot collide.""" + digest = hashlib.sha256() + for entry in publish.items: + for part in (entry.topic.encode(), entry.data.encode("ascii")): + digest.update(len(part).to_bytes(8, "big")) + digest.update(part) + return digest.hexdigest() + + +def _producer_key(publish: PublishInput) -> tuple[str, str]: + # A batch is one topic's, so its first entry names the stream. + topic = publish.items[0].topic if publish.items else "" + return publish.publisher_id, topic + + +def _decode_payload(data: str) -> Payload: + return Payload.FromString(base64.b64decode(data)) + + +def _publish_id(publish: PublishInput, content: str) -> str: + # One id per (producer, sequence, content): a retry after a lost reply + # is answered from the first outcome without reaching the workflow, and a + # divergent repeat is a new Update the workflow gets to refuse. + key = f"{publish.publisher_id}\0{publish.sequence}\0{content}".encode() + return f"streams-publish-{hashlib.sha256(key).hexdigest()[:32]}" + + +@dataclass +class _PublishResult: + """The publish Update's answer: where the batch's last record sits.""" + + run_id: str + last_offset: int + + +@dataclass +class _Held: + """What the workflow keeps of a producer's most recent accepted batch.""" + + sequence: int + content: str + last_offset: int + + class _InstanceStream: """A view over the shipped stream object of the running workflow instance. @@ -215,6 +296,12 @@ class _InstanceStream: def __init__(self, stream: WorkflowStream | None = None) -> None: self.stream = WorkflowStream() if stream is None else stream + # Per producer and topic, the most recent batch taken by the publish + # Update, so a repeat is answered and a divergent one refused. A + # producer's identity is one per stream, as on every provider, so + # the same identity on two topics is two producers. Workflow state, + # rebuilt on replay like the log itself. + self._producers: dict[tuple[str, str], _Held] = {} if workflow.get_query_handler(_TAIL_QUERY) is None: # The poll Update stops answering once the workflow is closing, # and a reader between polls at that moment would lose what the @@ -225,6 +312,58 @@ def __init__(self, stream: WorkflowStream | None = None) -> None: workflow.set_query_handler(_LATEST_QUERY, self._latest) if workflow.get_query_handler(_START_QUERY) is None: workflow.set_query_handler(_START_QUERY, self._start) + if workflow.get_update_handler(_PUBLISH_UPDATE) is None: + workflow.set_update_handler( + _PUBLISH_UPDATE, self._publish, validator=self._validate_publish + ) + + def _validate_publish(self, publish: PublishInput) -> None: + """Refuse a conflicting batch before the Update is accepted. + + Refused here rather than in the handler so the refusal writes no + event: the caller learns of it and the log is untouched. + + Raises: + ApplicationError: Typed ``PRODUCER_CONFLICT_ERROR_TYPE``. The + batch repeats this producer's most recent sequence with + other content, or names a sequence behind it. + """ + held = self._producers.get(_producer_key(publish)) + if held is None: + return + if publish.sequence < held.sequence: + raise ApplicationError( + f"producer {publish.publisher_id!r} already wrote past sequence " + f"{publish.sequence}; only its most recent append can be repeated", + type=PRODUCER_CONFLICT_ERROR_TYPE, + ) + if publish.sequence == held.sequence and _content(publish) != held.content: + raise ApplicationError( + f"producer {publish.publisher_id!r} repeated sequence " + f"{publish.sequence} with different content", + type=PRODUCER_CONFLICT_ERROR_TYPE, + ) + + def _publish(self, publish: PublishInput) -> _PublishResult: + """Take one outside batch into the log and answer where it landed. + + A repeat of the producer's most recent batch, which the validator + let through because its content matches, is answered with the + original position and writes nothing. The Update keeps its own + table rather than the shipped Signal handler's, which is one per + publisher across topics; a producer stays on one transport for its + life, so the two never judge the same batch. + """ + run_id = workflow.info().run_id + key = _producer_key(publish) + held = self._producers.get(key) + if held is not None and publish.sequence == held.sequence: + return _PublishResult(run_id=run_id, last_offset=held.last_offset) + for entry in publish.items: + self.stream.topic(entry.topic).publish(_decode_payload(entry.data)) + last = self.stream.next_offset - 1 + self._producers[key] = _Held(publish.sequence, _content(publish), last) + return _PublishResult(run_id=run_id, last_offset=last) def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: """One page of ``topic``'s items at or past ``from_offset``. @@ -315,6 +454,10 @@ def _instance() -> _InstanceStream: # id: the SDK rebuilds an evicted workflow from history as a new object, # and a map would hand that object the stale log with its unregistered # handlers and the records of a task that failed. + handler = workflow.get_update_handler(_PUBLISH_UPDATE) + registered = getattr(handler, "__self__", None) + if isinstance(registered, _InstanceStream): + return registered stream = _registered_stream() return _InstanceStream() if stream is None else _InstanceStream(stream) @@ -433,18 +576,19 @@ async def on_workflow_finish(self) -> None: class WorkflowStreamsProducer(Generic[T]): - """Appends by sending the shipped publish Signal directly. + """Appends through the publish Update, or the shipped publish Signal. Direct rather than through ``WorkflowStreamClient`` because the interface owns the publisher identity: it must be ``producer#attempt`` for the shipped dedupe to drop a retry and pass a new generation, and the client would use its own random id. - Sequences are committed only after the server accepted the Signal. A - batch whose Signal raised stays pending and goes out again under the - same signal sequence, either when the caller retries the same values or - ahead of whatever the caller sends next, so an ambiguous failure writes - the batch once and loses nothing. + Sequences are committed only after the server accepted the batch. A + batch whose send raised stays pending and goes out again under the same + sequence, either when the caller retries the same values or ahead of + whatever the caller sends next, so an ambiguous failure writes the batch + once and loses nothing. A batch the workflow refused is dropped from the + pending slot: the refusal is the answer. """ def __init__( @@ -456,12 +600,16 @@ def __init__( attempt: int, *, item_topic: str | None = None, + transport: PublishTransport = "update", + retry_cooldown: timedelta = timedelta(milliseconds=100), ) -> None: """Bind this producer to ``topic`` on the workflow behind ``handle``. ``item_topic`` is the name the log files the records under when it differs from the name the records carry, as an activity's reserved - topics do. + topics do. ``transport`` is how batches travel; ``"update"`` falls + back to ``"signal"`` on a workflow that has no publish Update, after + one retry ``retry_cooldown`` later. """ self._handle = handle self._converter = converter @@ -469,10 +617,18 @@ def __init__( self._item_topic = topic if item_topic is None else item_topic self._producer_id = producer_id self._attempt = attempt + self._transport: PublishTransport = transport + self._retry_cooldown = retry_cooldown # One-based, because zero on the wire says the producer does not # number its records and this one does. self._sequence = 1 self._pending: tuple[list[PublishEntry], int] | None = None + self._last: Cursor | None = None + + @property + def transport(self) -> PublishTransport: + """How this producer's batches travel now: the Signal once it fell back.""" + return self._transport @property def producer_id(self) -> str: @@ -493,16 +649,23 @@ def _publisher_id(self) -> str: ) async def append(self, *values: T) -> Cursor | None: - """Append ``values`` through the shipped publish Signal. + """Append ``values`` and return the cursor of the batch's last record. + + A repeat of this producer's most recent append with the same content + is answered with the position the original landed at; an empty call + returns the position of the last record this producer wrote, or + ``None`` before its first. ``None`` is also the answer on the Signal + transport, which learns positions at read time, so a caller that + wants to follow from now asks :meth:`WorkflowStreamsHandle.latest`. - Always ``None``: this transport learns positions at read time, so a - caller that wants to follow from now asks - :meth:`WorkflowStreamsHandle.latest`. + Raises: + StreamProducerError: The workflow holds this sequence with other + content, or the producer already wrote past it. + StreamNotFoundError: The workflow does not exist or has closed. """ if not values: - return None - await self._send([(RecordKind.DATA, value) for value in values]) - return None + return self._last + return await self._send([(RecordKind.DATA, value) for value in values]) async def finish(self) -> None: """Write ``FINISH`` for this producer on this topic.""" @@ -529,33 +692,104 @@ def _entries( ) return entries, sequence - async def _send(self, batch: list[tuple[RecordKind, Any]]) -> None: + async def _send(self, batch: list[tuple[RecordKind, Any]]) -> Cursor | None: entries, next_sequence = self._entries(batch) if self._pending is not None and self._pending[0] != entries: - # The caller moved on from a batch whose Signal raised. It goes - # first, under the signal sequence it already had, so a copy the - # server did accept is dropped and one it never saw lands. The - # new batch is then renumbered behind it. - await self._signal(*self._pending) + # The caller moved on from a batch whose send raised. It goes + # first, under the sequence it already had, so a copy the server + # did accept is answered or dropped and one it never saw lands. + # The new batch is then renumbered behind it. + await self._deliver(*self._pending) entries, next_sequence = self._entries(batch) - await self._signal(entries, next_sequence) + return await self._deliver(entries, next_sequence) - async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None: + async def _deliver( + self, entries: list[PublishEntry], next_sequence: int + ) -> Cursor | None: # The dedupe sequence is where this producer's records end, not how - # many signals it has sent. The two differ once a retry batches its + # many batches it has sent. The two differ once a retry batches its # records differently from the send it is repeating, and a counter of - # signals then either drops a batch of new records or lets records + # batches then either drops a batch of new records or lets records # that are already there through a second time. self._pending = (entries, next_sequence) + publish = PublishInput( + items=entries, publisher_id=self._publisher_id, sequence=next_sequence + ) + if self._transport == "signal": + landed = await self._signal(publish) + else: + landed = await self._update(publish) + self._sequence = next_sequence + self._pending = None + if landed is not None: + self._last = landed + return landed + + async def _update(self, publish: PublishInput) -> Cursor | None: + content = _content(publish) + retried = False + while True: + try: + answer = await self._handle.execute_update( + _PUBLISH_UPDATE, + publish, + id=_publish_id(publish, content), + result_type=_PublishResult, + ) + except WorkflowUpdateFailedError as error: + cause = error.cause + cause_type = getattr(cause, "type", None) + if cause_type == PRODUCER_CONFLICT_ERROR_TYPE: + # Final for this batch: there is nothing to send again. + self._pending = None + raise StreamProducerError(str(cause)) from error + if cause_type == _UPDATE_OUTLIVED_RUN: + # Accepted, then the run closed before the handler ran: + # the same answer a Signal to a closed run gets. + raise StreamNotFoundError( + f"workflow {self._handle.id!r} closed before taking the " + "batch, so its stream cannot be appended to" + ) from error + if _HANDLER_NOT_FOUND not in str(cause): + raise StreamError( + f"the publish update on workflow {self._handle.id!r} " + f"failed: {error}" + ) from error + if not retried: + # A rejection comes back with the completion of the task + # that made it, so a retry reaches a later task, and the + # start hook registers the handler on the first one. + retried = True + await asyncio.sleep(self._retry_cooldown.total_seconds()) + continue + # Rejected across a task boundary: the workflow's worker + # predates the publish Update. Its Signal takes the batch, + # and every later one from this producer. + logger.info( + "workflow %r has no publish update; appending by Signal", + self._handle.id, + ) + self._transport = "signal" + return await self._signal(publish) + except RPCError as error: + if _run_is_closing(error): + # Where a Signal would be carried to the successor, the + # Update is refused; the retry reaches the run that + # takes over, since the handle follows the chain. + await asyncio.sleep(self._retry_cooldown.total_seconds()) + continue + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._handle.id!r} was not found, so its stream " + "cannot be appended to" + ) from error + raise + return _cursor(answer.run_id, answer.last_offset) + + async def _signal(self, publish: PublishInput) -> Cursor | None: + # Always None: a Signal has no response to carry the position. try: - await self._handle.signal( - PUBLISH_SIGNAL_NAME, - PublishInput( - items=entries, - publisher_id=self._publisher_id, - sequence=next_sequence, - ), - ) + await self._handle.signal(PUBLISH_SIGNAL_NAME, publish) except RPCError as error: if error.status == RPCStatusCode.NOT_FOUND: raise StreamNotFoundError( @@ -563,8 +797,7 @@ async def _signal(self, entries: list[PublishEntry], next_sequence: int) -> None "cannot be appended to" ) from error raise - self._sequence = next_sequence - self._pending = None + return None class WorkflowStreamsHandle: @@ -576,12 +809,14 @@ def __init__( workflow_id: str, run_id: str | None, poll_cooldown: timedelta, + publish_transport: PublishTransport = "update", ) -> None: """Address ``workflow_id``'s log, pinned to ``run_id`` when one is given.""" self._client = client self._workflow_id = workflow_id self._run_id = run_id self._poll_cooldown = poll_cooldown + self._publish_transport: PublishTransport = publish_transport self._converter = client.data_converter.payload_converter def read( @@ -782,6 +1017,11 @@ async def _poll( f"the poll update on workflow {self._workflow_id!r} failed: {error}" ) from error except RPCError as error: + if _run_is_closing(error): + # Continuing as new, or completing: the next attempt + # learns how it closed, the way a draining poll does. + await asyncio.sleep(cooldown) + continue # The run closed and its poll Update went with it, or the # workflow does not exist; the caller describes to tell which. if error.status != RPCStatusCode.NOT_FOUND: @@ -1006,6 +1246,8 @@ def producer( producer_id, attempt, item_topic=wire_topic, + transport=self._publish_transport, + retry_cooldown=self._poll_cooldown, ) @@ -1034,9 +1276,10 @@ def __init__( run_id: str | None, activity_id: str, poll_cooldown: timedelta, + publish_transport: PublishTransport = "update", ) -> None: """Address ``activity_id``'s streams inside ``workflow_id``'s log.""" - super().__init__(client, workflow_id, run_id, poll_cooldown) + super().__init__(client, workflow_id, run_id, poll_cooldown, publish_transport) self._activity_id = activity_id def _wire_topic(self, topic: str) -> str: @@ -1105,15 +1348,28 @@ class WorkflowStreamsProvider(ProviderPlugin): """The provider over the shipped Workflow Streams transport.""" def __init__( - self, *, poll_cooldown: timedelta = timedelta(milliseconds=100) + self, + *, + poll_cooldown: timedelta = timedelta(milliseconds=100), + publish_transport: PublishTransport = "update", ) -> None: """Create the provider. Args: poll_cooldown: How long an outside reader that is caught up waits - between polls. Backlogs drain at full speed regardless. + between polls. Backlogs drain at full speed regardless. Also + how long a producer waits before retrying a publish Update + the workflow's first task rejected. + publish_transport: How an outside producer's batches reach the + workflow. ``"update"``, the default, answers each batch with + its position and refuses a conflicting repeat; on a workflow + whose worker predates the publish Update the producer falls + back to the Signal by itself. ``"signal"`` is the transport + of the first release and skips the detection. Both cost one + Action per batch. """ self._poll_cooldown = poll_cooldown + self._publish_transport: PublishTransport = publish_transport def workflow_provider(self) -> _WSWorkflowProvider: """The workflow half, over the running instance's shipped stream object.""" @@ -1123,7 +1379,9 @@ def get_stream_handle( self, client: Client, workflow_id: str, *, run_id: str | None = None ) -> WorkflowStreamsHandle: """A handle on ``workflow_id``'s log; without ``run_id`` it follows the chain.""" - return WorkflowStreamsHandle(client, workflow_id, run_id, self._poll_cooldown) + return WorkflowStreamsHandle( + client, workflow_id, run_id, self._poll_cooldown, self._publish_transport + ) def get_activity_stream_handle( self, @@ -1151,7 +1409,12 @@ def get_activity_stream_handle( "workflow hosts this activity's" ) return WorkflowStreamsActivityHandle( - client, workflow_id, run_id, activity_id, self._poll_cooldown + client, + workflow_id, + run_id, + activity_id, + self._poll_cooldown, + self._publish_transport, ) async def close(self) -> None: diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 942595170..e8959dfc7 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -166,15 +166,12 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: TruncatingStreamHost.truncate, args=[topic, keep] ) + # A publish is an Update, so the workflow answers with the position + # and refuses a divergent repeat: both capabilities hold here. yield ProviderCase( "workflow_streams", provider, client, - reports_positions=False, - # A publish is a Signal, so the dedupe decision is taken in the - # workflow with nowhere to report it. See the module docstring of - # the provider. - detects_divergent_retries=False, host=host, task_queue=worker.task_queue, truncate=truncate, diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index f37fccd8a..e5d713971 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -1,11 +1,12 @@ """Conformance for the workflow_streams provider on the test environment's server. Runs the interface loop over the shipped Option 0 transport: an outside -producer appends through the publish Signal, the workflow reads and -republishes through its own state, and an outside reader follows the poll -Update while the run is open and the tail Query once it has closed. The -outside-surface cases shared by every provider run from -``test_streams_conformance``; this file covers what the transport adds. +producer appends through the publish Update, or the shipped publish Signal +where a workflow has no Update, the workflow reads and republishes through +its own state, and an outside reader follows the poll Update while the run is +open and the tail Query once it has closed. The outside-surface cases shared +by every provider run from ``test_streams_conformance``; this file covers +what the transport adds. """ from __future__ import annotations @@ -26,17 +27,21 @@ Client, WorkflowExecutionStatus, WorkflowQueryFailedError, + WorkflowUpdateFailedError, ) from temporalio.contrib.workflow_streams import PublishInput, WorkflowStream from temporalio.converter import DataConverter +from temporalio.exceptions import ApplicationError from temporalio.service import RPCError, RPCStatusCode from temporalio.streams import ( BEGINNING, END, + Cursor, RecordKind, StreamCursorError, StreamError, StreamNotFoundError, + StreamProducerError, StreamUnsupportedError, Supersession, ) @@ -47,6 +52,7 @@ WorkflowStreamsProducer, WorkflowStreamsProvider, _InstanceStream, + _PublishResult, ) from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND from tests.helpers import new_worker @@ -150,21 +156,32 @@ async def test_retried_producer_dedupes_and_new_attempt_supersedes( stream = provider.get_stream_handle(client, workflow_id) first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) - # Positions are learnt at read time on this transport. - assert await first.append({"n": 1}) is None - # The retry of the same attempt re-sends its first batch. + # The publish Update answers with where the batch landed. + landed = await first.append({"n": 1}) + assert landed is not None + # The retry of the same attempt re-sends its first batch and is + # answered with the same position. retry = stream.producer(topic=INPUTS, producer_id="model", attempt=1) - await retry.append({"n": 1}) + assert await retry.append({"n": 1}) == landed second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) await second.append({"n": 2}) records = await take(stream.read(topic=INPUTS, result_type=dict), 3) assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 + assert records[0].cursor == landed assert records[1].kind is RecordKind.SUPERSEDED assert records[1].supersession == Supersession("model", 1, 2) assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 assert all(r.topic == INPUTS for r in records) + # A sequence behind the producer's most recent one is stale, and a + # refusal is typed rather than a silent drop. + await second.append({"n": 3}) + stale = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + with pytest.raises(StreamProducerError, match="most recent"): + await stale.append({"n": "other"}) + assert stale.transport == "update" + await second.finish() await handle.signal(EchoLoop.release) await handle.result() @@ -668,6 +685,48 @@ async def test_a_live_read_without_a_run_id_follows_a_reset( ] +@workflow.defn +class ShippedOnly: + """Holds the shipped stream object alone, as a workflow on a worker without the provider.""" + + def __init__(self) -> None: + self._stream = WorkflowStream() + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def test_a_workflow_without_the_publish_update_is_appended_to_by_signal( + client: Client, provider: WorkflowStreamsProvider +): + # The worker has no provider, so the workflow serves the shipped Signal + # and nothing else. The producer learns that from two rejections and + # falls back, and the batches land in the shipped log all the same. + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, ShippedOnly) as worker: + handle = await client.start_workflow( + ShippedOnly.run, id=workflow_id, task_queue=worker.task_queue + ) + producer = provider.get_stream_handle(client, workflow_id).producer( + topic=INPUTS, producer_id="model", attempt=1 + ) + assert await producer.append({"n": 1}) is None + assert producer.transport == "signal" + await producer.append({"n": 2}, {"n": 3}) + assert ( + await handle.query("__temporal_workflow_stream_offset", result_type=int) + == 3 + ) + await handle.signal(ShippedOnly.release) + await handle.result() + + async def test_a_running_workflow_without_the_provider_fails_the_read_clearly( client: Client, provider: WorkflowStreamsProvider ): @@ -735,18 +794,25 @@ def _sequences(sent: list[PublishInput]) -> list[tuple[int, list[int]]]: ] -def _producer(handle: _FlakyHandle) -> WorkflowStreamsProducer: +def _producer(handle: Any, transport: Any = "signal") -> WorkflowStreamsProducer: return WorkflowStreamsProducer( - handle, DataConverter.default.payload_converter, INPUTS, "model", 1 + handle, + DataConverter.default.payload_converter, + INPUTS, + "model", + 1, + transport=transport, + retry_cooldown=timedelta(0), ) async def test_a_retried_append_after_an_ambiguous_failure_writes_once(): + # The Signal transport: what a producer that fell back to it does. handle = _FlakyHandle() producer = _producer(handle) with pytest.raises(ConnectionResetError): await producer.append({"n": 1}) - await producer.append({"n": 1}) + assert await producer.append({"n": 1}) is None await producer.append({"n": 2}) # The retry carries the same signal sequence and the same record # sequence as the failed send, so the shipped dedupe drops the copy; the @@ -766,6 +832,99 @@ async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) +class _UpdateHandle: + """A workflow handle serving the publish Update as the server and workflow would. + + Answers a repeated update id from the first outcome, the way the server + does, positions each new batch at the head of a log, the way the + workflow does, and loses the reply of the first call when told to. + Without a handler it rejects every Update the way a workflow whose + worker predates it does, and takes Signals. + """ + + id = "update" + + def __init__(self, *, lose_first_reply: bool = False, handler: bool = True): + self.sent: list[tuple[str, PublishInput]] = [] + self.signalled: list[PublishInput] = [] + self._outcomes: dict[str, _PublishResult] = {} + self._head = 0 + self._lose = lose_first_reply + self._handler = handler + + async def execute_update( + self, name: str, arg: PublishInput, *, id: str, result_type: Any + ) -> _PublishResult: + del name, result_type + self.sent.append((id, arg)) + if not self._handler: + raise WorkflowUpdateFailedError( + ApplicationError( + f"Update handler for 'x' {QUERY_HANDLER_NOT_FOUND}, known updates: []" + ) + ) + answer = self._outcomes.get(id) + if answer is None: + answer = _PublishResult("run", self._head + len(arg.items) - 1) + self._head += len(arg.items) + self._outcomes[id] = answer + if self._lose: + self._lose = False + raise ConnectionResetError("the server took the update, the reply was lost") + return answer + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.signalled.append(arg) + + +def _at(offset: int) -> Cursor: + return Cursor(f"workflow_streams:run:{offset}") + + +async def test_a_retried_update_after_a_lost_reply_is_answered_from_its_id(): + handle = _UpdateHandle(lose_first_reply=True) + producer = _producer(handle, "update") + assert await producer.append() is None + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + # The retry is the same Update: same producer, sequence and content make + # the same id, so the server answers with the outcome it already holds + # and the log takes the batch once. The batch after it is a new one. + assert await producer.append({"n": 1}) == _at(0) + assert await producer.append({"n": 2}) == _at(1) + assert await producer.append() == _at(1) + ids = [id for id, _ in handle.sent] + assert ids[0] == ids[1] != ids[2] + assert _sequences([arg for _, arg in handle.sent]) == [(2, [1]), (2, [1]), (3, [2])] + + +async def test_a_batch_whose_update_failed_goes_out_before_the_next_one(): + handle = _UpdateHandle(lose_first_reply=True) + producer = _producer(handle, "update") + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + # The pending batch lands first, at offset 0, so the new one follows it. + assert await producer.append({"n": 2}, {"n": 3}) == _at(2) + await producer.finish() + sent = [arg for _, arg in handle.sent] + assert _sequences(sent) == [(2, [1]), (2, [1]), (4, [2, 3]), (5, [4])] + assert _wires(sent[-1])[0].kind == int(RecordKind.FINISH) + + +async def test_a_producer_falls_back_to_the_signal_without_a_publish_update(): + handle = _UpdateHandle(handler=False) + producer = _producer(handle, "update") + # Rejected twice across a task boundary means the workflow's worker + # predates the Update; the batch goes by Signal, and so does every later + # one, without asking again. + assert await producer.append({"n": 1}) is None + assert producer.transport == "signal" + assert await producer.append({"n": 2}) is None + assert len(handle.sent) == 2 + assert _sequences(handle.signalled) == [(2, [1]), (3, [2])] + + class _Description: run_id = "the-only-run" status = WorkflowExecutionStatus.COMPLETED @@ -955,7 +1114,14 @@ async def test_a_workflow_reader_starts_at_end_or_the_newest_records( result = asyncio.ensure_future(handle.result()) if start == "end": for _ in range(150): - await producer.append({"n": "new"}) + try: + await producer.append({"n": "new"}) + except StreamNotFoundError: + # The Update returns once the workflow took the + # batch, and reading it is what closes the workflow, + # so the next append can find it gone before the + # result has been noticed. + break done, _ = await asyncio.wait({result}, timeout=0.2) if done: break From 3db4601b2cea6e7bffb8025f0ed800aa33dd42b2 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:42:09 -0700 Subject: [PATCH 29/32] Absorbed refs, the body hash and the standalone refusal on Workflow Streams. Handles name their stream as a StreamRef and refuse close(), standalone streams are refused with a housing workflow noted as the option, and each record carries the plaintext hash of its body, which the workflow now matches a repeated batch by. Bodies meet the codec and external storage at the transport's envelope, since the workflow thread reads them from state. --- .../streams/providers/workflow_streams.py | 116 +++++++++++++++++- tests/streams/conftest.py | 5 + tests/streams/test_streams_conformance.py | 13 +- .../streams/test_workflow_streams_provider.py | 108 +++++++++++++++- 4 files changed, 237 insertions(+), 5 deletions(-) diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 0ec67b870..2b7036941 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -58,7 +58,27 @@ provider reserves ``activity/`` in the owner's map. The record itself carries the plain name. A standalone activity has no workflow to host a log, so its streams are refused, and a workflow's own topic may not start - with the reserved prefix. + with the reserved prefix. A standalone stream, one with no owner at all, + is refused too: there is no workflow whose state could hold it. A housing + workflow per stream id is the design option for that, not built here. +- A record's body meets the client's data converter at the transport's + envelope rather than one record at a time. A batch travels as a Signal or + Update argument and comes back as an Update or Query result, and the SDK + runs the payload codec and external storage over those the way it does + over every payload it sends, off the workflow thread, so a batch above the + storage threshold is offloaded as a claim and a codec protects it in + History. The bodies inside are left as the payload converter produced + them, because the workflow thread reads them straight out of its state + and could not decode a codec's output or redeem a claim there. So the + worker's converter has to match the clients', as it does for every other + payload, and a client with a converter of its own cannot read another's + records. What each record does carry is the plaintext hash of its body + under ``temporal.io/content-hash``, stamped by the producer before the + envelope is encoded, and the workflow matches a repeated batch by those + hashes rather than by the bytes. +- A handle names its stream as a :class:`temporalio.streams.StreamRef` with + ``ref()``, a workflow's or an activity's, and ``close()`` refuses, since + an owned stream ends with its owner. """ from __future__ import annotations @@ -100,6 +120,7 @@ from temporalio.converter import PayloadConverter from temporalio.exceptions import ApplicationError from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import CONTENT_HASH_KEY, content_hash from temporalio.streams._errors import ( StreamCursorError, StreamError, @@ -116,6 +137,7 @@ StreamRecord, check_read_start, ) +from temporalio.streams._ref import StreamRef from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( RecordDecoder, @@ -243,10 +265,23 @@ def _entry_data(payload: Payload) -> str: def _content(publish: PublishInput) -> str: - """A digest of one batch's content, length-delimited so a resplit cannot collide.""" + """A digest of one batch's content, length-delimited so a resplit cannot collide. + + Each record counts by the plaintext hash stamped on it under + ``CONTENT_HASH_KEY``, so a body's encoding never enters the identity a + repeat is matched by; a record without a body, such as ``FINISH``, counts + by its bytes. + """ digest = hashlib.sha256() for entry in publish.items: - for part in (entry.topic.encode(), entry.data.encode("ascii")): + record = WireRecord.FromString(_decode_payload(entry.data).data) + stamped = record.metadata.get(CONTENT_HASH_KEY) + identity = ( + stamped.data + if stamped is not None + else record.SerializeToString(deterministic=True) + ) + for part in (entry.topic.encode(), identity): digest.update(len(part).to_bytes(8, "big")) digest.update(part) return digest.hexdigest() @@ -687,6 +722,16 @@ def _entries( sequence=sequence, ) sequence += 1 + if wire.HasField("body"): + # The plaintext hash the workflow matches a repeat by, taken + # as the converter produced the body, before the transport's + # codec meets the envelope. + wire.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload( + metadata={"encoding": _ENCODING}, + data=content_hash(wire.body).encode(), + ) + ) entries.append( PublishEntry(topic=self._item_topic, data=_entry_data(_wrap(wire))) ) @@ -1250,6 +1295,19 @@ def producer( retry_cooldown=self._poll_cooldown, ) + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this workflow's stream, pinned as this handle is.""" + return StreamRef.for_workflow( + self._workflow_id, run_id=self._run_id, topic=topic + ) + + async def close(self) -> None: + """Refuse: an owned stream ends with its owner, not by a caller.""" + raise ValueError( + "only a standalone stream can be closed; this handle is on an owned " + "stream, which ends when its workflow or activity does" + ) + class WorkflowStreamsActivityHandle(WorkflowStreamsHandle): """The streams one activity of a workflow owns, kept in that workflow's log. @@ -1285,6 +1343,15 @@ def __init__( def _wire_topic(self, topic: str) -> str: return _activity_topic(self._activity_id, topic) + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this activity's streams, through its workflow.""" + return StreamRef.for_activity( + self._activity_id, + workflow_id=self._workflow_id, + run_id=self._run_id, + topic=topic, + ) + async def _read( self, topic: str, @@ -1417,5 +1484,48 @@ def get_activity_stream_handle( self._publish_transport, ) + async def create_standalone_stream( + self, + client: Client, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> WorkflowStreamsHandle: + """Refused: this provider has no store for a stream without an owner. + + Every log here is a running workflow's state, served by that + workflow's handlers, and a standalone stream has no workflow. A + housing workflow started for the stream id, with the retention policy + as its state and ``close()`` as a Signal, would be one way to offer + it on this transport; it is a design option for a later round, not + built here. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the workflow_streams provider does not host standalone streams: every " + "log is a running workflow's state, and a stream without an owner has " + "no workflow" + ) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> WorkflowStreamsHandle: + """Refused: this provider has no store for a stream without an owner. + + See :meth:`create_standalone_stream`. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the workflow_streams provider does not host standalone streams: every " + "log is a running workflow's state, and a stream without an owner has " + "no workflow" + ) + async def close(self) -> None: """Nothing to release: the provider holds no connection of its own.""" diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 45e09bfd5..a9699ab94 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -15,6 +15,11 @@ 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", + "encodes_bodies: the case needs the outside path to run each body through " + "the client's data converter", + ) config.addinivalue_line( "markers", "standalone_activities: the case needs the streams of an activity outside " diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 6f9dd793e..ed3b27610 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -88,6 +88,9 @@ class ProviderCase: """``append()`` returns where the records landed.""" detects_divergent_retries: bool = True """``append()`` compares a repeat's content with what it already holds.""" + encodes_bodies: bool = True + """The outside path runs each body through the client's data converter, so a + client with a converter of its own reads and writes another's records.""" host: Callable[[str], Awaitable[None]] | None = None """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" task_queue: str | None = None @@ -249,11 +252,16 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: ) # A publish is an Update, so the workflow answers with the position - # and refuses a divergent repeat: both capabilities hold here. + # and refuses a divergent repeat: both capabilities hold here. Bodies + # meet the codec and external storage at the transport's envelope, + # which the worker's converter has to match, so a case whose client + # carries a converter of its own is skipped; see the provider's + # module docstring. yield ProviderCase( "workflow_streams", provider, client, + encodes_bodies=False, host=host, task_queue=worker.task_queue, truncate=truncate, @@ -270,6 +278,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, "detects_divergent_retries": lambda case: case.detects_divergent_retries, + "encodes_bodies": lambda case: case.encodes_bodies, "truncates": lambda case: case.truncate is not None, } @@ -722,6 +731,7 @@ async def test_an_owned_stream_cannot_be_closed_by_a_handle(case: ProviderCase): await stream.close() +@pytest.mark.encodes_bodies async def test_a_body_above_the_threshold_is_offloaded_and_read_back( case: ProviderCase, client: Client ): @@ -747,6 +757,7 @@ async def test_a_body_above_the_threshold_is_offloaded_and_read_back( @pytest.mark.detects_divergent_retries +@pytest.mark.encodes_bodies async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( case: ProviderCase, client: Client ): diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index e5d713971..320a8bf03 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -13,6 +13,7 @@ import asyncio import base64 +import dataclasses import uuid from collections.abc import AsyncIterator from datetime import timedelta @@ -30,11 +31,13 @@ WorkflowUpdateFailedError, ) from temporalio.contrib.workflow_streams import PublishInput, WorkflowStream -from temporalio.converter import DataConverter +from temporalio.converter import DataConverter, ExternalStorage from temporalio.exceptions import ApplicationError from temporalio.service import RPCError, RPCStatusCode from temporalio.streams import ( BEGINNING, + CONTENT_HASH_KEY, + DEFAULT_TOPIC, END, Cursor, RecordKind, @@ -42,10 +45,13 @@ StreamError, StreamNotFoundError, StreamProducerError, + StreamRef, StreamUnsupportedError, Supersession, + content_hash, ) from temporalio.streams._wire import WireRecord +from temporalio.streams.providers import workflow_streams from temporalio.streams.providers.workflow_streams import ( WorkflowStreamsActivityHandle, WorkflowStreamsHandle, @@ -1143,6 +1149,106 @@ async def test_a_standalone_activity_stream_is_refused( ) +async def test_a_standalone_stream_is_refused( + client: Client, provider: WorkflowStreamsProvider +): + # Every log is a running workflow's state; a stream with no owner has no + # workflow to live in. Both the create and the lookup say so. + with pytest.raises(StreamUnsupportedError, match="standalone"): + await provider.create_standalone_stream(client, "shared") + with pytest.raises(StreamUnsupportedError, match="standalone"): + provider.get_standalone_stream_handle(client, "shared") + + +async def test_a_handle_names_its_stream_as_a_ref( + client: Client, provider: WorkflowStreamsProvider +): + workflow_stream = provider.get_stream_handle(client, "wf", run_id="run-1") + assert workflow_stream.ref(topic=INPUTS) == StreamRef.for_workflow( + "wf", run_id="run-1", topic=INPUTS + ) + own = provider.get_activity_stream_handle(client, "act", workflow_id="wf") + assert own.ref(topic=TOKENS) == StreamRef.for_activity( + "act", workflow_id="wf", topic=TOKENS + ) + assert own.ref().topic == DEFAULT_TOPIC + # An owned stream ends with its owner; only a standalone one closes. + with pytest.raises(ValueError, match="standalone"): + await own.close() + + +def test_a_record_carries_the_plaintext_hash_the_workflow_dedupes_by(): + handle = _UpdateHandle() + producer = _producer(handle, "update") + entries, _ = producer._entries( + [(RecordKind.DATA, {"n": 1}), (RecordKind.FINISH, None)] + ) # pyright: ignore[reportPrivateUsage] + data, finish = _wires(PublishInput(items=entries)) + # A DATA record is stamped with the hash of its converted body, where the + # workflow can read it without the body; FINISH has nothing to hash. + assert data.metadata[CONTENT_HASH_KEY].data.decode() == content_hash(data.body) + assert CONTENT_HASH_KEY not in finish.metadata + + # The workflow's identity for a batch is those hashes, so a body whose + # bytes a codec changed is still the same batch, and a different value + # is not. + def batch(*values: dict) -> PublishInput: + made, _ = _producer(handle, "update")._entries( # pyright: ignore[reportPrivateUsage] + [(RecordKind.DATA, value) for value in values] + ) + return PublishInput(items=made, publisher_id="model#1", sequence=3) + + same, recoded, other = batch({"n": 1}), batch({"n": 1}), batch({"n": 2}) + record = _wires(recoded)[0] + record.body.data = b"\x00" + record.body.data + recoded.items[0].data = base64.b64encode( + Payload( + metadata={"encoding": b"binary/plain"}, data=record.SerializeToString() + ).SerializeToString() + ).decode("ascii") + assert workflow_streams._content(same) == workflow_streams._content(recoded) # pyright: ignore[reportPrivateUsage] + assert workflow_streams._content(same) != workflow_streams._content(other) # pyright: ignore[reportPrivateUsage] + + +async def test_external_storage_applies_at_the_envelope_the_workflow_reads_through( + client: Client, provider: WorkflowStreamsProvider +): + # Worker and clients share one converter, as a deployment's do. A batch + # above the threshold leaves as a claim on the Update argument, the + # worker redeems it, the workflow reads the value, and the outside + # reader gets the poll response the same way. + from tests.streams.test_streams_conformance import RecordingDriver + + driver = RecordingDriver() + converter = dataclasses.replace( + DataConverter.default, + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=512), + ) + config = client.config() + config["data_converter"] = converter + shared = Client(**config) + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(shared, EchoLoop, plugins=[provider]) as worker: + handle = await shared.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(shared, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + assert driver.stored == 0 + await producer.append({"n": 2, "blob": "x" * 4096}) + assert driver.stored == 1 + await producer.finish() + + records = await take(stream.read(topic=DECISIONS, result_type=dict), 3) + assert [r.value.get("echo") for r in records[:2]] == [1, 2] + assert records[2].kind is RecordKind.FINISH + assert driver.retrieved >= 1 + + await handle.signal(EchoLoop.release) + assert await handle.result() == 2 + + TOKENS = "tokens" From b9ee5456bd3e9780b94b9166d30d24151d4f61dc Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:08:23 -0700 Subject: [PATCH 30/32] Declared the standalone byte bound and age trim absent on Workflow Streams. The provider hosts no standalone streams, so neither retention policy exists here and the retention case skips with the rest of them. --- tests/streams/test_streams_conformance.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 5bd2efcd3..39c41fc8a 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -322,8 +322,10 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: task_queue=worker.task_queue, truncate=truncate, # Every log is a running workflow's state; a stream with no owner - # has no workflow to live in. + # has no workflow to live in, so neither of its policies exists. hosts_standalone_streams=False, + bounds_standalone_bytes=False, + trims_open_stream_by_age=False, ) for handle in hosts.values(): await handle.terminate() From cb6c7734ae9b2283ea9f1c089c964c2108ebafbd Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 18:54:47 -0700 Subject: [PATCH 31/32] Registered the provider's handlers before the first task's Updates run. The SDK handles a task's Signals and Updates ahead of the workflow function, so a publish or poll Update that arrived with the first task was rejected and, on a worker with no cache, on every rebuilt instance too. The start hook now runs when the instance's loop first turns, and Workflow Streams binds the shipped stream object lazily so a workflow's own stays. --- temporalio/streams/_provider.py | 9 +- .../streams/providers/workflow_streams.py | 111 ++++++++++++++---- temporalio/worker/_workflow.py | 28 ++++- tests/streams/test_stream_hooks.py | 42 ++++--- .../streams/test_workflow_streams_provider.py | 29 +++-- 5 files changed, 163 insertions(+), 56 deletions(-) diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py index 340cea87a..8982bf87b 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/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py index 2b7036941..8a5ed5daa 100644 --- a/temporalio/streams/providers/workflow_streams.py +++ b/temporalio/streams/providers/workflow_streams.py @@ -52,6 +52,11 @@ through the handler the shipped class registers on it. An evicted and rebuilt workflow gets its own, so a task that failed leaks nothing into the next attempt's log and a replayed run does not see records twice. + The provider's handlers are registered as the instance is initialised, + before the first task's Signals and Updates are applied, so a publish or + poll that arrives with that task finds them; the stream object is bound + on first use, adopting one the workflow built in its ``__init__`` or + constructing the provider's own. - An activity a workflow scheduled keeps its own streams inside that workflow's log, under the reserved topic ``activity//`` with ``%`` and ``/`` in the id percent-encoded, the way the native @@ -322,15 +327,28 @@ class _Held: last_offset: int -class _InstanceStream: - """A view over the shipped stream object of the running workflow instance. +class _Shipped: + """Constructs the shipped stream object, which insists on a caller named ``__init__``.""" + + def __init__(self) -> None: + self.stream = WorkflowStream() - A separate class because ``WorkflowStream`` insists on being constructed - from a method named ``__init__``. + +class _InstanceStream: + """The provider's handlers on one workflow instance, over the shipped stream object. + + Built when the instance is initialised, before the workflow's own + ``__init__`` and before any Signal or Update of the first task is + applied, so the handlers are found by whatever arrives with that task. + The shipped stream object is bound on first use rather than here: a + workflow migrating from the contrib feature constructs its own in + ``__init__``, which the shipped class refuses to do twice, so this + adopts that one when it exists and constructs the provider's own when + nothing has, from a handler that may write or from the workflow half. """ - def __init__(self, stream: WorkflowStream | None = None) -> None: - self.stream = WorkflowStream() if stream is None else stream + def __init__(self) -> None: + self._stream: WorkflowStream | None = None # Per producer and topic, the most recent batch taken by the publish # Update, so a repeat is answered and a divergent one refused. A # producer's identity is one per stream, as on every provider, so @@ -351,6 +369,38 @@ def __init__(self, stream: WorkflowStream | None = None) -> None: workflow.set_update_handler( _PUBLISH_UPDATE, self._publish, validator=self._validate_publish ) + if workflow.get_update_handler(POLL_UPDATE_NAME) is None: + # Stands in for the shipped poll handler until a stream object is + # bound, whose constructor then registers the real one over it. + workflow.set_update_handler( + POLL_UPDATE_NAME, self._poll, validator=self._validate_poll + ) + + @property + def held(self) -> WorkflowStream | None: + """The shipped stream object, if the workflow or a handler has bound one. + + Looks one up and never constructs, so a Query may ask. + """ + if self._stream is None: + self._stream = _registered_stream() + return self._stream + + @property + def stream(self) -> WorkflowStream: + """The shipped stream object, adopting the workflow's own or constructing the provider's.""" + held = self.held + if held is None: + held = self._stream = _Shipped().stream + return held + + def _validate_poll(self, payload: PollInput) -> None: + held = self.held + if held is not None: + held._validate_poll(payload) # pyright: ignore[reportPrivateUsage] + + async def _poll(self, payload: PollInput) -> PollResult: + return await self.stream._on_poll(payload) # pyright: ignore[reportPrivateUsage] def _validate_publish(self, publish: PublishInput) -> None: """Refuse a conflicting batch before the Update is accepted. @@ -394,9 +444,10 @@ def _publish(self, publish: PublishInput) -> _PublishResult: held = self._producers.get(key) if held is not None and publish.sequence == held.sequence: return _PublishResult(run_id=run_id, last_offset=held.last_offset) + stream = self.stream for entry in publish.items: - self.stream.topic(entry.topic).publish(_decode_payload(entry.data)) - last = self.stream.next_offset - 1 + stream.topic(entry.topic).publish(_decode_payload(entry.data)) + last = stream.next_offset - 1 self._producers[key] = _Held(publish.sequence, _content(publish), last) return _PublishResult(run_id=run_id, last_offset=last) @@ -407,11 +458,20 @@ def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: response has to fit the server's blob limit and a log the reader only wants one topic of can be much larger than that. """ + stream = self.held + if stream is None: + # Nothing has been bound, so nothing has been published. + return { + "items": [], + "next_offset": 0, + "more_ready": False, + "base_offset": 0, + } items: list[dict[str, Any]] = [] size = 0 - next_offset = self.stream.next_offset + next_offset = stream.next_offset more_ready = False - held = self.stream.items_from(from_offset) + held = stream.items_from(from_offset) for offset, item_topic, payload in held: if item_topic != topic: continue @@ -437,7 +497,10 @@ def _latest(self, topic: str) -> int: answer about one topic. Scanning here costs one Query rather than shipping the log to the caller to find the same thing. """ - for offset, item_topic, _ in reversed(self.stream.items_from(0)): + stream = self.held + if stream is None: + return -1 + for offset, item_topic, _ in reversed(stream.items_from(0)): if item_topic == topic: return offset return -1 @@ -447,7 +510,8 @@ def _start(self, topic: str, last_n: int) -> int: Zero ``last_n`` asks for the head, which is where ``END`` starts. """ - return start_offset(self.stream, topic, last_n) + stream = self.held + return 0 if stream is None else start_offset(stream, topic, last_n) def start_offset(stream: WorkflowStream, topic: str, last_n: int) -> int: @@ -493,8 +557,7 @@ def _instance() -> _InstanceStream: registered = getattr(handler, "__self__", None) if isinstance(registered, _InstanceStream): return registered - stream = _registered_stream() - return _InstanceStream() if stream is None else _InstanceStream(stream) + return _InstanceStream() class _WSReadSource: @@ -594,19 +657,23 @@ def open_writer(self, topic: str) -> WriteSink: return _WSWriteSink(self._own_stream(), topic) def on_workflow_start(self) -> None: - # Registered on the first task, because an outside reader can poll - # before workflow code has opened anything, and an Update with no - # handler yet is rejected rather than held. A poll that arrives in - # that first task still runs ahead of this hook; the reader retries it. - self._own_stream() + # Called as the instance is initialised, so the handlers exist before + # the first task's Updates are evaluated: an outside publish or poll + # that arrives with that task is served rather than rejected. The + # stream object itself is bound later, once the workflow's own + # __init__ has had its chance to construct one. + _instance() async def on_workflow_finish(self) -> None: # An Option 0 stream dies with its run, and a parked long-poll Update # would otherwise hold completion open. Same recipe the shipped - # feature documents before a return or a continue-as-new. - if self._stream is None: + # feature documents before a return or a continue-as-new. Asked of + # the instance rather than this object's cache, because a poll can + # bind the stream before workflow code touches it. + held = _instance().held + if held is None: return - self._stream.detach_pollers() + held.detach_pollers() await workflow.wait_condition(workflow.all_handlers_finished) diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 79f5ca6e8..f4001b32e 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -13,7 +13,7 @@ from dataclasses import dataclass from datetime import timezone from types import TracebackType -from typing import Any +from typing import Any, cast import temporalio.api.common.v1 import temporalio.bridge.proto.common @@ -41,6 +41,7 @@ Interceptor, WorkflowInboundInterceptor, WorkflowInterceptorClassInput, + WorkflowOutboundInterceptor, ) from ._workflow_instance import ( _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS, @@ -62,9 +63,14 @@ 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. + 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 @@ -72,10 +78,22 @@ class _StreamHooksInterceptor(WorkflowInboundInterceptor): 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 - provider.on_workflow_start() try: result = await self.next.execute_workflow(input) except GeneratorExit: diff --git a/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py index a547daf3e..7908dfc57 100644 --- a/tests/streams/test_stream_hooks.py +++ b/tests/streams/test_stream_hooks.py @@ -10,7 +10,7 @@ from __future__ import annotations import asyncio -from typing import Any +from typing import Any, cast import pytest @@ -58,6 +58,11 @@ def workflow_streams(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.""" @@ -65,6 +70,9 @@ class _Body(WorkflowInboundInterceptor): 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: @@ -84,6 +92,18 @@ async def _unused_run_fn() -> None: _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() @@ -99,40 +119,34 @@ def _evicting(fake: _Provider) -> None: async def test_the_finish_hook_runs_on_return(provider: _Provider): - assert ( - await _StreamHooksInterceptor(_Body("done")).execute_workflow(_INPUT) == "done" - ) + 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 _StreamHooksInterceptor(_Body(RuntimeError("boom"))).execute_workflow( - _INPUT - ) + 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 _StreamHooksInterceptor(_Body(error)).execute_workflow(_INPUT) + 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 _StreamHooksInterceptor(_Body(asyncio.CancelledError())).execute_workflow( - _INPUT - ) + 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 = _StreamHooksInterceptor(_Body(_PARK)).execute_workflow(_INPUT) + 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) @@ -143,9 +157,7 @@ async def test_the_finish_hook_does_not_run_when_the_coroutine_is_collected( 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 - ) + await (await _hooked(_Body(asyncio.CancelledError()))).execute_workflow(_INPUT) assert provider.calls == ["start"] diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py index 320a8bf03..218c55fdc 100644 --- a/tests/streams/test_workflow_streams_provider.py +++ b/tests/streams/test_workflow_streams_provider.py @@ -373,7 +373,7 @@ def items_from(self, offset: int) -> list[Any]: return [item for item in items if item[0] >= offset] instance = _InstanceStream.__new__(_InstanceStream) - instance.stream = _Log() # type: ignore[assignment] + instance._stream = _Log() # type: ignore[assignment] # pyright: ignore[reportPrivateUsage] seen: list[int] = [] offset, more = 0, True @@ -590,17 +590,24 @@ async def _reset_at_last_completed_task( if event.HasField("workflow_task_completed_event_attributes"): completion_id = event.event_id assert completion_id - answer = await client.workflow_service.reset_workflow_execution( - ResetWorkflowExecutionRequest( - namespace=client.namespace, - workflow_execution=WorkflowExecution( - workflow_id=workflow_id, run_id=run_id - ), - reason="re-run from the last completed task", - workflow_task_finish_event_id=completion_id, - request_id=uuid.uuid4().hex, + try: + answer = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=run_id + ), + reason="re-run from the last completed task", + workflow_task_finish_event_id=completion_id, + request_id=uuid.uuid4().hex, + ) ) - ) + except RPCError as error: + if error.status != RPCStatusCode.UNIMPLEMENTED: + raise + # The Java time-skipping test server has no reset; the real server + # does, and this never skips there. + pytest.skip("this test server does not implement ResetWorkflowExecution") return answer.run_id From 8ed252371eb7b36570178f679484abd6016d8860 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 18:57:54 -0700 Subject: [PATCH 32/32] Declared the byte-cap refusal absent on Workflow Streams. The provider hosts no standalone streams, so there is no byte cap to refuse an append past; the flag says so next to the other three. --- tests/streams/test_streams_conformance.py | 1 + 1 file changed, 1 insertion(+) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 1968a5080..993cba69b 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -332,6 +332,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: hosts_standalone_streams=False, bounds_standalone_bytes=False, trims_open_stream_by_age=False, + refuses_appends_past_byte_cap=False, ) for handle in hosts.values(): await handle.terminate()