From c6a50b8cfb5851d611d60375506ddb0cc2fca6a6 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:21:41 -0700 Subject: [PATCH 1/3] Replayed workflows that read native streams. History records only the offsets each task consumed, so the replayer fetches those records from the stream service, or from slices carried with an exported history, and hands them to the replay. --- temporalio/bridge/src/worker.rs | 13 +- temporalio/client/_workflow.py | 63 +++++++- temporalio/streams/providers/native.py | 79 ++++++++-- temporalio/worker/_replayer.py | 210 +++++++++++++++++++++++++ temporalio/worker/_stream_ranges.py | 173 ++++++++++++++++++++ temporalio/worker/_worker.py | 3 + temporalio/worker/_workflow.py | 15 +- 7 files changed, 537 insertions(+), 19 deletions(-) create mode 100644 temporalio/worker/_stream_ranges.py diff --git a/temporalio/bridge/src/worker.rs b/temporalio/bridge/src/worker.rs index e74698e94..44e5a2005 100644 --- a/temporalio/bridge/src/worker.rs +++ b/temporalio/bridge/src/worker.rs @@ -14,6 +14,7 @@ use temporalio_common::protos::coresdk::{ nexus::NexusTaskCompletion, ActivityHeartbeat, ActivityTaskCompletion, }; use temporalio_common::protos::temporal::api::history::v1::History; +use temporalio_common::protos::temporal::api::stream::v1::StreamSlice; use temporalio_common::protos::temporal::api::worker::v1::{PluginInfo, StorageDriverInfo}; use temporalio_sdk_core::replay::{HistoryForReplay, ReplayWorkerInput}; use temporalio_sdk_core::{ @@ -950,14 +951,24 @@ impl HistoryPusher { #[pymethods] impl HistoryPusher { + /// Feed one history to the replay worker. `stream_slices` are serialized + /// `temporal.api.stream.v1.StreamSlice` messages carrying the records the + /// history's completed tasks consumed, which History itself never holds. + #[pyo3(signature = (workflow_id, history_proto, stream_slices = Vec::new()))] fn push_history<'p>( &self, py: Python<'p>, workflow_id: &str, history_proto: &Bound<'_, PyBytes>, + stream_slices: Vec>, ) -> PyResult> { let history = History::decode(history_proto.as_bytes()) .map_err(|err| PyValueError::new_err(format!("Invalid proto: {err}")))?; + let slices = stream_slices + .iter() + .map(|slice| StreamSlice::decode(slice.as_bytes())) + .collect::, _>>() + .map_err(|err| PyValueError::new_err(format!("Invalid stream slice proto: {err}")))?; let wfid = workflow_id.to_string(); let tx = if let Some(tx) = self.tx.as_ref() { tx.clone() @@ -968,7 +979,7 @@ impl HistoryPusher { }; // We accept this doesn't have logging/tracing self.runtime.future_into_py(py, async move { - tx.send(HistoryForReplay::new(history, wfid)) + tx.send(HistoryForReplay::new(history, wfid).with_stream_slices(slices)) .await .map_err(|_| { PyRuntimeError::new_err( diff --git a/temporalio/client/_workflow.py b/temporalio/client/_workflow.py index 9fbc60138..c2e57cc92 100644 --- a/temporalio/client/_workflow.py +++ b/temporalio/client/_workflow.py @@ -4,6 +4,7 @@ import asyncio import functools +import json import warnings from asyncio import Future from collections.abc import ( @@ -31,6 +32,7 @@ import temporalio.api.common.v1 import temporalio.api.enums.v1 import temporalio.api.history.v1 +import temporalio.api.stream.v1 import temporalio.api.update.v1 import temporalio.api.workflow.v1 import temporalio.api.workflowservice.v1 @@ -1698,6 +1700,20 @@ class WorkflowHistory: events: Sequence[temporalio.api.history.v1.HistoryEvent] """History events for the workflow.""" + stream_slices: Sequence[temporalio.api.stream.v1.StreamSlice] = () + """The stream records the workflow's completed tasks consumed, if captured. + + History records only the offsets each Workflow Task consumed from a + server-side stream; the records live in the stream. A history whose tasks + consumed records and that carries none here cannot be replayed without the + server that still holds the stream. One slice per recorded range, tagged + with the ``WorkflowTaskCompleted`` event that recorded it, in the shape the + server puts on a poll response. + :py:meth:`temporalio.worker.Replayer.fetch_stream_slices` fills it while + the stream is retained, and :py:meth:`to_json` and :py:meth:`from_json` + carry it, so an exported history file is the whole replay input. + """ + @property def run_id(self) -> str: """Run ID extracted from the first event.""" @@ -1719,31 +1735,62 @@ def from_json(workflow_id: str, history: str | dict[str, Any]) -> WorkflowHistor Args: workflow_id: The workflow's ID history: A string or parsed-to-dict representation of workflow - history + history. A ``streamSlices`` list beside ``events``, as + :py:meth:`to_json` writes one, is read into + :py:attr:`stream_slices`. Returns: Workflow history """ parsed = _history_from_json(history) - return WorkflowHistory(workflow_id, parsed.events) + raw: Any = [] + if isinstance(history, dict): + raw = history.get("streamSlices") or history.get("stream_slices") or [] + elif '"streamSlices"' in history or '"stream_slices"' in history: + # Read again only when the text mentions them. Almost every + # history has none, and handing the parsed dict to + # _history_from_json instead would make it deep-copy an export + # that can be very large. + decoded = json.loads(history) + if isinstance(decoded, dict): + raw = decoded.get("streamSlices") or decoded.get("stream_slices") or [] + slices: list[temporalio.api.stream.v1.StreamSlice] = [ + google.protobuf.json_format.ParseDict( + entry, + temporalio.api.stream.v1.StreamSlice(), + ignore_unknown_fields=True, + ) + for entry in raw + ] + return WorkflowHistory(workflow_id, parsed.events, slices) def to_json(self) -> str: """Convert this history to JSON. - Note, this does not include the workflow ID. + Note, this does not include the workflow ID. The stream slices, when + there are any, are written as a ``streamSlices`` list beside the + events; without them the output is the history proto's JSON alone. """ - return google.protobuf.json_format.MessageToJson( - temporalio.api.history.v1.History(events=self.events) - ) + if not self.stream_slices: + return google.protobuf.json_format.MessageToJson( + temporalio.api.history.v1.History(events=self.events) + ) + return json.dumps(self.to_json_dict(), indent=2) def to_json_dict(self) -> dict[str, Any]: """Convert this history to JSON-compatible dict. - Note, this does not include the workflow ID. + Note, this does not include the workflow ID. See :py:meth:`to_json` + for how the stream slices are carried. """ - return google.protobuf.json_format.MessageToDict( + out = google.protobuf.json_format.MessageToDict( temporalio.api.history.v1.History(events=self.events) ) + if self.stream_slices: + out["streamSlices"] = [ + google.protobuf.json_format.MessageToDict(s) for s in self.stream_slices + ] + return out @dataclass diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py index 95bb866a3..8779c3465 100644 --- a/temporalio/streams/providers/native.py +++ b/temporalio/streams/providers/native.py @@ -21,7 +21,11 @@ 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. +from the run's close event who came next; with a run id it is pinned. A run +that was reset is followed too: its close event does not name the run reset +from it, describe does, and that run's streams continue the base run's offset +space where it inherited a subscription, so the read resumes at the floor the +stream reports rather than at zero. Prototype support for AI-198. It needs a server built from that branch and reaches the stream service on a channel of its own, opened with the client's @@ -411,6 +415,10 @@ def read( 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. + + The chain is followed across continue-as-new and across a reset: a + run reset from a closed one is read next, from the floor its stream + reports. A handle pinned to a run that was reset ends with that run. """ check_read_start(after, last) topic, result_type = resolve_topic(topic, result_type) @@ -469,10 +477,10 @@ async def _read( break if self._run_id is not None: return - successor = await self._successor(run_id) - if successor is None: + following = await self._successor(topic, run_id) + if following is None: return - run_id, offset = successor, 0 + run_id, offset = following @staticmethod def _previous( @@ -581,31 +589,84 @@ async def _predecessor(self, run_id: str) -> str | None: 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 + if attributes.continued_execution_run_id: + return attributes.continued_execution_run_id + # A reset run's start event is the base run's, copied, and the + # original run id it carries is kept across resets, so it names + # the run the chain of resets began from. + original = attributes.original_execution_run_id + if original and original != run_id: + return original + return 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: + async def _successor(self, topic: str, run_id: str) -> tuple[str, int] | None: + """The run that carries on after ``run_id``, and where its stream starts. + + A continue-as-new names its successor in the close event, and the + successor's streams start at zero. A reset does not: the base run is + closed with no word of the reset in its own History, and only describe + names the run reset from it. That run reads on streams of its own that + continue the base run's offset space where it inherited a subscription, + and such a stream refuses any offset below its floor, so the read + resumes at the floor the stream reports. + """ handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) events = handle.fetch_history_events( event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT ) + closed = False try: async for event in events: + closed = True 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 + if not attributes.new_execution_run_id: + return None + return attributes.new_execution_run_id, 0 except RPCError as error: if error.status != RPCStatusCode.NOT_FOUND: raise - return None + return None + if not closed: + return None + reset_run = await self._reset_run(run_id) + if reset_run is None: + return None + return reset_run, await self._floor(topic, reset_run) + + async def _reset_run(self, run_id: str) -> str | None: + """The run ``run_id`` was reset into, which only describe reports.""" + handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + try: + description = await handle.describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + return None + raise + extended = description.raw_description.workflow_extended_info + return extended.reset_run_id or None + + async def _floor(self, topic: str, run_id: str) -> int: + """The first offset a run's stream holds. + + A stream the run inherited a subscription to exists from the reset on + and starts at the inherited offset. One the base run only published to + is created on the run's first publish, at zero, and does not exist + before that. + """ + try: + return (await self._stream(topic, run_id).describe()).base_offset + except StreamNotFoundError: + return 0 class NativeActivityStreamHandle(NativeStreamHandle): @@ -685,7 +746,7 @@ async def _first_run(self) -> str: async def _predecessor(self, run_id: str) -> str | None: return None - async def _successor(self, run_id: str) -> str | None: + async def _successor(self, topic: str, run_id: str) -> tuple[str, int] | None: return None diff --git a/temporalio/worker/_replayer.py b/temporalio/worker/_replayer.py index c7ae8d59d..38304be48 100644 --- a/temporalio/worker/_replayer.py +++ b/temporalio/worker/_replayer.py @@ -8,10 +8,13 @@ from collections.abc import AsyncIterator, Mapping, Sequence from contextlib import AbstractAsyncContextManager, asynccontextmanager from dataclasses import dataclass +from typing import Any from typing_extensions import TypedDict +import temporalio.api.enums.v1 import temporalio.api.history.v1 +import temporalio.api.stream.v1 import temporalio.bridge.proto.workflow_activation import temporalio.bridge.worker import temporalio.client @@ -23,6 +26,7 @@ from ..common import HeaderCodecBehavior from ._interceptor import Interceptor +from ._stream_ranges import fetch_range from ._worker import load_default_build_id from ._workflow import _WorkflowWorker from ._workflow_instance import ( @@ -58,6 +62,7 @@ def __init__( disable_safe_workflow_eviction: bool = False, header_codec_behavior: HeaderCodecBehavior = HeaderCodecBehavior.NO_CODEC, stream_provider: temporalio.streams.StreamProvider | None = None, + stream_client: temporalio.client.Client | None = None, ) -> None: """Create a replayer to replay workflows from history. @@ -66,6 +71,30 @@ def __init__( the replayer that were passed to the worker when the workflow originally ran, ``stream_provider`` included when the workflow used streams. + A workflow that read a server-side stream cannot be replayed from its + history alone. History records the offsets each Workflow Task consumed + and never the records; on a live task the server re-supplies them from + the stream. Three cases: + + * The history carries the records in + :py:attr:`temporalio.client.WorkflowHistory.stream_slices`, put + there by :py:meth:`fetch_stream_slices` while the stream was + retained and carried by ``to_json()`` and ``from_json()``. They are + handed to the replay and no server is contacted. + * The history carries none and ``stream_client`` is given: a client + connected to the server that still holds the streams. The replayer + fetches every range the completed tasks recorded from the stream + service and hands the records to the replay, so the workflow sees + what it saw the first time. A range the stream no longer holds fails + that replay with :py:class:`temporalio.streams.StreamNotFoundError`. + The stream service is reached on a channel of its own, opened with + the client's connection settings (target, TLS, API key, headers, + retries), as :py:class:`temporalio.client_stream.Connection` + describes. The replayer closes the channel it opened when the + replay is finished. + * The history carries none and there is no client: replaying it fails + and the message names both remedies. + Note, unlike the worker, for the replayer the workflow_task_executor will default to a new thread pool executor with no max_workers set that will be shared across all replay calls and never explicitly shut down. @@ -89,6 +118,7 @@ def __init__( disable_safe_workflow_eviction=disable_safe_workflow_eviction, header_codec_behavior=header_codec_behavior, stream_provider=stream_provider, + stream_client=stream_client, ) self._initial_config = self._config.copy() self._default_workflow_logic_flags = set(_DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS) @@ -110,6 +140,31 @@ def _set_default_workflow_logic_flag( else: self._default_workflow_logic_flags.discard(flag) + @staticmethod + async def fetch_stream_slices( + client: temporalio.client.Client, history: temporalio.client.WorkflowHistory + ) -> temporalio.client.WorkflowHistory: + """Return ``history`` with the stream records its tasks consumed attached. + + History records only the offsets each Workflow Task consumed from a + server-side stream, so an export cannot be replayed without the + server that still holds the stream. This fetches every recorded range + from the stream service through ``client`` while the stream is + retained and returns a copy of ``history`` with the records in + :py:attr:`temporalio.client.WorkflowHistory.stream_slices`. Its + ``to_json()`` then carries them, and a replayer given the result, or a + ``from_json()`` of it, needs no server. Ranges recorded before a reset + point are fetched from the run the workflow was reset from. + + Raises :py:class:`temporalio.streams.StreamNotFoundError` for a range + the stream no longer holds. + """ + return temporalio.client.WorkflowHistory( + history.workflow_id, + history.events, + await _stream_slices(client, history), + ) + def config(self, *, active_config: bool = False) -> ReplayerConfig: """Config, as a dictionary, used to create this replayer. @@ -213,6 +268,7 @@ async def _workflow_replay_iterator( pusher = None workflow_worker_task = None bridge_worker_scope = None + stream_client = self._config.get("stream_client") try: last_replay_failure: Exception | None @@ -368,6 +424,33 @@ def on_eviction_hook( # Yield iterator async def replay_iterator() -> AsyncIterator[WorkflowReplayResult]: async for history in histories: + # The records a consuming workflow read are not in its + # history unless they were captured into it. Otherwise + # fetch them from the stream service, or report here what + # is missing, rather than let Core fail the first task on + # input it was never given. + stream_slices: list[bytes] = [] + if history.stream_slices: + stream_slices = [ + s.SerializeToString() for s in history.stream_slices + ] + elif stream_client is not None: + try: + fetched = await _stream_slices(stream_client, history) + except temporalio.streams.StreamNotFoundError as err: + yield WorkflowReplayResult( + history=history, replay_failure=err + ) + continue + stream_slices = [s.SerializeToString() for s in fetched] + else: + missing = _stream_records_missing(history) + if missing is not None: + yield WorkflowReplayResult( + history=history, replay_failure=RuntimeError(missing) + ) + continue + # Clear last complete and push history last_replay_complete.clear() await pusher.push_history( @@ -375,6 +458,7 @@ async def replay_iterator() -> AsyncIterator[WorkflowReplayResult]: temporalio.api.history.v1.History( events=history.events ).SerializeToString(), + stream_slices, ) # Wait for worker error or last replay to complete. This @@ -401,6 +485,13 @@ async def replay_iterator() -> AsyncIterator[WorkflowReplayResult]: yield replay_iterator() finally: + # The stream channel is opened by the fetch, on this loop, and is + # the replayer's to close: nothing else in the process asked for + # it, and leaving it open outlives the replay it served. + if stream_client is not None: + from temporalio.client_stream import close_shared_clients, shared_key + + await close_shared_clients(shared_key(stream_client)) # Close the pusher if pusher is not None: pusher.close() @@ -420,6 +511,118 @@ async def replay_iterator() -> AsyncIterator[WorkflowReplayResult]: logger.warning("Failed to finalize shutdown", exc_info=True) +def _consumed_ranges( + history: temporalio.client.WorkflowHistory, +) -> list[tuple[int, temporalio.api.stream.v1.StreamRange]]: + """Every range a completed task recorded, with the id of the event that recorded it.""" + out: list[tuple[int, temporalio.api.stream.v1.StreamRange]] = [] + for event in history.events: + if not event.HasField("workflow_task_completed_event_attributes"): + continue + attributes = event.workflow_task_completed_event_attributes + for consumed in attributes.consumed_stream_ranges: + out.append((event.event_id, consumed)) + return out + + +def _eras( + history: temporalio.client.WorkflowHistory, +) -> list[tuple[str, list[tuple[int, temporalio.api.stream.v1.StreamRange]]]]: + """The recorded ranges grouped by the run whose streams hold them. + + A reset copies the base run's history into the new run, so the ranges + before the reset point were consumed from the base run's streams, and the + run they belong to is only named by the ``WorkflowTaskFailed`` event that + marks the reset point, after them in the history. Ranges after the last + reset point belong to the run itself, which that event names as well; a + history with no reset point belongs to the run its start event names. This + is the split the server makes when it re-supplies a cache miss. + """ + reset = temporalio.api.enums.v1.WorkflowTaskFailedCause.WORKFLOW_TASK_FAILED_CAUSE_RESET_WORKFLOW + eras: list[tuple[str, list[tuple[int, temporalio.api.stream.v1.StreamRange]]]] = [] + pending: list[tuple[int, temporalio.api.stream.v1.StreamRange]] = [] + own_run_id = history.run_id + for event in history.events: + if event.HasField("workflow_task_failed_event_attributes"): + failed = event.workflow_task_failed_event_attributes + if failed.cause == reset and failed.base_run_id: + eras.append((failed.base_run_id, pending)) + pending = [] + own_run_id = failed.new_run_id or own_run_id + continue + if not event.HasField("workflow_task_completed_event_attributes"): + continue + attributes = event.workflow_task_completed_event_attributes + for consumed in attributes.consumed_stream_ranges: + pending.append((event.event_id, consumed)) + eras.append((own_run_id, pending)) + return eras + + +def _stream_records_missing(history: temporalio.client.WorkflowHistory) -> str | None: + """Why this history cannot be replayed without a stream client, or ``None``.""" + ranges = [r for _, r in _consumed_ranges(history) if r.to_offset > r.from_offset] + if not ranges: + return None + streams = sorted({r.stream_id for r in ranges}) + return ( + f"workflow {history.workflow_id!r} consumed records from stream(s) " + f"{', '.join(repr(s) for s in streams)} in {len(ranges)} task(s), and History " + "records only the offsets. Either create the Replayer with stream_client= (a " + "Client connected to the server that still holds the streams) so the records " + "can be fetched from the stream service, or replay a history exported with " + "them: Replayer.fetch_stream_slices(client, history) attaches the records " + "while the stream is retained and WorkflowHistory.to_json() carries them." + ) + + +async def _stream_slices( + client: temporalio.client.Client, history: temporalio.client.WorkflowHistory +) -> list[temporalio.api.stream.v1.StreamSlice]: + """Fetch the records the history's tasks consumed, as ``StreamSlice`` messages. + + The shape is the one the server puts on a poll response when it re-supplies + a cache miss: one slice per recorded range, tagged with the completion that + recorded it, an empty range included, fetched from the run whose stream + holds it (the run reset from, for a range recorded before a reset point). + Raises :py:class:`temporalio.streams.StreamNotFoundError` for a range the + stream no longer holds. + """ + if not _consumed_ranges(history): + return [] + # Imported here: the stream client needs the grpc extra, which a replayer + # without a stream client never touches. + from temporalio.client_stream import shared_client + + streams = shared_client(client) + # The handle that served each run's stream, so later ranges of the same + # stream go straight to it. + served_by: dict[tuple[str, str], Any] = {} + slices: list[temporalio.api.stream.v1.StreamSlice] = [] + for run_id, ranges in _eras(history): + for event_id, consumed in ranges: + stream_slice = temporalio.api.stream.v1.StreamSlice( + stream_id=consumed.stream_id, + run_id=run_id, + from_offset=consumed.from_offset, + to_offset=consumed.to_offset, + workflow_task_completed_event_id=event_id, + ) + if consumed.to_offset > consumed.from_offset: + handle, records, owner_run_id = await fetch_range( + streams, + history.workflow_id, + run_id, + consumed, + served_by.get((run_id, consumed.stream_id)), + ) + served_by[(run_id, consumed.stream_id)] = handle + stream_slice.run_id = owner_run_id or run_id + stream_slice.records.extend(records) + slices.append(stream_slice) + return slices + + class ReplayerConfig(TypedDict, total=False): """TypedDict of config originally passed to :py:class:`Replayer`.""" @@ -438,6 +641,7 @@ class ReplayerConfig(TypedDict, total=False): disable_safe_workflow_eviction: bool header_codec_behavior: HeaderCodecBehavior stream_provider: temporalio.streams.StreamProvider | None + stream_client: temporalio.client.Client | None @dataclass(frozen=True) @@ -455,6 +659,12 @@ class WorkflowReplayResult: :py:class:`temporalio.workflow.NondeterminismError` was encountered during replay - likely indicating your workflow code is incompatible with the history. + + A workflow that read a server-side stream reports here, before any task + runs, when its records could not be fetched: a + :py:class:`temporalio.streams.StreamNotFoundError` when the stream no + longer holds a range a task consumed, or a ``RuntimeError`` when the + replayer was given no ``stream_client`` to fetch them with. """ diff --git a/temporalio/worker/_stream_ranges.py b/temporalio/worker/_stream_ranges.py new file mode 100644 index 000000000..6bf35dbef --- /dev/null +++ b/temporalio/worker/_stream_ranges.py @@ -0,0 +1,173 @@ +"""Reading a recorded stream range back from the stream service. + +History records the offsets each Workflow Task consumed from a server-side +stream and never the records. Two paths have to read those offsets back: the +replayer, for a history it was handed without its records, and a live worker +handed a task whose re-supplied ranges stop short of what History recorded. +Both resolve a subscribed name the way the server does and read exactly the +recorded offsets, so the workflow sees what it saw the first time. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +import temporalio.api.stream.v1 +import temporalio.bridge.proto.workflow_activation +import temporalio.service +import temporalio.streams + +if TYPE_CHECKING: + import temporalio.client + +__all__ = ["fetch_range", "fill_short_stream_ranges", "read_range"] + + +async def fill_short_stream_ranges( + activation: temporalio.bridge.proto.workflow_activation.WorkflowActivation, + workflow_id: str, + client: temporalio.client.Client, +) -> None: + """Fetch the records a short ``DeliverStreamRecords`` job lacks, before the activation runs. + + The server re-supplies a replaying workflow's recorded ranges on the task + that carries them, within a budget. A job whose records stop short of its + range is that budget's remainder: the offsets are recorded in History and + the stream still holds them, so they are read from the stream service the + way the replayer reads them for an exported history, and the job is made + whole before the workflow runs on it. A job that covers its range, or that + records an empty observation, is left as it is. The fetch reads from the + activation's own run; a range recorded before a reset point lives on the + base run, which the job does not name yet. + + Raises: + temporalio.streams.StreamNotFoundError: The stream no longer holds the + missing offsets. The task fails rather than run on less input. + """ + short = [ + job.deliver_stream_records + for job in activation.jobs + if job.HasField("deliver_stream_records") + and len(job.deliver_stream_records.records) + < job.deliver_stream_records.to_offset - job.deliver_stream_records.from_offset + ] + if not short: + return + # Imported here: the stream client needs the grpc extra, which a worker + # that never sees a short range never touches. + from temporalio.client_stream import shared_client + + streams = shared_client(client) + served_by: dict[str, Any] = {} + for job in short: + missing = temporalio.api.stream.v1.StreamRange( + stream_id=job.stream_id, + from_offset=job.from_offset + len(job.records), + to_offset=job.to_offset, + ) + handle, records, _ = await fetch_range( + streams, + workflow_id, + activation.run_id, + missing, + served_by.get(job.stream_id), + ) + served_by[job.stream_id] = handle + job.records.extend(records) + + +async def fetch_range( + streams: Any, + workflow_id: str, + run_id: str, + consumed: temporalio.api.stream.v1.StreamRange, + known: Any, +) -> tuple[Any, list[temporalio.api.stream.v1.StreamRecord], str]: + """The records at exactly the recorded range, with the handle that served them. + + A subscribed name is resolved as the server resolves it: a stream the + workflow owns by that name first, else a standalone stream by that id. The + owned stream cannot be asked whether it exists, since a name nobody wrote + reads as empty, so the owned stream is probed for the range's first record + and the standalone one is tried when it has nothing there. ``known`` is the + handle that served this stream before, when there was one. + """ + gone = ( + f"stream {consumed.stream_id!r} no longer holds offsets " + f"[{consumed.from_offset}, {consumed.to_offset}) that a completed task of " + f"workflow {workflow_id!r} run {run_id!r} consumed; its records cannot be " + "replayed" + ) + candidates = ( + [known] + if known is not None + else [ + streams.workflow_stream( + workflow_id, consumed.stream_id, owner_run_id=run_id + ), + streams.get(consumed.stream_id), + ] + ) + last: Exception | None = None + for handle in candidates: + try: + records, owner_run_id = await read_range(handle, consumed, gone) + except temporalio.streams.StreamNotFoundError as error: + last = error + continue + if records is not None: + return handle, records, owner_run_id + raise temporalio.streams.StreamNotFoundError(gone) from last + + +async def read_range( + handle: Any, consumed: temporalio.api.stream.v1.StreamRange, gone: str +) -> tuple[list[temporalio.api.stream.v1.StreamRecord] | None, str]: + """Read ``[from_offset, to_offset)`` from one stream. + + ``None`` when the stream has nothing at the range's first offset, which is + how a name that is not this stream's reads; a stream that has the start but + not the rest, or refuses the offset as truncated or past its head, raises + :class:`temporalio.streams.StreamNotFoundError` with ``gone`` as the reason. + """ + records: list[temporalio.api.stream.v1.StreamRecord] = [] + owner_run_id = "" + offset = consumed.from_offset + while offset < consumed.to_offset: + try: + page = await handle.poll( + from_offset=offset, + max_records=consumed.to_offset - offset, + wait=False, + ) + except temporalio.streams.StreamNotFoundError: + if not records: + return None, "" + raise + except temporalio.streams.StreamCursorError as error: + # Below the truncation floor: the records a task consumed are gone. + raise temporalio.streams.StreamNotFoundError(gone) from error + except temporalio.service.RPCError as error: + # Past the head, or a refusal the server does not type: it refuses + # the offset rather than answering short. + if error.status in ( + temporalio.service.RPCStatusCode.FAILED_PRECONDITION, + temporalio.service.RPCStatusCode.INVALID_ARGUMENT, + temporalio.service.RPCStatusCode.OUT_OF_RANGE, + ): + raise temporalio.streams.StreamNotFoundError(gone) from error + raise + owner_run_id = page.run_id or owner_run_id + if not page.entries: + if not records: + return None, "" + raise temporalio.streams.StreamNotFoundError(gone) + for entry in page.entries: + if entry.offset != offset or offset >= consumed.to_offset: + raise temporalio.streams.StreamNotFoundError( + f"{gone}: the stream answered offset {entry.offset} where " + f"{offset} was due" + ) + records.append(entry.record) + offset += 1 + return records, owner_run_id diff --git a/temporalio/worker/_worker.py b/temporalio/worker/_worker.py index fade69c4d..2f6ff3792 100644 --- a/temporalio/worker/_worker.py +++ b/temporalio/worker/_worker.py @@ -579,6 +579,9 @@ def check_activity(activity: str): != HeaderCodecBehavior.NO_CODEC, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, stream_provider=stream_provider, + stream_client=( + config["client"] if stream_provider is not None else None # type: ignore[reportTypedDictNotRequiredAccess] + ), ) tuner = config.get("tuner") diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index f4001b32e..40beaa65a 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -13,7 +13,7 @@ from dataclasses import dataclass from datetime import timezone from types import TracebackType -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast import temporalio.api.common.v1 import temporalio.bridge.proto.common @@ -43,6 +43,7 @@ WorkflowInterceptorClassInput, WorkflowOutboundInterceptor, ) +from ._stream_ranges import fill_short_stream_ranges from ._workflow_instance import ( _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS, PatchActivationInput, @@ -53,6 +54,9 @@ _WorkflowLogicFlag, ) +if TYPE_CHECKING: + import temporalio.client + logger = logging.getLogger(__name__) # Set to true to log all activations and completions @@ -158,6 +162,7 @@ def __init__( max_workflow_task_external_storage_concurrency: int, default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] | None = None, stream_provider: temporalio.streams.StreamProvider | None = None, + stream_client: temporalio.client.Client | None = None, ) -> None: # Debug mode is enabled if specified or if the TEMPORAL_DEBUG env var is truthy debug_mode = debug_mode or bool(os.environ.get("TEMPORAL_DEBUG")) @@ -221,6 +226,9 @@ def __init__( # Innermost, so the lifecycle hooks bracket the workflow function # itself, after every user interceptor has done its own setup. self._interceptor_classes.append(_StreamHooksInterceptor) + # For the records a task's re-supplied ranges leave out, which the + # stream service still holds. + self._stream_client = stream_client self._workflow_failure_exception_types = workflow_failure_exception_types self._patch_activation_callback = patch_activation_callback @@ -415,6 +423,11 @@ async def _handle_activation( "Cache already exists for activation with initialize job" ) + # Before the bodies are decoded, so what is fetched is decoded with + # the rest, and before the workflow runs on the range. + if self._stream_client is not None: + await fill_short_stream_ranges(act, workflow_id, self._stream_client) + workflow_context = temporalio.converter.WorkflowSerializationContext( namespace=self._namespace, workflow_id=workflow_id, From bedc0e9619f0ce573c2d570bac725adb1fdd25b0 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:21:41 -0700 Subject: [PATCH 2/3] Covered replay of native stream reads. Re-supply of short ranges, the replayer against a live stream, reset runs and offline histories with stream slices. --- tests/worker/test_replayer.py | 46 ++ tests/worker/test_stream_resupply.py | 150 +++++ tests/worker/test_workflow_stream_e2e.py | 710 ++++++++++++++++++++++- 3 files changed, 899 insertions(+), 7 deletions(-) create mode 100644 tests/worker/test_stream_resupply.py diff --git a/tests/worker/test_replayer.py b/tests/worker/test_replayer.py index 52fed7416..ddce04826 100644 --- a/tests/worker/test_replayer.py +++ b/tests/worker/test_replayer.py @@ -8,8 +8,12 @@ from pathlib import Path from typing import Any +import google.protobuf.json_format import pytest +import temporalio.api.common.v1 +import temporalio.api.history.v1 +import temporalio.api.stream.v1 import temporalio.worker._workflow_instance from temporalio import activity, workflow from temporalio.client import Client, WorkflowFailureError, WorkflowHistory @@ -121,6 +125,48 @@ async def test_replayer_workflow_complete(client: Client) -> None: ) +def test_workflow_history_json_carries_stream_slices_only_when_present() -> None: + """The JSON shape is unchanged without slices, and round-trips them when present.""" + with Path(__file__).with_name("test_replayer_complete_history.json").open("r") as f: + history = WorkflowHistory.from_json("fake", f.read()) + assert list(history.stream_slices) == [] + # Byte for byte what the history proto's JSON was before slices existed. + assert history.to_json() == google.protobuf.json_format.MessageToJson( + temporalio.api.history.v1.History(events=history.events) + ) + assert "streamSlices" not in history.to_json_dict() + + record = temporalio.api.stream.v1.StreamRecord( + body=temporalio.api.common.v1.Payload( + data=b'{"n": 1}', metadata={"encoding": b"json/plain"} + ), + topic="inputs", + kind=temporalio.api.stream.v1.StreamRecordKind.STREAM_RECORD_KIND_DATA, + producer_id="model", + attempt=1, + sequence=0, + ) + stream_slice = temporalio.api.stream.v1.StreamSlice( + stream_id="inputs", + run_id="run-1", + from_offset=0, + to_offset=1, + records=[record], + workflow_task_completed_event_id=4, + ) + bundle = WorkflowHistory("fake", history.events, [stream_slice]) + + text = bundle.to_json() + assert "streamSlices" in text + restored = WorkflowHistory.from_json("fake", text) + assert list(restored.stream_slices) == [stream_slice] + assert list(restored.events) == list(history.events) + # The dict form reads back the same, and a plain export stays plain. + from_dict = WorkflowHistory.from_json("fake", bundle.to_json_dict()) + assert list(from_dict.stream_slices) == [stream_slice] + assert WorkflowHistory("fake", restored.events).to_json() == history.to_json() + + @pytest.mark.skipif(sys.version_info < (3, 12), reason="Skipping for < 3.12") async def test_replayer_workflow_complete_json() -> None: # See `test_replayer_workflow_complete` for full skip description. diff --git a/tests/worker/test_stream_resupply.py b/tests/worker/test_stream_resupply.py new file mode 100644 index 000000000..f7dcf3f78 --- /dev/null +++ b/tests/worker/test_stream_resupply.py @@ -0,0 +1,150 @@ +"""A task whose re-supplied ranges stop short is made whole before it runs. + +The server re-supplies a replaying workflow's recorded ranges within a budget. +These pin the worker's half: a ``DeliverStreamRecords`` job with fewer records +than its range spans has the rest fetched from the stream service, resolved as +the server resolves a subscribed name, while a whole job and an empty +observation are left alone and a range the stream no longer holds fails the +task rather than let it run on less input. +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import pytest + +from temporalio import client_stream +from temporalio.api.common.v1 import Payload +from temporalio.api.stream.v1 import StreamRecord +from temporalio.bridge.proto.workflow_activation import WorkflowActivation +from temporalio.client_stream import Page, StreamEntry +from temporalio.streams import StreamNotFoundError +from temporalio.worker._stream_ranges import fill_short_stream_ranges + +# The fill-in only hands the client to the shared-channel lookup, which the +# tests replace. +_CLIENT: Any = SimpleNamespace() + + +def _record(offset: int) -> StreamRecord: + return StreamRecord(topic="inputs", body=Payload(data=f"r{offset}".encode())) + + +class _FakeStream: + """One stream's records by offset, answering polls as the service does.""" + + def __init__(self, held: dict[int, StreamRecord]) -> None: + self.held = held + self.polls: list[tuple[int, int]] = [] + + async def poll(self, *, from_offset: int, max_records: int, wait: bool) -> Page: + assert not wait + self.polls.append((from_offset, max_records)) + if not self.held: + raise StreamNotFoundError("no such stream") + entries = [ + StreamEntry(record=self.held[o], offset=o) + for o in range(from_offset, from_offset + max_records) + if o in self.held + ] + head = max(self.held) + 1 + return Page( + entries=entries, + next_offset=from_offset + len(entries), + head_offset=head, + closed=False, + ) + + +class _FakeStreams: + """The stream client: an owned stream per (workflow, name) and standalone ones by id.""" + + def __init__( + self, owned: dict[str, _FakeStream], standalone: dict[str, _FakeStream] + ) -> None: + self.owned = owned + self.standalone = standalone + + def workflow_stream( + self, workflow_id: str, name: str, *, owner_run_id: str + ) -> _FakeStream: + assert (workflow_id, owner_run_id) == ("wf", "run") + return self.owned.get(name, _FakeStream({})) + + def get(self, stream_id: str) -> _FakeStream: + return self.standalone.get(stream_id, _FakeStream({})) + + +def _activation(*jobs: tuple[str, int, int, int]) -> WorkflowActivation: + """An activation with one delivery per ``(stream, from, to, records present)``.""" + act = WorkflowActivation(run_id="run") + for stream_id, from_offset, to_offset, present in jobs: + job = act.jobs.add() + job.deliver_stream_records.stream_id = stream_id + job.deliver_stream_records.from_offset = from_offset + job.deliver_stream_records.to_offset = to_offset + job.deliver_stream_records.records.extend( + _record(o) for o in range(from_offset, from_offset + present) + ) + return act + + +@pytest.fixture +def streams(monkeypatch: pytest.MonkeyPatch) -> _FakeStreams: + fake = _FakeStreams( + owned={"inputs": _FakeStream({o: _record(o) for o in range(10)})}, + standalone={"shared": _FakeStream({o: _record(o) for o in range(5)})}, + ) + monkeypatch.setattr(client_stream, "shared_client", lambda _client: fake) + return fake + + +def _bodies(act: WorkflowActivation, index: int = 0) -> list[bytes]: + return [r.body.data for r in act.jobs[index].deliver_stream_records.records] + + +async def test_a_short_job_is_filled_from_the_owned_stream( + streams: _FakeStreams, +) -> None: + act = _activation(("inputs", 2, 7, 2)) + await fill_short_stream_ranges(act, "wf", _CLIENT) + assert _bodies(act) == [b"r2", b"r3", b"r4", b"r5", b"r6"] + # Only the missing tail was asked for. + assert streams.owned["inputs"].polls == [(4, 3)] + + +@pytest.mark.usefixtures("streams") +async def test_a_name_the_workflow_does_not_own_is_a_standalone_stream() -> None: + act = _activation(("shared", 0, 5, 1)) + await fill_short_stream_ranges(act, "wf", _CLIENT) + assert _bodies(act) == [b"r0", b"r1", b"r2", b"r3", b"r4"] + + +async def test_whole_and_empty_jobs_are_left_alone(streams: _FakeStreams) -> None: + act = _activation(("inputs", 0, 3, 3), ("inputs", 3, 3, 0)) + await fill_short_stream_ranges(act, "wf", _CLIENT) + assert _bodies(act, 0) == [b"r0", b"r1", b"r2"] + assert _bodies(act, 1) == [] + assert streams.owned["inputs"].polls == [] + + +@pytest.mark.usefixtures("streams") +async def test_a_range_the_stream_no_longer_holds_fails_loudly() -> None: + act = _activation(("inputs", 8, 12, 1)) + with pytest.raises( + StreamNotFoundError, match="no longer holds offsets \\[9, 12\\)" + ): + await fill_short_stream_ranges(act, "wf", _CLIENT) + + +async def test_no_client_is_needed_when_nothing_is_short( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def refuse(_client: Any) -> Any: + raise AssertionError("no channel should be opened") + + monkeypatch.setattr(client_stream, "shared_client", refuse) + act = _activation(("inputs", 0, 2, 2)) + await fill_short_stream_ranges(act, "wf", _CLIENT) diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py index e2ca2d9da..b05fd653a 100644 --- a/tests/worker/test_workflow_stream_e2e.py +++ b/tests/worker/test_workflow_stream_e2e.py @@ -17,19 +17,25 @@ import asyncio import os import uuid +from collections.abc import Sequence +from datetime import timedelta 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.api.common.v1 import Payload, WorkflowExecution +from temporalio.api.enums.v1 import EventType, WorkflowTaskFailedCause +from temporalio.api.history.v1 import HistoryEvent +from temporalio.api.stream.v1 import StreamRange, StreamRecord +from temporalio.api.workflowservice.v1 import ResetWorkflowExecutionRequest +from temporalio.client import Client, WorkflowExecutionStatus, WorkflowHistory from temporalio.client_stream import StreamClient -from temporalio.streams import END, RecordKind +from temporalio.converter import DataConverter, PayloadCodec +from temporalio.streams import END, RecordKind, StreamNotFoundError from temporalio.streams.providers.native import NativeStreams -from temporalio.worker import Worker +from temporalio.worker import Replayer, Worker +from temporalio.workflow import NondeterminismError from tests.streams.test_streams_conformance import take TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") @@ -69,11 +75,19 @@ async def _event_counts(client: Client, workflow_id: str) -> dict[Any, int]: class ContractLoop: """Reads ``inputs``, publishes a decision per value, reports control records.""" + def __init__(self) -> None: + self._trace: list[dict[str, Any]] = [] + + @workflow.query + def trace(self) -> list[dict[str, Any]]: + """What the loop has decided so far, for a query against replayed state.""" + return self._trace + @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]] = [] + trace = self._trace async for record in inputs: if record.kind is RecordKind.SUPERSEDED: assert record.supersession is not None @@ -386,6 +400,688 @@ async def test_a_cached_workflow_consumes_across_sticky_tasks() -> None: await streams.close() +class _EncryptingCodec(PayloadCodec): + """A codec whose output never contains its input. + + The SDK's own test codec marks payloads as encrypted but leaves the bytes + as they are, so it cannot show that a stored body is not plaintext. This + one keeps its metadata convention and runs a repeating-key XOR over the + serialized payload, which is enough for the plaintext markers the tests + look for to be absent from what the server stores. + """ + + _KEY = b"ai198" + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [ + Payload( + metadata={"encoding": b"binary/encrypted"}, + data=self._xor(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/encrypted": + out.append(p) + continue + out.append(Payload.FromString(self._xor(p.data))) + return out + + def _xor(self, data: bytes) -> bytes: + key = self._KEY + return bytes(b ^ key[i % len(key)] for i, b in enumerate(data)) + + +async def _drive_contract_loop(client: Client) -> WorkflowHistory: + """Run ``ContractLoop`` to completion on ``client`` and return its history.""" + task_queue = "codec-tq-" + uuid.uuid4().hex[:8] + workflow_id = "codec-wf-" + uuid.uuid4().hex[:8] + async with Worker(client, task_queue=task_queue, workflows=[ContractLoop]): + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + await producer.finish() + trace = await asyncio.wait_for(handle.result(), 60) + # The workflow decoded what the outside producer appended. + assert trace == [ + {"kind": "decision", "n": 1}, + {"kind": "decision", "n": 2}, + {"kind": "finish", "producer": "model"}, + ] + # The outside read decodes what the workflow published. + decisions = [ + (r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + assert decisions == [ + (RecordKind.DATA, {"decided": 1}), + (RecordKind.DATA, {"decided": 2}), + (RecordKind.FINISH, None), + ] + return await handle.fetch_history() + + +async def test_a_codec_encodes_records_on_both_halves() -> None: + """A payload codec on the client covers the workflow's records and an outside producer's. + + The worker's payload visitor runs the codec over the bodies a workflow + publishes and receives; the outside half applies the client's codec to each + body it sends and reads. Read raw, without the codec, every stored body is + ciphertext, and each side still reads the other's records in the clear. + """ + plain = await _connect() + config = plain.config() + config["data_converter"] = DataConverter(payload_codec=_EncryptingCodec()) + provider = NativeStreams() + config["plugins"] = [provider] + client = Client(**config) + raw = StreamClient.connect(TARGET or "") + try: + history = await _drive_contract_loop(client) + run_id = history.run_id + + # What the server holds, read through the stream service with no codec. + inputs = await raw.workflow_stream( + history.workflow_id, INPUTS, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + decisions = await raw.workflow_stream( + history.workflow_id, DECISIONS, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + + # The outside producer's two records, stored encoded. + stored_inputs = [e.record for e in inputs.entries if e.record.HasField("body")] + assert len(stored_inputs) == 2 + for record in stored_inputs: + assert record.body.metadata["encoding"] == b"binary/encrypted" + assert b'"n"' not in record.body.data + # The workflow's two decisions, stored encoded by the worker's visitor. + stored_decisions = [ + e.record for e in decisions.entries if e.record.HasField("body") + ] + assert len(stored_decisions) == 2 + for record in stored_decisions: + assert record.body.metadata["encoding"] == b"binary/encrypted" + assert b"decided" not in record.body.data + finally: + await raw.close() + await provider.close() + + +def _consumed_ranges(history: WorkflowHistory) -> list[StreamRange]: + return [ + consumed + for event in history.events + if event.HasField("workflow_task_completed_event_attributes") + for consumed in event.workflow_task_completed_event_attributes.consumed_stream_ranges + ] + + +def _copy_of( + history: WorkflowHistory, workflow_id: str | None = None +) -> WorkflowHistory: + events: list[HistoryEvent] = [] + for event in history.events: + copied = HistoryEvent() + copied.CopyFrom(event) + events.append(copied) + return WorkflowHistory(workflow_id or history.workflow_id, events) + + +async def test_a_replayer_with_a_client_replays_a_consuming_workflow() -> None: + """Rule 2 through the ``Replayer``: the records come back from the stream service. + + History holds the offsets each task consumed and never the records, so the + replayer is given a client to the server that still holds the streams and + fetches every recorded range before pushing the history. The same delivery + path as a live cache miss then hands each range to the task that consumed + it, and the reissued publishes match their events. + """ + provider = NativeStreams() + client = await _connect(provider) + try: + history = await _drive_contract_loop(client) + # Something was consumed, or the replay would prove nothing. + assert any(r.to_offset > r.from_offset for r in _consumed_ranges(history)) + + replayer = Replayer( + workflows=[ContractLoop], plugins=[provider], stream_client=client + ) + result = await replayer.replay_workflow(history) + assert result.replay_failure is None + finally: + await provider.close() + + +async def test_a_replayer_fails_a_tampered_range_as_nondeterministic() -> None: + """A history whose recorded ranges were changed replays on other input and fails. + + Every recorded range is emptied. The task that read two records and + published two decisions is replayed with nothing to read, so it issues no + publish, and Core finds the recorded publish event with no command for it. + """ + provider = NativeStreams() + client = await _connect(provider) + try: + history = await _drive_contract_loop(client) + tampered = _copy_of(history) + for consumed in _consumed_ranges(tampered): + consumed.to_offset = consumed.from_offset + + replayer = Replayer( + workflows=[ContractLoop], plugins=[provider], stream_client=client + ) + with pytest.raises(NondeterminismError): + await replayer.replay_workflow(tampered) + finally: + await provider.close() + + +async def test_a_replayer_fails_loudly_when_the_stream_is_gone() -> None: + """A range the stream service no longer serves is a ``StreamNotFoundError``, not a replay.""" + provider = NativeStreams() + client = await _connect(provider) + try: + history = await _drive_contract_loop(client) + # The same history under a workflow id that owns no stream: the + # records it consumed are nowhere to be fetched from. + gone = _copy_of(history, workflow_id="gone-wf-" + uuid.uuid4().hex[:8]) + + replayer = Replayer( + workflows=[ContractLoop], plugins=[provider], stream_client=client + ) + with pytest.raises(StreamNotFoundError, match="cannot be replayed"): + await replayer.replay_workflow(gone) + finally: + await provider.close() + + +async def test_a_replayer_without_a_client_says_what_it_needs() -> None: + """Without a stream client a consuming workflow's history is refused, with the remedy.""" + provider = NativeStreams() + client = await _connect(provider) + try: + history = await _drive_contract_loop(client) + replayer = Replayer(workflows=[ContractLoop], plugins=[provider]) + + with pytest.raises(RuntimeError, match="stream_client=") as raised: + await replayer.replay_workflow(history) + assert f"'{INPUTS}'" in str(raised.value) + + # The aggregating call reports it per run rather than aborting. + async def histories(): + yield history + + results = await replayer.replay_workflows( + histories(), raise_on_replay_failure=False + ) + assert isinstance(results.replay_failures[history.run_id], RuntimeError) + finally: + await provider.close() + + +async def test_a_replayer_fetches_a_standalone_stream_by_its_id() -> None: + """A subscribed name that is no stream of the workflow's is a standalone stream's id. + + The server resolves a subscription the same way, an owned stream by that + name first, so the replayer has to look in both places to hand the replay + the records the workflow actually read. + """ + client = await _connect() + streams = StreamClient.connect(TARGET or "") + task_queue = "replay-src-tq-" + uuid.uuid4().hex[:8] + workflow_id = "replay-src-wf-" + uuid.uuid4().hex[:8] + stream_id = "replay-src-" + uuid.uuid4().hex[:8] + tokens = ["a1", "a2", "b1"] + 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(tokens)], + id=workflow_id, + task_queue=task_queue, + ) + await streams.get(stream_id).append(*[_record(t.encode()) for t in tokens]) + assert await asyncio.wait_for(handle.result(), timeout=60) == tokens + history = await handle.fetch_history() + + result = await Replayer( + workflows=[ConsumeAcrossTasks], stream_client=client + ).replay_workflow(history) + assert result.replay_failure is None + finally: + await streams.close() + + +TWO_DECISIONS = [{"kind": "decision", "n": 1}, {"kind": "decision", "n": 2}] + + +async def _feed_two(client: Client, workflow_id: str) -> Any: + """Append two inputs and wait until both decisions are out, so their tasks have completed.""" + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + await take(stream.read(topic=DECISIONS, result_type=dict), 2, timeout=60) + return producer + + +async def test_a_query_against_a_cold_worker_answers_from_replayed_state() -> None: + """A query task built through matching carries the recorded ranges. + + With the cache off every task is a full replay. A query dispatched to a + worker that holds nothing has to rebuild the run from History, and the + records the run read are not in it, so the server attaches them to the + query task the way it does to a task after a cache miss. Without them the + replay would have nothing to read and the query would fail. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "query-cold-tq-" + uuid.uuid4().hex[:8] + workflow_id = "query-cold-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 + ) + producer = await _feed_two(client, workflow_id) + assert await handle.query(ContractLoop.trace) == TWO_DECISIONS + await producer.finish() + await asyncio.wait_for(handle.result(), 60) + finally: + await provider.close() + + +async def test_a_sticky_query_after_an_eviction_still_answers() -> None: + """A query sent to the sticky queue of a run the worker evicted still answers. + + The sticky task carries the history since the last task and no records, + which a worker that lost the run cannot use. Core sees that the history it + fetched itself records a consumed range it was sent no records for, and + lets the query go unanswered rather than answer it from the wrong state. + The server's sticky attempt then times out, stickiness is reset, and the + query is dispatched again on the normal queue, which carries the records. + The sticky timeout is shortened so the test does not wait the default out. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "query-sticky-tq-" + uuid.uuid4().hex[:8] + first_id = "query-sticky-a-" + uuid.uuid4().hex[:8] + second_id = "query-sticky-b-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[ContractLoop], + max_cached_workflows=1, + sticky_queue_schedule_to_start_timeout=timedelta(seconds=2), + ): + first = await client.start_workflow( + ContractLoop.run, id=first_id, task_queue=task_queue + ) + first_producer = await _feed_two(client, first_id) + # A second run on a one-slot cache pushes the first out of it. + second = await client.start_workflow( + ContractLoop.run, id=second_id, task_queue=task_queue + ) + second_producer = await _feed_two(client, second_id) + + answer = await first.query( + ContractLoop.trace, rpc_timeout=timedelta(seconds=60) + ) + assert answer == TWO_DECISIONS + + await first_producer.finish() + await second_producer.finish() + await asyncio.wait_for(first.result(), 60) + await asyncio.wait_for(second.result(), 60) + finally: + await provider.close() + + +async def _collect( + client: Client, workflow_id: str, topic: str, run_id: str | None +) -> list[tuple[Any, str, int]]: + """Every record on ``topic`` as ``(value, run, offset)``, the run and offset from the cursor.""" + handle = client.get_stream_handle(workflow_id, run_id=run_id) + out: list[tuple[Any, str, int]] = [] + async for record in handle.read(topic=topic, result_type=dict): + # native:: + _, run, offset = record.cursor.token.split(":") + out.append((record.value, run, int(offset))) + return out + + +async def test_a_reset_run_is_followed_and_replayed() -> None: + """A run reset from a consuming one carries its subscriptions on, and readers follow. + + The reset re-runs the task named by the reset point, so the base run's + history is copied up to that task and the ranges recorded in the copy are + what the reset run replays, from the base run's streams. The inherited + ``inputs`` stream is the reset run's own from the inherited cursor on: the + server seeds it with the input the re-run task had consumed, at the offset + it held, so the re-run task reads it again and the next input follows. The + ``decisions`` stream, which the base run only published to, starts at zero. + A handle without a run id follows the base run into the reset run and starts + each stream at the floor it reports; a handle pinned to the base run ends + with it; the ``Replayer`` fetches each era of the reset run's history from + the run whose stream holds it. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "reset-tq-" + uuid.uuid4().hex[:8] + workflow_id = "reset-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[ContractLoop]): + base = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=task_queue + ) + base_run = base.result_run_id + assert base_run + producer = await _feed_two(client, workflow_id) + # A third input in a task of its own, which is the task the reset + # re-runs: the reset run is seeded with it and consumes it again. + await producer.append({"n": 3}) + base_stream = client.get_stream_handle(workflow_id, run_id=base_run) + await take( + base_stream.read(topic=DECISIONS, result_type=dict), 3, timeout=60 + ) + + # Reset to the completion of the last consuming task, while the base + # run waits for more input. + completion_id = 0 + async for event in client.get_workflow_handle( + workflow_id, run_id=base_run + ).fetch_history_events(): + if event.HasField("workflow_task_completed_event_attributes"): + completed = event.workflow_task_completed_event_attributes + if any( + r.to_offset > r.from_offset + for r in completed.consumed_stream_ranges + ): + completion_id = event.event_id + assert completion_id + reset = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=base_run + ), + reason="re-run the last consuming task", + workflow_task_finish_event_id=completion_id, + request_id=uuid.uuid4().hex, + ) + ) + reset_run = reset.run_id + assert reset_run and reset_run != base_run + + # A fresh producer pins to the current run, the reset run, whose + # inherited inputs stream continues at the inherited offset. + continued = client.get_stream_handle(workflow_id).producer( + topic=INPUTS, producer_id="model2", attempt=1 + ) + await continued.append({"n": 4}) + await continued.finish() + trace = await asyncio.wait_for( + client.get_workflow_handle(workflow_id, run_id=reset_run).result(), + 60, + ) + # The first two decisions were replayed from the base run's stream; + # the third input came back with the task the reset re-ran. + assert trace == TWO_DECISIONS + [ + {"kind": "decision", "n": 3}, + {"kind": "decision", "n": 4}, + {"kind": "finish", "producer": "model2"}, + ] + + # The base run was terminated by the reset, and describe is the one + # place that names the run it was reset into. + described = await client.get_workflow_handle( + workflow_id, run_id=base_run + ).describe() + assert described.status == WorkflowExecutionStatus.TERMINATED + extended = described.raw_description.workflow_extended_info + assert extended.reset_run_id == reset_run + + # A reader pinned to the base run ends with the base run. + pinned = await asyncio.wait_for( + _collect(client, workflow_id, DECISIONS, base_run), 30 + ) + assert pinned == [ + ({"decided": 1}, base_run, 0), + ({"decided": 2}, base_run, 1), + ({"decided": 3}, base_run, 2), + ] + + # A chain-following reader crosses from the base run into the reset + # run on both topics, with no gap and no refusal: the published one + # restarts at zero, the inherited one continues at the floor. + decisions = await asyncio.wait_for( + _collect(client, workflow_id, DECISIONS, None), 30 + ) + assert decisions == [ + ({"decided": 1}, base_run, 0), + ({"decided": 2}, base_run, 1), + ({"decided": 3}, base_run, 2), + ({"decided": 3}, reset_run, 0), + ({"decided": 4}, reset_run, 1), + (None, reset_run, 2), + ] + # The seeded input sits at the offset it held in the base run. + inputs = await asyncio.wait_for( + _collect(client, workflow_id, INPUTS, None), 30 + ) + assert inputs == [ + ({"n": 1}, base_run, 0), + ({"n": 2}, base_run, 1), + ({"n": 3}, base_run, 2), + ({"n": 3}, reset_run, 2), + ({"n": 4}, reset_run, 3), + (None, reset_run, 4), + ] + latest = await client.get_stream_handle(workflow_id).latest(topic=INPUTS) + assert latest.token == f"native:{reset_run}:4" + + # Pinned to the reset run, BEGINNING is the floor its inherited + # stream starts at. Offset zero, which it never held, is refused. + from_floor = await asyncio.wait_for( + _collect(client, workflow_id, INPUTS, reset_run), 30 + ) + assert from_floor == [ + ({"n": 3}, reset_run, 2), + ({"n": 4}, reset_run, 3), + (None, reset_run, 4), + ] + + # The reset run's history: the base run's events, the reset marker + # naming both runs, then its own. The replayer fetches the first era + # from the base run's stream and the rest from the reset run's. + history = await client.get_workflow_handle( + workflow_id, run_id=reset_run + ).fetch_history() + reset_cause = WorkflowTaskFailedCause.WORKFLOW_TASK_FAILED_CAUSE_RESET_WORKFLOW + markers = [ + (failed.base_run_id, failed.new_run_id) + for event in history.events + if event.HasField("workflow_task_failed_event_attributes") + for failed in [event.workflow_task_failed_event_attributes] + if failed.cause == reset_cause + ] + assert markers == [(base_run, reset_run)] + result = await Replayer( + workflows=[ContractLoop], plugins=[provider], stream_client=client + ).replay_workflow(history) + assert result.replay_failure is None + + # Exported with its records, the reset run's history replays with no + # server: the slices come from both runs' streams. + bundle = await Replayer.fetch_stream_slices(client, history) + assert {s.run_id for s in bundle.stream_slices if s.records} == { + base_run, + reset_run, + } + restored = WorkflowHistory.from_json(workflow_id, bundle.to_json()) + offline = await Replayer( + workflows=[ContractLoop], plugins=[provider] + ).replay_workflow(restored) + assert offline.replay_failure is None + finally: + await provider.close() + + +async def test_an_exported_history_replays_offline_with_its_records() -> None: + """A history exported with its stream records is the whole replay input. + + ``fetch_stream_slices`` captures the records while the stream is retained, + ``to_json`` writes them beside the events, ``from_json`` reads them back, + and a replayer with no client replays the result. A plain export carries + none and is refused with both remedies named. + """ + provider = NativeStreams() + client = await _connect(provider) + try: + history = await _drive_contract_loop(client) + assert "streamSlices" not in history.to_json() + + bundle = await Replayer.fetch_stream_slices(client, history) + assert list(bundle.events) == list(history.events) + assert any(s.records for s in bundle.stream_slices) + text = bundle.to_json() + assert "streamSlices" in text + restored = WorkflowHistory.from_json(history.workflow_id, text) + assert list(restored.stream_slices) == list(bundle.stream_slices) + assert list(restored.events) == list(history.events) + + offline = Replayer(workflows=[ContractLoop], plugins=[provider]) + result = await offline.replay_workflow(restored) + assert result.replay_failure is None + + # With the records stripped it is a plain export again. + stripped = WorkflowHistory(restored.workflow_id, restored.events) + with pytest.raises(RuntimeError, match="stream_client=") as raised: + await offline.replay_workflow(stripped) + assert "fetch_stream_slices" in str(raised.value) + finally: + await provider.close() + + +async def test_a_chain_of_two_resets_is_followed_to_its_end() -> None: + """A run reset, and that run reset again: what a chain-following read walks. + + Resetting the same run twice is not a thing the server allows; the first + reset terminates it and the second is refused. So "reset more than once" + means a chain, base into A into B, and each run names only the one after + it. A chain-following read walks it one link at a time and stops at the + end, which is also the answer to whether it can loop: a reset always + makes a run that has never been reset itself. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "reset2-tq-" + uuid.uuid4().hex[:8] + workflow_id = "reset2-wf-" + uuid.uuid4().hex[:8] + + async def reset(run_id: str, reason: str) -> str: + # The last completed task of the run, whatever it did: a reset run + # picks up from there, and what matters here is the chain it makes. + completion_id = 0 + async for event in client.get_workflow_handle( + workflow_id, run_id=run_id + ).fetch_history_events(): + if event.HasField("workflow_task_completed_event_attributes"): + completion_id = event.event_id + assert completion_id + answer = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=run_id + ), + reason=reason, + workflow_task_finish_event_id=completion_id, + request_id=uuid.uuid4().hex, + ) + ) + return answer.run_id + + async def reset_run_of(run_id: str) -> str: + described = await client.get_workflow_handle( + workflow_id, run_id=run_id + ).describe() + return described.raw_description.workflow_extended_info.reset_run_id + + try: + async with Worker(client, task_queue=task_queue, workflows=[ContractLoop]): + base = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=task_queue + ) + base_run = base.result_run_id + assert base_run + await _feed_two(client, workflow_id) + base_stream = client.get_stream_handle(workflow_id, run_id=base_run) + await take( + base_stream.read(topic=DECISIONS, result_type=dict), 2, timeout=60 + ) + + first = await reset(base_run, "the first reset") + second = await reset(first, "the second reset") + assert len({base_run, first, second}) == 3 + + # Each link names the one after it and nothing else. The last has + # not been reset, so the walk ends there rather than going round. + assert await reset_run_of(base_run) == first + assert await reset_run_of(first) == second + assert not await reset_run_of(second) + + # Let the end of the chain finish so a read can reach an end. + carry_on = client.get_stream_handle(workflow_id).producer( + topic=INPUTS, producer_id="model2", attempt=1 + ) + await carry_on.append({"n": 4}) + await carry_on.finish() + await asyncio.wait_for( + client.get_workflow_handle(workflow_id, run_id=second).result(), 60 + ) + + # The walk itself: one link at a time, ending at the run that was + # not reset. A middle run may have published nothing, so this is + # asked of the provider rather than read off the records. + handle = client.get_stream_handle(workflow_id) + walked = [base_run] + while True: + onward = await handle._successor(DECISIONS, walked[-1]) # type: ignore[attr-defined] + if onward is None: + break + assert onward[0] not in walked, "the chain must not double back" + walked.append(onward[0]) + assert walked == [base_run, first, second] + + # And a chain-following read ends, having crossed the same runs + # and repeated none of them. + decisions = await asyncio.wait_for( + _collect(client, workflow_id, DECISIONS, None), 60 + ) + seen: list[str] = [] + for _, run, _ in decisions: + if not seen or seen[-1] != run: + seen.append(run) + assert seen == sorted(set(seen), key=walked.index) + assert seen[0] == base_run and seen[-1] == second + finally: + await provider.close() + + @workflow.defn class StartsWhenTold: """Subscribes to ``inputs`` at a start given by a signal and returns what it read.""" From 93f03ea697d190c35c66c666b40b1d7a1acf46b7 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:21:41 -0700 Subject: [PATCH 3/3] Added a changelog entry for server-side streams. --- CHANGELOG.md | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 8f8a2a95a..35129f595 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -66,6 +66,27 @@ to include examples, links to docs, or any other relevant information. converter, so a payload codec and external storage apply to them. `temporalio.streams.providers.memory.MemoryStreams` is the in-memory reference provider the conformance tests run against. +- **Experimental**: server-side streams. A workflow publishes to a stream it + owns with a command the server applies in its Workflow Task's commit, and + reads the ranges the server delivers on its Workflow Tasks, through + `temporalio.workflow.append_stream_records`, `subscribe_stream` and + `read_stream_records`. `temporalio.client_stream` and + `temporalio.contrib.server_streams` reach the same stream from outside a + workflow, and `temporalio.streams.providers.native.NativeStreams` puts it + behind the shared stream interface with one owned stream per topic. Requires + a server that serves the stream service. `Replayer(stream_client=)` replays a + workflow that read such a stream while the server still holds it: History + records only the offsets each task consumed, so the replayer fetches the + records from the stream service and hands them to the replay with the + history. A range the stream no longer holds fails the replay with + `StreamNotFoundError`. A handle without a run id follows a workflow reset as + it follows a continue-as-new, reading the reset run from the floor its stream + reports, and the replayer fetches the ranges recorded before a reset point + from the run the workflow was reset from. For offline replay, + `Replayer.fetch_stream_slices(client, history)` attaches the records to a + `WorkflowHistory` while the stream is retained, `to_json()` and `from_json()` + carry them as `streamSlices` beside the events, and a history that carries + them replays with no server. - **Experimental**: `temporalio.streams.providers.workflow_streams.WorkflowStreamsProvider` serves the stream interface over the shipped Workflow Streams transport as a worker plugin, so a workflow reads and publishes through