From 246f3c0d30ca13e37369bfdd24aa4cd7359b9ab8 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:07:19 -0700 Subject: [PATCH 01/27] Added the workflow half of a server-side stream. A task's publishes on one stream are held until it completes, so they become one command and one History event however many there were, and a failed task publishes nothing. Delivered ranges are buffered because the server records a range as consumed and never sends it again. --- temporalio/worker/_workflow_instance.py | 224 ++++++++++++++++++ temporalio/workflow/__init__.py | 8 + temporalio/workflow/_context.py | 97 ++++++++ tests/worker/test_workflow_stream.py | 288 ++++++++++++++++++++++++ 4 files changed, 617 insertions(+) create mode 100644 tests/worker/test_workflow_stream.py diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 26f796e57..9b1e1ea0b 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -47,6 +47,7 @@ import temporalio.api.common.v1 import temporalio.api.enums.v1 import temporalio.api.sdk.v1 +import temporalio.api.stream.v1 import temporalio.bridge.proto.activity_result import temporalio.bridge.proto.child_workflow import temporalio.bridge.proto.common @@ -272,6 +273,128 @@ def create_instance(self, det: WorkflowInstanceDetails) -> WorkflowInstance: _ExceptionHandler: TypeAlias = Callable[[asyncio.AbstractEventLoop, _Context], Any] +# Match the server's per-batch limits. A record over its limit is refused where +# it is published, because a rejected command would be reissued on every +# replay; a task's records are split into commands that fit the batch limits. +_MAX_STREAM_RECORDS_PER_BATCH = 1000 +_MAX_STREAM_RECORD_BYTES = 1 << 20 +_MAX_STREAM_BATCH_BYTES = 2 << 20 + + +def _stream_batches( + records: Sequence[temporalio.api.stream.v1.StreamRecord], +) -> Iterator[list[temporalio.api.stream.v1.StreamRecord]]: + """Split one task's records for a stream into batches the server accepts.""" + batch: list[temporalio.api.stream.v1.StreamRecord] = [] + size = 0 + for record in records: + record_size = record.ByteSize() + if batch and ( + len(batch) >= _MAX_STREAM_RECORDS_PER_BATCH + or size + record_size > _MAX_STREAM_BATCH_BYTES + ): + yield batch + batch, size = [], 0 + batch.append(record) + size += record_size + if batch: + yield batch + + +def _is_completion_command( + command: temporalio.bridge.proto.workflow_commands.WorkflowCommand, +) -> bool: + return ( + command.HasField("complete_workflow_execution") + or command.HasField("continue_as_new_workflow_execution") + or command.HasField("fail_workflow_execution") + or command.HasField("cancel_workflow_execution") + ) + + +class _StreamBuffer: + """Holds the stream ranges delivered to a workflow so far. + + Delivery is driven by the server, not by whether workflow code happens to be + reading. A range arrives once, is recorded in History as consumed, and is + never sent again, so anything not yet read has to be kept here rather than + dropped. + """ + + def __init__(self, stream_id: str = "") -> None: + self._stream_id = stream_id + self._records: list[temporalio.workflow.DeliveredStreamRecord] = [] + self._waiters: list[asyncio.Future] = [] + # Where the next range has to start. Unknown until the first one + # arrives, because a subscription may start wherever the stream is and + # the server is the one that resolves that. + self._next_offset: int | None = None + + def extend( + self, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + from_offset: int = 0, + to_offset: int | None = None, + ) -> None: + if to_offset is None: + to_offset = from_offset + len(records) + # A range is recorded as consumed once and never resent, so one that + # repeats, skips or mis-sizes would hand the workflow duplicate or + # shifted records with nothing to say so. Failing the task is what + # makes the fault visible. + if to_offset - from_offset != len(records): + raise RuntimeError( + f"stream {self._stream_id!r} delivered {len(records)} records " + f"for offsets [{from_offset}, {to_offset})" + ) + if self._next_offset is not None and from_offset != self._next_offset: + raise RuntimeError( + f"stream {self._stream_id!r} delivered offsets [{from_offset}, " + f"{to_offset}) but the last range ended at {self._next_offset}" + ) + self._next_offset = to_offset + # An empty range still counts as a delivery, but there is nothing to + # hand a reader, so only a non-empty one wakes anyone. + if not records: + return + # Offsets are dense inside a delivered range and the range arrives in + # order, so counting from its start is the position rather than an + # estimate of it. The per-record field is not on the activation, and a + # reader that has to resume elsewhere needs a position it can name. + for index, record in enumerate(records): + # Copied so the buffer outlives the activation that carried it. + kept = temporalio.api.stream.v1.StreamRecord() + kept.CopyFrom(record) + self._records.append( + temporalio.workflow.DeliveredStreamRecord( + record=kept, offset=from_offset + index + ) + ) + waiters, self._waiters = self._waiters, [] + for waiter in waiters: + if not waiter.done(): + waiter.set_result(None) + + def take(self) -> list[temporalio.workflow.DeliveredStreamRecord]: + taken, self._records = self._records, [] + return taken + + def put_back( + self, records: Sequence[temporalio.workflow.DeliveredStreamRecord] + ) -> None: + """Return an unread tail to the front of the buffer.""" + self._records[:0] = records + + def wait_future(self) -> asyncio.Future: + loop = asyncio.get_event_loop() + fut = loop.create_future() + self._waiters.append(fut) + return fut + + def __len__(self) -> int: + return len(self._records) + + class _WorkflowInstanceImpl( # type: ignore[reportImplicitAbstractClass] WorkflowInstance, temporalio.workflow._Runtime, asyncio.AbstractEventLoop ): @@ -401,6 +524,16 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: str, list[temporalio.bridge.proto.workflow_activation.SignalWorkflow] ] = {} + # Stream ranges delivered to this workflow, keyed by stream id. Ranges + # arrive whether or not anything is reading yet, because the server has + # already recorded them as consumed and will not send them again. + self._stream_buffers: dict[str, _StreamBuffer] = {} + # Records this task's publishes append, by stream, until the task + # completes and they become commands. + self._stream_appends: dict[ + str, list[temporalio.api.stream.v1.StreamRecord] + ] = {} + # When we evict, we have to mark the workflow as deleting so we don't # add any commands and we swallow exceptions on tear down self._deleting = False @@ -461,6 +594,9 @@ def activate( temporalio.bridge.proto.workflow_completion.WorkflowActivationCompletion() ) self._current_completion.successful.SetInParent() + # A failed task's publishes never became commands, so nothing carries + # over into this one. + self._stream_appends = {} self._current_activation_error: Exception | None = None self._deployment_version_for_current_task = ( @@ -552,6 +688,9 @@ def activate( ) activation_err = None + if activation_err is None and not self._deleting: + self._flush_stream_appends() + # Apply versioning behavior if one was established if self._versioning_behavior: self._current_completion.successful.versioning_behavior = ( @@ -616,6 +755,8 @@ def _apply( ) -> None: if job.HasField("cancel_workflow"): self._apply_cancel_workflow(job.cancel_workflow) + elif job.HasField("deliver_stream_records"): + self._apply_deliver_stream_records(job.deliver_stream_records) elif job.HasField("do_update"): self._apply_do_update(job.do_update) elif job.HasField("fire_timer"): @@ -1141,6 +1282,19 @@ def _apply_resolve_signal_external_workflow( else: fut.set_result(None) + def _apply_deliver_stream_records( + self, + job: temporalio.bridge.proto.workflow_activation.DeliverStreamRecords, + ) -> None: + buffer = self._stream_buffers.get(job.stream_id) + if buffer is None: + # Nothing subscribed. The range is already recorded as consumed and + # will not be sent again, so buffering it is the only way a + # subscription made later in the same task still sees it. + buffer = _StreamBuffer(job.stream_id) + self._stream_buffers[job.stream_id] = buffer + buffer.extend(job.records, job.from_offset, job.to_offset) + def _apply_signal_workflow( self, job: temporalio.bridge.proto.workflow_activation.SignalWorkflow ) -> None: @@ -1305,6 +1459,76 @@ def workflow_get_current_deployment_version( def get_info(self) -> temporalio.workflow.Info: return self._info + def workflow_subscribe_stream(self, stream_id: str, start_offset: int) -> None: + # Reissued on every replay, so the buffer has to exist before the first + # range arrives and the command has to be harmless the second time. A + # repeat subscription leaves the server-side cursor where it is. + self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)) + command = self._add_command() + command.subscribe_stream.stream_id = stream_id + command.subscribe_stream.start_offset = start_offset + + def workflow_append_stream_records( + self, + stream_id: str, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + ) -> None: + self._assert_not_read_only("append stream records") + if not records: + raise ValueError("append_stream_records needs at least one record") + kept: list[temporalio.api.stream.v1.StreamRecord] = [] + for record in records: + if record.ByteSize() > _MAX_STREAM_RECORD_BYTES: + raise ValueError( + f"a stream record is limited to {_MAX_STREAM_RECORD_BYTES} " + f"bytes, got {record.ByteSize()}" + ) + copy = temporalio.api.stream.v1.StreamRecord() + copy.CopyFrom(record) + # The workflow is the producer here, whatever the caller set. + copy.producer_id = "" + kept.append(copy) + # Held until the task completes, so a task's publishes on one stream + # become one command and one History event however many there were. + self._stream_appends.setdefault(stream_id, []).extend(kept) + + def _flush_stream_appends(self) -> None: + appends, self._stream_appends = self._stream_appends, {} + if not appends: + return + commands = self._current_completion.successful.commands + # Ahead of any command that ends the run, because the server accepts + # nothing after one of those. + insert_at = len(commands) + for index, command in enumerate(commands): + if _is_completion_command(command): + insert_at = index + break + for stream_id, records in appends.items(): + for batch in _stream_batches(records): + command = temporalio.bridge.proto.workflow_commands.WorkflowCommand() + command.append_stream_records.stream_id = stream_id + command.append_stream_records.records.extend(batch) + commands.insert(insert_at, command) + insert_at += 1 + + async def workflow_read_stream_records( + self, stream_id: str, max_records: int + ) -> list[temporalio.workflow.DeliveredStreamRecord]: + # Ranges arrive on Workflow Tasks, and a query activation carries none, + # so without this the read waits on a future nothing can resolve and the + # query times out with nothing to say why. + self._assert_not_read_only("read stream") + buffer = self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)) + while not len(buffer): + await buffer.wait_future() + taken = buffer.take() + if max_records and len(taken) > max_records: + # Put the tail back rather than dropping it: nothing will resend it. + buffer.put_back(taken[max_records:]) + taken = taken[:max_records] + return taken + def workflow_get_current_history_length(self) -> int: return self._current_history_length diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index f63135bf4..18d04cb0c 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -58,6 +58,7 @@ wait, ) from ._context import ( + DeliveredStreamRecord, Info, ParentInfo, RootInfo, @@ -65,6 +66,7 @@ _current_update_info, _Runtime, _set_current_update_info, + append_stream_records, cancellation_reason, current_update_info, deprecate_patch, @@ -86,9 +88,11 @@ payload_converter, random, random_seed, + read_stream_records, register_random_seed_callback, set_current_details, sleep, + subscribe_stream, time, time_ns, upsert_memo, @@ -233,6 +237,10 @@ "upsert_search_attributes", "uuid4", "uuid7", + "DeliveredStreamRecord", + "read_stream_records", + "append_stream_records", + "subscribe_stream", "wait_condition", "DynamicWorkflowConfig", "defn", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 25523836a..8b6c07782 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -14,6 +14,7 @@ from nexusrpc import InputT, OutputT import temporalio.api.common.v1 +import temporalio.api.stream.v1 import temporalio.common import temporalio.converter @@ -313,6 +314,21 @@ def workflow_get_current_deployment_version( self, ) -> temporalio.common.WorkerDeploymentVersion | None: ... + @abstractmethod + def workflow_subscribe_stream(self, stream_id: str, start_offset: int) -> None: ... + + @abstractmethod + def workflow_append_stream_records( + self, + stream_id: str, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + ) -> None: ... + + @abstractmethod + async def workflow_read_stream_records( + self, stream_id: str, max_records: int + ) -> list[DeliveredStreamRecord]: ... + @abstractmethod def workflow_get_current_history_length(self) -> int: ... @@ -949,6 +965,87 @@ async def sleep(duration: float | timedelta, *, summary: str | None = None) -> N ) +def subscribe_stream(stream_id: str, *, start_offset: int = 0) -> None: + """Subscribe this workflow to a server-side stream. + + From here on its Workflow Tasks carry the ranges it has not consumed yet, + and :func:`read_stream_records` returns them. Safe to call again: a second + subscription to a stream this run already consumes does not move its + cursor, though it does write one event. Calling it on every replay is + harmless because replay matches the command to the event already recorded. + + Only the stream id and start offset go to the server. The rest of the + stream's addressing is resolved there, because a workflow cannot look it up + without doing I/O and a value it carried would be a reading rather than a + fact. A name this workflow has not written yet names a stream it owns, and + subscribing creates it. + + Args: + stream_id: Stream to consume: the name of one this workflow owns, or + the id of a standalone stream. + start_offset: Where to start. Negative means from wherever the stream is + when the subscription is registered; the server resolves that once + and records it, so replay does not resolve it again. + """ + _Runtime.current().workflow_subscribe_stream(stream_id, start_offset) + + +def append_stream_records( + records: Sequence[temporalio.api.stream.v1.StreamRecord], + *, + stream_id: str = "", +) -> None: + """Publish records to a server-side stream this workflow owns. + + Returns at once. The records become one command when this Workflow Task + completes, so they are visible when the task is accepted and never if it + fails. Their bodies go to the stream's own log rather than into History, + which gets one fixed-size event naming the offset range, so a task that + publishes a thousand records costs History the same as one that publishes + one. Readers do not have to exist yet, and adding one costs the writer + nothing. + + Args: + records: Records to append, in order. The server stores each with an + empty ``producer_id``, because the workflow is the producer. + stream_id: Stream to publish to. Empty means the workflow's default + output stream. + + Raises: + ValueError: ``records`` is empty or one of them is over the server's + per-record size limit. + """ + _Runtime.current().workflow_append_stream_records(stream_id, records) + + +@dataclass(frozen=True) +class DeliveredStreamRecord: + """One record a consuming workflow was given, with where it sat.""" + + record: temporalio.api.stream.v1.StreamRecord + offset: int + """Its position in the whole stream, which is what a reader resumes from.""" + + +async def read_stream_records( + stream_id: str, *, max_records: int = 0 +) -> list[DeliveredStreamRecord]: + """Read the next records of a server-side stream this workflow consumes. + + Waits until at least one record is available. Ranges arrive on Workflow + Tasks, and only the offsets they covered are written to History, so this is + deterministic on replay: the server re-supplies the same ranges by reading + the stream again. Subscribe first with :func:`subscribe_stream`; this only + reads what has already been delivered to this workflow. + + Args: + stream_id: Stream to read from. + max_records: Most records to return at once, or 0 for everything + available. + """ + return await _Runtime.current().workflow_read_stream_records(stream_id, max_records) + + async def wait_condition( fn: Callable[[], bool], *, diff --git a/tests/worker/test_workflow_stream.py b/tests/worker/test_workflow_stream.py new file mode 100644 index 000000000..9083e74a9 --- /dev/null +++ b/tests/worker/test_workflow_stream.py @@ -0,0 +1,288 @@ +"""In-workflow consumption and publication of a server-side stream. + +Ranges arrive on Workflow Tasks and only the offsets they covered are written to +History, so replay is served by the server reading the stream again. These tests +drive the buffering, the workflow-facing read and the per-task publish directly, +which is the part this SDK owns; the delivery decision itself lives in the +server and sdk-core. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +import temporalio.api.common.v1 +import temporalio.api.stream.v1 as api_stream +from temporalio.bridge.proto.workflow_commands import WorkflowCommand +from temporalio.worker._workflow_instance import ( + _MAX_STREAM_BATCH_BYTES, + _MAX_STREAM_RECORD_BYTES, + _MAX_STREAM_RECORDS_PER_BATCH, + _StreamBuffer, + _WorkflowInstanceImpl, +) +from temporalio.workflow import ReadOnlyContextError + + +def record(body: bytes, topic: str = "") -> api_stream.StreamRecord: + return api_stream.StreamRecord( + body=temporalio.api.common.v1.Payload(data=body), topic=topic + ) + + +def bodies(delivered: list[Any]) -> list[bytes]: + return [item.record.body.data for item in delivered] + + +async def test_buffer_hands_over_in_order() -> None: + buffer = _StreamBuffer() + buffer.extend([record(b"one"), record(b"two")]) + + assert bodies(buffer.take()) == [b"one", b"two"] + assert len(buffer) == 0 + + +# A reader that arrives before the data must not miss it, and one that arrives +# after must not block: the range is delivered once and never resent. +async def test_buffer_wakes_a_waiting_reader() -> None: + buffer = _StreamBuffer() + waiter = buffer.wait_future() + assert not waiter.done() + + buffer.extend([record(b"late")]) + await asyncio.wait_for(waiter, timeout=1) + assert bodies(buffer.take()) == [b"late"] + + +# An empty range is still a delivery the server recorded, but there is nothing +# to hand a reader, so it must not wake one into returning nothing. +async def test_empty_range_does_not_wake_a_reader() -> None: + buffer = _StreamBuffer() + waiter = buffer.wait_future() + + buffer.extend([]) + + assert not waiter.done() + assert len(buffer) == 0 + + +async def test_buffer_keeps_data_delivered_before_anyone_reads() -> None: + buffer = _StreamBuffer() + buffer.extend([record(b"early")]) + + # No waiting: the data is already there. + assert len(buffer) == 1 + waiter = buffer.wait_future() + assert not waiter.done(), "a fresh waiter is only resolved by new data" + assert bodies(buffer.take()) == [b"early"] + + +class _ReadOnlyStub: + """Enough of the workflow instance to drive the real read. + + The cap lives inside ``workflow_read_stream_records``, so a test that + reimplements it proves nothing about the code that ships. Borrowing the + method off the real class is what keeps the test on the shipped path. + """ + + workflow_read_stream_records = _WorkflowInstanceImpl.workflow_read_stream_records + + def __init__(self, read_only: bool = False) -> None: + self._stream_buffers: dict[str, _StreamBuffer] = {} + self._read_only = read_only + + def _assert_not_read_only(self, action: str) -> None: + if self._read_only: + raise ReadOnlyContextError(f"cannot {action} in a read-only context") + + async def read(self, stream: str, max_records: int) -> list[bytes]: + instance: Any = self + return bodies( + await _WorkflowInstanceImpl.workflow_read_stream_records( + instance, stream, max_records + ) + ) + + +@pytest.mark.parametrize("max_records", [1, 2, 5]) +async def test_read_respects_a_cap_without_losing_the_tail(max_records: int) -> None: + stub = _ReadOnlyStub() + expected = [b"a", b"b", b"c"] + stub._stream_buffers["s"] = _StreamBuffer() + stub._stream_buffers["s"].extend([record(b) for b in expected]) + + got = await stub.read("s", max_records) + + assert got == expected[:max_records] + # Whatever the cap left behind has to still be there: nothing resends it. + assert len(stub._stream_buffers["s"]) == max(0, len(expected) - max_records) + + if len(stub._stream_buffers["s"]): + rest = await stub.read("s", 0) + assert got + rest == expected + else: + assert got == expected + + +# A query activation carries no ranges, so a read there would wait on a future +# nothing can resolve and the query would time out saying nothing. +async def test_read_is_refused_in_a_read_only_context() -> None: + stub = _ReadOnlyStub(read_only=True) + stub._stream_buffers["s"] = _StreamBuffer() + stub._stream_buffers["s"].extend([record(b"a")]) + + with pytest.raises(ReadOnlyContextError, match="read stream"): + await stub.read("s", 0) + # Refused before the buffer was touched, so the range is still there for + # the task that is allowed to read it. + assert len(stub._stream_buffers["s"]) == 1 + + +# A range is recorded as consumed once and never resent, so the buffer is the +# only place a repeated, skipped or mis-sized delivery can still be noticed. +async def test_ranges_have_to_abut_the_last_one() -> None: + buffer = _StreamBuffer("s") + buffer.extend([record(b"a"), record(b"b")], 0, 2) + # An empty range moves the expectation too: the server recorded it. + buffer.extend([], 2, 2) + buffer.extend([record(b"c")], 2, 3) + assert [item.offset for item in buffer.take()] == [0, 1, 2] + + with pytest.raises(RuntimeError, match=r"\[2, 3\).*ended at 3"): + buffer.extend([record(b"c")], 2, 3) + with pytest.raises(RuntimeError, match=r"\[5, 6\).*ended at 3"): + buffer.extend([record(b"f")], 5, 6) + + +async def test_a_range_has_to_carry_as_many_records_as_it_spans() -> None: + buffer = _StreamBuffer("s") + with pytest.raises(RuntimeError, match=r"2 records for offsets \[0, 1\)"): + buffer.extend([record(b"a"), record(b"b")], 0, 1) + assert len(buffer) == 0 + + +class _Completion: + def __init__(self) -> None: + self.commands: list[WorkflowCommand] = [] + + +class _Successful: + def __init__(self) -> None: + self.successful = _Completion() + + +class _CommandStub: + """Drives the real publish path and keeps the commands it issued. + + The completion's command list is a plain list here, which supports the + same ``insert`` the protobuf container does. + """ + + workflow_append_stream_records = ( + _WorkflowInstanceImpl.workflow_append_stream_records + ) + _flush_stream_appends = _WorkflowInstanceImpl._flush_stream_appends + + def __init__(self) -> None: + self._stream_appends: dict[str, list[api_stream.StreamRecord]] = {} + self._current_completion = _Successful() + + def _assert_not_read_only(self, _action: str) -> None: + pass + + @property + def commands(self) -> list[WorkflowCommand]: + return self._current_completion.successful.commands + + def publish(self, *records: api_stream.StreamRecord, stream_id: str = "") -> None: + instance: Any = self + _WorkflowInstanceImpl.workflow_append_stream_records( + instance, stream_id, list(records) + ) + + def flush(self) -> None: + instance: Any = self + _WorkflowInstanceImpl._flush_stream_appends(instance) + + +# A task's publishes on one stream become one command, whatever their number, +# because the event the command produces is what bounds a workflow's History. +def test_a_tasks_publishes_on_one_stream_become_one_command() -> None: + stub = _CommandStub() + stub.publish(record(b"a"), record(b"b")) + stub.publish(record(b"c")) + stub.publish(record(b"d"), stream_id="other") + assert stub.commands == [] + + stub.flush() + + by_stream = { + command.append_stream_records.stream_id: [ + r.body.data for r in command.append_stream_records.records + ] + for command in stub.commands + } + assert by_stream == {"": [b"a", b"b", b"c"], "other": [b"d"]} + # Flushed once: a second flush has nothing left to say. + stub.flush() + assert len(stub.commands) == 2 + + +# The workflow is the producer of what it publishes, whatever the caller set. +def test_the_workflows_records_carry_no_producer() -> None: + stub = _CommandStub() + stub.publish(api_stream.StreamRecord(producer_id="someone", attempt=3)) + stub.flush() + assert stub.commands[0].append_stream_records.records[0].producer_id == "" + + +# The server accepts nothing after a command that ends the run, so the +# publishes have to go ahead of it. +def test_publishes_are_flushed_ahead_of_the_completion_command() -> None: + stub = _CommandStub() + stub.publish(record(b"a")) + done = WorkflowCommand() + done.complete_workflow_execution.SetInParent() + stub.commands.append(done) + + stub.flush() + + assert [c.WhichOneof("variant") for c in stub.commands] == [ + "append_stream_records", + "complete_workflow_execution", + ] + + +# The server refuses an oversized record, and a refused command is reissued on +# every replay, so the limit has to be applied before the record is buffered. +def test_a_record_over_the_server_limit_is_refused_before_the_command() -> None: + stub = _CommandStub() + stub.publish(record(b"x" * (_MAX_STREAM_RECORD_BYTES - 16))) + with pytest.raises(ValueError, match=f"{_MAX_STREAM_RECORD_BYTES} bytes"): + stub.publish(record(b"x" * (_MAX_STREAM_RECORD_BYTES + 1))) + stub.flush() + assert len(stub.commands) == 1 + + +# A task that publishes more than one batch holds is split into commands the +# server accepts, by count and by bytes, rather than refused. +def test_a_task_over_the_batch_limits_is_split_into_commands() -> None: + stub = _CommandStub() + stub.publish(*(record(b"x") for _ in range(_MAX_STREAM_RECORDS_PER_BATCH + 1))) + stub.flush() + assert [len(c.append_stream_records.records) for c in stub.commands] == [ + _MAX_STREAM_RECORDS_PER_BATCH, + 1, + ] + + stub = _CommandStub() + # Two records that fit one batch together, then one whose framing alone + # overflows the room they leave. + big = record(b"x" * (_MAX_STREAM_RECORD_BYTES - 16)) + room = _MAX_STREAM_BATCH_BYTES - 2 * big.ByteSize() + stub.publish(big, big, record(b"y" * room)) + stub.flush() + assert [len(c.append_stream_records.records) for c in stub.commands] == [2, 1] From c76d1763244a9e8fb45fa081d5696b6036650afc Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:07:25 -0700 Subject: [PATCH 02/27] Added the native provider with one owned stream per topic. The workflow subscribes to the topic by name and publishes with the batched command; outside code appends and long-polls the same stream, following the run chain by cursor unless pinned. --- temporalio/streams/providers/native.py | 463 ++++++++++++++++++++++ tests/streams/test_streams_conformance.py | 53 ++- tests/worker/test_workflow_stream_e2e.py | 386 ++++++++++++++++++ 3 files changed, 901 insertions(+), 1 deletion(-) create mode 100644 temporalio/streams/providers/native.py create mode 100644 tests/worker/test_workflow_stream_e2e.py diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py new file mode 100644 index 000000000..9c0cb3dbd --- /dev/null +++ b/temporalio/streams/providers/native.py @@ -0,0 +1,463 @@ +"""The server-side (native) provider. + +Streams live on the Temporal server, beside the workflow that owns them. A +topic is one owned stream named after the topic, created by whoever touches it +first: the workflow publishes to it with a command the server applies in the +transaction that accepts the Workflow Task, subscribes to it by name and reads +the ranges the server delivers on its Workflow Tasks; outside code appends and +reads through the stream service, and the workflow's records and an outside +producer's land in one log in the order the server accepted them. + +A cursor names the run as well as the offset, because an owned stream belongs +to one run and a successor's starts over at zero. A handle without a run id +reads run after run, learning from the poll that a run's stream is closed and +from the run's close event who came next; with a run id it is pinned. + +Prototype support for AI-198. It needs a server built from that branch and +opens its own gRPC channel to it, because sdk-core does not know the stream +service yet, which is also why it does not support TLS or API keys. +""" + +from __future__ import annotations + +import logging +from collections.abc import AsyncGenerator +from typing import Any, Generic, TypeVar + +from temporalio import workflow +from temporalio.client import Client, WorkflowHistoryEventFilterType +from temporalio.client_stream import ( + StreamClient, + WorkflowStreamHandle, + close_shared_clients, + shared_client, +) +from temporalio.converter import PayloadCodec, PayloadConverter +from temporalio.service import RPCError, RPCStatusCode +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, + cursor_position, + mint_cursor, + producer_identity, + to_wire, +) +from temporalio.streams.providers import ProviderPlugin + +__all__ = ["NativeProducer", "NativeStreamHandle", "NativeStreams"] + +T = TypeVar("T") + +_PROVIDER = "native" + +logger = logging.getLogger(__name__) + + +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 " + "native provider" + ) from None + + +async def _encode_body(codec: PayloadCodec | None, record: WireRecord) -> WireRecord: + # The worker's payload visitor runs a codec over the bodies a workflow + # publishes and receives; the outside half has no such pass, so it applies + # the client's codec here or the two sides would not agree. + if codec is None or not record.HasField("body"): + return record + encoded = await codec.encode([record.body]) + record.body.CopyFrom(encoded[0]) + return record + + +async def _decode_body(codec: PayloadCodec | None, record: WireRecord) -> WireRecord: + if codec is None or not record.HasField("body"): + return record + decoded = await codec.decode([record.body]) + record.body.CopyFrom(decoded[0]) + return record + + +class _NativeReadSource: + """One subscription of the running workflow, fed by delivered ranges.""" + + def __init__(self, stream_id: str, run_id: str) -> None: + self._stream_id = stream_id + self._run_id = run_id + self._closed = False + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + if self._closed: + raise StopAsyncIteration + delivered = await workflow.read_stream_records(self._stream_id) + return [(_cursor(self._run_id, item.offset), item.record) for item in delivered] + + def close(self) -> None: + self._closed = True + + +class _NativeWriteSink: + def __init__(self, topic: str) -> None: + self._topic = topic + + def publish(self, record: WireRecord) -> None: + # Held by the runtime until the task completes, when the task's + # records on this topic become one command the server applies with + # the task: rule 1 through the server's own commit. + workflow.append_stream_records([record], stream_id=self._topic) + + +class _NativeWorkflowProvider: + """The workflow half: the server's commands and delivered ranges.""" + + 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 stream is its own" + ) + start = named[1] + 1 + workflow.subscribe_stream(topic, start_offset=start) + return _NativeReadSource(topic, run_id) + + def open_writer(self, topic: str) -> WriteSink: + _require_topic(topic) + return _NativeWriteSink(topic) + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +class NativeProducer(Generic[T]): + """Appends to a topic from outside workflow code. + + Every append is visible as soon as the server accepts it. That is the + point for an activity streaming model output, and it is why an activity + carries its own identity: the retry of a failed attempt has no commit + boundary to sort it out afterwards. + """ + + def __init__( + self, + handle: WorkflowStreamHandle, + pin: Any, + codec: PayloadCodec | None, + converter: PayloadConverter, + topic: str, + producer_id: str, + attempt: int, + ) -> None: + """Bind this producer to ``topic`` on the stream ``handle`` names.""" + self._handle = handle + self._pin = pin + self._codec = codec + self._converter = converter + self._topic = topic + self._producer_id = producer_id + self._attempt = attempt + self._sequence = 0 + self._last = BEGINNING + + @property + def producer_id(self) -> str: + """Who this producer writes as.""" + return self._producer_id + + @property + def attempt(self) -> int: + """The generation this producer is writing.""" + return self._attempt + + @property + def _writer(self) -> str: + # The server dedupes on this and the sequence. The attempt is part of + # it so a retried append is dropped while a new generation writing + # different words at the same sequence is not. + return ( + f"{self._producer_id}#{self._attempt}" + if self._attempt + else self._producer_id + ) + + async def append(self, *values: T) -> Cursor: + """Append ``values`` and return the cursor of the last record as stored. + + A repeat returns where the original landed, because the server + answers a deduplicated batch with the original offsets; an empty call + returns the position of this producer's last record. + """ + if not values: + return self._last + return await self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.DATA, + value=value, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence + index, + ) + for index, value in enumerate(values) + ] + ) + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic.""" + await self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.FINISH, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + ) + ] + ) + + async def _write(self, records: list[WireRecord]) -> Cursor: + # Pinned before the first write, so every cursor this producer hands + # out names the run its records landed in. + if not self._handle.owner_run_id: + self._handle.pin(await self._pin()) + for record in records: + await _encode_body(self._codec, record) + appended = await self._handle.append( + *records, producer_id=self._writer, sequence=self._sequence + ) + self._sequence += len(records) + self._last = _cursor(self._handle.owner_run_id, appended.next_offset - 1) + return self._last + + +class NativeStreamHandle: + """One workflow's topics from outside, over the stream service.""" + + def __init__(self, client: Client, workflow_id: str, run_id: str | None) -> None: + """Address ``workflow_id``'s topics, pinned to ``run_id`` when one is given.""" + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._converter = client.data_converter.payload_converter + self._codec = client.data_converter.payload_codec + self._streams: StreamClient | None = None + + def _service(self) -> StreamClient: + # Resolved on first use, because the shared channel belongs to the + # running loop and a handle may be made before there is one. + if self._streams is None: + self._streams = shared_client( + self._client.service_client.config.target_host, + self._client.namespace, + ) + return self._streams + + def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: + return self._service().workflow_stream( + self._workflow_id, topic, owner_run_id=run_id + ) + + def read( + self, + *, + 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.""" + 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) + 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 + ) + 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: + stream = self._stream(topic, run_id) + while True: + page = await stream.poll(from_offset=offset) + for entry in page.entries: + record = await _decode_body(self._codec, entry.record) + for out in decoder.decode(_cursor(run_id, entry.offset), record): + yield out + offset = page.next_offset + # On a pinned stream the server reports the run's end as closed, + # and a closed stream is finished once its head is delivered. + if page.closed and offset >= page.head_offset: + break + if self._run_id is not None: + return + successor = await self._successor(run_id) + if successor is None: + return + run_id, offset = successor, 0 + + async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: + """The cursor of the newest record on ``topic``, naming the run it was read from. + + 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. + """ + topic, _ = resolve_topic(topic) + run_id = self._run_id or await self._current_run() + try: + head = (await self._stream(topic, run_id).describe()).head_offset + except StreamNotFoundError: + # A topic nobody has written yet does not exist on the server, + # which is the same answer as an empty one. + head = 0 + if head > 0: + return _cursor(run_id, head - 1) + if self._run_id is None and await self._predecessor(run_id) is not None: + return _cursor(run_id, -1) + return BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any], + producer_id: str = "", + attempt: int = 0, + ) -> NativeProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return NativeProducer( + self._stream(topic, self._run_id or ""), + self._current_run, + self._codec, + self._converter, + topic, + producer_id, + attempt, + ) + + async def _current_run(self) -> str: + try: + description = await self._client.get_workflow_handle( + self._workflow_id + ).describe() + 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 description.run_id is not None + return description.run_id + + async def _first_run(self) -> str: + """The oldest retained run of the chain, walking back from the latest.""" + run_id = await self._current_run() + 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: + handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + try: + async for event in handle.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, run_id: str) -> str | None: + handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + events = handle.fetch_history_events( + event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT + ) + 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 + return None + + +class NativeStreams(ProviderPlugin): + """The server-side provider. + + Takes no options: the streams are on the server the client is already + connected to. Construct one, pass it to the worker as a plugin and open + handles from it anywhere else; :meth:`close` releases the channels this + process opened to the stream service. + """ + + def workflow_provider(self) -> _NativeWorkflowProvider: + """The workflow half, over the server's commands and delivered ranges.""" + return _NativeWorkflowProvider() + + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> NativeStreamHandle: + """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain.""" + return NativeStreamHandle(client, workflow_id, run_id) + + async def close(self) -> None: + """Close the channels this process opened to the stream service.""" + await close_shared_clients() diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 9002ea16a..3bdc93e1a 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -21,6 +21,7 @@ from __future__ import annotations import asyncio +import os import uuid from collections.abc import AsyncIterator, Awaitable, Callable from dataclasses import dataclass @@ -28,8 +29,9 @@ 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.native import NativeStreams +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,9 +95,56 @@ 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 _native_case(client: Client) -> AsyncIterator[ProviderCase]: + # The store is a server built from the stream-carrying branch, which the + # test environment's own server is not; TEMPORAL_ADDRESS names it. + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = NativeStreams() + # 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) 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("native", provider, client, host=host) + for handle in hosts.values(): + await handle.terminate() + await provider.close() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { "memory": _memory_case } +if os.environ.get("STREAMS_LIVE") == "native": + SETUPS["native"] = _native_case _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py new file mode 100644 index 000000000..1d974b401 --- /dev/null +++ b/tests/worker/test_workflow_stream_e2e.py @@ -0,0 +1,386 @@ +"""The native provider inside a real workflow, against a server that has streams. + +The two rules the memory provider cannot keep are the measurement here: a +publish commits with its Workflow Task and never lands if the task fails, and +a read is a recorded observation the server re-supplies on replay. Needs a +Temporal server built from the AI-198 branch, because neither the stream +service nor the commands exist on a released one: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/worker/test_workflow_stream_e2e.py + +Skipped otherwise, rather than passing against a server that has no idea what a +stream is. +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.enums.v1 import EventType +from temporalio.api.stream.v1 import StreamRecord +from temporalio.client import Client +from temporalio.client_stream import StreamClient +from temporalio.streams import RecordKind +from temporalio.streams.providers.native import NativeStreams +from temporalio.worker import Worker +from tests.streams.test_streams_conformance import take + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream commands", +) + +INPUTS = "inputs" +DECISIONS = "decisions" + +EVENT_STREAM_SUBSCRIBED = EventType.EVENT_TYPE_WORKFLOW_STREAM_SUBSCRIBED +EVENT_STREAM_RECORDS_APPENDED = EventType.EVENT_TYPE_WORKFLOW_STREAM_RECORDS_APPENDED + + +async def _connect(provider: NativeStreams | None = None) -> Client: + # Registered once, on the client: the worker inherits it and the tests + # open handles through client.get_stream_handle. The contrib tests below + # need no provider. + return await Client.connect(TARGET or "", plugins=[provider] if provider else []) + + +async def _event_counts(client: Client, workflow_id: str) -> dict[Any, int]: + counts = { + EVENT_STREAM_RECORDS_APPENDED: 0, + EVENT_STREAM_SUBSCRIBED: 0, + EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED: 0, + } + async for event in client.get_workflow_handle(workflow_id).fetch_history_events(): + if event.event_type in counts: + counts[event.event_type] += 1 + return counts + + +@workflow.defn +class ContractLoop: + """Reads ``inputs``, publishes a decision per value, reports control records.""" + + @workflow.run + async def run(self) -> list[dict[str, Any]]: + inputs = workflow.stream_reader(INPUTS, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "replaced": record.supersession.previous_attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decisions.publish({"decided": record.value["n"]}) + trace.append({"kind": "decision", "n": record.value["n"]}) + decisions.finish() + return trace + + +async def test_the_interface_loop_runs_on_the_server_with_a_cold_cache() -> None: + """Rule 2 on the native provider: every task replays from the server. + + With the cache off, each Workflow Task rebuilds the workflow from History + and the server re-supplies the ranges earlier tasks consumed, so the loop + completing at all means the same records came back in the same order. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "loop-tq-" + uuid.uuid4().hex[:8] + workflow_id = "loop-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[ContractLoop], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"n": 1}, {"n": 2}) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"n": 3}) + await second.finish() + trace = await asyncio.wait_for(handle.result(), 60) + + assert trace == [ + {"kind": "decision", "n": 1}, + {"kind": "decision", "n": 2}, + {"kind": "superseded", "replaced": 1}, + {"kind": "decision", "n": 3}, + {"kind": "finish", "producer": "model"}, + ] + + # The read ends by itself: the workflow is closed and the tail + # delivered, with the workflow's own records carrying no producer. + async def read_everything() -> list[Any]: + return [ + (r.producer_id, r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + assert await asyncio.wait_for(read_everything(), 60) == [ + ("", RecordKind.DATA, {"decided": 1}), + ("", RecordKind.DATA, {"decided": 2}), + ("", RecordKind.DATA, {"retracting_attempt": 1}), + ("", RecordKind.DATA, {"decided": 3}), + ("", RecordKind.FINISH, None), + ] + counts = await _event_counts(client, workflow_id) + assert counts[EVENT_STREAM_SUBSCRIBED] == 1 + # Several tasks published, one event each; the loop spanned more than + # one task or the cold cache proved nothing. + assert counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] >= 2 + assert ( + 1 + <= counts[EVENT_STREAM_RECORDS_APPENDED] + <= counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] + ) + finally: + await provider.close() + + +# Run ids whose first workflow task already failed, shared with the workflow +# thread so the retry can tell it is the retry. Outside the sandbox on +# purpose: the sandbox re-imports this module per run and would hide the set. +_failed_once: set[str] = set() + + +@workflow.defn(sandboxed=False) +class PublishThenFail: + """Publishes, then fails its first workflow task; the retry publishes again.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + run_id = workflow.info().run_id + committed = run_id in _failed_once + decisions.publish({"committed": committed}) + if not committed: + _failed_once.add(run_id) + raise RuntimeError("the first task fails after publishing") + decisions.finish() + + +async def test_a_failed_task_publishes_nothing() -> None: + """Rule 1 on the native provider: the server applies the command with the task.""" + provider = NativeStreams() + client = await _connect(provider) + task_queue = "fail-tq-" + uuid.uuid4().hex[:8] + workflow_id = "fail-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[PublishThenFail], + ): + handle = await client.start_workflow( + PublishThenFail.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + records = await take( + stream.read(topic=DECISIONS, result_type=dict), 2, timeout=60 + ) + await handle.result() + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"committed": True}), + (RecordKind.FINISH, None), + ] + finally: + await provider.close() + + +@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() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "relay-tq-" + uuid.uuid4().hex[:8] + workflow_id = "relay-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Relay]): + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + + async def read_everything() -> list[Any]: + return [ + (r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + records = await asyncio.wait_for(read_everything(), 60) + await handle.result() + # The chain is followed: the successor's record arrives on the same + # read, each cursor names its run, and the read ends with the chain. + assert records == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + # Pinned to the last run, a handle sees that run's stream alone. + last_run = (await client.get_workflow_handle(workflow_id).describe()).run_id + pinned = client.get_stream_handle(workflow_id, run_id=last_run) + only_last = [ + r.value async for r in pinned.read(topic=DECISIONS, result_type=dict) + ] + assert only_last == [{"run": 1}, None] + finally: + await provider.close() + + +@workflow.defn +class PublishAndRead: + """Publishes through the low-level surface and reads its own records back.""" + + @workflow.run + async def run(self) -> list[str]: + workflow.append_stream_records( + [_record(b"alpha", "progress"), _record(b"beta", "progress")], + stream_id="output", + ) + workflow.append_stream_records([_record(b"gamma")], stream_id="output") + # A name this workflow has not written yet still names a stream it + # owns, so subscribing creates the one the later publish lands in. + workflow.subscribe_stream("output", start_offset=0) + + received: list[str] = [] + while len(received) < 3: + for item in await workflow.read_stream_records("output"): + received.append(item.record.body.data.decode()) + return received + + +def _record(body: bytes, topic: str = "") -> StreamRecord: + return StreamRecord( + body=Payload(data=body, metadata={"encoding": b"binary/plain"}), topic=topic + ) + + +async def test_a_tasks_publishes_become_one_event() -> None: + client = await _connect() + task_queue = "publish-tq-" + uuid.uuid4().hex[:8] + workflow_id = "publish-wf-" + uuid.uuid4().hex[:8] + + async with Worker( + client, + task_queue=task_queue, + workflows=[PublishAndRead], + # Every task after the first is a replay, so completing at all means the + # reissued publish matched the event the first run wrote. + max_cached_workflows=0, + ): + result = await asyncio.wait_for( + client.execute_workflow( + PublishAndRead.run, id=workflow_id, task_queue=task_queue + ), + timeout=60, + ) + assert result == ["alpha", "beta", "gamma"] + + counts = await _event_counts(client, workflow_id) + # Two calls in one task carrying three records, so one event: the event is + # per task and stream, which is what makes publishing often free. + assert counts[EVENT_STREAM_RECORDS_APPENDED] == 1 + assert counts[EVENT_STREAM_SUBSCRIBED] == 1 + + +@workflow.defn +class ConsumeAcrossTasks: + """Reads a standalone stream over several Workflow Tasks, then reports what it saw. + + The turn structure is the point: each read that finds nothing blocks, + which ends a Workflow Task, so the run spans several. With the cache on, + every task after the first is sticky. + """ + + @workflow.run + async def run(self, stream_id: str, expected: int) -> list[str]: + workflow.subscribe_stream(stream_id, start_offset=0) + seen: list[str] = [] + while len(seen) < expected: + for item in await workflow.read_stream_records(stream_id): + seen.append(item.record.body.data.decode()) + return seen + + +async def test_a_cached_workflow_consumes_across_sticky_tasks() -> None: + """The workflow cache stays on, which is what a real worker does. + + The server sends no replay slice for a sticky task, while the sticky + history still carries the previous task's consumed range, and the two + together have to agree on every task after the first consumed range. + """ + client = await _connect() + streams = StreamClient.connect(TARGET or "") + task_queue = "sticky-tq-" + uuid.uuid4().hex[:8] + workflow_id = "sticky-wf-" + uuid.uuid4().hex[:8] + stream_id = "sticky-src-" + uuid.uuid4().hex[:8] + + batches = [["a1", "a2"], ["b1"], ["c1", "c2", "c3"]] + expected = [tok for batch in batches for tok in batch] + + try: + await streams.create(stream_id) + async with Worker( + client, task_queue=task_queue, workflows=[ConsumeAcrossTasks] + ): + handle = await client.start_workflow( + ConsumeAcrossTasks.run, + args=[stream_id, len(expected)], + id=workflow_id, + task_queue=task_queue, + ) + + # Spaced out so the workflow drains, blocks and ends a task between + # them. Without the gap the appends coalesce into one task and the + # sticky path is never taken. + for batch in batches: + await streams.get(stream_id).append( + *[_record(t.encode()) for t in batch] + ) + await asyncio.sleep(0.4) + + assert await asyncio.wait_for(handle.result(), timeout=60) == expected + + counts = await _event_counts(client, workflow_id) + assert counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] >= 3, ( + "sticky path not exercised" + ) + finally: + await streams.close() From cf6a38c4fa61c8a673896354a0d7419d1b27159f Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:07:25 -0700 Subject: [PATCH 03/27] Added the Workflow Streams surface over a server-side stream. The same API as temporalio.contrib.workflow_streams on a Temporal-owned log, so an application swaps the import and keeps its code. The LangGraph docstring qualifies its stream references, because this module gives both names a second definition pydoctor cannot resolve. --- temporalio/contrib/langgraph/_plugin.py | 8 +- temporalio/contrib/server_streams/__init__.py | 424 ++++++++++++++++++ tests/contrib/test_server_streams.py | 161 +++++++ 3 files changed, 590 insertions(+), 3 deletions(-) create mode 100644 temporalio/contrib/server_streams/__init__.py create mode 100644 tests/contrib/test_server_streams.py diff --git a/temporalio/contrib/langgraph/_plugin.py b/temporalio/contrib/langgraph/_plugin.py index 03881ca2a..1819c646d 100644 --- a/temporalio/contrib/langgraph/_plugin.py +++ b/temporalio/contrib/langgraph/_plugin.py @@ -104,11 +104,13 @@ class LangGraphPlugin(SimplePlugin): inherited default of either form. streaming_topic: When set, ``langgraph.config.get_stream_writer()`` inside a node publishes to this topic on the workflow's - :class:`WorkflowStream`. The workflow must construct - ``WorkflowStream()`` in its ``@workflow.init`` (the plugin's + :class:`temporalio.contrib.workflow_streams.WorkflowStream`. The + workflow must construct ``WorkflowStream()`` in its + ``@workflow.init`` (the plugin's interceptor verifies this on workflow start). Nodes with ``execute_in='activity'`` publish through - :class:`WorkflowStreamClient` (signal); nodes with + :class:`temporalio.contrib.workflow_streams.WorkflowStreamClient` + (signal); nodes with ``execute_in='workflow'`` publish synchronously to the in-workflow stream (no signal). streaming_batch_interval: How often the activity-side stream diff --git a/temporalio/contrib/server_streams/__init__.py b/temporalio/contrib/server_streams/__init__.py new file mode 100644 index 000000000..95eb48d7e --- /dev/null +++ b/temporalio/contrib/server_streams/__init__.py @@ -0,0 +1,424 @@ +"""Server-side streams behind the Workflow Streams API. + +This is the same surface as :mod:`temporalio.contrib.workflow_streams`, backed +by a Temporal-owned log instead of by Signals and Updates. An application swaps +the import and keeps its code: publishing from a Workflow is still a plain +call, publishing from an Activity is still a buffered handle, and a consumer +still subscribes by topic from an offset. + +What changes is underneath. A publish is a Workflow Command whose payload never +enters History, so History gets one fixed-size event per Workflow Task rather +than a Signal per batch. A consumer reads the log directly rather than +long-polling an Update, so there is no per-Workflow limit on how many can read +at once, and a closed Workflow stays readable until its stream's retention +expires. + +As in the shipped feature, a topic here is a label on a record in the +Workflow's one default stream, which a consumer filters on. The provider in +:mod:`temporalio.streams.providers.native` keeps one stream per topic instead; +the two do not share a log. + +Prototype support for AI-198. It needs a server built from that branch, and it +opens its own gRPC channel because sdk-core does not know the stream service +yet, which is also why it does not support TLS or API keys. +""" + +from __future__ import annotations + +import asyncio +import builtins +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from dataclasses import dataclass +from datetime import timedelta +from typing import Any, Generic, TypeVar, overload + +from temporalio import activity, workflow +from temporalio.api.stream.v1 import StreamRecord, StreamRecordKind +from temporalio.client import Client, WorkflowExecutionDescription +from temporalio.client_stream import StreamClient, WorkflowStreamHandle, shared_client +from temporalio.common import RawValue +from temporalio.converter import PayloadCodec, PayloadConverter + +__all__ = [ + "RawPage", + "TopicHandle", + "WorkflowStream", + "WorkflowStreamClient", + "WorkflowStreamItem", + "WorkflowTopicHandle", +] + +T = TypeVar("T") + +DEFAULT_BATCH_INTERVAL = timedelta(milliseconds=50) + + +@dataclass +class RawPage: + """One read whose items are still the stored records. + + ``closed`` with ``next_offset >= head_offset`` is the end of the stream. + The bodies are as the server holds them, codec included, for a caller + that forwards records rather than using them. + """ + + items: list[WorkflowStreamItem[StreamRecord]] + next_offset: int + head_offset: int + closed: bool + + +@dataclass +class WorkflowStreamItem(Generic[T]): + """One item read from a workflow's stream. + + ``offset`` is where the item sits in the whole stream, so it is what a + consumer hands back to resume. A topic filter leaves gaps in it. + """ + + topic: str + data: T + offset: int = 0 + + +def _record(converter: PayloadConverter, topic: str, value: Any) -> StreamRecord: + """The record one published value becomes. + + The body is the value's payload, so the encoding metadata the consumer + needs to decode into a type travels with it and a + :class:`temporalio.common.RawValue` passes through pre-encoded. + """ + record = StreamRecord(topic=topic, kind=StreamRecordKind.STREAM_RECORD_KIND_DATA) + record.body.CopyFrom(converter.to_payloads([value])[0]) + return record + + +def _decode( + converter: PayloadConverter, record: StreamRecord, as_type: type | None +) -> Any: + if as_type is None: + return converter.from_payloads([record.body])[0] + return converter.from_payloads([record.body], [as_type])[0] + + +class WorkflowTopicHandle(Generic[T]): + """A topic on the stream the running Workflow owns.""" + + def __init__(self, topic: str, value_type: type[T]) -> None: + """Prefer :meth:`WorkflowStream.topic`.""" + self._name = topic + self._type = value_type + + @property + def name(self) -> str: + """The topic name this handle is bound to.""" + return self._name + + @property + def type(self) -> type[T]: + """The value type this handle is bound to.""" + return self._type + + def publish(self, value: T | RawValue) -> None: + """Append ``value`` to the Workflow's stream on this topic. + + Returns at once. There is nothing to await: the Workflow Task's + publishes become one command the server applies in the task's own + commit, so it costs this Workflow no round trip and no extra + transition. The Worker's payload codec applies to the body as it does + to any other payload the Workflow sends. + """ + workflow.append_stream_records( + [_record(workflow.payload_converter(), self._name, value)] + ) + + +class WorkflowStream: + """The stream the running Workflow owns, from inside it. + + Construct in ``@workflow.init``. Unlike the Signals-and-Updates + implementation this holds no state of its own: the log lives on the server, + so there is nothing here for replay to reconstruct. + """ + + def __init__(self, prior_state: Any = None) -> None: + """Take the same argument as the Workflow Streams version and ignore it. + + That version carried the log across a continue-as-new, because the log + was Workflow state. Here it is not, so there is nothing to carry. + """ + self._prior_state = prior_state + + @overload + def topic(self, name: str) -> WorkflowTopicHandle[Any]: ... + + @overload + def topic(self, name: str, *, type: type[T]) -> WorkflowTopicHandle[T]: ... + + def topic(self, name: str, *, type: type = object) -> WorkflowTopicHandle[Any]: + """Bind a topic on this Workflow's stream.""" + return WorkflowTopicHandle(name, type) + + +class TopicHandle(Generic[T]): + """A topic on a Workflow's stream, from outside that Workflow.""" + + def __init__( + self, client: WorkflowStreamClient, topic: str, value_type: type[T] + ) -> None: + """Prefer :meth:`WorkflowStreamClient.topic`.""" + self._client = client + self._name = topic + self._type = value_type + + @property + def name(self) -> str: + """The topic name this handle is bound to.""" + return self._name + + @property + def type(self) -> type[T]: + """The value type this handle is bound to.""" + return self._type + + def publish(self, value: T | RawValue, *, force_flush: bool = False) -> None: + """Buffer ``value`` for the next flush. + + Buffered rather than sent, because an append costs one transition on + the owning execution whatever its size. A token at a time would pay + that per token. + """ + self._client._buffer(self._name, value) + if force_flush: + self._client._flush_soon() + + def subscribe( + self, + *, + from_offset: int = 0, + # Spelled out because `type` in this class body is the property below. + result_type: builtins.type | None = None, + poll_cooldown: timedelta | None = None, + ) -> AsyncIterator[WorkflowStreamItem[T]]: + """Read this topic from ``from_offset`` onwards.""" + return self._client.subscribe( + topics=[self._name], + from_offset=from_offset, + result_type=result_type or self._type, + poll_cooldown=poll_cooldown, + ) + + +class WorkflowStreamClient: + """Publishes to and reads from a Workflow's stream, from outside it.""" + + def __init__( + self, + handle: WorkflowStreamHandle, + converter: PayloadConverter, + batch_interval: timedelta = DEFAULT_BATCH_INTERVAL, + *, + codec: PayloadCodec | None = None, + describe: Callable[[], Awaitable[WorkflowExecutionDescription]] | None = None, + ) -> None: + """Prefer :meth:`create` or :meth:`from_within_activity`. + + ``codec`` is applied to every body this client sends and receives, so + a namespace whose payloads are encoded agrees with the Worker, whose + payload visitor applies the same codec to the Workflow's publishes. + ``describe`` is how an unpinned handle learns which run it follows; + see :meth:`WorkflowStreamHandle.pin`. + """ + self._handle = handle + self._converter = converter + self._codec = codec + self._batch_interval = batch_interval + self._describe = describe + self._buffered: list[tuple[str, Any]] = [] + self._flusher: asyncio.Task[None] | None = None + self._wake = asyncio.Event() + self._closed = False + + @classmethod + def create( + cls, + client: Client, + workflow_id: str, + *, + owner_run_id: str = "", + batch_interval: timedelta = DEFAULT_BATCH_INTERVAL, + ) -> WorkflowStreamClient: + """Open the stream owned by ``workflow_id``. + + Without ``owner_run_id`` the current run is looked up on the first + call and the handle pinned to it, so a reader following across a + continue-as-new sees the run end rather than being moved to the + successor's stream at a stale offset. + """ + return cls( + _stream_client(client).workflow_stream( + workflow_id, owner_run_id=owner_run_id + ), + client.data_converter.payload_converter, + batch_interval, + codec=client.data_converter.payload_codec, + describe=client.get_workflow_handle(workflow_id).describe, + ) + + @classmethod + def from_within_activity( + cls, *, batch_interval: timedelta = DEFAULT_BATCH_INTERVAL + ) -> WorkflowStreamClient: + """Open the stream owned by the Workflow that scheduled this Activity.""" + info = activity.info() + if info.workflow_id is None: + raise RuntimeError( + "no Workflow stream to open: this Activity was not started by a " + "Workflow" + ) + # The Activity's output belongs to the run that scheduled it, and the + # Activity already knows which run that is. + return cls.create( + activity.client(), + info.workflow_id, + owner_run_id=info.workflow_run_id or "", + batch_interval=batch_interval, + ) + + async def __aenter__(self) -> WorkflowStreamClient: + """Start the background flusher.""" + self._flusher = asyncio.create_task(self._run_flusher()) + return self + + async def __aexit__(self, *_exc: object) -> None: + """Drain what is buffered before letting the caller go. + + An Activity that returned with a batch still buffered would have + reported work its readers never saw. The flusher is asked to stop + rather than cancelled: a cancel landing inside its append would + unwind with the batch it had already taken off the buffer. + """ + self._closed = True + self._wake.set() + if self._flusher is not None: + await self._flusher + self._flusher = None + await self.flush() + + async def _pin(self) -> None: + if self._handle.owner_run_id or self._describe is None: + return + self._handle.pin((await self._describe()).run_id) + + @overload + def topic(self, name: str) -> TopicHandle[Any]: ... + + @overload + def topic(self, name: str, *, type: type[T]) -> TopicHandle[T]: ... + + def topic(self, name: str, *, type: type = object) -> TopicHandle[Any]: + """Bind a topic on this Workflow's stream.""" + return TopicHandle(self, name, type) + + async def get_offset(self) -> int: + """Where the stream currently ends. + + A reader that wants only what comes next starts here. + """ + await self._pin() + return (await self._handle.describe()).head_offset + + async def subscribe( + self, + *, + topics: Sequence[str] = (), + from_offset: int = 0, + result_type: type | None = None, + poll_cooldown: timedelta | None = None, + ) -> AsyncIterator[WorkflowStreamItem[Any]]: + """Yield items from ``from_offset`` as they arrive. + + ``poll_cooldown`` is accepted and ignored. It paced a client that had + to re-ask; the server parks this read until something arrives. + """ + del poll_cooldown + await self._pin() + async for entry in self._handle.follow(from_offset=from_offset, topics=topics): + record = await self._decoded(entry.record) + yield WorkflowStreamItem( + topic=record.topic, + data=_decode(self._converter, record, result_type), + offset=entry.offset, + ) + + async def poll_raw( + self, + *, + topics: Sequence[str] = (), + from_offset: int = 0, + wait: bool = True, + ) -> RawPage: + """One read, with the records left as they were stored. + + For a caller that forwards items on rather than using them. A gateway + would only have to encode again what this decoded. + """ + await self._pin() + page = await self._handle.poll( + from_offset=from_offset, topics=topics, wait=wait + ) + return RawPage( + items=[ + WorkflowStreamItem( + topic=entry.record.topic, data=entry.record, offset=entry.offset + ) + for entry in page.entries + ], + next_offset=page.next_offset, + head_offset=page.head_offset, + closed=page.closed, + ) + + async def flush(self) -> None: + """Append everything buffered as one batch.""" + pending, self._buffered = self._buffered, [] + if not pending: + return + await self._pin() + records = [ + await self._encoded(_record(self._converter, topic, value)) + for topic, value in pending + ] + await self._handle.append(*records) + + def _buffer(self, topic: str, value: Any) -> None: + self._buffered.append((topic, value)) + + def _flush_soon(self) -> None: + self._wake.set() + + async def _run_flusher(self) -> None: + """Append on a fixed cadence, so a slow producer still gets delivered.""" + while not self._closed: + try: + await asyncio.wait_for( + self._wake.wait(), self._batch_interval.total_seconds() + ) + except asyncio.TimeoutError: + pass + self._wake.clear() + await self.flush() + + async def _encoded(self, record: StreamRecord) -> StreamRecord: + if self._codec is not None and record.HasField("body"): + record.body.CopyFrom((await self._codec.encode([record.body]))[0]) + return record + + async def _decoded(self, record: StreamRecord) -> StreamRecord: + if self._codec is not None and record.HasField("body"): + record.body.CopyFrom((await self._codec.decode([record.body]))[0]) + return record + + +def _stream_client(client: Client) -> StreamClient: + return shared_client(client.service_client.config.target_host, client.namespace) diff --git a/tests/contrib/test_server_streams.py b/tests/contrib/test_server_streams.py new file mode 100644 index 000000000..540f7f7d7 --- /dev/null +++ b/tests/contrib/test_server_streams.py @@ -0,0 +1,161 @@ +"""The Workflow Streams surface over a server-side stream. + +Both producers and the consumer, since the point of the surface is that an +application does not have to know which of them wrote a given item. Needs a +Temporal server built from the AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/contrib/test_server_streams.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from dataclasses import dataclass +from datetime import timedelta + +import pytest + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.contrib.server_streams import WorkflowStream, WorkflowStreamClient +from temporalio.worker import Worker + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream service", +) + +TOPIC = "turn_events" + + +@dataclass +class Event: + source: str + text: str + + +@activity.defn +async def emit_from_activity(count: int) -> None: + async with WorkflowStreamClient.from_within_activity() as client: + events = client.topic(TOPIC, type=Event) + for i in range(count): + events.publish(Event(source="activity", text=f"token {i}")) + + +@workflow.defn +class Emitting: + def __init__(self) -> None: + self._events = WorkflowStream().topic(TOPIC, type=Event) + + @workflow.run + async def run(self, count: int) -> None: + self._events.publish(Event(source="workflow", text="turn started")) + await workflow.execute_activity( + emit_from_activity, + count, + start_to_close_timeout=timedelta(seconds=30), + ) + self._events.publish(Event(source="workflow", text="turn ended")) + + +async def test_both_producers_reach_one_subscriber() -> None: + client = await Client.connect(TARGET or "") + task_queue = "ss-tq-" + uuid.uuid4().hex[:8] + wf_id = "ss-wf-" + uuid.uuid4().hex[:8] + tokens = 4 + + seen: list[Event] = [] + offsets: list[int] = [] + + async with Worker( + client, + task_queue=task_queue, + workflows=[Emitting], + activities=[emit_from_activity], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + Emitting.run, tokens, id=wf_id, task_queue=task_queue + ) + + # Subscribed before the Workflow has published anything, which is what + # a consumer attaching to a session does. + stream = WorkflowStreamClient.create(client, wf_id) + + async def read() -> None: + async for item in stream.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ): + seen.append(item.data) + offsets.append(item.offset) + if len(seen) == tokens + 2: + return + + reading = asyncio.ensure_future(read()) + await asyncio.wait_for(handle.result(), timeout=60) + await asyncio.wait_for(reading, timeout=60) + + # The Workflow's own publishes bracket the Activity's, and both are on one + # log in the order the server took them. + assert [e.source for e in seen] == ["workflow"] + ["activity"] * tokens + [ + "workflow" + ] + assert seen[0].text == "turn started" + assert seen[-1].text == "turn ended" + assert offsets == list(range(tokens + 2)) + + # A consumer that arrives after the fact reads the same thing, and is not + # left tailing: the Workflow has ended, so its stream is finished and the + # subscription ends on its own. + late = WorkflowStreamClient.create(client, wf_id) + assert await late.get_offset() == tokens + 2 + replayed = [ + item.data.text + async for item in late.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ) + ] + assert replayed == [e.text for e in seen] + + +async def test_a_reader_resumes_from_an_offset_it_was_given() -> None: + client = await Client.connect(TARGET or "") + task_queue = "ss-resume-tq-" + uuid.uuid4().hex[:8] + wf_id = "ss-resume-wf-" + uuid.uuid4().hex[:8] + + async with Worker( + client, + task_queue=task_queue, + workflows=[Emitting], + activities=[emit_from_activity], + max_cached_workflows=0, + ): + await client.execute_workflow(Emitting.run, 3, id=wf_id, task_queue=task_queue) + + stream = WorkflowStreamClient.create(client, wf_id) + first = [ + item + async for item in stream.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ) + ] + assert len(first) == 5 + + # Resuming past the second item skips exactly the two before it, so the + # offset a reader was handed is the position it means. + resumed = [ + item.data.text + async for item in stream.subscribe( + topics=[TOPIC], from_offset=first[1].offset + 1, result_type=Event + ) + ] + assert resumed == [item.data.text for item in first[2:]] + + # The raw page carries the stored records themselves. + raw = await stream.poll_raw(topics=[TOPIC], from_offset=0, wait=False) + assert [item.data.topic for item in raw.items] == [TOPIC] * 5 + assert raw.closed and raw.next_offset == raw.head_offset == 5 From f2511f7d1a1bad49252e218650a28c9d3e850cef Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:46:43 -0700 Subject: [PATCH 04/27] Gave the contrib client a producer identity and a pending batch. The module says an application swaps the import and keeps its code, but the append carried no identity and took the batch off the buffer before the await, so a retried Activity wrote twice and a failed append lost what it held. --- temporalio/contrib/server_streams/__init__.py | 64 +++++++++++--- tests/contrib/test_server_streams_flush.py | 86 +++++++++++++++++++ 2 files changed, 136 insertions(+), 14 deletions(-) create mode 100644 tests/contrib/test_server_streams_flush.py diff --git a/temporalio/contrib/server_streams/__init__.py b/temporalio/contrib/server_streams/__init__.py index 95eb48d7e..0f6c5142d 100644 --- a/temporalio/contrib/server_streams/__init__.py +++ b/temporalio/contrib/server_streams/__init__.py @@ -4,7 +4,9 @@ by a Temporal-owned log instead of by Signals and Updates. An application swaps the import and keeps its code: publishing from a Workflow is still a plain call, publishing from an Activity is still a buffered handle, and a consumer -still subscribes by topic from an offset. +still subscribes by topic from an offset. An Activity's appends carry its own +id and attempt, so a retried Activity's repeat is deduplicated by the server +rather than written twice. What changes is underneath. A publish is a Workflow Command whose payload never enters History, so History gets one fixed-size event per Workflow Task rather @@ -128,7 +130,7 @@ def publish(self, value: T | RawValue) -> None: transition. The Worker's payload codec applies to the body as it does to any other payload the Workflow sends. """ - workflow.append_stream_records( + workflow._append_stream_records( [_record(workflow.payload_converter(), self._name, value)] ) @@ -220,6 +222,7 @@ def __init__( *, codec: PayloadCodec | None = None, describe: Callable[[], Awaitable[WorkflowExecutionDescription]] | None = None, + producer_id: str = "", ) -> None: """Prefer :meth:`create` or :meth:`from_within_activity`. @@ -227,13 +230,18 @@ def __init__( a namespace whose payloads are encoded agrees with the Worker, whose payload visitor applies the same codec to the Workflow's publishes. ``describe`` is how an unpinned handle learns which run it follows; - see :meth:`WorkflowStreamHandle.pin`. + see :meth:`WorkflowStreamHandle.pin`. ``producer_id`` is who the + appends are written as; without one they are at-least-once, because + the server has nothing to deduplicate a retry against. """ self._handle = handle self._converter = converter self._codec = codec self._batch_interval = batch_interval self._describe = describe + self._producer_id = producer_id + self._sequence = 0 + self._pending: tuple[list[StreamRecord], int] | None = None self._buffered: list[tuple[str, Any]] = [] self._flusher: asyncio.Task[None] | None = None self._wake = asyncio.Event() @@ -247,13 +255,16 @@ def create( *, owner_run_id: str = "", batch_interval: timedelta = DEFAULT_BATCH_INTERVAL, + producer_id: str = "", ) -> WorkflowStreamClient: """Open the stream owned by ``workflow_id``. Without ``owner_run_id`` the current run is looked up on the first call and the handle pinned to it, so a reader following across a continue-as-new sees the run end rather than being moved to the - successor's stream at a stale offset. + successor's stream at a stale offset. Without ``producer_id`` the + appends are at-least-once; inside an Activity, + :meth:`from_within_activity` supplies one. """ return cls( _stream_client(client).workflow_stream( @@ -263,6 +274,7 @@ def create( batch_interval, codec=client.data_converter.payload_codec, describe=client.get_workflow_handle(workflow_id).describe, + producer_id=producer_id, ) @classmethod @@ -277,12 +289,14 @@ def from_within_activity( "Workflow" ) # The Activity's output belongs to the run that scheduled it, and the - # Activity already knows which run that is. + # Activity already knows which run that is. Its id and attempt are + # also what lets the server drop a batch a retried Activity re-sends. return cls.create( activity.client(), info.workflow_id, owner_run_id=info.workflow_run_id or "", batch_interval=batch_interval, + producer_id=f"{info.activity_id}#{info.attempt}", ) async def __aenter__(self) -> WorkflowStreamClient: @@ -380,16 +394,38 @@ async def poll_raw( ) async def flush(self) -> None: - """Append everything buffered as one batch.""" - pending, self._buffered = self._buffered, [] - if not pending: - return + """Append everything buffered as one batch. + + A batch whose append failed stays pending and goes out again on the + next flush under the sequence it already had, so an append the server + did accept is deduplicated and one it never saw still lands. Nothing + comes off the buffer until there is a batch to replace it with. + + Raises: + temporalio.streams.StreamProducerError: The server holds this + producer's sequence with different content. + """ + if self._pending is not None: + records, sequence = self._pending + else: + if not self._buffered: + return + await self._pin() + # Encoded before the buffer is cleared, so a converter failure + # leaves the values where the caller can still see them. + records = [ + await self._encoded(_record(self._converter, topic, value)) + for topic, value in self._buffered + ] + sequence = self._sequence + self._buffered = [] + self._pending = (records, sequence) await self._pin() - records = [ - await self._encoded(_record(self._converter, topic, value)) - for topic, value in pending - ] - await self._handle.append(*records) + await self._handle.append( + *records, producer_id=self._producer_id, sequence=sequence + ) + self._sequence = sequence + len(records) + self._pending = None def _buffer(self, topic: str, value: Any) -> None: self._buffered.append((topic, value)) diff --git a/tests/contrib/test_server_streams_flush.py b/tests/contrib/test_server_streams_flush.py new file mode 100644 index 000000000..23a72106e --- /dev/null +++ b/tests/contrib/test_server_streams_flush.py @@ -0,0 +1,86 @@ +"""What the buffered client does with a batch, without a server. + +The parity claim in the module docstring is that an application swaps the +import and keeps its code, so the two things the shipped client does on every +flush have to hold here too: the append carries a producer identity, and a +batch whose append failed is kept for the retry rather than lost with it. +""" + +from __future__ import annotations + +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio.api.stream.v1 import StreamRecord +from temporalio.contrib.server_streams import WorkflowStreamClient +from temporalio.converter import DataConverter + + +class _Handle: + """A stream handle that records its appends and can fail the first one.""" + + owner_run_id = "run" + + def __init__(self, *, fail_first: bool = False) -> None: + self.sent: list[tuple[list[StreamRecord], str, int]] = [] + self._fail_first = fail_first + + def pin(self, run_id: str) -> None: + del run_id + + async def append( + self, *records: StreamRecord, producer_id: str = "", sequence: int = 0 + ) -> Any: + self.sent.append((list(records), producer_id, sequence)) + if self._fail_first: + self._fail_first = False + raise ConnectionResetError("the server took it, the reply was lost") + return None + + +def _client(handle: _Handle) -> WorkflowStreamClient: + return WorkflowStreamClient( + handle, # type: ignore[arg-type] + DataConverter.default.payload_converter, + timedelta(seconds=60), + producer_id="act#1", + ) + + +def _bodies(sent: tuple[list[StreamRecord], str, int]) -> list[bytes]: + return [record.body.data for record in sent[0]] + + +async def test_an_append_carries_who_wrote_it_and_where_it_sits() -> None: + handle = _Handle() + client = _client(handle) + client.topic("t").publish({"n": 1}) + client.topic("t").publish({"n": 2}) + await client.flush() + client.topic("t").publish({"n": 3}) + await client.flush() + + # Without an identity the append is at-least-once, which is what the + # shipped client refuses to be. + assert [(who, seq) for _, who, seq in handle.sent] == [("act#1", 0), ("act#1", 2)] + + +async def test_a_batch_whose_append_failed_goes_out_again() -> None: + handle = _Handle(fail_first=True) + client = _client(handle) + client.topic("t").publish({"n": 1}) + with pytest.raises(ConnectionResetError): + await client.flush() + # Taken off the buffer before the await, it would have gone with the + # failure. It goes out again under the sequence it already had, so a copy + # the server did accept is deduplicated and one it never saw lands. + await client.flush() + assert len(handle.sent) == 2 + assert _bodies(handle.sent[0]) == _bodies(handle.sent[1]) + assert [seq for _, _, seq in handle.sent] == [0, 0] + + client.topic("t").publish({"n": 2}) + await client.flush() + assert handle.sent[2][2] == 1, "the next batch continues past the first" From e3b42303b6e3d611b3f5b2118749dbbb4df44eac Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:47:16 -0700 Subject: [PATCH 05/27] Closed only the stream channels a provider opened. One channel is shared per loop, target and namespace, so closing a provider took out channels another one in the same process was still reading through. --- temporalio/client_stream.py | 24 ++++++++---- temporalio/streams/providers/native.py | 40 ++++++++++++++++---- tests/test_client_stream_sharing.py | 52 ++++++++++++++++++++++++++ 3 files changed, 102 insertions(+), 14 deletions(-) create mode 100644 tests/test_client_stream_sharing.py diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index a043b8e03..6e0d47227 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -639,12 +639,22 @@ def shared_client(target_host: str, namespace: str) -> StreamClient: return existing -async def close_shared_clients() -> None: - """Close every shared client this loop opened. - - For a process that is done with streams, and for tests, which open a - loop per case and would otherwise leave a channel behind on each. +async def close_shared_clients(*keys: tuple[str, str]) -> None: + """Close the shared clients this loop opened for ``keys``, or all of them. + + A provider closes the ones it opened, named by ``(target host, + namespace)``: another provider on the same loop may still be reading + through a channel of its own, and taking that out from under it is not + this one's to do. With no keys it closes every one, which is what a + process finished with streams wants, and what a test that opened a loop + of its own wants. """ - per_loop = _shared.pop(asyncio.get_running_loop(), {}) - for client in per_loop.values(): + loop = asyncio.get_running_loop() + if not keys: + per_loop = _shared.pop(loop, {}) + closing = list(per_loop.values()) + else: + per_loop = _shared.get(loop, {}) + closing = [per_loop.pop(key) for key in keys if key in per_loop] + for client in closing: await client.close() diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 9c0cb3dbd..6ec1de705 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -267,11 +267,23 @@ async def _write(self, records: list[WireRecord]) -> Cursor: class NativeStreamHandle: """One workflow's topics from outside, over the stream service.""" - def __init__(self, client: Client, workflow_id: str, run_id: str | None) -> None: - """Address ``workflow_id``'s topics, pinned to ``run_id`` when one is given.""" + def __init__( + self, + client: Client, + workflow_id: str, + run_id: str | None, + *, + opened: set[tuple[str, str]] | None = None, + ) -> None: + """Address ``workflow_id``'s topics, pinned to ``run_id`` when one is given. + + ``opened`` is where this handle records the shared channel it used, so + the provider that made it closes that one and no other. + """ self._client = client self._workflow_id = workflow_id self._run_id = run_id + self._opened = set() if opened is None else opened self._converter = client.data_converter.payload_converter self._codec = client.data_converter.payload_codec self._streams: StreamClient | None = None @@ -280,10 +292,12 @@ def _service(self) -> StreamClient: # Resolved on first use, because the shared channel belongs to the # running loop and a handle may be made before there is one. if self._streams is None: - self._streams = shared_client( + key = ( self._client.service_client.config.target_host, self._client.namespace, ) + self._streams = shared_client(*key) + self._opened.add(key) return self._streams def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: @@ -445,9 +459,15 @@ class NativeStreams(ProviderPlugin): Takes no options: the streams are on the server the client is already connected to. Construct one, pass it to the worker as a plugin and open handles from it anywhere else; :meth:`close` releases the channels this - process opened to the stream service. + provider opened to the stream service. """ + def __init__(self) -> None: + """Create the provider.""" + # What this provider's handles opened, so closing it leaves another + # provider's channels on the same loop alone. + self._opened: set[tuple[str, str]] = set() + def workflow_provider(self) -> _NativeWorkflowProvider: """The workflow half, over the server's commands and delivered ranges.""" return _NativeWorkflowProvider() @@ -456,8 +476,14 @@ def get_stream_handle( self, client: Client, workflow_id: str, *, run_id: str | None = None ) -> NativeStreamHandle: """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain.""" - return NativeStreamHandle(client, workflow_id, run_id) + return NativeStreamHandle(client, workflow_id, run_id, opened=self._opened) async def close(self) -> None: - """Close the channels this process opened to the stream service.""" - await close_shared_clients() + """Close the channels this provider opened to the stream service. + + The application calls this; no worker or client owns the provider's + lifetime, because one provider serves the workers built from a client + and every handle opened outside them. + """ + await close_shared_clients(*self._opened) + self._opened.clear() diff --git a/tests/test_client_stream_sharing.py b/tests/test_client_stream_sharing.py new file mode 100644 index 000000000..793ba122a --- /dev/null +++ b/tests/test_client_stream_sharing.py @@ -0,0 +1,52 @@ +"""Who owns a shared channel, and who may close it. + +One channel is shared per loop, target and namespace, so two providers in one +process can be using the same one. Closing a provider has to leave the other's +channels alone. +""" + +from __future__ import annotations + +import asyncio + +from temporalio import client_stream +from temporalio.streams.providers.native import NativeStreams + + +class _FakeClient: + """Stands in for a StreamClient, which would want a real channel.""" + + def __init__(self) -> None: + self.closed = False + + async def close(self) -> None: + self.closed = True + + +def _put(target: str, namespace: str) -> _FakeClient: + fake = _FakeClient() + per_loop = client_stream._shared.setdefault(asyncio.get_running_loop(), {}) + per_loop[(target, namespace)] = fake # type: ignore[assignment] + return fake + + +async def test_closing_named_clients_leaves_the_others_open() -> None: + mine = _put("host-a:7233", "ns") + theirs = _put("host-b:7233", "ns") + await client_stream.close_shared_clients(("host-a:7233", "ns")) + assert mine.closed + assert not theirs.closed, "another provider is still reading through it" + await client_stream.close_shared_clients() + assert theirs.closed + + +async def test_a_provider_closes_only_what_its_own_handles_opened() -> None: + mine = _put("host-a:7233", "ns") + theirs = _put("host-b:7233", "ns") + + provider = NativeStreams() + provider._opened.add(("host-a:7233", "ns")) + await provider.close() + assert mine.closed + assert not theirs.closed + await client_stream.close_shared_clients() From a223504ce78779e638cb9891652568ffa7d79fa2 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:47:53 -0700 Subject: [PATCH 06/27] Named the way out of a stream continuity failure. The check fails the Workflow Task and keeps failing it, because the range is recorded as consumed and is never sent again, so the message has to say what a reader can do about it. The copied batch limits now say what breaks if the server's differ. --- temporalio/worker/_workflow_instance.py | 26 +++++++++++++++++++++++-- tests/worker/test_workflow_stream.py | 9 +++++++++ 2 files changed, 33 insertions(+), 2 deletions(-) diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 9b1e1ea0b..9503e2e5c 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -276,6 +276,27 @@ def create_instance(self, det: WorkflowInstanceDetails) -> WorkflowInstance: # Match the server's per-batch limits. A record over its limit is refused where # it is published, because a rejected command would be reissued on every # replay; a task's records are split into commands that fit the batch limits. +# +# Copied rather than learned: the activation does not carry them and the +# server does not report them, so this is a second copy of a number somebody +# else owns. If the server lowers one, or makes it per namespace, the split +# here stops fitting and the command is rejected on every replay, which is +# the failure the split exists to avoid. Carrying them on the activation is +# what would fix that, and it needs a Core and server change. +_STREAM_CONTINUITY_REMEDY = ( + "This fails the Workflow Task and will keep failing it, because the range " + "is recorded as consumed and will not be sent again. Reset the workflow to " + "before the subscription to start its stream reading over, or terminate it " + "if its output is no longer wanted." +) + +_STREAM_CONTINUITY_REMEDY = ( + "This fails the Workflow Task and will keep failing it, because the range " + "is recorded as consumed and will not be sent again. Reset the workflow to " + "before the subscription to start its stream reading over, or terminate it " + "if its output is no longer wanted." +) + _MAX_STREAM_RECORDS_PER_BATCH = 1000 _MAX_STREAM_RECORD_BYTES = 1 << 20 _MAX_STREAM_BATCH_BYTES = 2 << 20 @@ -345,12 +366,13 @@ def extend( if to_offset - from_offset != len(records): raise RuntimeError( f"stream {self._stream_id!r} delivered {len(records)} records " - f"for offsets [{from_offset}, {to_offset})" + f"for offsets [{from_offset}, {to_offset}). {_STREAM_CONTINUITY_REMEDY}" ) if self._next_offset is not None and from_offset != self._next_offset: raise RuntimeError( f"stream {self._stream_id!r} delivered offsets [{from_offset}, " - f"{to_offset}) but the last range ended at {self._next_offset}" + f"{to_offset}) but the last range ended at {self._next_offset}. " + f"{_STREAM_CONTINUITY_REMEDY}" ) self._next_offset = to_offset # An empty range still counts as a delivery, but there is nothing to diff --git a/tests/worker/test_workflow_stream.py b/tests/worker/test_workflow_stream.py index 9083e74a9..568e6adaf 100644 --- a/tests/worker/test_workflow_stream.py +++ b/tests/worker/test_workflow_stream.py @@ -286,3 +286,12 @@ def test_a_task_over_the_batch_limits_is_split_into_commands() -> None: stub.publish(big, big, record(b"y" * room)) stub.flush() assert [len(c.append_stream_records.records) for c in stub.commands] == [2, 1] + + +async def test_the_continuity_failure_says_what_to_do_about_it() -> None: + buffer = _StreamBuffer("s") + buffer.extend([record(b"one")], 0, 1) + with pytest.raises(RuntimeError) as failed: + buffer.extend([record(b"two")], 5, 6) + # The task fails and keeps failing, so the message has to name the way out. + assert "Reset the workflow" in str(failed.value) From 2b3577fc489ae55817fa69461668313d14467d02 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:48:12 -0700 Subject: [PATCH 07/27] Stopped holding records for a reader that closed. There is no unsubscribe command, so the server keeps delivering for the life of the run and the buffer grew for a reader nobody would read again. The raw stream API that takes stream ids and protos is private now: the typed surface beside it is the one an application should reach for. --- temporalio/streams/providers/native.py | 21 +++++++-- temporalio/worker/_workflow_instance.py | 51 ++++++++++++++------- temporalio/workflow/__init__.py | 25 +++++++---- temporalio/workflow/_context.py | 35 +++++++++++---- tests/worker/test_workflow_stream.py | 56 ++++++++++++++++++++++++ tests/worker/test_workflow_stream_e2e.py | 12 ++--- 6 files changed, 159 insertions(+), 41 deletions(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 6ec1de705..83aeacfa7 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -113,11 +113,26 @@ def __init__(self, stream_id: str, run_id: str) -> None: async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: if self._closed: raise StopAsyncIteration - delivered = await workflow.read_stream_records(self._stream_id) + delivered = await workflow._read_stream_records(self._stream_id) + if self._closed: + # Closed while this was parked; the buffer woke it with nothing. + raise StopAsyncIteration return [(_cursor(self._run_id, item.offset), item.record) for item in delivered] def close(self) -> None: + """Stop reading, and stop keeping what the server keeps delivering. + + The server has no unsubscribe command, so ranges keep arriving on + every Workflow Task for the life of the run. What this ends is the + reading and the keeping: nothing further is held for this stream, so + a run that closes a reader early does not grow for the rest of its + life. The subscription itself, and the delivery it costs each task, + stay until the run ends. + """ + if self._closed: + return self._closed = True + workflow._close_stream_records(self._stream_id) class _NativeWriteSink: @@ -128,7 +143,7 @@ def publish(self, record: WireRecord) -> None: # Held by the runtime until the task completes, when the task's # records on this topic become one command the server applies with # the task: rule 1 through the server's own commit. - workflow.append_stream_records([record], stream_id=self._topic) + workflow._append_stream_records([record], stream_id=self._topic) class _NativeWorkflowProvider: @@ -145,7 +160,7 @@ def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: f"cursor {after.token!r} names another run; a run's stream is its own" ) start = named[1] + 1 - workflow.subscribe_stream(topic, start_offset=start) + workflow._subscribe_stream(topic, start_offset=start) return _NativeReadSource(topic, run_id) def open_writer(self, topic: str) -> WriteSink: diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 9503e2e5c..69182d192 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -290,13 +290,6 @@ def create_instance(self, det: WorkflowInstanceDetails) -> WorkflowInstance: "if its output is no longer wanted." ) -_STREAM_CONTINUITY_REMEDY = ( - "This fails the Workflow Task and will keep failing it, because the range " - "is recorded as consumed and will not be sent again. Reset the workflow to " - "before the subscription to start its stream reading over, or terminate it " - "if its output is no longer wanted." -) - _MAX_STREAM_RECORDS_PER_BATCH = 1000 _MAX_STREAM_RECORD_BYTES = 1 << 20 _MAX_STREAM_BATCH_BYTES = 2 << 20 @@ -344,12 +337,35 @@ class _StreamBuffer: def __init__(self, stream_id: str = "") -> None: self._stream_id = stream_id - self._records: list[temporalio.workflow.DeliveredStreamRecord] = [] + self._records: list[temporalio.workflow._DeliveredStreamRecord] = [] self._waiters: list[asyncio.Future] = [] # Where the next range has to start. Unknown until the first one # arrives, because a subscription may start wherever the stream is and # the server is the one that resolves that. self._next_offset: int | None = None + self._closed = False + + @property + def closed(self) -> bool: + """Whether workflow code has said it wants no more of this stream.""" + return self._closed + + def close(self) -> None: + """Stop keeping what arrives, and let go of what is held. Idempotent. + + There is no unsubscribe command, so the server keeps delivering for + the life of the run. Holding those records would grow the instance + without bound for a reader nobody will read again. Dropping them is + replay-safe because the close happens at the same point of the same + workflow code every time, so the same ranges are dropped; continuity + is still tracked, so a range that repeats or skips is still caught. + """ + self._closed = True + self._records = [] + waiters, self._waiters = self._waiters, [] + for waiter in waiters: + if not waiter.done(): + waiter.set_result(None) def extend( self, @@ -376,8 +392,10 @@ def extend( ) self._next_offset = to_offset # An empty range still counts as a delivery, but there is nothing to - # hand a reader, so only a non-empty one wakes anyone. - if not records: + # hand a reader, so only a non-empty one wakes anyone. A closed buffer + # counts the range and keeps nothing: continuity is still checked + # above, and nobody is left to read what it held. + if not records or self._closed: return # Offsets are dense inside a delivered range and the range arrives in # order, so counting from its start is the position rather than an @@ -388,7 +406,7 @@ def extend( kept = temporalio.api.stream.v1.StreamRecord() kept.CopyFrom(record) self._records.append( - temporalio.workflow.DeliveredStreamRecord( + temporalio.workflow._DeliveredStreamRecord( record=kept, offset=from_offset + index ) ) @@ -397,12 +415,12 @@ def extend( if not waiter.done(): waiter.set_result(None) - def take(self) -> list[temporalio.workflow.DeliveredStreamRecord]: + def take(self) -> list[temporalio.workflow._DeliveredStreamRecord]: taken, self._records = self._records, [] return taken def put_back( - self, records: Sequence[temporalio.workflow.DeliveredStreamRecord] + self, records: Sequence[temporalio.workflow._DeliveredStreamRecord] ) -> None: """Return an unread tail to the front of the buffer.""" self._records[:0] = records @@ -1534,15 +1552,18 @@ def _flush_stream_appends(self) -> None: commands.insert(insert_at, command) insert_at += 1 + def workflow_close_stream_records(self, stream_id: str) -> None: + self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)).close() + async def workflow_read_stream_records( self, stream_id: str, max_records: int - ) -> list[temporalio.workflow.DeliveredStreamRecord]: + ) -> list[temporalio.workflow._DeliveredStreamRecord]: # Ranges arrive on Workflow Tasks, and a query activation carries none, # so without this the read waits on a future nothing can resolve and the # query times out with nothing to say why. self._assert_not_read_only("read stream") buffer = self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)) - while not len(buffer): + while not len(buffer) and not buffer.closed: await buffer.wait_future() taken = buffer.take() if max_records and len(taken) > max_records: diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index 18d04cb0c..68822fd49 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -57,8 +57,7 @@ as_completed, wait, ) -from ._context import ( - DeliveredStreamRecord, +from ._context import ( # noqa: F401 Info, ParentInfo, RootInfo, @@ -66,7 +65,6 @@ _current_update_info, _Runtime, _set_current_update_info, - append_stream_records, cancellation_reason, current_update_info, deprecate_patch, @@ -88,11 +86,9 @@ payload_converter, random, random_seed, - read_stream_records, register_random_seed_callback, set_current_details, sleep, - subscribe_stream, time, time_ns, upsert_memo, @@ -101,6 +97,21 @@ uuid7, wait_condition, ) +from ._context import ( + _append_stream_records as _append_stream_records, +) +from ._context import ( + _close_stream_records as _close_stream_records, +) +from ._context import ( + _DeliveredStreamRecord as _DeliveredStreamRecord, +) +from ._context import ( + _read_stream_records as _read_stream_records, +) +from ._context import ( + _subscribe_stream as _subscribe_stream, +) from ._definition import ( DynamicWorkflowConfig, _Definition, @@ -237,10 +248,6 @@ "upsert_search_attributes", "uuid4", "uuid7", - "DeliveredStreamRecord", - "read_stream_records", - "append_stream_records", - "subscribe_stream", "wait_condition", "DynamicWorkflowConfig", "defn", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 8b6c07782..2b9c0a8f8 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -324,10 +324,13 @@ def workflow_append_stream_records( records: Sequence[temporalio.api.stream.v1.StreamRecord], ) -> None: ... + @abstractmethod + def workflow_close_stream_records(self, stream_id: str) -> None: ... + @abstractmethod async def workflow_read_stream_records( self, stream_id: str, max_records: int - ) -> list[DeliveredStreamRecord]: ... + ) -> list[_DeliveredStreamRecord]: ... @abstractmethod def workflow_get_current_history_length(self) -> int: ... @@ -965,11 +968,13 @@ async def sleep(duration: float | timedelta, *, summary: str | None = None) -> N ) -def subscribe_stream(stream_id: str, *, start_offset: int = 0) -> None: +def _subscribe_stream( # type: ignore[reportUnusedFunction] + stream_id: str, *, start_offset: int = 0 +) -> None: """Subscribe this workflow to a server-side stream. From here on its Workflow Tasks carry the ranges it has not consumed yet, - and :func:`read_stream_records` returns them. Safe to call again: a second + and :func:`_read_stream_records` returns them. Safe to call again: a second subscription to a stream this run already consumes does not move its cursor, though it does write one event. Calling it on every replay is harmless because replay matches the command to the event already recorded. @@ -990,7 +995,7 @@ def subscribe_stream(stream_id: str, *, start_offset: int = 0) -> None: _Runtime.current().workflow_subscribe_stream(stream_id, start_offset) -def append_stream_records( +def _append_stream_records( # type: ignore[reportUnusedFunction] records: Sequence[temporalio.api.stream.v1.StreamRecord], *, stream_id: str = "", @@ -1019,7 +1024,7 @@ def append_stream_records( @dataclass(frozen=True) -class DeliveredStreamRecord: +class _DeliveredStreamRecord: """One record a consuming workflow was given, with where it sat.""" record: temporalio.api.stream.v1.StreamRecord @@ -1027,15 +1032,29 @@ class DeliveredStreamRecord: """Its position in the whole stream, which is what a reader resumes from.""" -async def read_stream_records( +def _close_stream_records(stream_id: str) -> None: # type: ignore[reportUnusedFunction] + """Say this workflow wants no more of ``stream_id``. + + There is no unsubscribe command, so the server keeps delivering for the + life of the run; this drops what arrives instead of holding it for a + reader that has gone. Deterministic on replay, because the same workflow + code closes at the same point and the same ranges are dropped. + + Args: + stream_id: Stream to stop keeping records for. + """ + _Runtime.current().workflow_close_stream_records(stream_id) + + +async def _read_stream_records( # type: ignore[reportUnusedFunction] stream_id: str, *, max_records: int = 0 -) -> list[DeliveredStreamRecord]: +) -> list[_DeliveredStreamRecord]: """Read the next records of a server-side stream this workflow consumes. Waits until at least one record is available. Ranges arrive on Workflow Tasks, and only the offsets they covered are written to History, so this is deterministic on replay: the server re-supplies the same ranges by reading - the stream again. Subscribe first with :func:`subscribe_stream`; this only + the stream again. Subscribe first with :func:`_subscribe_stream`; this only reads what has already been delivered to this workflow. Args: diff --git a/tests/worker/test_workflow_stream.py b/tests/worker/test_workflow_stream.py index 568e6adaf..676b29c90 100644 --- a/tests/worker/test_workflow_stream.py +++ b/tests/worker/test_workflow_stream.py @@ -288,6 +288,37 @@ def test_a_task_over_the_batch_limits_is_split_into_commands() -> None: assert [len(c.append_stream_records.records) for c in stub.commands] == [2, 1] +async def test_a_closed_buffer_keeps_nothing_and_still_checks_continuity() -> None: + # There is no unsubscribe command, so the server keeps delivering for the + # life of the run. A reader that closed would otherwise grow the instance + # for the rest of it. + buffer = _StreamBuffer("s") + buffer.extend([record(b"one")], 0, 1) + assert len(buffer) == 1 + + buffer.close() + assert buffer.closed + assert len(buffer) == 0, "what it held is let go of, not kept for nobody" + + buffer.extend([record(b"two"), record(b"three")], 1, 3) + assert len(buffer) == 0 + # Continuity is still tracked across what it dropped, so a range that + # repeats or skips is caught rather than passing unnoticed. + with pytest.raises(RuntimeError, match="last range ended at 3"): + buffer.extend([record(b"four")], 9, 10) + buffer.extend([record(b"four")], 3, 4) + + +async def test_closing_a_buffer_wakes_a_reader_parked_on_it() -> None: + buffer = _StreamBuffer("s") + waiter = buffer.wait_future() + buffer.close() + # Woken rather than left parked: the reader has to unwind, and nothing + # will ever arrive for it again. + await asyncio.wait_for(waiter, 5) + assert waiter.done() + + async def test_the_continuity_failure_says_what_to_do_about_it() -> None: buffer = _StreamBuffer("s") buffer.extend([record(b"one")], 0, 1) @@ -295,3 +326,28 @@ async def test_the_continuity_failure_says_what_to_do_about_it() -> None: buffer.extend([record(b"two")], 5, 6) # The task fails and keeps failing, so the message has to name the way out. assert "Reset the workflow" in str(failed.value) + + +def test_the_raw_workflow_stream_api_is_not_public() -> None: + # A second workflow surface taking raw stream ids and raw protos, beside + # the typed one, is not what an application should reach for. It stays + # reachable under its private name, which is what the provider and the + # contrib surface use. + import temporalio.workflow as wf + + for name in ( + "subscribe_stream", + "append_stream_records", + "read_stream_records", + "DeliveredStreamRecord", + ): + assert name not in wf.__all__ + assert not hasattr(wf, name) + for name in ( + "_subscribe_stream", + "_append_stream_records", + "_read_stream_records", + "_close_stream_records", + "_DeliveredStreamRecord", + ): + assert hasattr(wf, name) diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py index 1d974b401..694b90b71 100644 --- a/tests/worker/test_workflow_stream_e2e.py +++ b/tests/worker/test_workflow_stream_e2e.py @@ -270,18 +270,18 @@ class PublishAndRead: @workflow.run async def run(self) -> list[str]: - workflow.append_stream_records( + workflow._append_stream_records( [_record(b"alpha", "progress"), _record(b"beta", "progress")], stream_id="output", ) - workflow.append_stream_records([_record(b"gamma")], stream_id="output") + workflow._append_stream_records([_record(b"gamma")], stream_id="output") # A name this workflow has not written yet still names a stream it # owns, so subscribing creates the one the later publish lands in. - workflow.subscribe_stream("output", start_offset=0) + workflow._subscribe_stream("output", start_offset=0) received: list[str] = [] while len(received) < 3: - for item in await workflow.read_stream_records("output"): + for item in await workflow._read_stream_records("output"): received.append(item.record.body.data.decode()) return received @@ -331,10 +331,10 @@ class ConsumeAcrossTasks: @workflow.run async def run(self, stream_id: str, expected: int) -> list[str]: - workflow.subscribe_stream(stream_id, start_offset=0) + workflow._subscribe_stream(stream_id, start_offset=0) seen: list[str] = [] while len(seen) < expected: - for item in await workflow.read_stream_records(stream_id): + for item in await workflow._read_stream_records(stream_id): seen.append(item.record.body.data.decode()) return seen From 55eef1855d0fc1eedbf94541803eb5f8a7f465e6 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 15:51:13 -0700 Subject: [PATCH 08/27] Started the native producer's numbering at one. Zero on the wire says a producer does not number its records, and this one does, so it cannot start there. --- temporalio/streams/providers/native.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 83aeacfa7..c1595d039 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -201,7 +201,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._last = BEGINNING @property From 9160c82911217fb37351909075b6e37e71194251 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 16:05:03 -0700 Subject: [PATCH 09/27] Renamed the stream command fields the workflow writes. Core names these fields after what they carry, so the writes and the helpers that feed them read the same way. --- temporalio/streams/providers/native.py | 2 +- temporalio/worker/_workflow_instance.py | 18 +++++++++------- temporalio/workflow/_context.py | 26 ++++++++++++++---------- tests/worker/test_workflow_stream.py | 8 ++++---- tests/worker/test_workflow_stream_e2e.py | 4 ++-- 5 files changed, 33 insertions(+), 25 deletions(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index c1595d039..d0026b7ed 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -143,7 +143,7 @@ def publish(self, record: WireRecord) -> None: # Held by the runtime until the task completes, when the task's # records on this topic become one command the server applies with # the task: rule 1 through the server's own commit. - workflow._append_stream_records([record], stream_id=self._topic) + workflow._append_stream_records([record], stream_name=self._topic) class _NativeWorkflowProvider: diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 69182d192..fd65fd1cc 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -1499,18 +1499,22 @@ def workflow_get_current_deployment_version( def get_info(self) -> temporalio.workflow.Info: return self._info - def workflow_subscribe_stream(self, stream_id: str, start_offset: int) -> None: + def workflow_subscribe_stream( + self, stream_name_or_id: str, start_offset: int + ) -> None: # Reissued on every replay, so the buffer has to exist before the first # range arrives and the command has to be harmless the second time. A # repeat subscription leaves the server-side cursor where it is. - self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)) + self._stream_buffers.setdefault( + stream_name_or_id, _StreamBuffer(stream_name_or_id) + ) command = self._add_command() - command.subscribe_stream.stream_id = stream_id + command.subscribe_stream.stream_name_or_id = stream_name_or_id command.subscribe_stream.start_offset = start_offset def workflow_append_stream_records( self, - stream_id: str, + stream_name: str, records: Sequence[temporalio.api.stream.v1.StreamRecord], ) -> None: self._assert_not_read_only("append stream records") @@ -1530,7 +1534,7 @@ def workflow_append_stream_records( kept.append(copy) # Held until the task completes, so a task's publishes on one stream # become one command and one History event however many there were. - self._stream_appends.setdefault(stream_id, []).extend(kept) + self._stream_appends.setdefault(stream_name, []).extend(kept) def _flush_stream_appends(self) -> None: appends, self._stream_appends = self._stream_appends, {} @@ -1544,10 +1548,10 @@ def _flush_stream_appends(self) -> None: if _is_completion_command(command): insert_at = index break - for stream_id, records in appends.items(): + for stream_name, records in appends.items(): for batch in _stream_batches(records): command = temporalio.bridge.proto.workflow_commands.WorkflowCommand() - command.append_stream_records.stream_id = stream_id + command.append_stream_records.stream_name = stream_name command.append_stream_records.records.extend(batch) commands.insert(insert_at, command) insert_at += 1 diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index 2b9c0a8f8..8d09cb8d3 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -315,12 +315,14 @@ def workflow_get_current_deployment_version( ) -> temporalio.common.WorkerDeploymentVersion | None: ... @abstractmethod - def workflow_subscribe_stream(self, stream_id: str, start_offset: int) -> None: ... + def workflow_subscribe_stream( + self, stream_name_or_id: str, start_offset: int + ) -> None: ... @abstractmethod def workflow_append_stream_records( self, - stream_id: str, + stream_name: str, records: Sequence[temporalio.api.stream.v1.StreamRecord], ) -> None: ... @@ -969,7 +971,7 @@ async def sleep(duration: float | timedelta, *, summary: str | None = None) -> N def _subscribe_stream( # type: ignore[reportUnusedFunction] - stream_id: str, *, start_offset: int = 0 + stream_name_or_id: str, *, start_offset: int = 0 ) -> None: """Subscribe this workflow to a server-side stream. @@ -979,26 +981,27 @@ def _subscribe_stream( # type: ignore[reportUnusedFunction] cursor, though it does write one event. Calling it on every replay is harmless because replay matches the command to the event already recorded. - Only the stream id and start offset go to the server. The rest of the + Only the name or id and the start offset go to the server. The rest of the stream's addressing is resolved there, because a workflow cannot look it up without doing I/O and a value it carried would be a reading rather than a fact. A name this workflow has not written yet names a stream it owns, and subscribing creates it. Args: - stream_id: Stream to consume: the name of one this workflow owns, or - the id of a standalone stream. + stream_name_or_id: Stream to consume: the name of one this workflow + owns, or the id of a standalone stream. The server tries them in + that order. start_offset: Where to start. Negative means from wherever the stream is when the subscription is registered; the server resolves that once and records it, so replay does not resolve it again. """ - _Runtime.current().workflow_subscribe_stream(stream_id, start_offset) + _Runtime.current().workflow_subscribe_stream(stream_name_or_id, start_offset) def _append_stream_records( # type: ignore[reportUnusedFunction] records: Sequence[temporalio.api.stream.v1.StreamRecord], *, - stream_id: str = "", + stream_name: str = "", ) -> None: """Publish records to a server-side stream this workflow owns. @@ -1013,14 +1016,15 @@ def _append_stream_records( # type: ignore[reportUnusedFunction] Args: records: Records to append, in order. The server stores each with an empty ``producer_id``, because the workflow is the producer. - stream_id: Stream to publish to. Empty means the workflow's default - output stream. + stream_name: Name of a stream this workflow owns, created on first + use. Empty means the workflow's default output stream. A workflow + cannot append to a stream another execution owns. Raises: ValueError: ``records`` is empty or one of them is over the server's per-record size limit. """ - _Runtime.current().workflow_append_stream_records(stream_id, records) + _Runtime.current().workflow_append_stream_records(stream_name, records) @dataclass(frozen=True) diff --git a/tests/worker/test_workflow_stream.py b/tests/worker/test_workflow_stream.py index 676b29c90..b4dd6107b 100644 --- a/tests/worker/test_workflow_stream.py +++ b/tests/worker/test_workflow_stream.py @@ -197,10 +197,10 @@ def _assert_not_read_only(self, _action: str) -> None: def commands(self) -> list[WorkflowCommand]: return self._current_completion.successful.commands - def publish(self, *records: api_stream.StreamRecord, stream_id: str = "") -> None: + def publish(self, *records: api_stream.StreamRecord, stream_name: str = "") -> None: instance: Any = self _WorkflowInstanceImpl.workflow_append_stream_records( - instance, stream_id, list(records) + instance, stream_name, list(records) ) def flush(self) -> None: @@ -214,13 +214,13 @@ def test_a_tasks_publishes_on_one_stream_become_one_command() -> None: stub = _CommandStub() stub.publish(record(b"a"), record(b"b")) stub.publish(record(b"c")) - stub.publish(record(b"d"), stream_id="other") + stub.publish(record(b"d"), stream_name="other") assert stub.commands == [] stub.flush() by_stream = { - command.append_stream_records.stream_id: [ + command.append_stream_records.stream_name: [ r.body.data for r in command.append_stream_records.records ] for command in stub.commands diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py index 694b90b71..5dc15dd4f 100644 --- a/tests/worker/test_workflow_stream_e2e.py +++ b/tests/worker/test_workflow_stream_e2e.py @@ -272,9 +272,9 @@ class PublishAndRead: async def run(self) -> list[str]: workflow._append_stream_records( [_record(b"alpha", "progress"), _record(b"beta", "progress")], - stream_id="output", + stream_name="output", ) - workflow._append_stream_records([_record(b"gamma")], stream_id="output") + workflow._append_stream_records([_record(b"gamma")], stream_name="output") # A name this workflow has not written yet still names a stream it # owns, so subscribing creates the one the later publish lands in. workflow._subscribe_stream("output", start_offset=0) From a48b223af43151bf6a3e2ba9b1b558e9d1f865fa Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 14:48:45 -0700 Subject: [PATCH 10/27] Mapped the native default topic onto the server's default stream. The server resolves an unnamed stream to the same name as DEFAULT_TOPIC, so the native handle takes no topic and sends the name explicitly, which keeps a record's topic equal to its stream's name. --- temporalio/streams/providers/native.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index d0026b7ed..78e194c7b 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -6,7 +6,11 @@ transaction that accepts the Workflow Task, subscribes to it by name and reads the ranges the server delivers on its Workflow Tasks; outside code appends and reads through the stream service, and the workflow's records and an outside -producer's land in one log in the order the server accepted them. +producer's land in one log in the order the server accepted them. The default +topic, :data:`temporalio.streams.DEFAULT_TOPIC`, is the server's default +stream: the server resolves an unnamed stream to that same name, so the +provider sends the name explicitly and a record's topic and its stream's name +never differ. A cursor names the run as well as the offset, because an owned stream belongs to one run and a successor's starts over at zero. A handle without a run id @@ -325,7 +329,7 @@ def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: def read( self, *, - topic: str | StreamTopic[Any], + topic: str | StreamTopic[Any] | None = None, after: Cursor = BEGINNING, result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: @@ -374,7 +378,7 @@ async def _read( return run_id, offset = successor, 0 - async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: """The cursor of the newest record on ``topic``, naming the run it was read from. An empty topic on the chain's first run is the beginning of the @@ -398,7 +402,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, ) -> NativeProducer[Any]: From 6be84731eadebac9291e3bc17ae1a5ad19479958 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:01:48 -0700 Subject: [PATCH 11/27] Opened activity-owned streams on the native provider. A standalone activity is its own owner on the server and a workflow's activity is reached through its workflow; the handle pins one activity execution and never follows a chain, because a retry writes to the same stream. --- temporalio/streams/providers/native.py | 85 +++++++++++++++++++++++++- tests/streams/test_activity_streams.py | 19 ++++++ 2 files changed, 103 insertions(+), 1 deletion(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 78e194c7b..8bcf32c04 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -12,6 +12,12 @@ provider sends the name explicitly and a record's topic and its stream's name never differ. +An activity owns topics of its own, apart from its workflow's: a standalone +activity is its own owner on the server, and an activity a workflow scheduled +is addressed through that workflow. They are one stream per activity +execution, so a retry writes to the same stream, and the server ends them when +the activity reaches a terminal status. + A cursor names the run as well as the offset, because an owned stream belongs to one run and a successor's starts over at zero. A handle without a run id reads run after run, learning from the poll that a run's stream is closed and @@ -52,7 +58,12 @@ ) from temporalio.streams.providers import ProviderPlugin -__all__ = ["NativeProducer", "NativeStreamHandle", "NativeStreams"] +__all__ = [ + "NativeActivityStreamHandle", + "NativeProducer", + "NativeStreamHandle", + "NativeStreams", +] T = TypeVar("T") @@ -474,6 +485,60 @@ async def _successor(self, run_id: str) -> str | None: return None +class NativeActivityStreamHandle(NativeStreamHandle): + """The topics one activity owns, from outside, over the stream service. + + An activity's streams belong to one activity execution, not to a chain of + runs: a retry writes to the same stream and a read ends when the activity + reaches a terminal status. So the handle pins the execution on first use + and never follows a successor. A standalone activity is its own owner; an + activity a workflow scheduled is reached through that workflow's run. + """ + + def __init__( + self, + client: Client, + activity_id: str, + workflow_id: str | None, + run_id: str | None, + *, + opened: set[tuple[str, str]] | None = None, + ) -> None: + """Address ``activity_id``'s topics, pinned to ``run_id`` when one is given.""" + super().__init__(client, workflow_id or "", run_id, opened=opened) + self._activity_id = activity_id + + def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: + return self._service().activity_stream( + self._activity_id, topic, workflow_id=self._workflow_id, run_id=run_id + ) + + async def _current_run(self) -> str: + if self._workflow_id: + return await super()._current_run() + try: + description = await self._client.get_activity_handle( + self._activity_id + ).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"activity {self._activity_id!r} was not found" + ) from error + raise + assert description.activity_run_id is not None + return description.activity_run_id + + async def _first_run(self) -> str: + return await self._current_run() + + async def _predecessor(self, run_id: str) -> str | None: + return None + + async def _successor(self, run_id: str) -> str | None: + return None + + class NativeStreams(ProviderPlugin): """The server-side provider. @@ -499,6 +564,24 @@ def get_stream_handle( """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain.""" return NativeStreamHandle(client, workflow_id, run_id, opened=self._opened) + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> NativeActivityStreamHandle: + """A handle on the topics ``activity_id`` owns, apart from any workflow's. + + Without ``workflow_id`` the activity is a standalone one and ``run_id`` + pins its run; with one it is that workflow's activity and ``run_id`` + pins the workflow's run. + """ + return NativeActivityStreamHandle( + client, activity_id, workflow_id, run_id, opened=self._opened + ) + async def close(self) -> None: """Close the channels this provider opened to the stream service. diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index 8aa4a653a..62d8b924c 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -16,6 +16,7 @@ from __future__ import annotations import asyncio +import os import uuid from collections.abc import AsyncIterator, Callable from dataclasses import dataclass @@ -34,6 +35,7 @@ topic, ) from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.native import NativeStreams from temporalio.testing import WorkflowEnvironment from tests.helpers import new_worker @@ -57,9 +59,26 @@ async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider.reset() +async def _native_setup(client: Client) -> AsyncIterator[ActivitySetup]: + # The store is a server built from the stream-carrying branch, which the + # test environment's own server is not; TEMPORAL_ADDRESS names it. + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = NativeStreams() + config = client.config() + config["plugins"] = [provider] + yield ActivitySetup("native", provider, Client(**config)) + await provider.close() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ActivitySetup]]] = { "memory": _memory_setup } +if os.environ.get("STREAMS_LIVE") == "native": + SETUPS["native"] = _native_setup @pytest.fixture(params=sorted(SETUPS)) From c8cd5e84ef6f65b7df27cc84afe069d2b5395564 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 17:05:20 -0700 Subject: [PATCH 12/27] Mapped every read start onto the native provider. BEGINNING now asks the server for the earliest record rather than offset zero, which a stream with a raised floor refuses. END and last=N become the tail and last-N positions, resolved by the server on the first poll or at subscription and recorded, so replay never resolves them again. --- temporalio/streams/providers/native.py | 92 +++++++++++++++---- temporalio/worker/_workflow_instance.py | 6 +- temporalio/workflow/_context.py | 21 +++-- tests/streams/test_streams_conformance.py | 3 + tests/worker/test_workflow_stream_e2e.py | 105 +++++++++++++++++++++- 5 files changed, 200 insertions(+), 27 deletions(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 8bcf32c04..62a1c8b7a 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -35,8 +35,10 @@ from typing import Any, Generic, TypeVar from temporalio import workflow +from temporalio.api.stream.v1 import StreamStartPosition from temporalio.client import Client, WorkflowHistoryEventFilterType from temporalio.client_stream import ( + Page, StreamClient, WorkflowStreamHandle, close_shared_clients, @@ -46,7 +48,14 @@ from temporalio.service import RPCError, RPCStatusCode 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._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( RecordDecoder, @@ -164,18 +173,30 @@ def publish(self, record: WireRecord) -> None: class _NativeWorkflowProvider: """The workflow half: the server's commands and delivered ranges.""" - 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 - named = _position(after) - if named is not None: + # The server resolves the position when it registers the subscription + # and records the offset on the subscribed event, so replay never + # resolves it again. + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif after == BEGINNING: + start = StreamStartPosition(earliest=True) + else: + named = _position(after) + assert named is not None if named[0] != run_id: raise StreamCursorError( f"cursor {after.token!r} names another run; a run's stream is its own" ) - start = named[1] + 1 - workflow._subscribe_stream(topic, start_offset=start) + start = StreamStartPosition(offset=named[1] + 1) + workflow._subscribe_stream(topic, start=start) return _NativeReadSource(topic, run_id) def open_writer(self, topic: str) -> WriteSink: @@ -342,37 +363,63 @@ def read( *, 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 record the chain's first retained run + still holds. ``END`` and ``last=`` start on the current run, or the + pinned one: an earlier run of a chain has ended and holds neither the + tail nor the newest records. The server resolves each on the first + poll, in the read that serves it. + """ + 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) + start: StreamStartPosition | None = None + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif named is None: + start = StreamStartPosition(earliest=True) + return self._read(topic, named, start, after, result_type) async def _read( self, topic: str, named: tuple[str, int] | None, + start: StreamStartPosition | None, after: Cursor, result_type: type | None, ) -> AsyncGenerator[StreamRecord[Any], None]: - decoder = RecordDecoder( - self._converter, result_type, after=after, warn=logger.warning - ) + decoder: RecordDecoder | None = None + offset = 0 if named is not None: run_id, offset = named[0], named[1] + 1 + elif start is not None and start.WhichOneof("position") != "earliest": + run_id = self._run_id or await self._current_run() else: - run_id, offset = self._run_id or await self._first_run(), 0 + run_id = self._run_id or await self._first_run() while True: stream = self._stream(topic, run_id) while True: - page = await stream.poll(from_offset=offset) + page = await stream.poll(from_offset=offset, start=start) + if decoder is None: + decoder = RecordDecoder( + self._converter, + result_type, + after=self._previous(run_id, page, start, after), + warn=logger.warning, + ) + start = None for entry in page.entries: record = await _decode_body(self._codec, entry.record) for out in decoder.decode(_cursor(run_id, entry.offset), record): @@ -389,6 +436,21 @@ async def _read( return run_id, offset = successor, 0 + @staticmethod + def _previous( + run_id: str, page: Page, start: StreamStartPosition | None, after: Cursor + ) -> Cursor: + """The position before the first record a read yields. + + A synthesized record is positioned there. After ``END`` or ``last=`` + it is only known once the server resolved the start, from the first + page. It names a run so a chain-following resume stays on this one. + """ + if start is None or start.WhichOneof("position") == "earliest": + return after + first = page.entries[0].offset if page.entries else page.next_offset + return _cursor(run_id, first - 1) + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: """The cursor of the newest record on ``topic``, naming the run it was read from. diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index c31bc9dc4..86d4922d1 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -1510,7 +1510,9 @@ def get_info(self) -> temporalio.workflow.Info: return self._info def workflow_subscribe_stream( - self, stream_name_or_id: str, start_offset: int + self, + stream_name_or_id: str, + start: temporalio.api.stream.v1.StreamStartPosition, ) -> None: # Reissued on every replay, so the buffer has to exist before the first # range arrives and the command has to be harmless the second time. A @@ -1520,7 +1522,7 @@ def workflow_subscribe_stream( ) command = self._add_command() command.subscribe_stream.stream_name_or_id = stream_name_or_id - command.subscribe_stream.start_offset = start_offset + command.subscribe_stream.start_position.CopyFrom(start) def workflow_append_stream_records( self, diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index ef5c6016e..d9fe9b051 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -325,7 +325,9 @@ def workflow_get_current_deployment_version( @abstractmethod def workflow_subscribe_stream( - self, stream_name_or_id: str, start_offset: int + self, + stream_name_or_id: str, + start: temporalio.api.stream.v1.StreamStartPosition, ) -> None: ... @abstractmethod @@ -1049,7 +1051,9 @@ async def sleep( def _subscribe_stream( # type: ignore[reportUnusedFunction] - stream_name_or_id: str, *, start_offset: int = 0 + stream_name_or_id: str, + *, + start: temporalio.api.stream.v1.StreamStartPosition | None = None, ) -> None: """Subscribe this workflow to a server-side stream. @@ -1059,7 +1063,7 @@ def _subscribe_stream( # type: ignore[reportUnusedFunction] cursor, though it does write one event. Calling it on every replay is harmless because replay matches the command to the event already recorded. - Only the name or id and the start offset go to the server. The rest of the + Only the name or id and the start position go to the server. The rest of the stream's addressing is resolved there, because a workflow cannot look it up without doing I/O and a value it carried would be a reading rather than a fact. A name this workflow has not written yet names a stream it owns, and @@ -1069,11 +1073,14 @@ def _subscribe_stream( # type: ignore[reportUnusedFunction] stream_name_or_id: Stream to consume: the name of one this workflow owns, or the id of a standalone stream. The server tries them in that order. - start_offset: Where to start. Negative means from wherever the stream is - when the subscription is registered; the server resolves that once - and records it, so replay does not resolve it again. + start: Where to start: an absolute offset, the oldest record the + stream holds, the tail as of registration, or the last N records. + Omitted, it is the oldest record held. The server resolves it + once and records the offset, so replay does not resolve it again. """ - _Runtime.current().workflow_subscribe_stream(stream_name_or_id, start_offset) + if start is None: + start = temporalio.api.stream.v1.StreamStartPosition(earliest=True) + _Runtime.current().workflow_subscribe_stream(stream_name_or_id, start) def _append_stream_records( # type: ignore[reportUnusedFunction] diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index e6d544df8..52261c61c 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -144,6 +144,9 @@ async def host(workflow_id: str) -> None: StreamHost.run, id=workflow_id, task_queue=worker.task_queue ) + # No truncate hook: a stream a workflow owns has no truncation call. + # BEGINNING on a truncated stream is covered on the stream client, + # whose standalone streams can be truncated. yield ProviderCase("native", provider, client, host=host) for handle in hosts.values(): await handle.terminate() diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py index 5dc15dd4f..e2ca2d9da 100644 --- a/tests/worker/test_workflow_stream_e2e.py +++ b/tests/worker/test_workflow_stream_e2e.py @@ -27,7 +27,7 @@ from temporalio.api.stream.v1 import StreamRecord from temporalio.client import Client from temporalio.client_stream import StreamClient -from temporalio.streams import RecordKind +from temporalio.streams import END, RecordKind from temporalio.streams.providers.native import NativeStreams from temporalio.worker import Worker from tests.streams.test_streams_conformance import take @@ -277,7 +277,7 @@ async def run(self) -> list[str]: workflow._append_stream_records([_record(b"gamma")], stream_name="output") # A name this workflow has not written yet still names a stream it # owns, so subscribing creates the one the later publish lands in. - workflow._subscribe_stream("output", start_offset=0) + workflow._subscribe_stream("output") received: list[str] = [] while len(received) < 3: @@ -331,7 +331,7 @@ class ConsumeAcrossTasks: @workflow.run async def run(self, stream_id: str, expected: int) -> list[str]: - workflow._subscribe_stream(stream_id, start_offset=0) + workflow._subscribe_stream(stream_id) seen: list[str] = [] while len(seen) < expected: for item in await workflow._read_stream_records(stream_id): @@ -384,3 +384,102 @@ async def test_a_cached_workflow_consumes_across_sticky_tasks() -> None: ) finally: await streams.close() + + +@workflow.defn +class StartsWhenTold: + """Subscribes to ``inputs`` at a start given by a signal 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 _subscribed_offsets(client: Client, workflow_id: str) -> list[int]: + return [ + event.workflow_stream_subscribed_event_attributes.start_offset + async for event in client.get_workflow_handle( + workflow_id + ).fetch_history_events() + if event.event_type == EVENT_STREAM_SUBSCRIBED + ] + + +async def test_a_workflow_reader_starts_at_the_last_n_records_on_the_server() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "last-tq-" + uuid.uuid4().hex[:8] + workflow_id = "last-wf-" + uuid.uuid4().hex[:8] + try: + # Cold, so every task replays and the recorded start is what places + # the reader, not a second resolution. + async with Worker( + client, + task_queue=task_queue, + workflows=[StartsWhenTold], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + StartsWhenTold.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_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, "last") + assert await asyncio.wait_for(handle.result(), 60) == [3, 4] + assert await _subscribed_offsets(client, workflow_id) == [2] + finally: + await provider.close() + + +async def test_a_workflow_reader_at_end_skips_what_was_there_on_the_server() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "end-tq-" + uuid.uuid4().hex[:8] + workflow_id = "end-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[StartsWhenTold], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + StartsWhenTold.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "old"}, {"n": "old"}) + await handle.signal(StartsWhenTold.begin, "end") + result = asyncio.ensure_future(handle.result()) + # The subscription registers when the signalled task completes, + # which the test does not observe, so appends keep coming. + 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, 60) == ["new"] + offsets = await _subscribed_offsets(client, workflow_id) + assert len(offsets) == 1 and offsets[0] >= 2 + finally: + await provider.close() From 560f5a01edbef4216a9653e841625a1da263e229 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 29 Sep 2026 12:44:03 -0700 Subject: [PATCH 13/27] Retried throttled and unavailable stream calls as sdk-core does. The stream client has a channel of its own, so Core's retry never covered it and a reader under frontend rate limiting raised mid-read. An append without a producer id goes again only on a refusal the server sent before doing anything, because the server has nothing to deduplicate it against. --- temporalio/client_stream.py | 277 +++++++++++++++++++++++++++--- tests/test_client_stream_owner.py | 1 + tests/test_client_stream_retry.py | 257 +++++++++++++++++++++++++++ 3 files changed, 507 insertions(+), 28 deletions(-) create mode 100644 tests/test_client_stream_retry.py diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index 4364073d2..e2afc8297 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -26,11 +26,24 @@ :class:`temporalio.streams.StreamProducerError` when it refuses a producer sequence it already holds, and :class:`temporalio.service.RPCError` otherwise, never the transport's own exception type. + +A failure sdk-core would retry is retried here, on the same codes and with the +same default :class:`temporalio.service.RetryConfig`, because this channel is +not Core's and gets none of its retrying. ``RESOURCE_EXHAUSTED`` backs off +longer than the rest, as in Core, so a caller the server is throttling does not +add to the load. A call the server cannot tell from its own repeat, an append +without a producer id or a create, is retried only on ``RESOURCE_EXHAUSTED``, +which the server sends before it does anything. The budget is bounded; a caller +that wants a shorter one cancels, as with ``asyncio.timeout``, and the +cancellation lands whether an attempt or a wait between attempts is in progress. """ from __future__ import annotations import asyncio +import logging +import random +import time import weakref from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass @@ -42,9 +55,12 @@ from google.protobuf.message import Message import temporalio.api.streamservice.v1 as stream +from temporalio.api.common.v1 import GrpcStatus +from temporalio.api.enums.v1 import ResourceExhaustedCause +from temporalio.api.errordetails.v1 import ResourceExhaustedFailure from temporalio.api.stream.v1 import StreamRecord, StreamStartPosition from temporalio.api.streamservice.v1 import service_pb2_grpc -from temporalio.service import RPCError, RPCStatusCode +from temporalio.service import RetryConfig, RPCError, RPCStatusCode from temporalio.streams import StreamNotFoundError, StreamProducerError __all__ = [ @@ -60,12 +76,133 @@ _T = TypeVar("_T") +logger = logging.getLogger(__name__) + # The server refuses a producer sequence it already holds with a message and # no typed detail, so the phrase is the only thing to match on. Both refusals # it sends carry it: a repeat with different content, and one behind the # sequence it accepted last. _PRODUCER_CONFLICT = "producer sequence" +# The codes sdk-core retries. +_RETRYABLE = frozenset( + { + grpc.StatusCode.DATA_LOSS, + grpc.StatusCode.INTERNAL, + grpc.StatusCode.UNKNOWN, + grpc.StatusCode.RESOURCE_EXHAUSTED, + grpc.StatusCode.ABORTED, + grpc.StatusCode.OUT_OF_RANGE, + grpc.StatusCode.UNAVAILABLE, + } +) +# What the server sends before it does anything, so a call that cannot be told +# from its own repeat is still safe to make again on it. +_REFUSED = frozenset({grpc.StatusCode.RESOURCE_EXHAUSTED}) +# A message over the channel's limit comes back as RESOURCE_EXHAUSTED and is +# the same size every time. +_TOO_LARGE = ( + "grpc: received message larger than max", + "grpc: message after decompression larger than max", + "grpc: received message after decompression larger than max", +) +# The floor under a throttled call's wait, sdk-core's own. +_THROTTLE = RetryConfig( + initial_interval_millis=1000, + multiplier=2.0, + max_interval_millis=10000, + max_elapsed_time_millis=None, + max_retries=0, +) + + +class _Backoff: + """Exponential backoff over a :class:`RetryConfig`, with sdk-core's arithmetic.""" + + def __init__(self, config: RetryConfig) -> None: + self._config = config + self._started = time.monotonic() + self._interval = config.initial_interval_millis / 1000 + self._failures = 0 + + def next(self) -> float | None: + """Seconds to wait before the next attempt, or ``None`` once the budget is spent.""" + config = self._config + self._failures += 1 + if config.max_retries and self._failures >= config.max_retries: + return None + base = self._interval + self._interval = min( + base * config.multiplier, config.max_interval_millis / 1000 + ) + spread = base * config.randomization_factor + delay = max(base + random.uniform(-spread, spread), 0.0) + if config.max_elapsed_time_millis is not None and ( + time.monotonic() - self._started + delay + > config.max_elapsed_time_millis / 1000 + ): + return None + return delay + + +def _raw_status(error: grpc.aio.AioRpcError) -> bytes: + # The aio metadata iterates as (key, value) pairs at runtime, whatever + # shape the stubs give its items. + trailing: Any = error.trailing_metadata() + for item in trailing or (): + key, value = item[0], item[1] + if key == "grpc-status-details-bin" and isinstance(value, bytes): + return value + return b"" + + +def _exhausted_cause(error: grpc.aio.AioRpcError) -> int | None: + """The cause the server attached to a ``RESOURCE_EXHAUSTED``, when it attached one.""" + raw = _raw_status(error) + if not raw: + return None + status = GrpcStatus() + status.ParseFromString(raw) + for detail in status.details: + if detail.Is(ResourceExhaustedFailure.DESCRIPTOR): + failure = ResourceExhaustedFailure() + detail.Unpack(failure) + return failure.cause + return None + + +def _worth_waiting_out(error: grpc.aio.AioRpcError) -> bool: + """Whether a ``RESOURCE_EXHAUSTED`` is load the server will shed, rather than a limit.""" + if (error.details() or "").startswith(_TOO_LARGE): + return False + # A stream's budget refuses an append for as long as the stream is that + # full, which no wait changes. + return _exhausted_cause(error) != ( + ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_PERSISTENCE_STORAGE_LIMIT + ) + + +def _retry_after( + error: grpc.aio.AioRpcError, + backoff: _Backoff, + throttle: _Backoff, + *, + idempotent: bool, +) -> float | None: + """Seconds to wait before making the call again, or ``None`` to raise it.""" + code = error.code() + if code not in _RETRYABLE or (not idempotent and code not in _REFUSED): + return None + throttled = code is grpc.StatusCode.RESOURCE_EXHAUSTED + if throttled and not _worth_waiting_out(error): + return None + delay = backoff.next() + if delay is None: + return None + if throttled: + delay = max(delay, throttle.next() or 0.0) + return delay + @dataclass(frozen=True) class Appended: @@ -176,44 +313,93 @@ def _translate(error: grpc.aio.AioRpcError) -> Exception: # content or behind the one it accepted last. The caller asked to be # deduplicated and could not be, which is a condition of its own. return StreamProducerError(details) - raw = b"" - # The aio metadata iterates as (key, value) pairs at runtime, whatever - # shape the stubs give its items. - trailing: Any = error.trailing_metadata() - for item in trailing or (): - key, value = item[0], item[1] - if key == "grpc-status-details-bin" and isinstance(value, bytes): - raw = value - return RPCError(details, RPCStatusCode(code.value[0]), raw) - - -async def _call(method: Callable[[Any], Awaitable[_T]], request: Any) -> _T: - """Make one stub call, translating the transport's failure to the SDK's.""" - try: - return await method(request) - except grpc.aio.AioRpcError as error: - raise _translate(error) from error + return RPCError(details, RPCStatusCode(code.value[0]), _raw_status(error)) + + +async def _call( + method: Callable[[Any], Awaitable[_T]], + request: Any, + *, + retry_config: RetryConfig | None = None, + idempotent: bool = True, +) -> _T: + """Make one stub call, retried as sdk-core would, translating the failure to the SDK's. + + ``idempotent`` is false for a call the server cannot tell from its own + repeat, which is then retried only on a refusal the server sent before it + did anything. The request is the same object on every attempt, so a + numbered append is deduplicated by the server whichever attempt landed. + """ + config = retry_config or RetryConfig() + backoff = _Backoff(config) + throttle = _Backoff(_THROTTLE) + attempts = 0 + while True: + attempts += 1 + try: + return await method(request) + except grpc.aio.AioRpcError as error: + delay = _retry_after(error, backoff, throttle, idempotent=idempotent) + if delay is None: + raise _translate(error) from error + _log_retry(error, attempts, config) + await asyncio.sleep(delay) + + +def _log_retry(error: grpc.aio.AioRpcError, attempts: int, config: RetryConfig) -> None: + # Quiet at first and louder once half the budget is gone, as sdk-core does, + # so a single throttled call is not a warning but a struggling one is. + level = logging.DEBUG + if config.max_retries and attempts * 2 >= config.max_retries: + level = logging.WARNING + logger.log( + level, + "stream call failed with %s on attempt %d, retrying: %s", + error.code().name, + attempts, + error.details(), + ) class StreamClient: """Creates and opens streams on a namespace.""" - def __init__(self, channel: Any, namespace: str) -> None: - """Wrap an existing ``grpc.aio`` channel. Prefer :meth:`connect`.""" + def __init__( + self, + channel: Any, + namespace: str, + *, + retry_config: RetryConfig | None = None, + ) -> None: + """Wrap an existing ``grpc.aio`` channel. Prefer :meth:`connect`. + + ``retry_config`` is the policy every call made through this client + retries under; ``None`` is the SDK's default. + """ self._channel = channel self._namespace = namespace + self._retry_config = retry_config # The generated stub is typed for a synchronous channel. This client # drives it over ``grpc.aio``, where every call is awaited. self._stub: Any = service_pb2_grpc.StreamServiceStub(channel) @staticmethod - def connect(target_host: str, namespace: str = "default") -> StreamClient: + def connect( + target_host: str, + namespace: str = "default", + *, + retry_config: RetryConfig | None = None, + ) -> StreamClient: """Open a channel to a frontend. Separate from ``Client.connect`` because this does not share the connection the rest of the SDK uses. """ - return StreamClient(grpc.aio.insecure_channel(target_host), namespace) + return StreamClient( + grpc.aio.insecure_channel(target_host), + namespace, + retry_config=retry_config, + ) async def close(self) -> None: """Close the underlying channel.""" @@ -240,6 +426,8 @@ async def create( if max_items is not None: lifecycle.max_items = max_items + # A create that landed but was not answered would be refused as a + # repeat, so it goes again only on a refusal. response = await _call( self._stub.CreateStream, stream.CreateStreamRequest( @@ -249,17 +437,22 @@ async def create( lifecycle=lifecycle, ) ), + retry_config=self._retry_config, + idempotent=False, ) return StreamHandle( self._stub, self._namespace, stream_id, run_id=response.frontend_response.run_id, + retry_config=self._retry_config, ) def get(self, stream_id: str) -> StreamHandle: """Open an existing stream without a round trip.""" - return StreamHandle(self._stub, self._namespace, stream_id) + return StreamHandle( + self._stub, self._namespace, stream_id, retry_config=self._retry_config + ) def workflow_stream( self, workflow_id: str, name: str = "", *, owner_run_id: str = "" @@ -281,7 +474,12 @@ def workflow_stream( :meth:`WorkflowStreamHandle.pin` for why a follower wants that. """ return WorkflowStreamHandle( - self._stub, self._namespace, workflow_id, name, owner_run_id + self._stub, + self._namespace, + workflow_id, + name, + owner_run_id, + retry_config=self._retry_config, ) def activity_stream( @@ -309,6 +507,7 @@ def activity_stream( name, run_id, activity_id=activity_id, + retry_config=self._retry_config, ) @@ -316,7 +515,13 @@ class StreamHandle: """A handle to one standalone stream.""" def __init__( - self, stub: Any, namespace: str, stream_id: str, run_id: str = "" + self, + stub: Any, + namespace: str, + stream_id: str, + run_id: str = "", + *, + retry_config: RetryConfig | None = None, ) -> None: """Prefer :meth:`StreamClient.get` or :meth:`StreamClient.create`.""" self._stub = stub @@ -325,6 +530,7 @@ def __init__( # Passing this back saves the server resolving the current run on every # call, which is otherwise a persistence lookup per request. self._run_id = run_id + self._retry_config = retry_config @property def id(self) -> str: @@ -341,9 +547,11 @@ async def append( Supplying ``producer_id`` and ``sequence`` makes the append idempotent: a retry with the same pair returns the original offsets rather than - appending twice, and says so. Without them the append is - at-least-once, which is only the right trade when duplicates are - harmless. + appending twice, and says so. That is also what lets a failed append + be made again here on every code sdk-core retries. Without them the + append is at-least-once, which is only the right trade when duplicates + are harmless, and it is made again only on a refusal the server sent + before it did anything. """ response = await _call( self._stub.AddMessages, @@ -357,6 +565,8 @@ async def append( sequence=sequence, ) ), + retry_config=self._retry_config, + idempotent=bool(producer_id), ) return _appended(response.frontend_response) @@ -440,6 +650,7 @@ async def poll( wait_new_messages=wait, ) ), + retry_config=self._retry_config, ) return _page(response.frontend_response) @@ -459,6 +670,7 @@ async def truncate(self, new_base_offset: int) -> None: new_base_offset=new_base_offset, ) ), + retry_config=self._retry_config, ) async def finish_writing(self, producer_id: str) -> None: @@ -472,6 +684,7 @@ async def finish_writing(self, producer_id: str) -> None: producer_id=producer_id, ) ), + retry_config=self._retry_config, ) async def close(self) -> None: @@ -483,6 +696,7 @@ async def close(self) -> None: namespace=self._namespace, stream_id=self._id ) ), + retry_config=self._retry_config, ) async def describe(self) -> stream.StreamState: @@ -494,6 +708,7 @@ async def describe(self) -> stream.StreamState: namespace=self._namespace, stream_id=self._id ) ), + retry_config=self._retry_config, ) return response.frontend_response.state @@ -517,6 +732,7 @@ def __init__( owner_run_id: str = "", *, activity_id: str = "", + retry_config: RetryConfig | None = None, ) -> None: """Prefer :meth:`StreamClient.workflow_stream` or :meth:`StreamClient.activity_stream`.""" self._stub = stub @@ -525,6 +741,7 @@ def __init__( self._name = name self._owner_run_id = owner_run_id self._activity_id = activity_id + self._retry_config = retry_config @property def workflow_id(self) -> str: @@ -608,6 +825,8 @@ async def append( sequence=sequence, ) ), + retry_config=self._retry_config, + idempotent=bool(producer_id), ) return _appended(response.frontend_response) @@ -679,6 +898,7 @@ async def poll( wait_new_messages=wait, ) ), + retry_config=self._retry_config, ) return _page(response.frontend_response) @@ -697,6 +917,7 @@ async def describe(self) -> stream.StreamState: stream_name=self._name, ) ), + retry_config=self._retry_config, ) return response.frontend_response.state diff --git a/tests/test_client_stream_owner.py b/tests/test_client_stream_owner.py index a87f9d782..a19e8dbc8 100644 --- a/tests/test_client_stream_owner.py +++ b/tests/test_client_stream_owner.py @@ -39,6 +39,7 @@ def _client(recorder: _Recorder) -> StreamClient: client = StreamClient.__new__(StreamClient) client._stub = recorder client._namespace = "ns" + client._retry_config = None return client diff --git a/tests/test_client_stream_retry.py b/tests/test_client_stream_retry.py new file mode 100644 index 000000000..8baa2939b --- /dev/null +++ b/tests/test_client_stream_retry.py @@ -0,0 +1,257 @@ +"""What the stream client retries, and what it raises at once. + +The stream service is reached over a channel of the client's own, outside +sdk-core, so the retry Core gives every other call is reproduced here. These +pin its edges with a scripted stub: a throttled read or numbered append goes +again, an append the server could not tell from its repeat does not, a code +Core would not retry is raised at once, and cancelling the caller lands during +the wait between attempts. +""" + +from __future__ import annotations + +import asyncio +from collections import defaultdict +from types import SimpleNamespace +from typing import Any + +import grpc +import grpc.aio +import pytest + +import temporalio.api.streamservice.v1 as stream +import temporalio.converter +from temporalio import client_stream +from temporalio.api.common.v1 import GrpcStatus +from temporalio.api.enums.v1 import ResourceExhaustedCause +from temporalio.api.errordetails.v1 import ResourceExhaustedFailure +from temporalio.api.stream.v1 import StreamRecord +from temporalio.client_stream import WorkflowStreamHandle, _to_service +from temporalio.service import RetryConfig, RPCError, RPCStatusCode +from temporalio.streams import StreamNotFoundError +from temporalio.streams._record import RecordKind +from temporalio.streams._wire import to_wire +from temporalio.streams.providers.native import NativeStreamHandle + +FAST = RetryConfig( + initial_interval_millis=1, + max_interval_millis=2, + max_elapsed_time_millis=5000, + max_retries=5, +) + + +@pytest.fixture(autouse=True) +def fast_throttle(monkeypatch: pytest.MonkeyPatch) -> None: + # The floor under a throttled wait is a second, which is right for a + # caller and wrong for a test. + monkeypatch.setattr( + client_stream, + "_THROTTLE", + RetryConfig( + initial_interval_millis=1, + max_interval_millis=2, + max_elapsed_time_millis=None, + max_retries=0, + ), + ) + + +class _Stub: + """Answers each method from a script of exceptions and responses, in order.""" + + def __init__(self, **scripts: list[Any]) -> None: + self.calls: dict[str, int] = defaultdict(int) + self._scripts = scripts + + def __getattr__(self, method: str) -> Any: + script = self._scripts[method] + + async def call(_request: Any, **_: Any) -> Any: + self.calls[method] += 1 + outcome = script.pop(0) + if isinstance(outcome, BaseException): + raise outcome + return outcome + + return call + + +def _error( + code: grpc.StatusCode, + details: str = "", + *, + cause: ResourceExhaustedCause.ValueType | None = None, +) -> grpc.aio.AioRpcError: + trailing = grpc.aio.Metadata() + if cause is not None: + status = GrpcStatus(code=code.value[0], message=details) + status.details.add().Pack(ResourceExhaustedFailure(cause=cause)) + trailing.add("grpc-status-details-bin", status.SerializeToString()) + return grpc.aio.AioRpcError(code, grpc.aio.Metadata(), trailing, details=details) + + +def _throttled() -> grpc.aio.AioRpcError: + return _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "service rate limit exceeded", + cause=ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_RPS_LIMIT, + ) + + +def _page(*records: stream.StreamRecord) -> stream.PollWorkflowMessagesResponse: + return stream.PollWorkflowMessagesResponse( + frontend_response=stream.PollMessagesOutput( + records=list(records), + next_offset=len(records), + head_offset=len(records), + closed=True, + run_id="run", + ) + ) + + +def _appended() -> stream.AddWorkflowMessagesResponse: + return stream.AddWorkflowMessagesResponse( + frontend_response=stream.AddMessagesOutput( + first_offset=0, next_offset=1, count=1 + ) + ) + + +def _handle(stub: _Stub, retry: RetryConfig = FAST) -> WorkflowStreamHandle: + return WorkflowStreamHandle(stub, "ns", "wf", "topic", "run", retry_config=retry) + + +async def test_a_throttled_poll_is_read_again() -> None: + stub = _Stub(PollWorkflowMessages=[_throttled(), _page()]) + page = await _handle(stub).poll() + assert page.closed + assert stub.calls["PollWorkflowMessages"] == 2 + + +async def test_a_throttled_numbered_append_goes_again() -> None: + stub = _Stub(AddWorkflowMessages=[_throttled(), _appended()]) + appended = await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert appended.next_offset == 1 + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_a_numbered_append_survives_an_ambiguous_failure() -> None: + # The server holds the producer's sequence, so whichever attempt landed, + # the repeat comes back with the original offsets. + stub = _Stub(AddWorkflowMessages=[_error(grpc.StatusCode.UNAVAILABLE), _appended()]) + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_an_unnumbered_append_is_not_made_again_after_an_ambiguous_failure() -> ( + None +): + stub = _Stub(AddWorkflowMessages=[_error(grpc.StatusCode.UNAVAILABLE), _appended()]) + with pytest.raises(RPCError) as raised: + await _handle(stub).append(StreamRecord()) + assert raised.value.status == RPCStatusCode.UNAVAILABLE + assert stub.calls["AddWorkflowMessages"] == 1 + + +async def test_an_unnumbered_append_goes_again_after_a_refusal() -> None: + # Throttling is answered before the handler runs, so nothing landed. + stub = _Stub(AddWorkflowMessages=[_throttled(), _appended()]) + await _handle(stub).append(StreamRecord()) + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_a_code_core_would_not_retry_is_raised_at_once() -> None: + stub = _Stub( + PollWorkflowMessages=[_error(grpc.StatusCode.INVALID_ARGUMENT, "bad"), _page()] + ) + with pytest.raises(RPCError) as raised: + await _handle(stub).poll() + assert raised.value.status == RPCStatusCode.INVALID_ARGUMENT + assert stub.calls["PollWorkflowMessages"] == 1 + + stub = _Stub(PollWorkflowMessages=[_error(grpc.StatusCode.NOT_FOUND), _page()]) + with pytest.raises(StreamNotFoundError): + await _handle(stub).poll() + assert stub.calls["PollWorkflowMessages"] == 1 + + +async def test_the_budget_is_bounded() -> None: + stub = _Stub(PollWorkflowMessages=[_throttled() for _ in range(10)]) + with pytest.raises(RPCError) as raised: + await _handle( + stub, RetryConfig(initial_interval_millis=1, max_retries=3) + ).poll() + assert raised.value.status == RPCStatusCode.RESOURCE_EXHAUSTED + assert stub.calls["PollWorkflowMessages"] == 3 + + +async def test_a_full_stream_is_not_waited_out() -> None: + full = _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "stream holds 10 of its budget of 10 records", + cause=ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_PERSISTENCE_STORAGE_LIMIT, + ) + stub = _Stub(AddWorkflowMessages=[full, _appended()]) + with pytest.raises(RPCError) as raised: + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert raised.value.status == RPCStatusCode.RESOURCE_EXHAUSTED + assert stub.calls["AddWorkflowMessages"] == 1 + + too_large = _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "grpc: received message larger than max (5 vs. 4)", + ) + stub = _Stub(AddWorkflowMessages=[too_large, _appended()]) + with pytest.raises(RPCError): + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert stub.calls["AddWorkflowMessages"] == 1 + + +async def test_cancelling_the_caller_lands_during_the_wait( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + client_stream, + "_THROTTLE", + RetryConfig(initial_interval_millis=60_000, max_elapsed_time_millis=None), + ) + stub = _Stub(PollWorkflowMessages=[_throttled(), _page()]) + task = asyncio.ensure_future(_handle(stub).poll()) + for _ in range(10): + await asyncio.sleep(0) + assert stub.calls["PollWorkflowMessages"] == 1, "parked in the wait" + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert stub.calls["PollWorkflowMessages"] == 1 + + +async def test_a_native_read_survives_a_throttled_poll( + monkeypatch: pytest.MonkeyPatch, +) -> None: + converter = temporalio.converter.default() + record = _to_service( + to_wire( + converter.payload_converter, + topic="topic", + kind=RecordKind.DATA, + value="hello", + producer_id="p", + sequence=1, + ) + ) + stub = _Stub(PollWorkflowMessages=[_throttled(), _page(record)]) + client: Any = SimpleNamespace(data_converter=converter) + handle = NativeStreamHandle(client, "wf", "run") + + def open_stream(_topic: str, _run_id: str) -> WorkflowStreamHandle: + return _handle(stub) + + monkeypatch.setattr(handle, "_stream", open_stream) + + values = [item.value async for item in handle.read(topic="topic")] + + assert values == ["hello"] + assert stub.calls["PollWorkflowMessages"] == 2 From 5c3421e074600c2842ef075f3f414167695d9baa Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:13:48 -0700 Subject: [PATCH 14/27] Derived the stream channel from the client's connection. The stream service was reached on a bare insecure channel, so a client with TLS, an API key, default headers or its own retry policy lost all of them on that path. Connection reads the client's ConnectConfig and opens a grpcio channel with the same target, TLS material, bearer header, headers, keep-alive and proxy, and the shared client retries under the client's retry_config. --- temporalio/client_stream.py | 219 ++++++++++++++++-- temporalio/contrib/server_streams/__init__.py | 2 +- temporalio/streams/providers/native.py | 21 +- tests/test_client_stream_connection.py | 216 +++++++++++++++++ tests/test_client_stream_sharing.py | 27 ++- 5 files changed, 443 insertions(+), 42 deletions(-) create mode 100644 tests/test_client_stream_connection.py diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index e2afc8297..55e924eda 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -18,8 +18,11 @@ - **This client is for use outside a Workflow.** Workflow code publishes and consumes with ``workflow.append_stream_records`` and ``workflow.read_stream_records`` instead. -- **No TLS or API-key support**, for the same reason: the channel is built - here rather than by the machinery that normally handles that. +- **The channel mirrors the client's connection rather than sharing it.** + :class:`Connection` reads a :class:`temporalio.service.ConnectConfig` and + opens a ``grpcio`` channel with the same target, TLS material, API key, + headers and keep-alive, so a client connected to Temporal Cloud reaches the + stream service the same way. What it cannot mirror is noted on that class. A failed call raises :class:`temporalio.streams.StreamNotFoundError` when the server answers ``NOT_FOUND``, @@ -47,7 +50,7 @@ import weakref from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass -from typing import Any, TypeVar +from typing import TYPE_CHECKING, Any, TypeVar import google.protobuf.duration_pb2 import grpc @@ -60,11 +63,22 @@ from temporalio.api.errordetails.v1 import ResourceExhaustedFailure from temporalio.api.stream.v1 import StreamRecord, StreamStartPosition from temporalio.api.streamservice.v1 import service_pb2_grpc -from temporalio.service import RetryConfig, RPCError, RPCStatusCode +from temporalio.service import ( + ConnectConfig, + RetryConfig, + RPCError, + RPCStatusCode, + TLSConfig, + __version__, +) from temporalio.streams import StreamNotFoundError, StreamProducerError +if TYPE_CHECKING: + from temporalio.client import Client + __all__ = [ "Appended", + "Connection", "Page", "StreamClient", "StreamEntry", @@ -72,6 +86,7 @@ "WorkflowStreamHandle", "close_shared_clients", "shared_client", + "shared_key", ] _T = TypeVar("_T") @@ -204,6 +219,145 @@ def _retry_after( return delay +class _Headers(grpc.aio.UnaryUnaryClientInterceptor): + """Attaches the connection's headers to every call, as Core's interceptor does.""" + + def __init__(self, headers: Sequence[tuple[str, str | bytes]]) -> None: + self._headers = headers + + async def intercept_unary_unary( # type: ignore[override] + self, + continuation: Callable[[grpc.aio.ClientCallDetails, Any], Awaitable[Any]], + client_call_details: grpc.aio.ClientCallDetails, + request: Any, + ) -> Any: + # The aio metadata iterates as (key, value) pairs at runtime, whatever + # shape the stubs give its items. + given: Any = client_call_details.metadata + metadata = grpc.aio.Metadata(*(given or ())) + for key, value in self._headers: + # A header the caller set on the call wins over the connection's. + if key not in metadata: + metadata.add(key, value) + details = client_call_details._replace(metadata=metadata) # type: ignore[attr-defined] + return await continuation(details, request) + + +@dataclass(frozen=True) +class Connection: + """How a stream channel reaches a frontend, taken from a client's connection. + + :meth:`from_config` reads what ``Client.connect`` was given and this opens + a ``grpc.aio`` channel that behaves the same way: the target, TLS with the + same root CA, client certificate and key, the API key as a bearer + ``authorization`` header, the client's default headers and keep-alive. + Two clients with the same settings yield equal connections, which is what + lets them share one channel per namespace. + + Two things ``grpcio`` cannot express the way sdk-core does. It has one + override for both the TLS server name it sends and the name it verifies, + so ``verification_server_name`` takes that override when set and + ``domain`` otherwise, while ``domain`` alone still sets the HTTP/2 + authority. And it reads the settings once, when the channel is opened, so + an API key or header updated on the client afterwards reaches the stream + channel only through a new connection. + """ + + target: str + secure: bool + server_root_ca_cert: bytes | None = None + client_cert: bytes | None = None + client_private_key: bytes | None = None + server_name: str | None = None + authority: str | None = None + headers: tuple[tuple[str, str | bytes], ...] = () + keep_alive: tuple[int, int] | None = None + http_proxy: str | None = None + + @staticmethod + def from_config(config: ConnectConfig) -> Connection: + """Read a :class:`temporalio.service.ConnectConfig` the way the bridge does.""" + target = config.target_host + tls: TLSConfig | None = None + if "://" in target: + # The bridge still accepts a URL with a scheme; the scheme decides. + scheme, _, target = target.partition("://") + secure = scheme == "https" + if isinstance(config.tls, TLSConfig): + tls = config.tls + elif isinstance(config.tls, TLSConfig): + secure, tls = True, config.tls + elif config.tls: + secure = True + else: + # TLS is on by default when an API key is given and tls was left unset. + secure = config.tls is None and config.api_key is not None + + headers: list[tuple[str, str | bytes]] = [ + ("client-name", "temporal-python"), + ("client-version", __version__), + ] + given = {key.lower() for key in config.rpc_metadata} + if config.api_key is not None and "authorization" not in given: + headers.append(("authorization", f"Bearer {config.api_key}")) + headers.extend(config.rpc_metadata.items()) + + proxy = config.http_connect_proxy_config + http_proxy: str | None = None + if proxy is not None: + auth = ( + f"{proxy.basic_auth[0]}:{proxy.basic_auth[1]}@" + if proxy.basic_auth + else "" + ) + http_proxy = f"http://{auth}{proxy.target_host}" + + keep_alive = config.keep_alive_config + return Connection( + target=target, + secure=secure, + server_root_ca_cert=tls.server_root_ca_cert if tls else None, + client_cert=tls.client_cert if tls else None, + client_private_key=tls.client_private_key if tls else None, + server_name=((tls.verification_server_name or tls.domain) if tls else None), + authority=tls.domain if tls else None, + headers=tuple(headers), + keep_alive=( + (keep_alive.interval_millis, keep_alive.timeout_millis) + if keep_alive + else None + ), + http_proxy=http_proxy, + ) + + def channel(self) -> grpc.aio.Channel: + """Open a channel with these settings. Nothing is sent until the first call.""" + options: list[tuple[str, Any]] = [] + if self.keep_alive is not None: + options.append(("grpc.keepalive_time_ms", self.keep_alive[0])) + options.append(("grpc.keepalive_timeout_ms", self.keep_alive[1])) + if self.server_name: + options.append(("grpc.ssl_target_name_override", self.server_name)) + if self.authority: + options.append(("grpc.default_authority", self.authority)) + if self.http_proxy: + options.append(("grpc.http_proxy", self.http_proxy)) + # The stubs do not know the aio interceptor base as a ClientInterceptor. + interceptors: Any = [_Headers(self.headers)] if self.headers else None + if not self.secure: + return grpc.aio.insecure_channel( + self.target, options=options, interceptors=interceptors + ) + credentials = grpc.ssl_channel_credentials( + root_certificates=self.server_root_ca_cert, + private_key=self.client_private_key, + certificate_chain=self.client_cert, + ) + return grpc.aio.secure_channel( + self.target, credentials, options=options, interceptors=interceptors + ) + + @dataclass(frozen=True) class Appended: """Where one append landed. @@ -390,10 +544,11 @@ def connect( *, retry_config: RetryConfig | None = None, ) -> StreamClient: - """Open a channel to a frontend. + """Open a plaintext channel to a frontend, for a local server. Separate from ``Client.connect`` because this does not share the - connection the rest of the SDK uses. + connection the rest of the SDK uses; :meth:`for_connection` opens + one with a client's settings. """ return StreamClient( grpc.aio.insecure_channel(target_host), @@ -401,6 +556,20 @@ def connect( retry_config=retry_config, ) + @staticmethod + def for_connection( + connection: Connection, + namespace: str = "default", + *, + retry_config: RetryConfig | None = None, + ) -> StreamClient: + """Open a channel the way ``connection`` describes. + + ``retry_config`` left ``None`` is the SDK's default; pass the client's + own to retry as its other calls do. + """ + return StreamClient(connection.channel(), namespace, retry_config=retry_config) + async def close(self) -> None: """Close the underlying channel.""" await self._channel.close() @@ -941,36 +1110,48 @@ def _page(out: stream.PollMessagesOutput) -> Page: ) -# One channel per loop, target and namespace, shared by every handle in the +# One channel per loop, connection and namespace, shared by every handle in the # process. A channel is multiplexed and long lived, and callers open a handle # per subscription, which would otherwise be a connection per subscription. The # loop is the key because a grpc.aio channel belongs to the loop that made it, # and it is held weakly so a loop that is gone cannot lend its channel to a # successor that happens to reuse its id. +SharedKey = tuple[Connection, str] + _shared: weakref.WeakKeyDictionary[ - asyncio.AbstractEventLoop, dict[tuple[str, str], StreamClient] + asyncio.AbstractEventLoop, dict[SharedKey, StreamClient] ] = weakref.WeakKeyDictionary() -def shared_client(target_host: str, namespace: str) -> StreamClient: - """The process-wide client for ``target_host`` and ``namespace`` on this loop.""" +def shared_key(client: Client) -> SharedKey: + """What names the shared channel ``client`` reaches the stream service through.""" + return Connection.from_config(client.service_client.config), client.namespace + + +def shared_client(client: Client) -> StreamClient: + """The process-wide stream client on this loop for ``client``'s connection and namespace. + + Opened with the client's connection settings and its ``retry_config``, so + a call on it authenticates and retries as the client's other calls do. + """ per_loop = _shared.setdefault(asyncio.get_running_loop(), {}) - key = (target_host, namespace) + key = shared_key(client) existing = per_loop.get(key) if existing is None: - existing = per_loop[key] = StreamClient.connect(target_host, namespace) + existing = per_loop[key] = StreamClient.for_connection( + key[0], key[1], retry_config=client.service_client.config.retry_config + ) return existing -async def close_shared_clients(*keys: tuple[str, str]) -> None: +async def close_shared_clients(*keys: SharedKey) -> None: """Close the shared clients this loop opened for ``keys``, or all of them. - A provider closes the ones it opened, named by ``(target host, - namespace)``: another provider on the same loop may still be reading - through a channel of its own, and taking that out from under it is not - this one's to do. With no keys it closes every one, which is what a - process finished with streams wants, and what a test that opened a loop - of its own wants. + A provider closes the ones it opened, named by :func:`shared_key`: + another provider on the same loop may still be reading through a channel + of its own, and taking that out from under it is not this one's to do. + With no keys it closes every one, which is what a process finished with + streams wants, and what a test that opened a loop of its own wants. """ loop = asyncio.get_running_loop() if not keys: diff --git a/temporalio/contrib/server_streams/__init__.py b/temporalio/contrib/server_streams/__init__.py index 0f6c5142d..f6ac83556 100644 --- a/temporalio/contrib/server_streams/__init__.py +++ b/temporalio/contrib/server_streams/__init__.py @@ -457,4 +457,4 @@ async def _decoded(self, record: StreamRecord) -> StreamRecord: def _stream_client(client: Client) -> StreamClient: - return shared_client(client.service_client.config.target_host, client.namespace) + return shared_client(client) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 62a1c8b7a..751489727 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -24,8 +24,9 @@ from the run's close event who came next; with a run id it is pinned. Prototype support for AI-198. It needs a server built from that branch and -opens its own gRPC channel to it, because sdk-core does not know the stream -service yet, which is also why it does not support TLS or API keys. +reaches the stream service on a channel of its own, opened with the client's +connection settings (target, TLS, API key, headers, retries), because sdk-core +does not know the service yet. """ from __future__ import annotations @@ -39,10 +40,12 @@ from temporalio.client import Client, WorkflowHistoryEventFilterType from temporalio.client_stream import ( Page, + SharedKey, StreamClient, WorkflowStreamHandle, close_shared_clients, shared_client, + shared_key, ) from temporalio.converter import PayloadCodec, PayloadConverter from temporalio.service import RPCError, RPCStatusCode @@ -326,7 +329,7 @@ def __init__( workflow_id: str, run_id: str | None, *, - opened: set[tuple[str, str]] | None = None, + opened: set[SharedKey] | None = None, ) -> None: """Address ``workflow_id``'s topics, pinned to ``run_id`` when one is given. @@ -345,12 +348,8 @@ def _service(self) -> StreamClient: # Resolved on first use, because the shared channel belongs to the # running loop and a handle may be made before there is one. if self._streams is None: - key = ( - self._client.service_client.config.target_host, - self._client.namespace, - ) - self._streams = shared_client(*key) - self._opened.add(key) + self._streams = shared_client(self._client) + self._opened.add(shared_key(self._client)) return self._streams def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: @@ -564,7 +563,7 @@ def __init__( workflow_id: str | None, run_id: str | None, *, - opened: set[tuple[str, str]] | None = None, + opened: set[SharedKey] | None = None, ) -> None: """Address ``activity_id``'s topics, pinned to ``run_id`` when one is given.""" super().__init__(client, workflow_id or "", run_id, opened=opened) @@ -614,7 +613,7 @@ def __init__(self) -> None: """Create the provider.""" # What this provider's handles opened, so closing it leaves another # provider's channels on the same loop alone. - self._opened: set[tuple[str, str]] = set() + self._opened: set[SharedKey] = set() def workflow_provider(self) -> _NativeWorkflowProvider: """The workflow half, over the server's commands and delivered ranges.""" diff --git a/tests/test_client_stream_connection.py b/tests/test_client_stream_connection.py new file mode 100644 index 000000000..73e4c5314 --- /dev/null +++ b/tests/test_client_stream_connection.py @@ -0,0 +1,216 @@ +"""The stream channel is the client's connection, opened again with ``grpcio``. + +sdk-core does not know the stream service, so its channel cannot be shared. +What can be shared is the configuration: these pin that a TLS-configured +client yields a secure channel with the same material, that an API key rides +as the bearer header along with the client's other headers, and that the +client's ``retry_config`` is what the shared stream client retries under. +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import grpc +import grpc.aio +import pytest + +import temporalio.api.streamservice.v1 as stream +from temporalio import client_stream +from temporalio.client_stream import Connection, StreamClient +from temporalio.service import ( + ConnectConfig, + HttpConnectProxyConfig, + KeepAliveConfig, + RetryConfig, + TLSConfig, + __version__, +) + +# The name the vendored stubs call the service by, server-internal as it is. +_SERVICE = "temporal.server.chasm.lib.stream.proto.v1.StreamService" + + +def _fake_client(config: ConnectConfig, namespace: str = "ns") -> Any: + return SimpleNamespace( + service_client=SimpleNamespace(config=config), namespace=namespace + ) + + +def test_a_plaintext_client_yields_an_insecure_channel( + monkeypatch: pytest.MonkeyPatch, +) -> None: + opened: dict[str, Any] = {} + + def insecure_channel(target: str, **kwargs: Any) -> str: + opened.update(target=target, **kwargs) + return "channel" + + monkeypatch.setattr(grpc.aio, "insecure_channel", insecure_channel) + connection = Connection.from_config(ConnectConfig(target_host="localhost:7233")) + assert not connection.secure + assert connection.channel() == "channel" + assert opened["target"] == "localhost:7233" + keep_alive = KeepAliveConfig.default + assert ("grpc.keepalive_time_ms", keep_alive.interval_millis) in opened["options"] + assert ("grpc.keepalive_timeout_ms", keep_alive.timeout_millis) in opened["options"] + + +def test_a_tls_client_yields_a_secure_channel_with_its_credentials( + monkeypatch: pytest.MonkeyPatch, +) -> None: + made: dict[str, Any] = {} + opened: dict[str, Any] = {} + + def ssl_channel_credentials(**kwargs: Any) -> str: + made.update(kwargs) + return "credentials" + + def secure_channel(target: str, credentials: Any, **kwargs: Any) -> str: + opened.update(target=target, credentials=credentials, **kwargs) + return "channel" + + monkeypatch.setattr(grpc, "ssl_channel_credentials", ssl_channel_credentials) + monkeypatch.setattr(grpc.aio, "secure_channel", secure_channel) + + connection = Connection.from_config( + ConnectConfig( + target_host="cloud.example:7233", + tls=TLSConfig( + server_root_ca_cert=b"root", + client_cert=b"cert", + client_private_key=b"key", + domain="cloud.example", + verification_server_name="pinned.test", + ), + ) + ) + assert connection.secure + assert connection.channel() == "channel" + assert made == { + "root_certificates": b"root", + "private_key": b"key", + "certificate_chain": b"cert", + } + assert opened["target"] == "cloud.example:7233" + assert opened["credentials"] == "credentials" + assert ("grpc.ssl_target_name_override", "pinned.test") in opened["options"] + assert ("grpc.default_authority", "cloud.example") in opened["options"] + + +def test_tls_is_on_by_default_with_an_api_key_and_off_when_refused() -> None: + with_key = Connection.from_config( + ConnectConfig(target_host="cloud.example:7233", api_key="secret") + ) + assert with_key.secure + assert with_key.server_root_ca_cert is None, "system roots" + refused = Connection.from_config( + ConnectConfig(target_host="localhost:7233", api_key="secret", tls=False) + ) + assert not refused.secure + assert ("authorization", "Bearer secret") in refused.headers + + +def test_a_scheme_in_the_target_decides_and_is_dropped() -> None: + connection = Connection.from_config( + ConnectConfig(target_host="https://cloud.example:7233") + ) + assert connection.secure + assert connection.target == "cloud.example:7233" + + +def test_a_proxy_and_the_clients_own_authorization_carry_over() -> None: + connection = Connection.from_config( + ConnectConfig( + target_host="localhost:7233", + api_key="ignored", + tls=False, + rpc_metadata={"Authorization": "Custom token", "x-tenant": "t1"}, + http_connect_proxy_config=HttpConnectProxyConfig( + target_host="proxy:3128", basic_auth=("user", "pass") + ), + ) + ) + # Core leaves a caller's own authorization header alone. + assert ("Authorization", "Custom token") in connection.headers + assert ("x-tenant", "t1") in connection.headers + assert not any(key == "authorization" for key, _ in connection.headers) + assert connection.http_proxy == "http://user:pass@proxy:3128" + + +class _Recorder: + """Answers describe and keeps the headers it arrived with.""" + + def __init__(self) -> None: + self.metadata: dict[str, str | bytes] = {} + + async def describe( + self, _request: Any, context: Any + ) -> stream.DescribeStreamResponse: + self.metadata = dict(context.invocation_metadata() or ()) + return stream.DescribeStreamResponse( + frontend_response=stream.DescribeStreamOutput( + state=stream.StreamState(head_offset=3) + ) + ) + + def register(self, server: grpc.aio.Server) -> None: + # The generated registration wants the whole servicer; one method is + # enough to see the headers. + handler: Any = grpc.unary_unary_rpc_method_handler( + self.describe, + request_deserializer=stream.DescribeStreamRequest.FromString, + response_serializer=stream.DescribeStreamResponse.SerializeToString, + ) + server.add_generic_rpc_handlers( + ( + grpc.method_handlers_generic_handler( + _SERVICE, {"DescribeStream": handler} + ), + ) + ) + + +async def test_an_api_key_client_sends_the_header_on_every_call() -> None: + recorder = _Recorder() + server = grpc.aio.server() + recorder.register(server) + port = server.add_insecure_port("127.0.0.1:0") + await server.start() + config = ConnectConfig( + target_host=f"127.0.0.1:{port}", + api_key="secret", + tls=False, + rpc_metadata={"x-tenant": "t1"}, + ) + streams = StreamClient.for_connection(Connection.from_config(config), "ns") + try: + state = await streams.get("s").describe() + assert state.head_offset == 3 + assert recorder.metadata["authorization"] == "Bearer secret" + assert recorder.metadata["x-tenant"] == "t1" + assert recorder.metadata["client-name"] == "temporal-python" + assert recorder.metadata["client-version"] == __version__ + finally: + await streams.close() + await server.stop(None) + + +async def test_the_shared_client_retries_under_the_clients_config() -> None: + retry = RetryConfig(max_retries=3) + client = _fake_client( + ConnectConfig(target_host="localhost:7233", retry_config=retry), "ns" + ) + try: + shared = client_stream.shared_client(client) + assert shared._retry_config is retry + assert shared is client_stream.shared_client(client), "one channel per key" + # A client with other credentials to the same host is not the same channel. + other = _fake_client( + ConnectConfig(target_host="localhost:7233", api_key="k", tls=False), "ns" + ) + assert client_stream.shared_client(other) is not shared + assert client_stream.shared_key(other) != client_stream.shared_key(client) + finally: + await client_stream.close_shared_clients() diff --git a/tests/test_client_stream_sharing.py b/tests/test_client_stream_sharing.py index 793ba122a..55fd3e7e3 100644 --- a/tests/test_client_stream_sharing.py +++ b/tests/test_client_stream_sharing.py @@ -1,8 +1,8 @@ """Who owns a shared channel, and who may close it. -One channel is shared per loop, target and namespace, so two providers in one -process can be using the same one. Closing a provider has to leave the other's -channels alone. +One channel is shared per loop, connection and namespace, so two providers in +one process can be using the same one. Closing a provider has to leave the +other's channels alone. """ from __future__ import annotations @@ -10,6 +10,7 @@ import asyncio from temporalio import client_stream +from temporalio.client_stream import Connection, SharedKey from temporalio.streams.providers.native import NativeStreams @@ -23,17 +24,21 @@ async def close(self) -> None: self.closed = True -def _put(target: str, namespace: str) -> _FakeClient: +def _key(target: str, namespace: str) -> SharedKey: + return Connection(target=target, secure=False), namespace + + +def _put(key: SharedKey) -> _FakeClient: fake = _FakeClient() per_loop = client_stream._shared.setdefault(asyncio.get_running_loop(), {}) - per_loop[(target, namespace)] = fake # type: ignore[assignment] + per_loop[key] = fake # type: ignore[assignment] return fake async def test_closing_named_clients_leaves_the_others_open() -> None: - mine = _put("host-a:7233", "ns") - theirs = _put("host-b:7233", "ns") - await client_stream.close_shared_clients(("host-a:7233", "ns")) + mine = _put(_key("host-a:7233", "ns")) + theirs = _put(_key("host-b:7233", "ns")) + await client_stream.close_shared_clients(_key("host-a:7233", "ns")) assert mine.closed assert not theirs.closed, "another provider is still reading through it" await client_stream.close_shared_clients() @@ -41,11 +46,11 @@ async def test_closing_named_clients_leaves_the_others_open() -> None: async def test_a_provider_closes_only_what_its_own_handles_opened() -> None: - mine = _put("host-a:7233", "ns") - theirs = _put("host-b:7233", "ns") + mine = _put(_key("host-a:7233", "ns")) + theirs = _put(_key("host-b:7233", "ns")) provider = NativeStreams() - provider._opened.add(("host-a:7233", "ns")) + provider._opened.add(_key("host-a:7233", "ns")) await provider.close() assert mine.closed assert not theirs.closed From 47f75c17f13ef29ca2f23e1272d5f1d1f7e75723 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:16:38 -0700 Subject: [PATCH 15/27] Mapped the server's typed refusals to the interface errors. A divergent or stale producer repeat is a StreamProducerError and a poll below the retention floor is a StreamCursorError, the fix the spec lists as pending. The server names each refusal with a reason token at the front of the message, so one function matches the status code and the token, keeps the phrases an older server sends, and leaves every other failure an RPCError. --- temporalio/client_stream.py | 69 +++++++++++++++++----- tests/test_client_stream_errors.py | 95 ++++++++++++++++++++++++++++++ 2 files changed, 148 insertions(+), 16 deletions(-) create mode 100644 tests/test_client_stream_errors.py diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index 55e924eda..babab6fc1 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -27,8 +27,10 @@ A failed call raises :class:`temporalio.streams.StreamNotFoundError` when the server answers ``NOT_FOUND``, :class:`temporalio.streams.StreamProducerError` when it refuses a producer -sequence it already holds, and :class:`temporalio.service.RPCError` -otherwise, never the transport's own exception type. +sequence it already holds, :class:`temporalio.streams.StreamCursorError` when +it refuses a read below the retention floor, and +:class:`temporalio.service.RPCError` otherwise, never the transport's own +exception type. :func:`translate_error` is the one place that decides. A failure sdk-core would retry is retried here, on the same codes and with the same default :class:`temporalio.service.RetryConfig`, because this channel is @@ -71,7 +73,11 @@ TLSConfig, __version__, ) -from temporalio.streams import StreamNotFoundError, StreamProducerError +from temporalio.streams import ( + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) if TYPE_CHECKING: from temporalio.client import Client @@ -87,17 +93,28 @@ "close_shared_clients", "shared_client", "shared_key", + "translate_error", ] _T = TypeVar("_T") logger = logging.getLogger(__name__) -# The server refuses a producer sequence it already holds with a message and -# no typed detail, so the phrase is the only thing to match on. Both refusals -# it sends carry it: a repeat with different content, and one behind the -# sequence it accepted last. -_PRODUCER_CONFLICT = "producer sequence" +# A refusal the caller has to act on is a FAILED_PRECONDITION whose message +# begins with a reason token and ": ", since the service carries no typed +# detail for these yet. A repeat with different content and one behind the +# sequence the server accepted last are both a producer error; a read below +# the retention floor is a cursor error. +_REASON_SEPARATOR = ": " +_REASONS: dict[str, type[Exception]] = { + "STREAM_PRODUCER_CONFLICT": StreamProducerError, + "STREAM_PRODUCER_STALE_SEQUENCE": StreamProducerError, + "STREAM_CURSOR_BELOW_FLOOR": StreamCursorError, +} +# The phrases a server built before the tokens existed sends for the same +# refusals, so a reader of either server gets the typed error. +_PRODUCER_PHRASE = "producer sequence" +_CURSOR_PHRASE = "below the stream's floor" # The codes sdk-core retries. _RETRYABLE = frozenset( @@ -457,17 +474,37 @@ def _to_public(record: stream.StreamRecord) -> StreamEntry: return StreamEntry(record=out, offset=record.offset) -def _translate(error: grpc.aio.AioRpcError) -> Exception: - code = error.code() - details = error.details() or code.name +def translate_error( + code: grpc.StatusCode, details: str, raw_status: bytes = b"" +) -> Exception: + """The SDK error for one failed call, from its status code and message. + + ``NOT_FOUND`` is :class:`temporalio.streams.StreamNotFoundError`. A + ``FAILED_PRECONDITION`` whose message begins with a reason token is the + error the token names: a producer refusal, a repeat with different content + or one behind the sequence the server accepted last, is + :class:`temporalio.streams.StreamProducerError`, because the caller asked + to be deduplicated and could not be; a read below the retention floor is + :class:`temporalio.streams.StreamCursorError`. Everything else is + :class:`temporalio.service.RPCError` with the code and the raw status. + """ + details = details or code.name if code is grpc.StatusCode.NOT_FOUND: return StreamNotFoundError(details) - if code is grpc.StatusCode.INVALID_ARGUMENT and _PRODUCER_CONFLICT in details: - # A producer sequence the store already holds, either with different - # content or behind the one it accepted last. The caller asked to be - # deduplicated and could not be, which is a condition of its own. + if code is grpc.StatusCode.FAILED_PRECONDITION: + token, separator, _ = details.partition(_REASON_SEPARATOR) + typed = _REASONS.get(token) if separator else None + if typed is not None: + return typed(details) + if _CURSOR_PHRASE in details: + return StreamCursorError(details) + if code is grpc.StatusCode.INVALID_ARGUMENT and _PRODUCER_PHRASE in details: return StreamProducerError(details) - return RPCError(details, RPCStatusCode(code.value[0]), _raw_status(error)) + return RPCError(details, RPCStatusCode(code.value[0]), raw_status) + + +def _translate(error: grpc.aio.AioRpcError) -> Exception: + return translate_error(error.code(), error.details() or "", _raw_status(error)) async def _call( diff --git a/tests/test_client_stream_errors.py b/tests/test_client_stream_errors.py new file mode 100644 index 000000000..ae0b452a6 --- /dev/null +++ b/tests/test_client_stream_errors.py @@ -0,0 +1,95 @@ +"""What a refused stream call raises. + +The service carries no typed detail for its refusals, so a typed one arrives +as a ``FAILED_PRECONDITION`` with a reason token at the front of the message. +These pin the mapping: each token to its error, the phrases an older server +sends to the same errors, and an unrelated message on the same code to a +plain ``RPCError`` that keeps the code. +""" + +from __future__ import annotations + +import grpc +import pytest + +from temporalio.client_stream import translate_error +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import ( + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) + + +@pytest.mark.parametrize( + "details", + [ + "STREAM_PRODUCER_CONFLICT: producer sequence 3 already used with different content", + "STREAM_PRODUCER_STALE_SEQUENCE: stale producer sequence 2, last accepted for " + 'producer "p" is 3', + ], +) +def test_a_producer_refusal_is_a_producer_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamProducerError) + assert str(error) == details + + +def test_an_older_servers_producer_refusal_is_a_producer_error() -> None: + # Before the tokens the refusal was an INVALID_ARGUMENT with the phrase. + details = 'stale producer sequence 2, last accepted for producer "p" is 3' + error = translate_error(grpc.StatusCode.INVALID_ARGUMENT, details) + assert isinstance(error, StreamProducerError) + + +@pytest.mark.parametrize( + "details", + [ + "STREAM_CURSOR_BELOW_FLOOR: offset 2 is below the stream's floor of 5", + "offset 2 is below the stream's floor of 5", + ], +) +def test_a_read_below_the_floor_is_a_cursor_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamCursorError) + assert str(error) == details + + +def test_not_found_is_a_not_found_error() -> None: + error = translate_error(grpc.StatusCode.NOT_FOUND, "no stream with id 's'") + assert isinstance(error, StreamNotFoundError) + + +def test_an_unrelated_message_keeps_its_code() -> None: + # The same codes with another message are the server's ordinary refusals. + for code, expected in [ + (grpc.StatusCode.INVALID_ARGUMENT, RPCStatusCode.INVALID_ARGUMENT), + (grpc.StatusCode.FAILED_PRECONDITION, RPCStatusCode.FAILED_PRECONDITION), + (grpc.StatusCode.UNAVAILABLE, RPCStatusCode.UNAVAILABLE), + ]: + error = translate_error(code, "stream is closed", b"raw") + assert type(error) is RPCError + assert error.status == expected + assert error.raw_grpc_status == b"raw" + assert str(error) == "stream is closed" + + +def test_a_token_needs_its_code_and_its_separator() -> None: + # A token on another code is prose, and so is one with no ": " after it. + error = translate_error( + grpc.StatusCode.INVALID_ARGUMENT, "STREAM_PRODUCER_CONFLICT: elsewhere" + ) + assert type(error) is RPCError + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, "STREAM_CURSOR_BELOW_FLOOR" + ) + assert type(error) is RPCError + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, "STREAM_SOMETHING_ELSE: unknown token" + ) + assert type(error) is RPCError + + +def test_an_empty_message_reads_as_the_code() -> None: + error = translate_error(grpc.StatusCode.UNAVAILABLE, "") + assert str(error) == "UNAVAILABLE" From defaf5c061ead0820a48e3e162e8e2c71fa76bd3 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:20:46 -0700 Subject: [PATCH 16/27] Stamped every native record with the hash of its plaintext body. The server deduplicated a producer's repeat on the encoded bytes, so a codec with a fresh nonce per call made every retry look divergent. The provider now puts the hex SHA-256 of the converted body under temporal.io/content-hash in the record's metadata, before the codec on the outside path and before the worker's payload pass on the workflow path; the server dedupes on it when present. The live nonce-codec case needs a server with that dedupe. --- temporalio/streams/providers/native.py | 29 ++- tests/streams/test_native_fingerprint.py | 116 +++++++++++ .../test_workflow_stream_payloads_e2e.py | 184 ++++++++++++++++++ 3 files changed, 326 insertions(+), 3 deletions(-) create mode 100644 tests/streams/test_native_fingerprint.py create mode 100644 tests/worker/test_workflow_stream_payloads_e2e.py diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 751489727..0a420c76c 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -31,11 +31,13 @@ from __future__ import annotations +import hashlib import logging from collections.abc import AsyncGenerator from typing import Any, Generic, TypeVar from temporalio import workflow +from temporalio.api.common.v1 import Payload from temporalio.api.stream.v1 import StreamStartPosition from temporalio.client import Client, WorkflowHistoryEventFilterType from temporalio.client_stream import ( @@ -71,6 +73,7 @@ from temporalio.streams.providers import ProviderPlugin __all__ = [ + "CONTENT_HASH_KEY", "NativeActivityStreamHandle", "NativeProducer", "NativeStreamHandle", @@ -81,6 +84,9 @@ _PROVIDER = "native" +CONTENT_HASH_KEY = "temporal.io/content-hash" +"""The record metadata key carrying the hex SHA-256 of the plaintext body.""" + logger = logging.getLogger(__name__) @@ -110,6 +116,22 @@ def _position(after: Cursor) -> tuple[str, int] | None: ) from None +def _fingerprint(record: WireRecord) -> WireRecord: + """Stamp ``record`` with the hash of its body as converted, before a codec or offload. + + The server deduplicates a producer's repeat on this when present, so a + codec that encrypts with a fresh nonce per call cannot turn a retry into + a divergent write. A record without a body has nothing to compare. + """ + if not record.HasField("body"): + return record + digest = hashlib.sha256(record.body.SerializeToString()).hexdigest() + record.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload(metadata={"encoding": b"binary/plain"}, data=digest.encode()) + ) + return record + + async def _encode_body(codec: PayloadCodec | None, record: WireRecord) -> WireRecord: # The worker's payload visitor runs a codec over the bodies a workflow # publishes and receives; the outside half has no such pass, so it applies @@ -169,8 +191,9 @@ def __init__(self, topic: str) -> None: def publish(self, record: WireRecord) -> None: # Held by the runtime until the task completes, when the task's # records on this topic become one command the server applies with - # the task: rule 1 through the server's own commit. - workflow._append_stream_records([record], stream_name=self._topic) + # the task: rule 1 through the server's own commit. The body is still + # plaintext here; the worker's payload pass runs after the task. + workflow._append_stream_records([_fingerprint(record)], stream_name=self._topic) class _NativeWorkflowProvider: @@ -311,7 +334,7 @@ async def _write(self, records: list[WireRecord]) -> Cursor: if not self._handle.owner_run_id: self._handle.pin(await self._pin()) for record in records: - await _encode_body(self._codec, record) + await _encode_body(self._codec, _fingerprint(record)) appended = await self._handle.append( *records, producer_id=self._writer, sequence=self._sequence ) diff --git a/tests/streams/test_native_fingerprint.py b/tests/streams/test_native_fingerprint.py new file mode 100644 index 000000000..0134b379a --- /dev/null +++ b/tests/streams/test_native_fingerprint.py @@ -0,0 +1,116 @@ +"""The native provider stamps every body with the hash of its plaintext. + +The server deduplicates a producer's repeat on that hash when it is present, +so a codec that encrypts with a fresh nonce per call cannot turn a retry into +a divergent write. These pin where the stamp is taken: over the body as the +payload converter produced it, before the codec on the outside path and before +the worker's payload pass on the workflow path. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Sequence +from typing import Any + +import pytest + +import temporalio.converter +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.stream.v1 import StreamRecord +from temporalio.client_stream import Appended +from temporalio.converter import PayloadCodec +from temporalio.streams.providers import native +from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeProducer + + +class _NonceCodec(PayloadCodec): + """Encodes to something different every time, like a nonce-based cipher.""" + + def __init__(self) -> None: + self.calls = 0 + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + self.calls += 1 + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=f"{self.calls}:".encode() + p.data, + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [Payload(data=p.data.split(b":", 1)[1]) for p in payloads] + + +class _Handle: + """Records what a producer appends and answers as the server would.""" + + def __init__(self) -> None: + self.owner_run_id = "run" + self.appended: list[StreamRecord] = [] + + async def append(self, *records: StreamRecord, **_: Any) -> Appended: + self.appended.extend(records) + return Appended(first_offset=0, next_offset=len(records), count=len(records)) + + +def _plaintext_hash(value: Any) -> bytes: + payload = temporalio.converter.default().payload_converter.to_payloads([value])[0] + return hashlib.sha256(payload.SerializeToString()).hexdigest().encode() + + +async def test_an_outside_append_is_stamped_before_the_codec() -> None: + handle: Any = _Handle() + codec = _NonceCodec() + converter = temporalio.converter.default().payload_converter + producer: NativeProducer[Any] = NativeProducer( + handle, None, codec, converter, "t", "p", 1 + ) + + await producer.append({"n": 1}) + await producer.append({"n": 1}) + + first, second = handle.appended + # The bodies went out differently encoded and the stamps agree anyway. + assert first.body.data != second.body.data + assert first.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 1}) + assert second.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 1}) + assert first.metadata[CONTENT_HASH_KEY].metadata["encoding"] == b"binary/plain" + + +async def test_a_finish_record_carries_no_stamp() -> None: + handle: Any = _Handle() + converter = temporalio.converter.default().payload_converter + producer: NativeProducer[Any] = NativeProducer( + handle, None, None, converter, "t", "p", 1 + ) + await producer.finish() + (record,) = handle.appended + assert CONTENT_HASH_KEY not in record.metadata + + +def test_a_workflow_publish_is_stamped_on_the_workflow_thread( + monkeypatch: pytest.MonkeyPatch, +) -> None: + staged: list[StreamRecord] = [] + + def stage(records: Sequence[StreamRecord], *, stream_name: str) -> None: + assert stream_name == "t" + staged.extend(records) + + monkeypatch.setattr(workflow, "_append_stream_records", stage) + converter = temporalio.converter.default().payload_converter + body = converter.to_payloads([{"n": 2}])[0] + record = StreamRecord(topic="t", body=body) + + native._NativeWriteSink("t").publish(record) + + (stamped,) = staged + assert stamped.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 2}) + # Two publishes of the same value stamp the same, which is what replay + # reissues. + native._NativeWriteSink("t").publish(StreamRecord(topic="t", body=body)) + assert staged[1].metadata[CONTENT_HASH_KEY] == stamped.metadata[CONTENT_HASH_KEY] diff --git a/tests/worker/test_workflow_stream_payloads_e2e.py b/tests/worker/test_workflow_stream_payloads_e2e.py new file mode 100644 index 000000000..97357a0d5 --- /dev/null +++ b/tests/worker/test_workflow_stream_payloads_e2e.py @@ -0,0 +1,184 @@ +"""How a body's encoding meets the server: retry identity and external storage. + +A payload codec and an ``ExternalStorage`` driver both change the bytes a body +is stored as. The server must still recognize a retry, and a driver must be +applied on both halves of the native provider. Needs a Temporal server built +from the AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/worker/test_workflow_stream_payloads_e2e.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from collections.abc import Sequence +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_stream import StreamClient +from temporalio.converter import DataConverter, PayloadCodec +from temporalio.streams import RecordKind +from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeStreams +from temporalio.worker import Worker +from tests.streams.test_streams_conformance import take + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream commands", +) + +INPUTS = "inputs" +DECISIONS = "decisions" + + +class _NonceCodec(PayloadCodec): + """Encodes to different bytes on every call, as a nonce-based cipher does. + + The plaintext is kept in the clear behind a counter so the test can read + the stored bytes back and see that two encodings of one value differ. + """ + + def __init__(self) -> None: + self.calls = 0 + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + self.calls += 1 + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=f"{self.calls}:".encode() + p.SerializeToString(), + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + out: list[Payload] = [] + for p in payloads: + if p.metadata.get("encoding", b"") != b"binary/nonce": + out.append(p) + continue + out.append(Payload.FromString(p.data.split(b":", 1)[1])) + return out + + +@workflow.defn +class Echo: + """Reads ``inputs`` and publishes each value on ``decisions``, until told to stop.""" + + def __init__(self) -> None: + self._stop = False + + @workflow.signal + def stop(self) -> None: + self._stop = True + + @workflow.run + async def run(self, count: int) -> list[Any]: + inputs = workflow.stream_reader(INPUTS, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) + seen: list[Any] = [] + async for value in inputs.values(): + decisions.publish({"echo": value}) + seen.append(value) + if len(seen) >= count: + break + decisions.finish() + await workflow.wait_condition(lambda: self._stop) + return seen + + +async def _connect(converter: DataConverter, provider: NativeStreams) -> Client: + plain = await Client.connect(TARGET or "") + config = plain.config() + config["data_converter"] = converter + config["plugins"] = [provider] + return Client(**config) + + +async def test_a_retried_append_under_a_nonce_codec_is_deduplicated() -> None: + """Retry identity is the plaintext hash, not the encoded bytes. + + Two producers with the same identity append the same value at the same + sequence, as a retry after a lost response would. The codec encodes each + differently, so the server sees two different bodies and one hash, and + writes the record once. + """ + provider = NativeStreams() + client = await _connect(DataConverter(payload_codec=_NonceCodec()), provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "nonce-tq-" + uuid.uuid4().hex[:8] + workflow_id = "nonce-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + first = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + retry = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + landed = await first.append({"n": 1}) + again = await retry.append({"n": 1}) + assert again == landed, "the repeat answers with the original position" + + echoed = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert [r.value for r in echoed] == [{"echo": {"n": 1}}] + await handle.signal(Echo.stop) + assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}] + + run_id = (await handle.describe()).run_id + assert run_id is not None + page = await raw.workflow_stream(workflow_id, INPUTS, owner_run_id=run_id).poll( + from_offset=0, wait=False + ) + bodies = [e.record for e in page.entries if e.record.HasField("body")] + assert len(bodies) == 1, "written once" + assert bodies[0].body.metadata["encoding"] == b"binary/nonce" + assert bodies[0].metadata[CONTENT_HASH_KEY].data + finally: + await raw.close() + await provider.close() + + +async def test_a_workflow_publish_carries_the_plaintext_hash() -> None: + """The workflow's own records are stamped too, from the workflow thread.""" + provider = NativeStreams() + client = await _connect(DataConverter.default, provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "stamp-tq-" + uuid.uuid4().hex[:8] + workflow_id = "stamp-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + await stream.producer(topic=INPUTS, producer_id="tool", attempt=1).append( + {"n": 7} + ) + await take(stream.read(topic=DECISIONS, result_type=dict), 1) + await handle.signal(Echo.stop) + await asyncio.wait_for(handle.result(), 60) + + run_id = (await handle.describe()).run_id + assert run_id is not None + page = await raw.workflow_stream( + workflow_id, DECISIONS, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + kinds = {e.record.kind for e in page.entries} + assert len(page.entries) == 2, "one decision and the finish" + data = [e.record for e in page.entries if e.record.HasField("body")] + assert len(data) == 1 and data[0].metadata[CONTENT_HASH_KEY].data + finish = [e.record for e in page.entries if not e.record.HasField("body")] + assert CONTENT_HASH_KEY not in finish[0].metadata + assert RecordKind.FINISH in {RecordKind(k) for k in kinds} + finally: + await raw.close() + await provider.close() From 0ee2aaa35d826be4aecea6006cab272cee5556a1 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:26:01 -0700 Subject: [PATCH 17/27] Routed native stream bodies through the client's data converter. The outside half applied the payload codec alone, so an ExternalStorage driver on the client never saw a stream body and a body the worker offloaded could not be read back. Encode and decode now run the codec and the external store in the worker's order, an offloaded body is stored under the owning execution, and the fingerprint is taken before either. --- temporalio/streams/providers/native.py | 81 +++++++++++++------ tests/streams/test_native_fingerprint.py | 18 +++-- .../test_workflow_stream_payloads_e2e.py | 70 +++++++++++++++- 3 files changed, 139 insertions(+), 30 deletions(-) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 0a420c76c..70a0de6a5 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -33,7 +33,7 @@ import hashlib import logging -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, Callable from typing import Any, Generic, TypeVar from temporalio import workflow @@ -49,7 +49,11 @@ shared_client, shared_key, ) -from temporalio.converter import PayloadCodec, PayloadConverter +from temporalio.converter import ( + DataConverter, + StorageDriverStoreContext, + StorageDriverWorkflowInfo, +) from temporalio.service import RPCError, RPCStatusCode from temporalio.streams._errors import StreamCursorError, StreamNotFoundError from temporalio.streams._provider import ReadSource, WriteSink @@ -132,21 +136,24 @@ def _fingerprint(record: WireRecord) -> WireRecord: return record -async def _encode_body(codec: PayloadCodec | None, record: WireRecord) -> WireRecord: - # The worker's payload visitor runs a codec over the bodies a workflow - # publishes and receives; the outside half has no such pass, so it applies - # the client's codec here or the two sides would not agree. - if codec is None or not record.HasField("body"): +async def _encode_body(converter: DataConverter, record: WireRecord) -> WireRecord: + # The worker's payload pass runs the codec and then the external store over + # the bodies a workflow publishes and receives; the outside half has no such + # pass, so it applies the client's data converter here in the same order, + # or the two sides would not agree. + if not record.HasField("body"): return record - encoded = await codec.encode([record.body]) - record.body.CopyFrom(encoded[0]) + encoded = await converter._encode_payload_sequence([record.body]) + stored = await converter._external_store_payload_sequence(encoded) + record.body.CopyFrom(stored[0]) return record -async def _decode_body(codec: PayloadCodec | None, record: WireRecord) -> WireRecord: - if codec is None or not record.HasField("body"): +async def _decode_body(converter: DataConverter, record: WireRecord) -> WireRecord: + if not record.HasField("body"): return record - decoded = await codec.decode([record.body]) + retrieved = await converter._external_retrieve_payload_sequence([record.body]) + decoded = await converter._decode_payload_sequence(retrieved) record.body.CopyFrom(decoded[0]) return record @@ -249,17 +256,24 @@ def __init__( self, handle: WorkflowStreamHandle, pin: Any, - codec: PayloadCodec | None, - converter: PayloadConverter, + converter: DataConverter, + store_target: Callable[[str], StorageDriverWorkflowInfo], topic: str, producer_id: str, attempt: int, ) -> None: - """Bind this producer to ``topic`` on the stream ``handle`` names.""" + """Bind this producer to ``topic`` on the stream ``handle`` names. + + ``converter`` is the client's data converter, applied to every body + as the worker applies it to a workflow's own records. ``store_target`` + names the execution an offloaded body is stored under, given the run + the producer pinned. + """ self._handle = handle self._pin = pin - self._codec = codec self._converter = converter + self._store_target = store_target + self._bound: DataConverter | None = None self._topic = topic self._producer_id = producer_id self._attempt = attempt @@ -301,7 +315,7 @@ async def append(self, *values: T) -> Cursor: return await self._write( [ to_wire( - self._converter, + self._converter.payload_converter, topic=self._topic, kind=RecordKind.DATA, value=value, @@ -318,7 +332,7 @@ async def finish(self) -> None: await self._write( [ to_wire( - self._converter, + self._converter.payload_converter, topic=self._topic, kind=RecordKind.FINISH, producer_id=self._producer_id, @@ -333,8 +347,16 @@ async def _write(self, records: list[WireRecord]) -> Cursor: # out names the run its records landed in. if not self._handle.owner_run_id: self._handle.pin(await self._pin()) + if self._bound is None: + # An offloaded body is stored under the execution that owns the + # stream, as the worker stores a workflow's own. + self._bound = self._converter._with_store_context( + StorageDriverStoreContext( + target=self._store_target(self._handle.owner_run_id) + ) + ) for record in records: - await _encode_body(self._codec, _fingerprint(record)) + await _encode_body(self._bound, _fingerprint(record)) appended = await self._handle.append( *records, producer_id=self._writer, sequence=self._sequence ) @@ -363,8 +385,8 @@ def __init__( self._workflow_id = workflow_id self._run_id = run_id self._opened = set() if opened is None else opened + self._data_converter = client.data_converter self._converter = client.data_converter.payload_converter - self._codec = client.data_converter.payload_codec self._streams: StreamClient | None = None def _service(self) -> StreamClient: @@ -443,7 +465,7 @@ async def _read( ) start = None for entry in page.entries: - record = await _decode_body(self._codec, entry.record) + record = await _decode_body(self._data_converter, entry.record) for out in decoder.decode(_cursor(run_id, entry.offset), record): yield out offset = page.next_offset @@ -507,13 +529,19 @@ def producer( return NativeProducer( self._stream(topic, self._run_id or ""), self._current_run, - self._codec, - self._converter, + self._data_converter, + self._store_target, topic, producer_id, attempt, ) + def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: + """The execution an offloaded body of this owner's stream is stored under.""" + return StorageDriverWorkflowInfo( + namespace=self._client.namespace, id=self._workflow_id, run_id=run_id + ) + async def _current_run(self) -> str: try: description = await self._client.get_workflow_handle( @@ -597,6 +625,13 @@ def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: self._activity_id, topic, workflow_id=self._workflow_id, run_id=run_id ) + def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: + # A workflow's activity stores under that workflow, as the worker does + # for its activities; a standalone activity has no workflow to name. + if self._workflow_id: + return super()._store_target(run_id) + return StorageDriverWorkflowInfo(namespace=self._client.namespace) + async def _current_run(self) -> str: if self._workflow_id: return await super()._current_run() diff --git a/tests/streams/test_native_fingerprint.py b/tests/streams/test_native_fingerprint.py index 0134b379a..a66c4c79f 100644 --- a/tests/streams/test_native_fingerprint.py +++ b/tests/streams/test_native_fingerprint.py @@ -20,7 +20,11 @@ from temporalio.api.common.v1 import Payload from temporalio.api.stream.v1 import StreamRecord from temporalio.client_stream import Appended -from temporalio.converter import PayloadCodec +from temporalio.converter import ( + DataConverter, + PayloadCodec, + StorageDriverWorkflowInfo, +) from temporalio.streams.providers import native from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeProducer @@ -62,12 +66,15 @@ def _plaintext_hash(value: Any) -> bytes: return hashlib.sha256(payload.SerializeToString()).hexdigest().encode() +def _target(_run_id: str) -> StorageDriverWorkflowInfo: + return StorageDriverWorkflowInfo(namespace="ns", id="wf") + + async def test_an_outside_append_is_stamped_before_the_codec() -> None: handle: Any = _Handle() - codec = _NonceCodec() - converter = temporalio.converter.default().payload_converter + converter = DataConverter(payload_codec=_NonceCodec()) producer: NativeProducer[Any] = NativeProducer( - handle, None, codec, converter, "t", "p", 1 + handle, None, converter, _target, "t", "p", 1 ) await producer.append({"n": 1}) @@ -83,9 +90,8 @@ async def test_an_outside_append_is_stamped_before_the_codec() -> None: async def test_a_finish_record_carries_no_stamp() -> None: handle: Any = _Handle() - converter = temporalio.converter.default().payload_converter producer: NativeProducer[Any] = NativeProducer( - handle, None, None, converter, "t", "p", 1 + handle, None, DataConverter.default, _target, "t", "p", 1 ) await producer.finish() (record,) = handle.appended diff --git a/tests/worker/test_workflow_stream_payloads_e2e.py b/tests/worker/test_workflow_stream_payloads_e2e.py index 97357a0d5..fd300d512 100644 --- a/tests/worker/test_workflow_stream_payloads_e2e.py +++ b/tests/worker/test_workflow_stream_payloads_e2e.py @@ -20,13 +20,20 @@ from temporalio import workflow from temporalio.api.common.v1 import Payload +from temporalio.api.sdk.v1.external_storage_pb2 import ExternalStorageReference from temporalio.client import Client from temporalio.client_stream import StreamClient -from temporalio.converter import DataConverter, PayloadCodec +from temporalio.converter import ( + DataConverter, + ExternalStorage, + JSONProtoPayloadConverter, + PayloadCodec, +) from temporalio.streams import RecordKind from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeStreams from temporalio.worker import Worker from tests.streams.test_streams_conformance import take +from tests.test_extstore import InMemoryTestDriver TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") @@ -182,3 +189,64 @@ async def test_a_workflow_publish_carries_the_plaintext_hash() -> None: finally: await raw.close() await provider.close() + + +async def test_external_storage_applies_on_both_halves() -> None: + """A ``StorageDriver`` on the client offloads stream bodies on both paths. + + Every body is over the threshold. The outside producer's append is stored + through the driver before it reaches the server; the workflow retrieves it + on its task, publishes, and the worker's payload pass offloads that too, off + the workflow thread, so the outside reader retrieves it. Read raw, the + server holds references on both topics and no plaintext. + """ + driver = InMemoryTestDriver() + converter = DataConverter( + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=0) + ) + provider = NativeStreams() + client = await _connect(converter, provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "offload-tq-" + uuid.uuid4().hex[:8] + workflow_id = "offload-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "plaintext-marker"}) + stores_after_append = driver._store_calls + assert stores_after_append >= 1, "the outside append offloaded" + + echoed = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert [r.value for r in echoed] == [{"echo": {"n": "plaintext-marker"}}] + await handle.signal(Echo.stop) + assert await asyncio.wait_for(handle.result(), 60) == [ + {"n": "plaintext-marker"} + ] + + run_id = (await handle.describe()).run_id + assert run_id is not None + for topic in (INPUTS, DECISIONS): + page = await raw.workflow_stream( + workflow_id, topic, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + bodies = [e.record.body for e in page.entries if e.record.HasField("body")] + assert len(bodies) == 1, topic + assert b"plaintext-marker" not in bodies[0].data, topic + reference = JSONProtoPayloadConverter().from_payload( + bodies[0], ExternalStorageReference + ) + assert reference.driver_name == driver.name(), topic + # The workflow's publish was offloaded by the worker, after the append. + assert driver._store_calls > stores_after_append + # Both halves retrieved: the worker on its task and the outside reader. + assert driver._retrieve_calls >= 2 + # The outside append was stored under the workflow that owns the stream. + targets = [ctx.target for ctx in driver._store_contexts if ctx.target] + assert any(t.id == workflow_id and t.run_id == run_id for t in targets) + finally: + await raw.close() + await provider.close() From 60fd7eca8466ffa2768893906082b553cf5797ea Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:28:35 -0700 Subject: [PATCH 18/27] Kept the declared content hash out of the worker's payload pass. The payload visitor walks a stream record's metadata payloads like any other, so a codec encrypted the hash a workflow publish declared and an external store could offload it. The server reads that value as sent and refuses one that is not hex, which would have failed the workflow's own append. The key and the digest now live in the streams wire module, where both halves and the worker read them. --- temporalio/streams/_wire.py | 18 +++++++++++ temporalio/streams/providers/native.py | 13 ++++---- temporalio/worker/_command_aware_visitor.py | 14 +++++++++ tests/streams/test_native_fingerprint.py | 31 ++++++++++++++++++- .../test_workflow_stream_payloads_e2e.py | 3 +- 5 files changed, 70 insertions(+), 9 deletions(-) diff --git a/temporalio/streams/_wire.py b/temporalio/streams/_wire.py index 4913cc308..7f9818ace 100644 --- a/temporalio/streams/_wire.py +++ b/temporalio/streams/_wire.py @@ -14,10 +14,12 @@ from __future__ import annotations +import hashlib from collections.abc import Callable from typing import Any import temporalio.converter +from temporalio.api.common.v1 import Payload from temporalio.api.stream.v1 import StreamRecord as WireRecord from temporalio.api.stream.v1 import StreamRecordKind from temporalio.streams._errors import StreamCursorError @@ -25,8 +27,10 @@ from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord __all__ = [ + "CONTENT_HASH_KEY", "RecordDecoder", "WireRecord", + "content_hash", "cursor_position", "from_wire", "mint_cursor", @@ -34,6 +38,20 @@ "to_wire", ] +CONTENT_HASH_KEY = "temporal.io/content-hash" +"""The record metadata key under which a producer declares its body's identity. + +The value is the lowercase hex SHA-256 of the body as the payload converter +produced it, before any codec or offload, so a store can tell a retry from a +divergent repeat whatever the encoding did to the bytes. It is read as sent: +nothing may encode or offload this one metadata payload. +""" + + +def content_hash(body: Payload) -> str: + """The lowercase hex SHA-256 of ``body`` as converted.""" + return hashlib.sha256(body.SerializeToString()).hexdigest() + def to_wire( converter: temporalio.converter.PayloadConverter, diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 70a0de6a5..8311cf0b1 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -31,7 +31,6 @@ from __future__ import annotations -import hashlib import logging from collections.abc import AsyncGenerator, Callable from typing import Any, Generic, TypeVar @@ -67,8 +66,10 @@ ) from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( + CONTENT_HASH_KEY, RecordDecoder, WireRecord, + content_hash, cursor_position, mint_cursor, producer_identity, @@ -77,7 +78,6 @@ from temporalio.streams.providers import ProviderPlugin __all__ = [ - "CONTENT_HASH_KEY", "NativeActivityStreamHandle", "NativeProducer", "NativeStreamHandle", @@ -88,9 +88,6 @@ _PROVIDER = "native" -CONTENT_HASH_KEY = "temporal.io/content-hash" -"""The record metadata key carrying the hex SHA-256 of the plaintext body.""" - logger = logging.getLogger(__name__) @@ -129,9 +126,11 @@ def _fingerprint(record: WireRecord) -> WireRecord: """ if not record.HasField("body"): return record - digest = hashlib.sha256(record.body.SerializeToString()).hexdigest() record.metadata[CONTENT_HASH_KEY].CopyFrom( - Payload(metadata={"encoding": b"binary/plain"}, data=digest.encode()) + Payload( + metadata={"encoding": b"binary/plain"}, + data=content_hash(record.body).encode(), + ) ) return record diff --git a/temporalio/worker/_command_aware_visitor.py b/temporalio/worker/_command_aware_visitor.py index 7c03c2cd4..bc3e3f35a 100644 --- a/temporalio/worker/_command_aware_visitor.py +++ b/temporalio/worker/_command_aware_visitor.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from temporalio.api.enums.v1.command_type_pb2 import CommandType +from temporalio.api.stream.v1 import StreamRecord from temporalio.bridge._visitor import PayloadVisitor from temporalio.bridge._visitor_functions import VisitorFunctions from temporalio.bridge.proto.workflow_activation.workflow_activation_pb2 import ( @@ -26,6 +27,7 @@ StartChildWorkflowExecution, WorkflowCommand, ) +from temporalio.streams._wire import CONTENT_HASH_KEY @dataclass(frozen=True) @@ -116,6 +118,18 @@ async def _visit_coresdk_workflow_commands_ScheduleNexusOperation( with current_command(CommandType.COMMAND_TYPE_SCHEDULE_NEXUS_OPERATION, o.seq): await super()._visit_coresdk_workflow_commands_ScheduleNexusOperation(fs, o) + async def _visit_temporal_api_stream_v1_StreamRecord( + self, fs: VisitorFunctions, o: StreamRecord + ) -> None: + if o.HasField("body"): + await self._visit_temporal_api_common_v1_Payload(fs, o.body) + for key, value in o.metadata.items(): + # The declared content hash is the server's dedupe identity and is + # read as sent, so neither the codec nor the store may rewrite it. + if key == CONTENT_HASH_KEY: + continue + await self._visit_temporal_api_common_v1_Payload(fs, value) + async def _visit_coresdk_workflow_commands_WorkflowCommand( self, fs: VisitorFunctions, o: WorkflowCommand ) -> None: diff --git a/tests/streams/test_native_fingerprint.py b/tests/streams/test_native_fingerprint.py index a66c4c79f..9cac71755 100644 --- a/tests/streams/test_native_fingerprint.py +++ b/tests/streams/test_native_fingerprint.py @@ -19,14 +19,17 @@ from temporalio import workflow from temporalio.api.common.v1 import Payload from temporalio.api.stream.v1 import StreamRecord +from temporalio.bridge.proto.workflow_completion import WorkflowActivationCompletion +from temporalio.bridge.worker import encode_completion from temporalio.client_stream import Appended from temporalio.converter import ( DataConverter, PayloadCodec, StorageDriverWorkflowInfo, ) +from temporalio.streams._wire import CONTENT_HASH_KEY from temporalio.streams.providers import native -from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeProducer +from temporalio.streams.providers.native import NativeProducer class _NonceCodec(PayloadCodec): @@ -120,3 +123,29 @@ def stage(records: Sequence[StreamRecord], *, stream_name: str) -> None: # reissues. native._NativeWriteSink("t").publish(StreamRecord(topic="t", body=body)) assert staged[1].metadata[CONTENT_HASH_KEY] == stamped.metadata[CONTENT_HASH_KEY] + + +async def test_the_workers_payload_pass_leaves_the_stamp_alone() -> None: + """The codec and the store rewrite the body and the other metadata, never the hash. + + The server reads the declared hash as sent and refuses a value that is not + hex, so a codec that touched it would fail the workflow's own append. + """ + converter = DataConverter(payload_codec=_NonceCodec()) + body = converter.payload_converter.to_payloads([{"n": 3}])[0] + record = native._fingerprint(StreamRecord(topic="t", body=body)) + record.metadata["note"].CopyFrom(Payload(data=b"plain")) + completion = WorkflowActivationCompletion() + command = completion.successful.commands.add() + command.append_stream_records.stream_name = "t" + command.append_stream_records.records.append(record) + + await encode_completion( + completion, converter, encode_headers=False, storage_concurrency_limit=1 + ) + + sent = completion.successful.commands[0].append_stream_records.records[0] + assert sent.body.metadata["encoding"] == b"binary/nonce" + assert sent.metadata["note"].metadata["encoding"] == b"binary/nonce" + assert sent.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 3}) + assert sent.metadata[CONTENT_HASH_KEY].metadata["encoding"] == b"binary/plain" diff --git a/tests/worker/test_workflow_stream_payloads_e2e.py b/tests/worker/test_workflow_stream_payloads_e2e.py index fd300d512..54ffe74ea 100644 --- a/tests/worker/test_workflow_stream_payloads_e2e.py +++ b/tests/worker/test_workflow_stream_payloads_e2e.py @@ -30,7 +30,8 @@ PayloadCodec, ) from temporalio.streams import RecordKind -from temporalio.streams.providers.native import CONTENT_HASH_KEY, NativeStreams +from temporalio.streams._wire import CONTENT_HASH_KEY +from temporalio.streams.providers.native import NativeStreams from temporalio.worker import Worker from tests.streams.test_streams_conformance import take from tests.test_extstore import InMemoryTestDriver From f37da2e073fae81dcfa7bf3829ac2b756a88e7a7 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:58:40 -0700 Subject: [PATCH 19/27] Hosted standalone streams and refs on the native provider. A stream with an id of its own is created and sealed over the existing CreateStream and CloseStream calls, read through a handle whose cursors name the stream where an owned stream's name a run, and a read on an id nobody has created yet relies on the server's parked poll. Every native handle answers ref(), the provider opens a StreamRef, and a create of an id that exists is answered from the stream's own policy since the server does not say ALREADY_EXISTS. The content hash and the body encoding now come from the shared streams helper the interface chain added, and an append on a sealed stream is a StreamClosedError. --- temporalio/client_stream.py | 32 +- temporalio/streams/_wire.py | 18 - temporalio/streams/providers/native.py | 389 ++++++++++++++++-- temporalio/worker/_command_aware_visitor.py | 2 +- tests/streams/test_native_fingerprint.py | 2 +- tests/streams/test_native_handles.py | 110 +++++ tests/streams/test_native_standalone_e2e.py | 95 +++++ tests/streams/test_streams_conformance.py | 44 +- .../test_workflow_stream_payloads_e2e.py | 3 +- 9 files changed, 612 insertions(+), 83 deletions(-) create mode 100644 tests/streams/test_native_handles.py create mode 100644 tests/streams/test_native_standalone_e2e.py diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index babab6fc1..8ae0f114c 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -28,9 +28,11 @@ server answers ``NOT_FOUND``, :class:`temporalio.streams.StreamProducerError` when it refuses a producer sequence it already holds, :class:`temporalio.streams.StreamCursorError` when -it refuses a read below the retention floor, and -:class:`temporalio.service.RPCError` otherwise, never the transport's own -exception type. :func:`translate_error` is the one place that decides. +it refuses a read below the retention floor, +:class:`temporalio.streams.StreamClosedError` when it refuses an append to a +sealed stream, and :class:`temporalio.service.RPCError` otherwise, never the +transport's own exception type. :func:`translate_error` is the one place that +decides. A failure sdk-core would retry is retried here, on the same codes and with the same default :class:`temporalio.service.RetryConfig`, because this channel is @@ -52,9 +54,9 @@ import weakref from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass +from datetime import timedelta from typing import TYPE_CHECKING, Any, TypeVar -import google.protobuf.duration_pb2 import grpc import grpc.aio from google.protobuf.message import Message @@ -74,6 +76,7 @@ __version__, ) from temporalio.streams import ( + StreamClosedError, StreamCursorError, StreamNotFoundError, StreamProducerError, @@ -112,9 +115,11 @@ "STREAM_CURSOR_BELOW_FLOOR": StreamCursorError, } # The phrases a server built before the tokens existed sends for the same -# refusals, so a reader of either server gets the typed error. +# refusals, so a reader of either server gets the typed error. An append on a +# sealed stream has no token yet and is matched on its whole message. _PRODUCER_PHRASE = "producer sequence" _CURSOR_PHRASE = "below the stream's floor" +_CLOSED_PHRASE = "stream is closed" # The codes sdk-core retries. _RETRYABLE = frozenset( @@ -498,6 +503,8 @@ def translate_error( return typed(details) if _CURSOR_PHRASE in details: return StreamCursorError(details) + if details == _CLOSED_PHRASE: + return StreamClosedError(details) if code is grpc.StatusCode.INVALID_ARGUMENT and _PRODUCER_PHRASE in details: return StreamProducerError(details) return RPCError(details, RPCStatusCode(code.value[0]), raw_status) @@ -615,20 +622,21 @@ async def create( self, stream_id: str, *, - retention: float | None = None, + retention: float | timedelta | None = None, max_items: int | None = None, ) -> StreamHandle: """Create a stream and return a handle to it. - ``retention`` is how long a closed stream stays readable, in seconds. - ``max_items`` caps how many records remain readable, dropping the - oldest, which bounds storage for a stream nobody truncates. + ``retention`` is how long a closed stream stays readable, in seconds + or as a ``timedelta``. ``max_items`` caps how many records remain + readable, dropping the oldest, which bounds storage for a stream + nobody truncates. """ lifecycle = stream.StreamLifecycle() if retention is not None: - lifecycle.retention.CopyFrom( - google.protobuf.duration_pb2.Duration(seconds=int(retention)) - ) + if not isinstance(retention, timedelta): + retention = timedelta(seconds=retention) + lifecycle.retention.FromTimedelta(retention) if max_items is not None: lifecycle.max_items = max_items diff --git a/temporalio/streams/_wire.py b/temporalio/streams/_wire.py index 7f9818ace..4913cc308 100644 --- a/temporalio/streams/_wire.py +++ b/temporalio/streams/_wire.py @@ -14,12 +14,10 @@ from __future__ import annotations -import hashlib from collections.abc import Callable from typing import Any import temporalio.converter -from temporalio.api.common.v1 import Payload from temporalio.api.stream.v1 import StreamRecord as WireRecord from temporalio.api.stream.v1 import StreamRecordKind from temporalio.streams._errors import StreamCursorError @@ -27,10 +25,8 @@ from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord __all__ = [ - "CONTENT_HASH_KEY", "RecordDecoder", "WireRecord", - "content_hash", "cursor_position", "from_wire", "mint_cursor", @@ -38,20 +34,6 @@ "to_wire", ] -CONTENT_HASH_KEY = "temporal.io/content-hash" -"""The record metadata key under which a producer declares its body's identity. - -The value is the lowercase hex SHA-256 of the body as the payload converter -produced it, before any codec or offload, so a store can tell a retry from a -divergent repeat whatever the encoding did to the bytes. It is read as sent: -nothing may encode or offload this one metadata payload. -""" - - -def content_hash(body: Payload) -> str: - """The lowercase hex SHA-256 of ``body`` as converted.""" - return hashlib.sha256(body.SerializeToString()).hexdigest() - def to_wire( converter: temporalio.converter.PayloadConverter, diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 8311cf0b1..6904c7cae 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -31,8 +31,11 @@ from __future__ import annotations +import asyncio import logging +import time from collections.abc import AsyncGenerator, Callable +from datetime import timedelta from typing import Any, Generic, TypeVar from temporalio import workflow @@ -40,6 +43,7 @@ from temporalio.api.stream.v1 import StreamStartPosition from temporalio.client import Client, WorkflowHistoryEventFilterType from temporalio.client_stream import ( + Appended, Page, SharedKey, StreamClient, @@ -48,14 +52,25 @@ shared_client, shared_key, ) +from temporalio.client_stream import StreamHandle as ServiceStreamHandle from temporalio.converter import ( DataConverter, StorageDriverStoreContext, StorageDriverWorkflowInfo, ) from temporalio.service import RPCError, RPCStatusCode -from temporalio.streams._errors import StreamCursorError, StreamNotFoundError -from temporalio.streams._provider import ReadSource, WriteSink +from temporalio.streams._body import ( + CONTENT_HASH_KEY, + content_hash, + decode_body, + encode_body, +) +from temporalio.streams._errors import ( + StreamCursorError, + StreamNotFoundError, + StreamUnsupportedError, +) +from temporalio.streams._provider import ReadSource, StreamHandle, WriteSink from temporalio.streams._record import ( BEGINNING, END, @@ -64,12 +79,11 @@ StreamRecord, check_read_start, ) +from temporalio.streams._ref import StreamRef, open_ref from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( - CONTENT_HASH_KEY, RecordDecoder, WireRecord, - content_hash, cursor_position, mint_cursor, producer_identity, @@ -80,6 +94,7 @@ __all__ = [ "NativeActivityStreamHandle", "NativeProducer", + "NativeStandaloneStreamHandle", "NativeStreamHandle", "NativeStreams", ] @@ -88,6 +103,10 @@ _PROVIDER = "native" +# Seconds between polls on a standalone stream that does not exist yet, when +# the server answers without parking. +_CREATE_WAIT_PACE = 1.0 + logger = logging.getLogger(__name__) @@ -97,6 +116,7 @@ def _require_topic(topic: str) -> None: def _cursor(run_id: str, offset: int) -> Cursor: + # A standalone stream has no run; its id takes the run's place. return mint_cursor(_PROVIDER, f"{run_id}:{offset}") @@ -118,11 +138,12 @@ def _position(after: Cursor) -> tuple[str, int] | None: def _fingerprint(record: WireRecord) -> WireRecord: - """Stamp ``record`` with the hash of its body as converted, before a codec or offload. + """Stamp ``record`` with the hash of its body as converted, before the worker's pass. - The server deduplicates a producer's repeat on this when present, so a - codec that encrypts with a fresh nonce per call cannot turn a retry into - a divergent write. A record without a body has nothing to compare. + The outside half gets the same stamp from :func:`encode_body`; a workflow's + own publish is encoded later, by the worker's payload pass, so the stamp + is taken here while the body is still what the converter produced. A + record without a body has nothing to compare. """ if not record.HasField("body"): return record @@ -135,28 +156,6 @@ def _fingerprint(record: WireRecord) -> WireRecord: return record -async def _encode_body(converter: DataConverter, record: WireRecord) -> WireRecord: - # The worker's payload pass runs the codec and then the external store over - # the bodies a workflow publishes and receives; the outside half has no such - # pass, so it applies the client's data converter here in the same order, - # or the two sides would not agree. - if not record.HasField("body"): - return record - encoded = await converter._encode_payload_sequence([record.body]) - stored = await converter._external_store_payload_sequence(encoded) - record.body.CopyFrom(stored[0]) - return record - - -async def _decode_body(converter: DataConverter, record: WireRecord) -> WireRecord: - if not record.HasField("body"): - return record - retrieved = await converter._external_retrieve_payload_sequence([record.body]) - decoded = await converter._decode_payload_sequence(retrieved) - record.body.CopyFrom(decoded[0]) - return record - - class _NativeReadSource: """One subscription of the running workflow, fed by delivered ranges.""" @@ -253,7 +252,7 @@ class NativeProducer(Generic[T]): def __init__( self, - handle: WorkflowStreamHandle, + handle: WorkflowStreamHandle | _StandaloneTarget, pin: Any, converter: DataConverter, store_target: Callable[[str], StorageDriverWorkflowInfo], @@ -264,9 +263,9 @@ def __init__( """Bind this producer to ``topic`` on the stream ``handle`` names. ``converter`` is the client's data converter, applied to every body - as the worker applies it to a workflow's own records. ``store_target`` - names the execution an offloaded body is stored under, given the run - the producer pinned. + as the worker applies it to a workflow's own records, with the + plaintext hash stamped first. ``store_target`` names the execution an + offloaded body is stored under, given the run the producer pinned. """ self._handle = handle self._pin = pin @@ -355,7 +354,7 @@ async def _write(self, records: list[WireRecord]) -> Cursor: ) ) for record in records: - await _encode_body(self._bound, _fingerprint(record)) + await encode_body(self._bound, record) appended = await self._handle.append( *records, producer_id=self._writer, sequence=self._sequence ) @@ -464,7 +463,7 @@ async def _read( ) start = None for entry in page.entries: - record = await _decode_body(self._data_converter, entry.record) + record = await decode_body(self._data_converter, entry.record) for out in decoder.decode(_cursor(run_id, entry.offset), record): yield out offset = page.next_offset @@ -535,6 +534,23 @@ def producer( attempt, ) + 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: a workflow's stream ends with the workflow. + + Raises: + ValueError: Always; only a standalone stream is closed by hand. + """ + raise ValueError( + "a workflow's stream ends with its workflow; only a standalone stream " + "can be closed" + ) + def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: """The execution an offloaded body of this owner's stream is stored under.""" return StorageDriverWorkflowInfo( @@ -631,6 +647,26 @@ def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: return super()._store_target(run_id) return StorageDriverWorkflowInfo(namespace=self._client.namespace) + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this activity's streams, pinned as this handle is.""" + return StreamRef.for_activity( + self._activity_id, + workflow_id=self._workflow_id or None, + run_id=self._run_id, + topic=topic, + ) + + async def close(self) -> None: + """Refuse: an activity's streams end with the activity. + + Raises: + ValueError: Always; only a standalone stream is closed by hand. + """ + raise ValueError( + "an activity's streams end with the activity; only a standalone stream " + "can be closed" + ) + async def _current_run(self) -> str: if self._workflow_id: return await super()._current_run() @@ -657,6 +693,201 @@ async def _successor(self, run_id: str) -> str | None: return None +class _StandaloneTarget: + """A standalone stream as the producer addresses an owned one. + + The producer pins a run before its first write and names it in every + cursor it hands out. A standalone stream has no run to pin, so its id + stands in that place from the start and nothing is ever resolved. + """ + + def __init__(self, handle: ServiceStreamHandle, stream_id: str) -> None: + self._handle = handle + self.owner_run_id = stream_id + + def pin(self, run_id: str) -> None: + self.owner_run_id = run_id + + async def append( + self, *records: WireRecord, producer_id: str = "", sequence: int = 0 + ) -> Appended: + return await self._handle.append( + *records, producer_id=producer_id, sequence=sequence + ) + + +class NativeStandaloneStreamHandle: + """One standalone stream's topics from outside, over the stream service. + + A standalone stream has an id of its own and no owner, so there is no + chain to follow and no run to pin; a cursor names the stream id where an + owned stream's names a run, and a cursor from another stream is refused. + Its topics share one log, so a topic read is the server's filter over it. + + A read on an id nobody has created yet parks on the server until the + stream appears, so a reader can attach before the producer's first write. + ``latest`` and a producer's append on such an id raise + :class:`temporalio.streams.StreamNotFoundError`: they have nothing to wait + on. The server keeps a policy's ``retention`` for the stream's records + after it is closed rather than trimming an open stream by age, and it has + no byte bound on a standalone stream's lifecycle, so the provider refuses + ``max_bytes``. + """ + + def __init__( + self, client: Client, stream_id: str, *, opened: set[SharedKey] | None = None + ) -> None: + """Address the standalone stream ``stream_id``, which must exist or be created later.""" + self._client = client + self._stream_id = stream_id + self._opened = set() if opened is None else opened + self._data_converter = client.data_converter + self._converter = client.data_converter.payload_converter + self._streams: StreamClient | None = None + + @property + def stream_id(self) -> str: + """The id of the stream this handle is on.""" + return self._stream_id + + def _service(self) -> StreamClient: + if self._streams is None: + self._streams = shared_client(self._client) + self._opened.add(shared_key(self._client)) + return self._streams + + def _stream(self) -> ServiceStreamHandle: + return self._service().get(self._stream_id) + + def read( + self, + *, + 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`` from where the read starts until the stream is closed and drained. + + The server resolves ``BEGINNING``, ``END`` and ``last=`` on the first + poll. On an id that does not exist yet the poll parks until the stream + is created, so the read is open before the first write. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + named = None if after == END else _position(after) + if named is not None and named[0] != self._stream_id: + raise StreamCursorError( + f"cursor {after.token!r} names another stream than {self._stream_id!r}" + ) + start: StreamStartPosition | None = None + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif named is None: + start = StreamStartPosition(earliest=True) + return self._read(topic, named, start, after, result_type) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + start: StreamStartPosition | None, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + stream = self._stream() + decoder: RecordDecoder | None = None + offset = named[1] + 1 if named is not None else 0 + while True: + asked = time.monotonic() + try: + page = await stream.poll( + from_offset=offset, start=start, topics=[topic] + ) + except StreamNotFoundError: + # The server parks a poll on an id nobody has created for its + # wait budget and answers NOT_FOUND when that runs out. The + # stream may still be created, so the read keeps waiting; one + # that has delivered before is gone for good. A server that + # answers at once does not park, so the wait is paced here. + if decoder is not None: + raise + if time.monotonic() - asked < _CREATE_WAIT_PACE: + await asyncio.sleep(_CREATE_WAIT_PACE) + continue + if decoder is None: + decoder = RecordDecoder( + self._converter, + result_type, + after=NativeStreamHandle._previous( + self._stream_id, page, start, after + ), + warn=logger.warning, + ) + start = None + for entry in page.entries: + record = await decode_body(self._data_converter, entry.record) + for out in decoder.decode( + _cursor(self._stream_id, entry.offset), record + ): + yield out + offset = page.next_offset + if page.closed and offset >= page.head_offset: + return + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The cursor of the newest record on the stream, or ``BEGINNING`` when empty. + + Topics share the stream's offsets, so the newest record may be on + another topic; a read after this cursor still yields only what lands + on ``topic`` later. + + Raises: + StreamNotFoundError: The stream does not exist. + """ + resolve_topic(topic) + head = (await self._stream().describe()).head_offset + return _cursor(self._stream_id, head - 1) if head > 0 else BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> NativeProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return NativeProducer( + _StandaloneTarget(self._stream(), self._stream_id), + self._no_run, + self._data_converter, + self._store_target, + topic, + producer_id, + attempt, + ) + + async def _no_run(self) -> str: + return self._stream_id + + def _store_target(self, _run_id: str) -> StorageDriverWorkflowInfo: + # No workflow owns the stream, so an offloaded body has only the + # namespace to be stored under. + return StorageDriverWorkflowInfo(namespace=self._client.namespace) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this stream.""" + return StreamRef.for_standalone(self._stream_id, topic=topic) + + async def close(self) -> None: + """Seal the stream: appends are refused from now on and the tail stays readable. Idempotent.""" + await self._stream().close() + + class NativeStreams(ProviderPlugin): """The server-side provider. @@ -677,11 +908,91 @@ def workflow_provider(self) -> _NativeWorkflowProvider: return _NativeWorkflowProvider() def get_stream_handle( - self, client: Client, workflow_id: str, *, run_id: str | None = None - ) -> NativeStreamHandle: - """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain.""" + self, + client: Client, + workflow_id: str | StreamRef, + *, + run_id: str | None = None, + ) -> StreamHandle: + """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain. + + A :class:`temporalio.streams.StreamRef` in place of the id opens the + stream it names, whatever its owner kind, with the ref's topic as the + handle's default. + """ + if isinstance(workflow_id, StreamRef): + if run_id is not None: + raise ValueError("a StreamRef names the run itself; pass no run_id") + return open_ref(self, client, workflow_id) return NativeStreamHandle(client, workflow_id, run_id, opened=self._opened) + 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, + ) -> NativeStandaloneStreamHandle: + """Create the standalone stream ``stream_id`` on the server and return a handle on it. + + ``max_records`` keeps the newest records and drops the oldest as they + are appended. ``retention`` is how long the records stay readable + after the stream is closed; the server does not trim an open stream + by age. ``max_bytes`` is refused: the server's lifecycle has no byte + bound. A create of an id that exists with the same policy returns a + handle on it; with another policy it is a ``ValueError``. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + if retention is not None and retention <= timedelta(0): + raise ValueError("retention must be positive") + if max_records is not None and max_records <= 0: + raise ValueError("max_records must be positive") + if max_bytes is not None: + if max_bytes <= 0: + raise ValueError("max_bytes must be positive") + raise StreamUnsupportedError( + "the native provider bounds a standalone stream by records and, " + "after it closes, by age; max_bytes is not on the server's lifecycle" + ) + streams = shared_client(client) + self._opened.add(shared_key(client)) + try: + await streams.create(stream_id, retention=retention, max_items=max_records) + except RPCError as error: + # The server answers a create of an id that exists with a generic + # failure rather than ALREADY_EXISTS; describe tells the two apart. + try: + held = (await streams.get(stream_id).describe()).lifecycle + except StreamNotFoundError: + raise error from None + # A bound left to the server's default is not a disagreement with + # whatever default the server filled in. + if (max_records is not None and held.max_items != max_records) or ( + retention is not None and held.retention.ToTimedelta() != retention + ): + raise ValueError( + f"stream {stream_id!r} exists with another policy: it keeps " + f"{held.max_items or 'all'} records for " + f"{held.retention.ToTimedelta()} after it closes" + ) from None + return NativeStandaloneStreamHandle(client, stream_id, opened=self._opened) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> NativeStandaloneStreamHandle: + """A handle on the standalone stream ``stream_id``. + + Nothing here creates the stream. A read on an id that does not exist + yet parks on the server until it is created; ``latest`` and a + producer's append raise :class:`temporalio.streams.StreamNotFoundError`. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + return NativeStandaloneStreamHandle(client, stream_id, opened=self._opened) + def get_activity_stream_handle( self, client: Client, diff --git a/temporalio/worker/_command_aware_visitor.py b/temporalio/worker/_command_aware_visitor.py index bc3e3f35a..fa6d5ba18 100644 --- a/temporalio/worker/_command_aware_visitor.py +++ b/temporalio/worker/_command_aware_visitor.py @@ -27,7 +27,7 @@ StartChildWorkflowExecution, WorkflowCommand, ) -from temporalio.streams._wire import CONTENT_HASH_KEY +from temporalio.streams._body import CONTENT_HASH_KEY @dataclass(frozen=True) diff --git a/tests/streams/test_native_fingerprint.py b/tests/streams/test_native_fingerprint.py index 9cac71755..fccc54aaa 100644 --- a/tests/streams/test_native_fingerprint.py +++ b/tests/streams/test_native_fingerprint.py @@ -27,7 +27,7 @@ PayloadCodec, StorageDriverWorkflowInfo, ) -from temporalio.streams._wire import CONTENT_HASH_KEY +from temporalio.streams import CONTENT_HASH_KEY from temporalio.streams.providers import native from temporalio.streams.providers.native import NativeProducer diff --git a/tests/streams/test_native_handles.py b/tests/streams/test_native_handles.py new file mode 100644 index 000000000..50b68c4db --- /dev/null +++ b/tests/streams/test_native_handles.py @@ -0,0 +1,110 @@ +"""What the native handles answer without a server. + +A ref names the owner as the handle addresses it, a close on an owned stream +is refused, a standalone handle refuses a cursor from another stream, and a +create refuses a policy the server cannot hold before any call is made. +""" + +from __future__ import annotations + +from datetime import timedelta +from types import SimpleNamespace +from typing import Any + +import pytest + +import temporalio.converter +from temporalio.streams import ( + DEFAULT_TOPIC, + StreamCursorError, + StreamRef, + StreamUnsupportedError, + topic, +) +from temporalio.streams.providers.native import ( + NativeActivityStreamHandle, + NativeStandaloneStreamHandle, + NativeStreamHandle, + NativeStreams, + _cursor, +) + +OUT = topic("out", dict) +_CLIENT: Any = SimpleNamespace( + data_converter=temporalio.converter.default(), namespace="ns" +) + + +def test_a_workflow_handle_refs_its_owner_as_it_is_pinned() -> None: + following = NativeStreamHandle(_CLIENT, "wf", None) + assert following.ref() == StreamRef.for_workflow("wf") + assert following.ref(topic=OUT) == StreamRef.for_workflow("wf", topic="out") + pinned = NativeStreamHandle(_CLIENT, "wf", "run-1") + assert pinned.ref(topic="a") == StreamRef.for_workflow( + "wf", run_id="run-1", topic="a" + ) + + +def test_an_activity_handle_refs_the_activity_and_its_workflow() -> None: + scheduled = NativeActivityStreamHandle(_CLIENT, "act", "wf", "run-1") + assert scheduled.ref() == StreamRef.for_activity( + "act", workflow_id="wf", run_id="run-1" + ) + standalone = NativeActivityStreamHandle(_CLIENT, "act", None, None) + assert standalone.ref(topic=OUT) == StreamRef.for_activity("act", topic="out") + assert standalone.ref().workflow_id is None + + +def test_a_standalone_handle_refs_its_id() -> None: + handle = NativeStandaloneStreamHandle(_CLIENT, "s1") + assert handle.stream_id == "s1" + assert handle.ref() == StreamRef.for_standalone("s1") + assert handle.ref().topic == DEFAULT_TOPIC + assert handle.ref(topic=OUT) == StreamRef.for_standalone("s1", topic="out") + + +async def test_only_a_standalone_stream_can_be_closed() -> None: + with pytest.raises(ValueError, match="standalone"): + await NativeStreamHandle(_CLIENT, "wf", None).close() + with pytest.raises(ValueError, match="standalone"): + await NativeActivityStreamHandle(_CLIENT, "act", "wf", None).close() + + +def test_a_standalone_handle_refuses_another_streams_cursor() -> None: + handle = NativeStandaloneStreamHandle(_CLIENT, "s1") + with pytest.raises(StreamCursorError, match="another stream"): + handle.read(topic=OUT, after=_cursor("s2", 3)) + with pytest.raises(StreamCursorError): + handle.read(topic=OUT, after=_cursor("", 3)) + + +def test_a_ref_opens_the_handle_it_names_on_the_provider() -> None: + provider = NativeStreams() + opened = provider.get_stream_handle( + _CLIENT, StreamRef.for_standalone("s1", topic="out") + ) + assert opened.ref() == StreamRef.for_standalone("s1", topic="out") + assert opened.ref(topic="b") == StreamRef.for_standalone("s1", topic="b") + workflow = provider.get_stream_handle( + _CLIENT, StreamRef.for_workflow("wf", run_id="r") + ) + assert workflow.ref() == StreamRef.for_workflow("wf", run_id="r") + with pytest.raises(ValueError, match="run_id"): + provider.get_stream_handle(_CLIENT, StreamRef.for_workflow("wf"), run_id="r") + + +async def test_a_create_refuses_a_policy_the_server_cannot_hold() -> None: + provider = NativeStreams() + for bad in ( + dict(max_records=0), + dict(max_bytes=-1), + dict(retention=timedelta(0)), + ): + with pytest.raises(ValueError): + await provider.create_standalone_stream(_CLIENT, "s1", **bad) # type: ignore[arg-type] + with pytest.raises(ValueError): + await provider.create_standalone_stream(_CLIENT, "") + with pytest.raises(StreamUnsupportedError, match="max_bytes"): + await provider.create_standalone_stream(_CLIENT, "s1", max_bytes=700) + with pytest.raises(ValueError): + provider.get_standalone_stream_handle(_CLIENT, "") diff --git a/tests/streams/test_native_standalone_e2e.py b/tests/streams/test_native_standalone_e2e.py new file mode 100644 index 000000000..2b9a43773 --- /dev/null +++ b/tests/streams/test_native_standalone_e2e.py @@ -0,0 +1,95 @@ +"""Standalone streams on the native provider against a live server. + +The conformance suite covers what every provider owes a standalone stream. +These pin what native adds on top: a read on an id nobody has created yet +parks on the server and delivers the first record once the stream exists, a +ref carries the stream to another client, and a create of an id that exists +is answered from the stream's own policy. Needs a server built from the +AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7433 uv run pytest tests/streams/test_native_standalone_e2e.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid + +import pytest + +from temporalio.client import Client +from temporalio.streams import StreamClosedError, StreamNotFoundError, StreamRef, topic +from temporalio.streams.providers.native import NativeStreams +from tests.streams.test_streams_conformance import take + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream service", +) + +OUT = topic("out", dict) + + +async def _connect(provider: NativeStreams) -> Client: + return await Client.connect(TARGET or "", plugins=[provider]) + + +async def test_a_read_on_a_stream_not_yet_created_waits_for_it() -> None: + provider = NativeStreams() + client = await _connect(provider) + stream_id = "late-" + uuid.uuid4().hex[:8] + try: + reader = client.get_stream_handle(stream_id=stream_id) + # Nothing to wait on for these: the stream does not exist. + with pytest.raises(StreamNotFoundError): + await reader.latest(topic=OUT) + with pytest.raises(StreamNotFoundError): + await reader.producer(topic=OUT, producer_id="early", attempt=1).append( + {"n": 0} + ) + + parked = asyncio.ensure_future(take(reader.read(topic=OUT), 1, timeout=30)) + await asyncio.sleep(0.5) + assert not parked.done(), "the read is parked, not failed" + + created = await client.create_stream(stream_id) + await created.producer(topic=OUT, producer_id="writer", attempt=1).append( + {"n": 1} + ) + assert [r.value for r in await parked] == [{"n": 1}] + finally: + await provider.close() + + +async def test_a_ref_carries_a_standalone_stream_to_another_client() -> None: + provider = NativeStreams() + client = await _connect(provider) + other_provider = NativeStreams() + other = await _connect(other_provider) + stream_id = "ref-" + uuid.uuid4().hex[:8] + try: + created = await client.create_stream(stream_id, max_records=10) + await created.producer(topic=OUT, producer_id="writer", attempt=1).append( + {"n": 1} + ) + ref = created.ref(topic=OUT) + assert ref == StreamRef.for_standalone(stream_id, topic="out") + + opened = other.get_stream_handle(ref) + records = await take(opened.read(), 1) + assert [r.value for r in records] == [{"n": 1}] + assert await opened.latest() == records[0].cursor + + await opened.close() + with pytest.raises(StreamClosedError): + await created.producer(topic=OUT, producer_id="writer", attempt=2).append( + {"n": 2} + ) + # The tail stays readable and the read ends on its own. + assert [r.value async for r in created.read(topic=OUT)] == [{"n": 1}] + finally: + await provider.close() + await other_provider.close() diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index aae7148ef..9cd49daf8 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -62,6 +62,7 @@ StreamProducerError, StreamProvider, StreamRef, + StreamUnsupportedError, Supersession, topic, ) @@ -100,6 +101,11 @@ class ProviderCase: waits_for_standalone_creation: bool = False """A read on a standalone stream id that does not exist yet parks until the first write instead of raising ``StreamNotFoundError``.""" + bounds_standalone_bytes: bool = True + """A standalone stream's policy can bound the bytes it keeps.""" + trims_open_stream_by_age: bool = True + """A standalone stream drops records older than ``retention`` while it is + open, rather than keeping them that long after it closes.""" async def open( self, @@ -273,7 +279,15 @@ async def host(workflow_id: str) -> None: # No truncate hook: a stream a workflow owns has no truncation call. # BEGINNING on a truncated stream is covered on the stream client, # whose standalone streams can be truncated. - yield ProviderCase("native", provider, client, host=host) + yield ProviderCase( + "native", + provider, + client, + host=host, + waits_for_standalone_creation=True, + bounds_standalone_bytes=False, + trims_open_stream_by_age=False, + ) for handle in hosts.values(): await handle.terminate() await provider.close() @@ -738,7 +752,9 @@ async def test_a_body_above_the_threshold_is_offloaded_and_read_back( external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=256), ) workflow_id = new_workflow_id() - stream = await case.open(workflow_id, client=_client_with(client, converter)) + stream = await case.open( + workflow_id, client=_client_with(case.client or client, converter) + ) producer = stream.producer(topic=OUT, producer_id="model", attempt=1) small = {"n": 1} large = {"blob": "x" * 1024} @@ -760,7 +776,9 @@ async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( codec = NonceCodec() converter = dataclasses.replace(DataConverter.default, payload_codec=codec) workflow_id = new_workflow_id() - stream = await case.open(workflow_id, client=_client_with(client, converter)) + stream = await case.open( + workflow_id, client=_client_with(case.client or client, converter) + ) first = stream.producer(topic=OUT, producer_id="model", attempt=1) landed = await first.append({"id": "r1"}) assert codec.encoded == 1 @@ -899,13 +917,19 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas kept = await take(by_count.read(topic=OUT), 2) assert [r.value for r in kept] == [{"n": 3}, {"n": 4}] - by_bytes = await case.create_stream(new_stream_id(), max_bytes=700) - producer = by_bytes.producer(topic=OUT, producer_id="writer", attempt=1) - for n in range(3): - await producer.append({"n": n, "blob": "x" * 500}) - kept = await take(by_bytes.read(topic=OUT), 1) - assert kept[0].value is not None and kept[0].value["n"] == 2 - + if case.bounds_standalone_bytes: + by_bytes = await case.create_stream(new_stream_id(), max_bytes=700) + producer = by_bytes.producer(topic=OUT, producer_id="writer", attempt=1) + for n in range(3): + await producer.append({"n": n, "blob": "x" * 500}) + kept = await take(by_bytes.read(topic=OUT), 1) + assert kept[0].value is not None and kept[0].value["n"] == 2 + else: + with pytest.raises(StreamUnsupportedError): + await case.create_stream(new_stream_id(), max_bytes=700) + + if not case.trims_open_stream_by_age: + return by_age = await case.create_stream( new_stream_id(), retention=timedelta(milliseconds=200) ) diff --git a/tests/worker/test_workflow_stream_payloads_e2e.py b/tests/worker/test_workflow_stream_payloads_e2e.py index 54ffe74ea..aa739fb1f 100644 --- a/tests/worker/test_workflow_stream_payloads_e2e.py +++ b/tests/worker/test_workflow_stream_payloads_e2e.py @@ -29,8 +29,7 @@ JSONProtoPayloadConverter, PayloadCodec, ) -from temporalio.streams import RecordKind -from temporalio.streams._wire import CONTENT_HASH_KEY +from temporalio.streams import CONTENT_HASH_KEY, RecordKind from temporalio.streams.providers.native import NativeStreams from temporalio.worker import Worker from tests.streams.test_streams_conformance import take From 3eb66641bf57efcbf42fe7e1042b682312cab50e Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:04:48 -0700 Subject: [PATCH 20/27] Mapped the STREAM_CLOSED reason token to StreamClosedError. The server is giving an append on a sealed stream a reason token of its own, so the mapping takes the token beside the three it already knew and keeps the whole-message match for a server built before it. --- temporalio/client_stream.py | 5 +++-- tests/test_client_stream_errors.py | 15 +++++++++++++-- 2 files changed, 16 insertions(+), 4 deletions(-) diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index 8ae0f114c..6b612dc3b 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -113,10 +113,11 @@ "STREAM_PRODUCER_CONFLICT": StreamProducerError, "STREAM_PRODUCER_STALE_SEQUENCE": StreamProducerError, "STREAM_CURSOR_BELOW_FLOOR": StreamCursorError, + "STREAM_CLOSED": StreamClosedError, } # The phrases a server built before the tokens existed sends for the same -# refusals, so a reader of either server gets the typed error. An append on a -# sealed stream has no token yet and is matched on its whole message. +# refusals, so a reader of either server gets the typed error. The sealed +# stream's message is matched whole. _PRODUCER_PHRASE = "producer sequence" _CURSOR_PHRASE = "below the stream's floor" _CLOSED_PHRASE = "stream is closed" diff --git a/tests/test_client_stream_errors.py b/tests/test_client_stream_errors.py index ae0b452a6..d317ced1f 100644 --- a/tests/test_client_stream_errors.py +++ b/tests/test_client_stream_errors.py @@ -15,6 +15,7 @@ from temporalio.client_stream import translate_error from temporalio.service import RPCError, RPCStatusCode from temporalio.streams import ( + StreamClosedError, StreamCursorError, StreamNotFoundError, StreamProducerError, @@ -55,6 +56,16 @@ def test_a_read_below_the_floor_is_a_cursor_error(details: str) -> None: assert str(error) == details +@pytest.mark.parametrize( + "details", + ["STREAM_CLOSED: stream is closed", "stream is closed"], +) +def test_an_append_on_a_sealed_stream_is_a_closed_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamClosedError) + assert str(error) == details + + def test_not_found_is_a_not_found_error() -> None: error = translate_error(grpc.StatusCode.NOT_FOUND, "no stream with id 's'") assert isinstance(error, StreamNotFoundError) @@ -67,11 +78,11 @@ def test_an_unrelated_message_keeps_its_code() -> None: (grpc.StatusCode.FAILED_PRECONDITION, RPCStatusCode.FAILED_PRECONDITION), (grpc.StatusCode.UNAVAILABLE, RPCStatusCode.UNAVAILABLE), ]: - error = translate_error(code, "stream is closed", b"raw") + error = translate_error(code, "no records to append", b"raw") assert type(error) is RPCError assert error.status == expected assert error.raw_grpc_status == b"raw" - assert str(error) == "stream is closed" + assert str(error) == "no records to append" def test_a_token_needs_its_code_and_its_separator() -> None: From 2b19d6eef5d2605f3b0b18ac23cf02857de355f3 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:11:54 -0700 Subject: [PATCH 21/27] Took the server's word on a create of a standalone stream that exists. The server is answering ALREADY_EXISTS for an id that exists with the same policy and a STREAM_POLICY_MISMATCH refusal for one that differs, so the provider returns the handle on the first and lets the second arrive as the ValueError the contract names. The describe comparison stays for a server that answers both with a generic failure. --- temporalio/client_stream.py | 3 +++ temporalio/streams/providers/native.py | 10 ++++++++-- tests/test_client_stream_errors.py | 8 ++++++++ 3 files changed, 19 insertions(+), 2 deletions(-) diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index 6b612dc3b..de0ed5c94 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -114,6 +114,9 @@ "STREAM_PRODUCER_STALE_SEQUENCE": StreamProducerError, "STREAM_CURSOR_BELOW_FLOOR": StreamCursorError, "STREAM_CLOSED": StreamClosedError, + # A create of an id that exists with another policy is the caller's + # mistake, which the interface contract spells as ValueError. + "STREAM_POLICY_MISMATCH": ValueError, } # The phrases a server built before the tokens existed sends for the same # refusals, so a reader of either server gets the typed error. The sealed diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 6904c7cae..5e6bd7f66 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -962,8 +962,14 @@ async def create_standalone_stream( try: await streams.create(stream_id, retention=retention, max_items=max_records) except RPCError as error: - # The server answers a create of an id that exists with a generic - # failure rather than ALREADY_EXISTS; describe tells the two apart. + # The server says ALREADY_EXISTS for an id that exists with this + # policy, and a policy that differs arrives typed as ValueError. An + # older server answers both with a generic failure, and describe + # tells the two apart. + if error.status == RPCStatusCode.ALREADY_EXISTS: + return NativeStandaloneStreamHandle( + client, stream_id, opened=self._opened + ) try: held = (await streams.get(stream_id).describe()).lifecycle except StreamNotFoundError: diff --git a/tests/test_client_stream_errors.py b/tests/test_client_stream_errors.py index d317ced1f..7b379f4c7 100644 --- a/tests/test_client_stream_errors.py +++ b/tests/test_client_stream_errors.py @@ -66,6 +66,14 @@ def test_an_append_on_a_sealed_stream_is_a_closed_error(details: str) -> None: assert str(error) == details +def test_a_create_with_another_policy_is_a_value_error() -> None: + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, + "STREAM_POLICY_MISMATCH: stream exists keeping 10 records, asked for 5", + ) + assert type(error) is ValueError + + def test_not_found_is_a_not_found_error() -> None: error = translate_error(grpc.StatusCode.NOT_FOUND, "no stream with id 's'") assert isinstance(error, StreamNotFoundError) From 20ae07a98591d7fc269f9165d9ef559755b1c8d0 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 18:30:26 -0700 Subject: [PATCH 22/27] Passed the byte cap through and trusted the server on a repeated create. The server's lifecycle now carries max_bytes, refuses an append that would cross it, reclaims an open standalone stream's records by age, and answers a create of an existing id with ALREADY_EXISTS or a STREAM_POLICY_MISMATCH refusal. The provider passes the cap through, drops the describe comparison, and the native conformance setup declares the byte cap as a refusal rather than a trim and waits for the store's age reclaim. --- temporalio/client_stream.py | 11 +++- temporalio/streams/providers/native.py | 69 ++++++++--------------- tests/streams/test_native_handles.py | 10 +--- tests/streams/test_streams_conformance.py | 38 +++++++++++-- 4 files changed, 66 insertions(+), 62 deletions(-) diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index de0ed5c94..a7b5f8d32 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -628,13 +628,16 @@ async def create( *, retention: float | timedelta | None = None, max_items: int | None = None, + max_bytes: int | None = None, ) -> StreamHandle: """Create a stream and return a handle to it. - ``retention`` is how long a closed stream stays readable, in seconds - or as a ``timedelta``. ``max_items`` caps how many records remain + ``retention`` is the age past which an open stream's records are + reclaimed and how long a closed stream stays readable, in seconds or + as a ``timedelta``. ``max_items`` caps how many records remain readable, dropping the oldest, which bounds storage for a stream - nobody truncates. + nobody truncates. ``max_bytes`` caps the bytes held; an append that + would cross it is refused rather than reclaiming anything. """ lifecycle = stream.StreamLifecycle() if retention is not None: @@ -643,6 +646,8 @@ async def create( lifecycle.retention.FromTimedelta(retention) if max_items is not None: lifecycle.max_items = max_items + if max_bytes is not None: + lifecycle.max_bytes = max_bytes # A create that landed but was not answered would be refused as a # repeat, so it goes again only on a refusal. diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 5e6bd7f66..95bb866a3 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -65,11 +65,7 @@ decode_body, encode_body, ) -from temporalio.streams._errors import ( - StreamCursorError, - StreamNotFoundError, - StreamUnsupportedError, -) +from temporalio.streams._errors import StreamCursorError, StreamNotFoundError from temporalio.streams._provider import ReadSource, StreamHandle, WriteSink from temporalio.streams._record import ( BEGINNING, @@ -728,10 +724,11 @@ class NativeStandaloneStreamHandle: stream appears, so a reader can attach before the producer's first write. ``latest`` and a producer's append on such an id raise :class:`temporalio.streams.StreamNotFoundError`: they have nothing to wait - on. The server keeps a policy's ``retention`` for the stream's records - after it is closed rather than trimming an open stream by age, and it has - no byte bound on a standalone stream's lifecycle, so the provider refuses - ``max_bytes``. + on. The policy is the server's lifecycle: ``retention`` is the age past + which an open stream's records are reclaimed and how long a closed one + stays readable, ``max_records`` drops the oldest records as new ones land, + and ``max_bytes`` refuses an append that would take the held bytes past it + rather than reclaiming anything. """ def __init__( @@ -937,12 +934,12 @@ async def create_standalone_stream( ) -> NativeStandaloneStreamHandle: """Create the standalone stream ``stream_id`` on the server and return a handle on it. - ``max_records`` keeps the newest records and drops the oldest as they - are appended. ``retention`` is how long the records stay readable - after the stream is closed; the server does not trim an open stream - by age. ``max_bytes`` is refused: the server's lifecycle has no byte - bound. A create of an id that exists with the same policy returns a - handle on it; with another policy it is a ``ValueError``. + The three bounds are the server's lifecycle: ``retention`` reclaims + records older than it on an open stream and times a closed one's + deletion, ``max_records`` drops the oldest records as new ones land, + and ``max_bytes`` refuses an append that would take the held bytes past + it. A create of an id that exists with the same policy returns a handle + on it; with another policy the server refuses it as a ``ValueError``. """ if not stream_id: raise ValueError("stream_id must not be empty") @@ -950,40 +947,22 @@ async def create_standalone_stream( raise ValueError("retention must be positive") if max_records is not None and max_records <= 0: raise ValueError("max_records must be positive") - if max_bytes is not None: - if max_bytes <= 0: - raise ValueError("max_bytes must be positive") - raise StreamUnsupportedError( - "the native provider bounds a standalone stream by records and, " - "after it closes, by age; max_bytes is not on the server's lifecycle" - ) + if max_bytes is not None and max_bytes <= 0: + raise ValueError("max_bytes must be positive") streams = shared_client(client) self._opened.add(shared_key(client)) try: - await streams.create(stream_id, retention=retention, max_items=max_records) + await streams.create( + stream_id, + retention=retention, + max_items=max_records, + max_bytes=max_bytes, + ) except RPCError as error: - # The server says ALREADY_EXISTS for an id that exists with this - # policy, and a policy that differs arrives typed as ValueError. An - # older server answers both with a generic failure, and describe - # tells the two apart. - if error.status == RPCStatusCode.ALREADY_EXISTS: - return NativeStandaloneStreamHandle( - client, stream_id, opened=self._opened - ) - try: - held = (await streams.get(stream_id).describe()).lifecycle - except StreamNotFoundError: - raise error from None - # A bound left to the server's default is not a disagreement with - # whatever default the server filled in. - if (max_records is not None and held.max_items != max_records) or ( - retention is not None and held.retention.ToTimedelta() != retention - ): - raise ValueError( - f"stream {stream_id!r} exists with another policy: it keeps " - f"{held.max_items or 'all'} records for " - f"{held.retention.ToTimedelta()} after it closes" - ) from None + # The id exists with this policy; a policy that differs arrives + # typed, as the ValueError the contract names. + if error.status != RPCStatusCode.ALREADY_EXISTS: + raise return NativeStandaloneStreamHandle(client, stream_id, opened=self._opened) def get_standalone_stream_handle( diff --git a/tests/streams/test_native_handles.py b/tests/streams/test_native_handles.py index 50b68c4db..318c30494 100644 --- a/tests/streams/test_native_handles.py +++ b/tests/streams/test_native_handles.py @@ -14,13 +14,7 @@ import pytest import temporalio.converter -from temporalio.streams import ( - DEFAULT_TOPIC, - StreamCursorError, - StreamRef, - StreamUnsupportedError, - topic, -) +from temporalio.streams import DEFAULT_TOPIC, StreamCursorError, StreamRef, topic from temporalio.streams.providers.native import ( NativeActivityStreamHandle, NativeStandaloneStreamHandle, @@ -104,7 +98,5 @@ async def test_a_create_refuses_a_policy_the_server_cannot_hold() -> None: await provider.create_standalone_stream(_CLIENT, "s1", **bad) # type: ignore[arg-type] with pytest.raises(ValueError): await provider.create_standalone_stream(_CLIENT, "") - with pytest.raises(StreamUnsupportedError, match="max_bytes"): - await provider.create_standalone_stream(_CLIENT, "s1", max_bytes=700) with pytest.raises(ValueError): provider.get_standalone_stream_handle(_CLIENT, "") diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index d6613a413..3c0d0c773 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -49,6 +49,7 @@ StorageDriverRetrieveContext, StorageDriverStoreContext, ) +from temporalio.service import RPCError from temporalio.streams import ( BEGINNING, DEFAULT_TOPIC, @@ -57,6 +58,7 @@ RecordKind, StreamClosedError, StreamCursorError, + StreamError, StreamHandle, StreamNotFoundError, StreamProducerError, @@ -103,6 +105,9 @@ class ProviderCase: the first write instead of raising ``StreamNotFoundError``.""" bounds_standalone_bytes: bool = True """A standalone stream's policy can bound the bytes it keeps.""" + refuses_appends_past_byte_cap: bool = False + """The byte bound refuses an append that would cross it, instead of + dropping the oldest records to make room.""" trims_open_stream_by_age: bool = True """A standalone stream drops records older than ``retention`` while it is open, rather than keeping them that long after it closes.""" @@ -291,8 +296,7 @@ async def host(workflow_id: str) -> None: client, host=host, waits_for_standalone_creation=True, - bounds_standalone_bytes=False, - trims_open_stream_by_age=False, + refuses_appends_past_byte_cap=True, ) for handle in hosts.values(): await handle.terminate() @@ -923,7 +927,19 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas kept = await take(by_count.read(topic=OUT), 2) assert [r.value for r in kept] == [{"n": 3}, {"n": 4}] - if case.bounds_standalone_bytes: + if case.bounds_standalone_bytes and case.refuses_appends_past_byte_cap: + # The cap counts the stored record, body and metadata included, so it + # holds one of these and not two. + by_bytes = await case.create_stream(new_stream_id(), max_bytes=1000) + producer = by_bytes.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": 0, "blob": "x" * 500}) + # The second record would take the stream past its cap, so the store + # refuses it and keeps what it holds. + with pytest.raises((RPCError, StreamError)): + await producer.append({"n": 1, "blob": "x" * 500}) + kept = await take(by_bytes.read(topic=OUT), 1) + assert kept[0].value is not None and kept[0].value["n"] == 0 + elif case.bounds_standalone_bytes: by_bytes = await case.create_stream(new_stream_id(), max_bytes=700) producer = by_bytes.producer(topic=OUT, producer_id="writer", attempt=1) for n in range(3): @@ -943,5 +959,17 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas await producer.append({"n": "old"}) await asyncio.sleep(0.3) await producer.append({"n": "new"}) - kept = await take(by_age.read(topic=OUT), 1) - assert [r.value for r in kept] == [{"n": "new"}] + # A store reclaims by age on its own schedule, and a coarse one may have + # aged the newer record out as well by the time it is looked at. What the + # policy decides is that the older record is no longer where BEGINNING + # starts. + kept: list = [] + for _ in range(50): + try: + kept = await take(by_age.read(topic=OUT), 1, timeout=1) + except asyncio.TimeoutError: + kept = [] + if not kept or kept[0].value == {"n": "new"}: + break + await asyncio.sleep(0.2) + assert [r.value for r in kept] in ([{"n": "new"}], []) From e80180e50aba6c83b906f06bf8cf41fdcf73ec74 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 18:32:36 -0700 Subject: [PATCH 23/27] Reused the retention case's result variable instead of redeclaring it. --- tests/streams/test_streams_conformance.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 3c0d0c773..7a2f6f435 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -963,7 +963,7 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas # aged the newer record out as well by the time it is looked at. What the # policy decides is that the older record is no longer where BEGINNING # starts. - kept: list = [] + kept = [] for _ in range(50): try: kept = await take(by_age.read(topic=OUT), 1, timeout=1) From c5b259c70cc4da93d7e84b331dc5ee57d2feafb1 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 19:29:33 -0700 Subject: [PATCH 24/27] Expected the typed cursor error from a read below the floor. This layer maps the STREAM_CURSOR_BELOW_FLOOR token to StreamCursorError, which is not an RPCError, so the live case written on the wire layer no longer matched the exception it got. --- tests/test_client_stream.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/test_client_stream.py b/tests/test_client_stream.py index 89abb0559..779d1b188 100644 --- a/tests/test_client_stream.py +++ b/tests/test_client_stream.py @@ -26,7 +26,11 @@ ) from temporalio.client_stream import StreamClient, StreamHandle from temporalio.service import RPCError -from temporalio.streams import StreamNotFoundError, StreamProducerError +from temporalio.streams import ( + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") @@ -230,7 +234,7 @@ async def test_earliest_reads_from_the_floor_of_a_truncated_stream( # Offset zero is what a reader with no position used to send, and a # truncated stream no longer holds it. - with pytest.raises(RPCError, match="truncated"): + with pytest.raises(StreamCursorError, match="truncated"): await stream.read(from_offset=0) entries, next_offset = await stream.read(start=StreamStartPosition(earliest=True)) assert data(entries) == [b"c", b"d"] From f235b47d4c9dd93d59313f8701200bb39ff56620 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:35:31 -0700 Subject: [PATCH 25/27] Followed a native stream through its channel on a live server. The two cases poll the channel stream_channel derives, one standalone and one owned by a workflow, and register a callback on the derived name. They need a server on which streams drive channels and skip otherwise. --- tests/streams/test_stream_channel_e2e.py | 182 +++++++++++++++++++++++ 1 file changed, 182 insertions(+) create mode 100644 tests/streams/test_stream_channel_e2e.py diff --git a/tests/streams/test_stream_channel_e2e.py b/tests/streams/test_stream_channel_e2e.py new file mode 100644 index 000000000..e2013c49f --- /dev/null +++ b/tests/streams/test_stream_channel_e2e.py @@ -0,0 +1,182 @@ +"""A native stream notifies the channel named by the stream, on a live server. + +Every append and the close of a native stream notify the channel +:func:`temporalio.client.stream_channel` derives from the stream's ref, so a +client follows a stream the way it follows an external one: by polling the +channel or registering a callback on it. The workflow's own consumption of a +stream is untouched by this and is covered elsewhere. Needs a server on which +streams drive channels, named with ``-E host:port``; skipped otherwise. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import uuid +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.client import ( + Callback, + ChannelAddress, + ChannelKind, + Client, + stream_channel, +) +from temporalio.service import RPCError +from temporalio.streams import StreamRef, topic +from temporalio.streams.providers.native import NativeStreams +from tests.helpers import assert_eventually, new_worker + +OUT = topic("out", dict) + +pytestmark = pytest.mark.needs_stream_channel_server + + +async def _native(client: Client, provider: NativeStreams) -> Client: + """A client on the same server as ``client`` with the native provider on it.""" + return await Client.connect( + client.service_client.config.target_host, + namespace=client.namespace, + plugins=[provider], + ) + + +def _closed(client: Client, notification: Any) -> bool: + if "closed" not in notification.metadata: + return False + return client.data_converter.payload_converter.from_payload( + notification.metadata["closed"], bool + ) + + +async def test_a_standalone_stream_notifies_the_channel_named_by_its_id( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + stream_id = f"chan-{uuid.uuid4().hex[:8]}" + try: + created = await native.create_stream(stream_id) + address = stream_channel(created.ref(topic=OUT)) + # A standalone stream's topics share one stream, so the topic is not in + # the name and the channel is an independent one. + assert address == ChannelAddress(f"stream/{stream_id}", None) + producer = created.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": 1}) + await producer.append({"n": 2}) + + async def two_changes() -> list[Any]: + # A standalone stream reaches its channel through a task of its + # own, so the notifications may trail the appends. + polled = await native.poll_channel(address.channel, wait=False) + assert [n.counter for n in polled] == [1, 2] + return polled + + polled = await assert_eventually(two_changes) + for notification in polled: + assert notification.channel == address.channel + assert notification.linked_to is None + # The position is the head after the change, as a native cursor + # names one: the stream id in place of a run, then the offset. + assert notification.position.decode().startswith(f"{stream_id}:") + assert not _closed(native, notification) + await created.close() + [third] = await native.poll_channel( + address.channel, after_counter=2, wait=timedelta(seconds=10) + ) + assert third.counter == 3 + assert _closed(native, third) + description = await native.describe_channel(address.channel) + assert description.kind == ChannelKind.INDEPENDENT + assert description.latest is not None and description.latest.counter == 3 + assert description.retained_count == 3 + # A callback registers on the derived name like on any channel. + callback = Callback(url="http://localhost:1/never-called", headers={}) + listener_id = await native.register_channel_listener(address.channel, callback) + description = await native.describe_channel(address.channel) + assert [listener.callback for listener in description.listeners] == [callback] + await native.unregister_channel_listener(address.channel, listener_id) + finally: + await provider.close() + + +@workflow.defn +class WriteOnNudge: + """Writes one record at the start and one more per nudge, until ``rounds``.""" + + def __init__(self) -> None: + self._nudges = 0 + self._written = 0 + + @workflow.run + async def run(self, rounds: int) -> int: + writer = workflow.stream_writer(OUT) + writer.publish({"n": 0}) + while self._written < rounds: + await workflow.wait_condition(lambda: self._nudges > self._written) + self._written += 1 + writer.publish({"n": self._written}) + return self._written + + @workflow.signal + def nudge(self) -> None: + self._nudges += 1 + + +async def test_a_workflow_stream_notifies_the_channel_linked_to_its_owner( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + worker = new_worker(native, WriteOnNudge) + running = asyncio.create_task(worker.run()) + # More rounds than nudges: the linked ring dies with the run, so the run + # stays open until the test has read it and is ended by hand. + handle = await native.start_workflow( + WriteOnNudge.run, 10, id=f"wf-{uuid.uuid4()}", task_queue=worker.task_queue + ) + try: + address = stream_channel(StreamRef.for_workflow(handle.id, topic=OUT)) + assert address == ChannelAddress("stream/out", handle.id) + assert address == stream_channel( + native.get_stream_handle(handle.id).ref(topic=OUT) + ) + + async def changes(count: int) -> list[Any]: + polled = await native.poll_channel( + address.channel, workflow_id=address.workflow_id, wait=False + ) + assert [n.counter for n in polled] == list(range(1, count + 1)) + return polled + + # The first task's publish is one append, so one notification. + await assert_eventually(lambda: changes(1), timeout=timedelta(seconds=30)) + await handle.signal(WriteOnNudge.nudge) + await assert_eventually(lambda: changes(2), timeout=timedelta(seconds=30)) + await handle.signal(WriteOnNudge.nudge) + polled = await assert_eventually( + lambda: changes(3), timeout=timedelta(seconds=30) + ) + run_id = handle.first_execution_run_id + for notification in polled: + assert notification.linked_to is not None + assert notification.linked_to.workflow_id == handle.id + assert notification.position.decode().startswith(f"{run_id}:") + description = await native.describe_channel( + address.channel, workflow_id=address.workflow_id + ) + assert description.kind == ChannelKind.LINKED + assert description.latest is not None and description.latest.counter == 3 + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + try: + await asyncio.wait_for(worker.shutdown(), 15) + except asyncio.TimeoutError: + running.cancel() + await asyncio.gather(running, return_exceptions=True) + await provider.close() From 1f7cee0cab82e611db6f2d8f1f6fb0fed4c910a7 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:43:18 -0700 Subject: [PATCH 26/27] Waited for each standalone append's notification before the next. A standalone stream hands its channel the latest change through one outstanding task, so a burst of appends arrives as its newest change. The case now observes every change rather than a fold. --- tests/streams/test_stream_channel_e2e.py | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/tests/streams/test_stream_channel_e2e.py b/tests/streams/test_stream_channel_e2e.py index e2013c49f..e698db6a1 100644 --- a/tests/streams/test_stream_channel_e2e.py +++ b/tests/streams/test_stream_channel_e2e.py @@ -66,17 +66,20 @@ async def test_a_standalone_stream_notifies_the_channel_named_by_its_id( # the name and the channel is an independent one. assert address == ChannelAddress(f"stream/{stream_id}", None) producer = created.producer(topic=OUT, producer_id="writer", attempt=1) - await producer.append({"n": 1}) - await producer.append({"n": 2}) - async def two_changes() -> list[Any]: - # A standalone stream reaches its channel through a task of its - # own, so the notifications may trail the appends. + async def changes(count: int) -> list[Any]: polled = await native.poll_channel(address.channel, wait=False) - assert [n.counter for n in polled] == [1, 2] + assert [n.counter for n in polled] == list(range(1, count + 1)) return polled - polled = await assert_eventually(two_changes) + # A standalone stream reaches its channel through a task of its own + # that hands over the latest change, so the notifications trail the + # appends and a burst arrives as its newest change. Each append waits + # for its notification so that every change is seen. + await producer.append({"n": 1}) + await assert_eventually(lambda: changes(1)) + await producer.append({"n": 2}) + polled = await assert_eventually(lambda: changes(2)) for notification in polled: assert notification.channel == address.channel assert notification.linked_to is None From 56ab5119d54bb831079dc44fc6ad9a0c3687dd68 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 18:13:22 -0700 Subject: [PATCH 27/27] Named a standalone activity's stream channel by execution on the native cases. The native e2e cases read the owner as an execution, and a new live case polls the channel linked to a standalone activity, which waits for the server layer that addresses a linked channel by execution. --- tests/streams/test_stream_channel_e2e.py | 75 +++++++++++++++++++++++- 1 file changed, 72 insertions(+), 3 deletions(-) diff --git a/tests/streams/test_stream_channel_e2e.py b/tests/streams/test_stream_channel_e2e.py index e698db6a1..6a632556a 100644 --- a/tests/streams/test_stream_channel_e2e.py +++ b/tests/streams/test_stream_channel_e2e.py @@ -18,7 +18,7 @@ import pytest -from temporalio import workflow +from temporalio import activity, workflow from temporalio.client import ( Callback, ChannelAddress, @@ -26,6 +26,7 @@ Client, stream_channel, ) +from temporalio.common import Execution, ExecutionType from temporalio.service import RPCError from temporalio.streams import StreamRef, topic from temporalio.streams.providers.native import NativeStreams @@ -144,7 +145,8 @@ async def test_a_workflow_stream_notifies_the_channel_linked_to_its_owner( ) try: address = stream_channel(StreamRef.for_workflow(handle.id, topic=OUT)) - assert address == ChannelAddress("stream/out", handle.id) + assert address == ChannelAddress("stream/out", Execution.workflow(handle.id)) + assert address.workflow_id == handle.id assert address == stream_channel( native.get_stream_handle(handle.id).ref(topic=OUT) ) @@ -167,7 +169,7 @@ async def changes(count: int) -> list[Any]: run_id = handle.first_execution_run_id for notification in polled: assert notification.linked_to is not None - assert notification.linked_to.workflow_id == handle.id + assert notification.linked_to.business_id == handle.id assert notification.position.decode().startswith(f"{run_id}:") description = await native.describe_channel( address.channel, workflow_id=address.workflow_id @@ -183,3 +185,70 @@ async def changes(count: int) -> list[Any]: running.cancel() await asyncio.gather(running, return_exceptions=True) await provider.close() + + +_released: dict[str, bool] = {} +"""Activity ids the test has let go, read by the activity in the same process.""" + + +@activity.defn +async def write_then_wait() -> str: + """Writes one record to the activity's own stream and lingers until released.""" + activity_id = activity.info().activity_id + producer = activity.stream_handle().producer(topic=OUT) + await producer.append({"n": 1}) + while not _released.get(activity_id): + activity.heartbeat() + await asyncio.sleep(0.2) + return activity_id + + +@pytest.mark.needs_execution_server +async def test_a_standalone_activity_stream_notifies_the_channel_linked_to_it( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + activity_id = f"act-{uuid.uuid4().hex[:8]}" + async with new_worker(native, activities=[write_then_wait]) as worker: + handle = await native.start_activity( + write_then_wait, + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=60), + ) + try: + # A standalone activity is an execution of its own, so its stream + # notifies a channel linked to it, under the topic's name alone. + address = stream_channel(StreamRef.for_activity(activity_id, topic=OUT)) + assert address == ChannelAddress( + "stream/out", Execution.activity(activity_id) + ) + assert address.workflow_id is None + + async def changes(count: int) -> list[Any]: + polled = await native.poll_channel( + address.channel, execution=address.execution, wait=False + ) + assert [n.counter for n in polled] == list(range(1, count + 1)) + return polled + + [first] = await assert_eventually( + lambda: changes(1), timeout=timedelta(seconds=30) + ) + assert first.linked_to is not None + assert (first.linked_to.type, first.linked_to.business_id) == ( + ExecutionType.ACTIVITY, + activity_id, + ) + description = await native.describe_channel( + address.channel, execution=address.execution + ) + assert description.kind == ChannelKind.LINKED + assert description.linked_to is not None + assert description.linked_to.business_id == activity_id + finally: + _released[activity_id] = True + with contextlib.suppress(Exception): + await asyncio.wait_for(handle.result(), 30) + await provider.close()