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 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, 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."""