diff --git a/CHANGELOG.md b/CHANGELOG.md index 5a234f1b1..4be5a64a9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -66,6 +66,16 @@ 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**: `temporalio.streams.providers.workflow_streams.WorkflowStreamsProvider` + serves the stream interface over the shipped Workflow Streams transport as a + worker plugin, so a workflow reads and publishes through + `temporalio.contrib.workflow_streams` without naming it. Records are the + `StreamRecord` proto inside the shipped item payload, and a handle without a + run id follows continue-as-new run by run and a reset into the run reset to. + An outside publish is an Update that answers with the batch's position and + refuses a conflicting repeat, falling back to the shipped Signal on a + workflow whose worker predates it. A workflow's activity keeps its own + streams in the workflow's log under `activity//`. ### Changed diff --git a/temporalio/contrib/workflow_streams/__init__.py b/temporalio/contrib/workflow_streams/__init__.py index 41f670f0c..c44bd25ea 100644 --- a/temporalio/contrib/workflow_streams/__init__.py +++ b/temporalio/contrib/workflow_streams/__init__.py @@ -13,12 +13,18 @@ """ from temporalio.contrib.workflow_streams._client import WorkflowStreamClient -from temporalio.contrib.workflow_streams._stream import WorkflowStream +from temporalio.contrib.workflow_streams._stream import ( + POLL_UPDATE_NAME, + PUBLISH_SIGNAL_NAME, + WorkflowStream, +) from temporalio.contrib.workflow_streams._topic_handle import ( TopicHandle, WorkflowTopicHandle, ) from temporalio.contrib.workflow_streams._types import ( + STREAM_DRAINING_ERROR_TYPE, + TRUNCATED_OFFSET_ERROR_TYPE, PollInput, PollResult, PublishEntry, @@ -29,6 +35,10 @@ ) __all__ = [ + "POLL_UPDATE_NAME", + "PUBLISH_SIGNAL_NAME", + "STREAM_DRAINING_ERROR_TYPE", + "TRUNCATED_OFFSET_ERROR_TYPE", "PollInput", "PollResult", "PublishEntry", diff --git a/temporalio/contrib/workflow_streams/_client.py b/temporalio/contrib/workflow_streams/_client.py index 605bf3f03..5c458bdab 100644 --- a/temporalio/contrib/workflow_streams/_client.py +++ b/temporalio/contrib/workflow_streams/_client.py @@ -325,6 +325,20 @@ def topic( self._topic_types[name] = bound return TopicHandle(self, name, bound) + @property + def handle(self) -> WorkflowHandle[Any, Any]: + """The workflow handle this client publishes to and polls. + + Re-targeted when :py:meth:`subscribe` follows a continue-as-new, so + read it when needed rather than caching it. + """ + return self._handle + + @property + def payload_converter(self) -> PayloadConverter: + """The sync payload converter used for per-item encode and decode.""" + return self._payload_converter() + async def flush(self) -> None: """Flush buffered (and pending) items and wait for server confirmation. diff --git a/temporalio/contrib/workflow_streams/_stream.py b/temporalio/contrib/workflow_streams/_stream.py index ae8608c3b..e1ebda15b 100644 --- a/temporalio/contrib/workflow_streams/_stream.py +++ b/temporalio/contrib/workflow_streams/_stream.py @@ -50,8 +50,22 @@ _WorkflowStreamWireItem, ) -_PUBLISH_SIGNAL = "__temporal_workflow_stream_publish" -_POLL_UPDATE = "__temporal_workflow_stream_poll" +PUBLISH_SIGNAL_NAME = "__temporal_workflow_stream_publish" +"""The signal :class:`WorkflowStream` registers for external publishes. + +Public so code that sends the signal itself, with its own publisher identity, +does not have to copy the name. +""" + +POLL_UPDATE_NAME = "__temporal_workflow_stream_poll" +"""The update :class:`WorkflowStream` registers for long polls. + +Public so code that drives the poll itself, with its own retry and +cancellation rules, does not have to copy the name. +""" + +_PUBLISH_SIGNAL = PUBLISH_SIGNAL_NAME +_POLL_UPDATE = POLL_UPDATE_NAME _OFFSET_QUERY = "__temporal_workflow_stream_offset" _MAX_POLL_RESPONSE_BYTES = 1_000_000 @@ -234,6 +248,25 @@ def topic( self._topic_types[name] = bound return WorkflowTopicHandle(self, name, bound) + @property + def next_offset(self) -> int: + """The global offset the next published item will receive.""" + return self._base_offset + len(self._log) + + def items_from(self, offset: int) -> list[tuple[int, str, Payload]]: + """Return ``(offset, topic, payload)`` for every item at or past ``offset``. + + Reads the log in place, so it is safe to call from a + :func:`temporalio.workflow.wait_condition` predicate. An ``offset`` + below the truncation base starts at the base instead; the offsets in + the result say where the items actually sit. + """ + start = max(offset, self._base_offset) - self._base_offset + return [ + (self._base_offset + index, item.topic, item.data) + for index, item in enumerate(self._log[start:], start) + ] + def get_state( self, *, publisher_ttl: timedelta = timedelta(seconds=900) ) -> WorkflowStreamState: diff --git a/temporalio/streams/providers/workflow_streams.py b/temporalio/streams/providers/workflow_streams.py new file mode 100644 index 000000000..8a5ed5daa --- /dev/null +++ b/temporalio/streams/providers/workflow_streams.py @@ -0,0 +1,1598 @@ +"""The provider over the shipped Workflow Streams transport (Option 0). + +Speaks the shipped contrib feature's wire format, the +``__temporal_workflow_stream_*`` Signal, Update and Query, so interface code +and existing Workflow Streams code share one log, and old histories replay. +Records live in the owning workflow's History, which is also this provider's +limit: the shipped caps (payloads in History, the Signal cap, bounded +subscribers) are transport properties and remain. + +The mapping, in one place: + +- A topic is the shipped topic of the same name in the workflow's one log. + A record rides as the item's ``Payload``: its data is the serialized + ``StreamRecord`` proto and its encoding is ``binary/plain``. The shipped + code stores and returns that ``Payload`` untouched, so the body's own + encoding never meets the transport. +- An outside publish is an Update, ``__temporal_streams_publish``, carrying + the shipped publish Signal's input. Producer identity dedupes through the + shipped publisher state, the publisher id being ``producer#attempt`` and + the sequence where the batch's records end, and the Update adds what a + Signal has no room for: a response. The workflow answers with the run and + offset the batch landed at, so ``append()`` returns a cursor, and it + refuses a repeat that carries *different* content at a sequence it holds, + or one behind its most recent, before accepting the Update, so the caller + gets :class:`temporalio.streams.StreamProducerError` and the log takes + nothing. Content is compared by a hash the workflow keeps per producer. + The Update's id is derived from producer, sequence and content, so a + retry after a lost reply is answered by the server from the first + outcome. The Action cost is unchanged: one Update per append batch in + place of one Signal. + + The Signal stays as the transport a worker that predates the Update + serves. A producer that is told twice, across a task boundary, that the + workflow has no publish Update falls back to it for the rest of its life + and returns ``None`` from ``append()``, since a Signal learns positions at + read time; ``publish_transport="signal"`` on the provider picks it from + the start. Records land in the same log either way, so a log written by + Signals reads the same. The shipped Signal handler's dedupe is one table + per publisher across topics, so on that transport one identity writing + two topics at the same sequence has its second batch dropped; the Update + keeps one per producer and topic, as the other providers do. +- A log belongs to one run and is not carried across continue-as-new, so a + cursor names the run as well as the offset. A handle without a run id + reads run after run: each log through the poll Update while its run is + open and through the tail Query once it has closed, then the successor's + from its first record. A run that was reset is followed too, into the run + describe names, at the position the read had reached: the reset run + rebuilt the base run's log up to the reset point by replay, so the items + before it sit at the same offsets. The reset itself is not reported to + the reader as a record; that is a wire change for a later round. +- The workflow-side stream object belongs to the workflow instance, found + through the handler the shipped class registers on it. An evicted and + rebuilt workflow gets its own, so a task that failed leaks nothing into + the next attempt's log and a replayed run does not see records twice. + The provider's handlers are registered as the instance is initialised, + before the first task's Signals and Updates are applied, so a publish or + poll that arrives with that task finds them; the stream object is bound + on first use, adopting one the workflow built in its ``__init__`` or + constructing the provider's own. +- An activity a workflow scheduled keeps its own streams inside that + workflow's log, under the reserved topic ``activity//`` + with ``%`` and ``/`` in the id percent-encoded, the way the native + provider reserves ``activity/`` in the owner's map. The record itself + carries the plain name. A standalone activity has no workflow to host a + log, so its streams are refused, and a workflow's own topic may not start + with the reserved prefix. A standalone stream, one with no owner at all, + is refused too: there is no workflow whose state could hold it. A housing + workflow per stream id is the design option for that, not built here. +- A record's body meets the client's data converter at the transport's + envelope rather than one record at a time. A batch travels as a Signal or + Update argument and comes back as an Update or Query result, and the SDK + runs the payload codec and external storage over those the way it does + over every payload it sends, off the workflow thread, so a batch above the + storage threshold is offloaded as a claim and a codec protects it in + History. The bodies inside are left as the payload converter produced + them, because the workflow thread reads them straight out of its state + and could not decode a codec's output or redeem a claim there. So the + worker's converter has to match the clients', as it does for every other + payload, and a client with a converter of its own cannot read another's + records. What each record does carry is the plaintext hash of its body + under ``temporal.io/content-hash``, stamped by the producer before the + envelope is encoded, and the workflow matches a repeated batch by those + hashes rather than by the bytes. +- A handle names its stream as a :class:`temporalio.streams.StreamRef` with + ``ref()``, a workflow's or an activity's, and ``close()`` refuses, since + an owned stream ends with its owner. +""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import logging +from collections.abc import AsyncGenerator +from dataclasses import dataclass +from datetime import timedelta +from typing import Any, Generic, Literal, TypeVar + +from google.protobuf.message import DecodeError + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.client import ( + Client, + WorkflowExecutionStatus, + WorkflowHandle, + WorkflowHistoryEventFilterType, + WorkflowQueryFailedError, + WorkflowUpdateFailedError, + WorkflowUpdateRPCTimeoutOrCancelledError, + WorkflowUpdateStage, +) +from temporalio.contrib.workflow_streams import ( + POLL_UPDATE_NAME, + PUBLISH_SIGNAL_NAME, + STREAM_DRAINING_ERROR_TYPE, + TRUNCATED_OFFSET_ERROR_TYPE, + PollInput, + PollResult, + PublishEntry, + PublishInput, + WorkflowStream, +) +from temporalio.converter import PayloadConverter +from temporalio.exceptions import ApplicationError +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import CONTENT_HASH_KEY, content_hash +from temporalio.streams._errors import ( + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, + StreamUnsupportedError, +) +from temporalio.streams._provider import ReadSource, WriteSink +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._ref import StreamRef +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import ( + RecordDecoder, + WireRecord, + cursor_position, + mint_cursor, + producer_identity, + to_wire, +) +from temporalio.streams.providers import ProviderPlugin +from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND + +__all__ = [ + "PRODUCER_CONFLICT_ERROR_TYPE", + "PublishTransport", + "WorkflowStreamsActivityHandle", + "WorkflowStreamsHandle", + "WorkflowStreamsProducer", + "WorkflowStreamsProvider", +] + +T = TypeVar("T") + +PublishTransport = Literal["update", "signal"] +"""How an outside producer's batches reach the workflow.""" + +PRODUCER_CONFLICT_ERROR_TYPE = "StreamProducerConflict" +"""The ``ApplicationError.type`` the publish Update refuses a producer conflict with. + +Public so a caller driving the Update itself can tell the refusal from any +other failure. +""" + +_PROVIDER = "workflow_streams" +_PUBLISH_UPDATE = "__temporal_streams_publish" +_TAIL_QUERY = "__temporal_streams_tail" +_LATEST_QUERY = "__temporal_streams_latest" +_START_QUERY = "__temporal_streams_start" +_ENCODING = b"binary/plain" +# The topics an activity owns live in its workflow's log under this prefix, +# so no workflow topic may start with it. +_ACTIVITY_PREFIX = "activity/" +# The same cap the shipped poll path answers under, because both are one +# response through the same server. +_MAX_TAIL_RESPONSE_BYTES = 1_000_000 +# The server's failure type for an accepted Update whose run closed before +# answering it: the poll's way of saying the run is over. +_UPDATE_OUTLIVED_RUN = "AcceptedUpdateCompletedWorkflow" +# The SDK's rejection of an Update no handler is registered for. It shares +# its wording with the Query one. +_HANDLER_NOT_FOUND = QUERY_HANDLER_NOT_FOUND + +logger = logging.getLogger(__name__) + + +def _run_is_closing(error: RPCError) -> bool: + """Whether the server refused an Update because the run is completing. + + The window between a run deciding to close, or continue as new, and its + close being recorded; the next attempt learns how it closed. + """ + return ( + error.status == RPCStatusCode.FAILED_PRECONDITION and "closing" in error.message + ) + + +def _require_topic(topic: str) -> None: + if not topic: + raise ValueError("topic must not be empty") + if topic.startswith(_ACTIVITY_PREFIX): + raise ValueError( + f"topic {topic!r} is reserved: names under {_ACTIVITY_PREFIX!r} hold the " + "streams of the workflow's activities on the workflow_streams provider" + ) + + +def _activity_topic(activity_id: str, topic: str) -> str: + """The reserved name ``topic`` of ``activity_id``'s streams takes in the log. + + The id is percent-encoded so an id holding ``/`` cannot be read as two + components; the name comes last, so it may hold anything. + """ + escaped = activity_id.replace("%", "%25").replace("/", "%2F") + return f"{_ACTIVITY_PREFIX}{escaped}/{topic}" + + +def _cursor(run_id: str, offset: int) -> Cursor: + return mint_cursor(_PROVIDER, f"{run_id}:{offset}") + + +def _position(after: Cursor) -> tuple[str, int] | None: + """The ``(run id, offset)`` a cursor of this provider names, or ``None`` for BEGINNING.""" + token = cursor_position(after, provider=_PROVIDER) + if token is None: + return None + run_id, _, offset = token.rpartition(":") + try: + if not run_id: + raise ValueError + return run_id, int(offset) + except ValueError: + raise StreamCursorError( + f"cursor {after.token!r} does not name a run and an offset on the " + "workflow_streams provider" + ) from None + + +def _wrap(record: WireRecord) -> Payload: + return Payload(metadata={"encoding": _ENCODING}, data=record.SerializeToString()) + + +def _unwrap(cursor: Cursor, payload: Payload, warn: Any) -> WireRecord | None: + try: + return WireRecord.FromString(payload.data) + except DecodeError as error: + # An item another publisher put on this topic, or a corrupt one: + # skip and say so, so one bad record cannot pin a reader. + warn("skipping stream record at %s: %s", cursor, error) + return None + + +def _entry_data(payload: Payload) -> str: + # The documented wire form of PublishEntry.data. + return base64.b64encode(payload.SerializeToString()).decode("ascii") + + +def _content(publish: PublishInput) -> str: + """A digest of one batch's content, length-delimited so a resplit cannot collide. + + Each record counts by the plaintext hash stamped on it under + ``CONTENT_HASH_KEY``, so a body's encoding never enters the identity a + repeat is matched by; a record without a body, such as ``FINISH``, counts + by its bytes. + """ + digest = hashlib.sha256() + for entry in publish.items: + record = WireRecord.FromString(_decode_payload(entry.data).data) + stamped = record.metadata.get(CONTENT_HASH_KEY) + identity = ( + stamped.data + if stamped is not None + else record.SerializeToString(deterministic=True) + ) + for part in (entry.topic.encode(), identity): + digest.update(len(part).to_bytes(8, "big")) + digest.update(part) + return digest.hexdigest() + + +def _producer_key(publish: PublishInput) -> tuple[str, str]: + # A batch is one topic's, so its first entry names the stream. + topic = publish.items[0].topic if publish.items else "" + return publish.publisher_id, topic + + +def _decode_payload(data: str) -> Payload: + return Payload.FromString(base64.b64decode(data)) + + +def _publish_id(publish: PublishInput, content: str) -> str: + # One id per (producer, sequence, content): a retry after a lost reply + # is answered from the first outcome without reaching the workflow, and a + # divergent repeat is a new Update the workflow gets to refuse. + key = f"{publish.publisher_id}\0{publish.sequence}\0{content}".encode() + return f"streams-publish-{hashlib.sha256(key).hexdigest()[:32]}" + + +@dataclass +class _PublishResult: + """The publish Update's answer: where the batch's last record sits.""" + + run_id: str + last_offset: int + + +@dataclass +class _Held: + """What the workflow keeps of a producer's most recent accepted batch.""" + + sequence: int + content: str + last_offset: int + + +class _Shipped: + """Constructs the shipped stream object, which insists on a caller named ``__init__``.""" + + def __init__(self) -> None: + self.stream = WorkflowStream() + + +class _InstanceStream: + """The provider's handlers on one workflow instance, over the shipped stream object. + + Built when the instance is initialised, before the workflow's own + ``__init__`` and before any Signal or Update of the first task is + applied, so the handlers are found by whatever arrives with that task. + The shipped stream object is bound on first use rather than here: a + workflow migrating from the contrib feature constructs its own in + ``__init__``, which the shipped class refuses to do twice, so this + adopts that one when it exists and constructs the provider's own when + nothing has, from a handler that may write or from the workflow half. + """ + + def __init__(self) -> None: + self._stream: WorkflowStream | None = None + # Per producer and topic, the most recent batch taken by the publish + # Update, so a repeat is answered and a divergent one refused. A + # producer's identity is one per stream, as on every provider, so + # the same identity on two topics is two producers. Workflow state, + # rebuilt on replay like the log itself. + self._producers: dict[tuple[str, str], _Held] = {} + if workflow.get_query_handler(_TAIL_QUERY) is None: + # The poll Update stops answering once the workflow is closing, + # and a reader between polls at that moment would lose what the + # final task published. The log is workflow state, so a Query + # still serves it after completion. + workflow.set_query_handler(_TAIL_QUERY, self._tail) + if workflow.get_query_handler(_LATEST_QUERY) is None: + workflow.set_query_handler(_LATEST_QUERY, self._latest) + if workflow.get_query_handler(_START_QUERY) is None: + workflow.set_query_handler(_START_QUERY, self._start) + if workflow.get_update_handler(_PUBLISH_UPDATE) is None: + workflow.set_update_handler( + _PUBLISH_UPDATE, self._publish, validator=self._validate_publish + ) + if workflow.get_update_handler(POLL_UPDATE_NAME) is None: + # Stands in for the shipped poll handler until a stream object is + # bound, whose constructor then registers the real one over it. + workflow.set_update_handler( + POLL_UPDATE_NAME, self._poll, validator=self._validate_poll + ) + + @property + def held(self) -> WorkflowStream | None: + """The shipped stream object, if the workflow or a handler has bound one. + + Looks one up and never constructs, so a Query may ask. + """ + if self._stream is None: + self._stream = _registered_stream() + return self._stream + + @property + def stream(self) -> WorkflowStream: + """The shipped stream object, adopting the workflow's own or constructing the provider's.""" + held = self.held + if held is None: + held = self._stream = _Shipped().stream + return held + + def _validate_poll(self, payload: PollInput) -> None: + held = self.held + if held is not None: + held._validate_poll(payload) # pyright: ignore[reportPrivateUsage] + + async def _poll(self, payload: PollInput) -> PollResult: + return await self.stream._on_poll(payload) # pyright: ignore[reportPrivateUsage] + + def _validate_publish(self, publish: PublishInput) -> None: + """Refuse a conflicting batch before the Update is accepted. + + Refused here rather than in the handler so the refusal writes no + event: the caller learns of it and the log is untouched. + + Raises: + ApplicationError: Typed ``PRODUCER_CONFLICT_ERROR_TYPE``. The + batch repeats this producer's most recent sequence with + other content, or names a sequence behind it. + """ + held = self._producers.get(_producer_key(publish)) + if held is None: + return + if publish.sequence < held.sequence: + raise ApplicationError( + f"producer {publish.publisher_id!r} already wrote past sequence " + f"{publish.sequence}; only its most recent append can be repeated", + type=PRODUCER_CONFLICT_ERROR_TYPE, + ) + if publish.sequence == held.sequence and _content(publish) != held.content: + raise ApplicationError( + f"producer {publish.publisher_id!r} repeated sequence " + f"{publish.sequence} with different content", + type=PRODUCER_CONFLICT_ERROR_TYPE, + ) + + def _publish(self, publish: PublishInput) -> _PublishResult: + """Take one outside batch into the log and answer where it landed. + + A repeat of the producer's most recent batch, which the validator + let through because its content matches, is answered with the + original position and writes nothing. The Update keeps its own + table rather than the shipped Signal handler's, which is one per + publisher across topics; a producer stays on one transport for its + life, so the two never judge the same batch. + """ + run_id = workflow.info().run_id + key = _producer_key(publish) + held = self._producers.get(key) + if held is not None and publish.sequence == held.sequence: + return _PublishResult(run_id=run_id, last_offset=held.last_offset) + stream = self.stream + for entry in publish.items: + stream.topic(entry.topic).publish(_decode_payload(entry.data)) + last = stream.next_offset - 1 + self._producers[key] = _Held(publish.sequence, _content(publish), last) + return _PublishResult(run_id=run_id, last_offset=last) + + def _tail(self, from_offset: int, topic: str) -> dict[str, Any]: + """One page of ``topic``'s items at or past ``from_offset``. + + Filtered and capped here rather than at the caller, because a Query + response has to fit the server's blob limit and a log the reader only + wants one topic of can be much larger than that. + """ + stream = self.held + if stream is None: + # Nothing has been bound, so nothing has been published. + return { + "items": [], + "next_offset": 0, + "more_ready": False, + "base_offset": 0, + } + items: list[dict[str, Any]] = [] + size = 0 + next_offset = stream.next_offset + more_ready = False + held = stream.items_from(from_offset) + for offset, item_topic, payload in held: + if item_topic != topic: + continue + data = base64.b64encode(payload.SerializeToString()).decode("ascii") + if items and size + len(data) > _MAX_TAIL_RESPONSE_BYTES: + next_offset, more_ready = offset, True + break + size += len(data) + items.append({"offset": offset, "topic": item_topic, "data": data}) + return { + "items": items, + "next_offset": next_offset, + "more_ready": more_ready, + # Where the retained log starts at or past ``from_offset``: a + # reader whose position is below it has fallen behind truncation. + "base_offset": held[0][0] if held else next_offset, + } + + def _latest(self, topic: str) -> int: + """The newest offset holding a record on ``topic``, or -1 when it has none. + + The log orders every topic together, so the head of the log is not an + answer about one topic. Scanning here costs one Query rather than + shipping the log to the caller to find the same thing. + """ + stream = self.held + if stream is None: + return -1 + for offset, item_topic, _ in reversed(stream.items_from(0)): + if item_topic == topic: + return offset + return -1 + + def _start(self, topic: str, last_n: int) -> int: + """Where a read starts: the log's head, or ``topic``'s newest ``last_n`` records. + + Zero ``last_n`` asks for the head, which is where ``END`` starts. + """ + stream = self.held + return 0 if stream is None else start_offset(stream, topic, last_n) + + +def start_offset(stream: WorkflowStream, topic: str, last_n: int) -> int: + """The log offset a read of ``topic`` starts at. + + ``last_n`` of zero is the head of the log, so only what is published next + is read. Otherwise it is the offset of ``topic``'s ``last_n``-th newest + item, or the oldest one the log holds when there are fewer. The log is + workflow state, so the workflow half answers the same on every replay. + """ + if last_n <= 0: + return stream.next_offset + items = stream.items_from(0) + seen = 0 + for offset, item_topic, _ in reversed(items): + if item_topic == topic: + seen += 1 + if seen == last_n: + return offset + return items[0][0] if items else stream.next_offset + + +def _registered_stream() -> WorkflowStream | None: + handler = workflow.get_signal_handler(PUBLISH_SIGNAL_NAME) + if handler is None: + return None + stream = getattr(handler, "__self__", None) + if not isinstance(stream, WorkflowStream): + raise RuntimeError( + f"the {PUBLISH_SIGNAL_NAME!r} signal on this workflow is handled by " + "something other than a WorkflowStream, so the workflow_streams " + "provider cannot share its log" + ) + return stream + + +def _instance() -> _InstanceStream: + # Found on the instance rather than in a process-level map keyed by run + # id: the SDK rebuilds an evicted workflow from history as a new object, + # and a map would hand that object the stale log with its unregistered + # handlers and the records of a task that failed. + handler = workflow.get_update_handler(_PUBLISH_UPDATE) + registered = getattr(handler, "__self__", None) + if isinstance(registered, _InstanceStream): + return registered + return _InstanceStream() + + +class _WSReadSource: + """Reads the signal-fed log the shipped feature keeps in workflow state.""" + + def __init__( + self, stream: WorkflowStream, topic: str, start: int, run_id: str + ) -> None: + self._stream = stream + self._topic = topic + self._offset = start + self._run_id = run_id + self._closed = False + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + while True: + if self._closed: + raise StopAsyncIteration + await workflow.wait_condition( + lambda: self._closed or self._stream.next_offset > self._offset + ) + if self._closed: + raise StopAsyncIteration + batch: list[tuple[Cursor, WireRecord]] = [] + for offset, topic, payload in self._stream.items_from(self._offset): + if topic != self._topic: + continue + cursor = _cursor(self._run_id, offset) + wire = _unwrap(cursor, payload, workflow.logger.warning) + if wire is not None: + batch.append((cursor, wire)) + self._offset = self._stream.next_offset + if batch: + return batch + + def close(self) -> None: + self._closed = True + + +class _WSWriteSink: + def __init__(self, stream: WorkflowStream, topic: str) -> None: + self._handle = stream.topic(topic) + + def publish(self, record: WireRecord) -> None: + # Appending to workflow state commits with the task, and a poll + # Update's result rides the same task completion, so a failed task + # leaks nothing: rule 1 through the shipped mechanics. + self._handle.publish(_wrap(record)) + + +class _WSWorkflowProvider: + """The workflow half: the shipped stream object of one workflow instance. + + Made per instance by the worker, so the stream object it captures on + first use is this instance's, and the finish hook lets go of that one + rather than of whatever the thread's handler lookup answers at the time. + """ + + def __init__(self) -> None: + self._stream: WorkflowStream | None = None + + def _own_stream(self) -> WorkflowStream: + if self._stream is None: + self._stream = _instance().stream + return self._stream + + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + check_read_start(after, last) + _require_topic(topic) + run_id = workflow.info().run_id + start = 0 + # The log is workflow state, so END and last= resolve against it here + # and land on the same offset on every replay. + if last is not None: + return _WSReadSource( + self._own_stream(), + topic, + start_offset(self._own_stream(), topic, last), + run_id, + ) + if after == END: + stream = self._own_stream() + return _WSReadSource(stream, topic, stream.next_offset, run_id) + named = _position(after) + if named is not None: + if named[0] != run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run; a run's log is its own" + ) + start = named[1] + 1 + return _WSReadSource(self._own_stream(), topic, start, run_id) + + def open_writer(self, topic: str) -> WriteSink: + _require_topic(topic) + return _WSWriteSink(self._own_stream(), topic) + + def on_workflow_start(self) -> None: + # Called as the instance is initialised, so the handlers exist before + # the first task's Updates are evaluated: an outside publish or poll + # that arrives with that task is served rather than rejected. The + # stream object itself is bound later, once the workflow's own + # __init__ has had its chance to construct one. + _instance() + + async def on_workflow_finish(self) -> None: + # An Option 0 stream dies with its run, and a parked long-poll Update + # would otherwise hold completion open. Same recipe the shipped + # feature documents before a return or a continue-as-new. Asked of + # the instance rather than this object's cache, because a poll can + # bind the stream before workflow code touches it. + held = _instance().held + if held is None: + return + held.detach_pollers() + await workflow.wait_condition(workflow.all_handlers_finished) + + +class WorkflowStreamsProducer(Generic[T]): + """Appends through the publish Update, or the shipped publish Signal. + + Direct rather than through ``WorkflowStreamClient`` because the interface + owns the publisher identity: it must be ``producer#attempt`` for the + shipped dedupe to drop a retry and pass a new generation, and the client + would use its own random id. + + Sequences are committed only after the server accepted the batch. A + batch whose send raised stays pending and goes out again under the same + sequence, either when the caller retries the same values or ahead of + whatever the caller sends next, so an ambiguous failure writes the batch + once and loses nothing. A batch the workflow refused is dropped from the + pending slot: the refusal is the answer. + """ + + def __init__( + self, + handle: Any, + converter: PayloadConverter, + topic: str, + producer_id: str, + attempt: int, + *, + item_topic: str | None = None, + transport: PublishTransport = "update", + retry_cooldown: timedelta = timedelta(milliseconds=100), + ) -> None: + """Bind this producer to ``topic`` on the workflow behind ``handle``. + + ``item_topic`` is the name the log files the records under when it + differs from the name the records carry, as an activity's reserved + topics do. ``transport`` is how batches travel; ``"update"`` falls + back to ``"signal"`` on a workflow that has no publish Update, after + one retry ``retry_cooldown`` later. + """ + self._handle = handle + self._converter = converter + self._topic = topic + self._item_topic = topic if item_topic is None else item_topic + self._producer_id = producer_id + self._attempt = attempt + self._transport: PublishTransport = transport + self._retry_cooldown = retry_cooldown + # One-based, because zero on the wire says the producer does not + # number its records and this one does. + self._sequence = 1 + self._pending: tuple[list[PublishEntry], int] | None = None + self._last: Cursor | None = None + + @property + def transport(self) -> PublishTransport: + """How this producer's batches travel now: the Signal once it fell back.""" + return self._transport + + @property + def producer_id(self) -> str: + """Who this producer writes as.""" + return self._producer_id + + @property + def attempt(self) -> int: + """The generation this producer is writing.""" + return self._attempt + + @property + def _publisher_id(self) -> str: + return ( + f"{self._producer_id}#{self._attempt}" + if self._attempt + else self._producer_id + ) + + async def append(self, *values: T) -> Cursor | None: + """Append ``values`` and return the cursor of the batch's last record. + + A repeat of this producer's most recent append with the same content + is answered with the position the original landed at; an empty call + returns the position of the last record this producer wrote, or + ``None`` before its first. ``None`` is also the answer on the Signal + transport, which learns positions at read time, so a caller that + wants to follow from now asks :meth:`WorkflowStreamsHandle.latest`. + + Raises: + StreamProducerError: The workflow holds this sequence with other + content, or the producer already wrote past it. + StreamNotFoundError: The workflow does not exist or has closed. + """ + if not values: + return self._last + return await self._send([(RecordKind.DATA, value) for value in values]) + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic.""" + await self._send([(RecordKind.FINISH, None)]) + + def _entries( + self, batch: list[tuple[RecordKind, Any]] + ) -> tuple[list[PublishEntry], int]: + sequence = self._sequence + entries = [] + for kind, value in batch: + wire = to_wire( + self._converter, + topic=self._topic, + kind=kind, + value=value, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=sequence, + ) + sequence += 1 + if wire.HasField("body"): + # The plaintext hash the workflow matches a repeat by, taken + # as the converter produced the body, before the transport's + # codec meets the envelope. + wire.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload( + metadata={"encoding": _ENCODING}, + data=content_hash(wire.body).encode(), + ) + ) + entries.append( + PublishEntry(topic=self._item_topic, data=_entry_data(_wrap(wire))) + ) + return entries, sequence + + async def _send(self, batch: list[tuple[RecordKind, Any]]) -> Cursor | None: + entries, next_sequence = self._entries(batch) + if self._pending is not None and self._pending[0] != entries: + # The caller moved on from a batch whose send raised. It goes + # first, under the sequence it already had, so a copy the server + # did accept is answered or dropped and one it never saw lands. + # The new batch is then renumbered behind it. + await self._deliver(*self._pending) + entries, next_sequence = self._entries(batch) + return await self._deliver(entries, next_sequence) + + async def _deliver( + self, entries: list[PublishEntry], next_sequence: int + ) -> Cursor | None: + # The dedupe sequence is where this producer's records end, not how + # many batches it has sent. The two differ once a retry batches its + # records differently from the send it is repeating, and a counter of + # batches then either drops a batch of new records or lets records + # that are already there through a second time. + self._pending = (entries, next_sequence) + publish = PublishInput( + items=entries, publisher_id=self._publisher_id, sequence=next_sequence + ) + if self._transport == "signal": + landed = await self._signal(publish) + else: + landed = await self._update(publish) + self._sequence = next_sequence + self._pending = None + if landed is not None: + self._last = landed + return landed + + async def _update(self, publish: PublishInput) -> Cursor | None: + content = _content(publish) + retried = False + while True: + try: + answer = await self._handle.execute_update( + _PUBLISH_UPDATE, + publish, + id=_publish_id(publish, content), + result_type=_PublishResult, + ) + except WorkflowUpdateFailedError as error: + cause = error.cause + cause_type = getattr(cause, "type", None) + if cause_type == PRODUCER_CONFLICT_ERROR_TYPE: + # Final for this batch: there is nothing to send again. + self._pending = None + raise StreamProducerError(str(cause)) from error + if cause_type == _UPDATE_OUTLIVED_RUN: + # Accepted, then the run closed before the handler ran: + # the same answer a Signal to a closed run gets. + raise StreamNotFoundError( + f"workflow {self._handle.id!r} closed before taking the " + "batch, so its stream cannot be appended to" + ) from error + if _HANDLER_NOT_FOUND not in str(cause): + raise StreamError( + f"the publish update on workflow {self._handle.id!r} " + f"failed: {error}" + ) from error + if not retried: + # A rejection comes back with the completion of the task + # that made it, so a retry reaches a later task, and the + # start hook registers the handler on the first one. + retried = True + await asyncio.sleep(self._retry_cooldown.total_seconds()) + continue + # Rejected across a task boundary: the workflow's worker + # predates the publish Update. Its Signal takes the batch, + # and every later one from this producer. + logger.info( + "workflow %r has no publish update; appending by Signal", + self._handle.id, + ) + self._transport = "signal" + return await self._signal(publish) + except RPCError as error: + if _run_is_closing(error): + # Where a Signal would be carried to the successor, the + # Update is refused; the retry reaches the run that + # takes over, since the handle follows the chain. + await asyncio.sleep(self._retry_cooldown.total_seconds()) + continue + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._handle.id!r} was not found, so its stream " + "cannot be appended to" + ) from error + raise + return _cursor(answer.run_id, answer.last_offset) + + async def _signal(self, publish: PublishInput) -> Cursor | None: + # Always None: a Signal has no response to carry the position. + try: + await self._handle.signal(PUBLISH_SIGNAL_NAME, publish) + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._handle.id!r} was not found, so its stream " + "cannot be appended to" + ) from error + raise + return None + + +class WorkflowStreamsHandle: + """One workflow's log from outside, through the shipped poll Update and a tail Query.""" + + def __init__( + self, + client: Client, + workflow_id: str, + run_id: str | None, + poll_cooldown: timedelta, + publish_transport: PublishTransport = "update", + ) -> None: + """Address ``workflow_id``'s log, pinned to ``run_id`` when one is given.""" + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._poll_cooldown = poll_cooldown + self._publish_transport: PublishTransport = publish_transport + self._converter = client.data_converter.payload_converter + + def read( + self, + *, + topic: str | StreamTopic[Any] | None = None, + after: Cursor = BEGINNING, + last: int | None = None, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` from where the read starts until the chain, or the pinned run, closes. + + ``BEGINNING`` is the oldest item the first retained run's log still + holds: the poll Update reads offset zero as the log's base. ``END`` + and ``last=`` start on the current run, or the pinned one, at an + offset its workflow answers by Query when the read starts. + + The chain is followed across continue-as-new, into the successor's + log from its first record, and across a reset, into the run describe + names at the position the read had reached. A handle pinned to a run + ends with that run either way. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + wire_topic = self._wire_topic(topic) + # Parsed here so a foreign cursor fails this call, not the first + # iteration of the generator. + named = None if after == END else _position(after) + if named is not None and self._run_id is not None and named[0] != self._run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run than this handle is pinned to" + ) + return self._read(wire_topic, named, after, last, result_type) + + def _wire_topic(self, topic: str) -> str: + """The name the log files ``topic`` under: the topic itself for a workflow's.""" + _require_topic(topic) + return topic + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + after: Cursor, + last: int | None, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + if named is not None: + run_id, offset = named[0], named[1] + 1 + elif last is not None or after == END: + run_id, offset = await self._start_on_current_run(topic, last or 0) + # A synthesized record is positioned before the first one read. + after = _cursor(run_id, offset - 1) + else: + run_id, offset = self._run_id or await self._first_run(), 0 + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + while True: + handle = self._handle(run_id) + next_offset = offset + polls = self._poll(handle, topic, offset) + try: + async for item_offset, payload in polls: + next_offset = item_offset + 1 + for record in self._records(decoder, run_id, item_offset, payload): + yield record + finally: + await polls.aclose() + status = await self._status(handle) + if status is None: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {run_id!r} was not found" + ) + if status == WorkflowExecutionStatus.RUNNING: + # The subscription ended early, on an RPC timeout for + # instance; the run is still open, so pick up where it left. + offset = next_offset + continue + # What landed after the last poll is still in workflow state, so + # the tail comes back by Query rather than being lost with the run. + tail_offset = next_offset + while True: + page, tail_offset, more = await self._tail(handle, topic, tail_offset) + for offset_, payload in page: + for record in self._records(decoder, run_id, offset_, payload): + yield record + if not more: + break + if self._run_id is not None: + return + following = await self._following(handle, status, tail_offset) + if following is None: + return + run_id, offset = following + + async def _following( + self, + handle: WorkflowHandle[Any, Any], + status: WorkflowExecutionStatus, + position: int, + ) -> tuple[str, int] | None: + """The run that carries on after ``handle``'s, and where to read it from. + + A continue-as-new names its successor in the close event, and the + successor's log starts over at zero. A reset does not: the base run is + closed with no word of it in its own History, and only describe names + the run reset from it. That run rebuilt its log by replaying the base + run's History up to the reset point, so the items before that point + sit at the same offsets in both logs, and the read carries on at the + position it reached. Past the reset point the two logs differ, and a + reader already there is not told: reporting the reset to consumers as + a record is a wire change for a later round. + """ + if status == WorkflowExecutionStatus.CONTINUED_AS_NEW: + successor = await self._successor(handle) + return None if successor is None else (successor, 0) + reset_run = await self._reset_run(handle) + return None if reset_run is None else (reset_run, position) + + async def _reset_run(self, handle: WorkflowHandle[Any, Any]) -> str | None: + """The run ``handle``'s run was reset into, which only describe reports.""" + description = await self._describe(handle) + if description is None: + return None + extended = description.raw_description.workflow_extended_info + return extended.reset_run_id or None + + async def _poll( + self, handle: WorkflowHandle[Any, Any], topic: str, offset: int + ) -> AsyncGenerator[tuple[int, Payload], None]: + """Drive the shipped poll Update on one run until that run closes. + + Written here rather than through ``WorkflowStreamClient.subscribe`` + for two reasons. That loop swallows a cancellation of the caller's + task, so a consumer's ``asyncio.timeout`` or task cancel around a read + would end the subscription and let the read resubscribe forever; here + the cancellation leaves ``read()`` as what it was. And it follows + continue-as-new with offsets this provider does not carry across + runs, which is the outer loop's job. + """ + cooldown = self._poll_cooldown.total_seconds() + unhandled = False + while True: + try: + update = await handle.start_update( + POLL_UPDATE_NAME, + PollInput(topics=[topic], from_offset=offset), + wait_for_stage=WorkflowUpdateStage.ACCEPTED, + result_type=PollResult, + ) + result = await update.result() + except WorkflowUpdateRPCTimeoutOrCancelledError as error: + if isinstance(error.__cause__, asyncio.CancelledError): + # The SDK wraps a cancelled await in this error; the + # consumer cancelled us, so that is what comes out. + raise asyncio.CancelledError() from error + # The RPC itself timed out while the run is still open. + continue + except WorkflowUpdateFailedError as error: + cause = getattr(error.cause, "type", None) + if cause == TRUNCATED_OFFSET_ERROR_TYPE: + # Restarting from the beginning would hand the caller + # records it already handled, and only the caller can + # decide to do that. + raise StreamCursorError( + f"offset {offset} of workflow {self._workflow_id!r} run " + f"{handle.run_id!r} is no longer retained" + ) from error + if cause == STREAM_DRAINING_ERROR_TYPE: + # Pollers are detached because the run is closing; the + # next attempt learns how it closed. + await asyncio.sleep(cooldown) + continue + if cause == _UPDATE_OUTLIVED_RUN: + return + if _HANDLER_NOT_FOUND in str(error.cause): + if await self._status(handle) != WorkflowExecutionStatus.RUNNING: + # The run closed with the rejecting task, as a run + # that continues as new on its first task does; the + # caller describes it and follows the chain. + return + if unhandled: + raise StreamError( + f"workflow {self._workflow_id!r} run {handle.run_id!r} " + "does not serve the poll update: the workflow_streams " + "provider is not installed on its worker, and no " + "WorkflowStream was constructed" + ) from error + # A rejection comes back with the completion of the task + # that made it, so a retry reaches a later task, and the + # start hook registers the handler on the first one. Only + # a second rejection means there is no handler to wait for. + unhandled = True + await asyncio.sleep(cooldown) + continue + raise StreamError( + f"the poll update on workflow {self._workflow_id!r} failed: {error}" + ) from error + except RPCError as error: + if _run_is_closing(error): + # Continuing as new, or completing: the next attempt + # learns how it closed, the way a draining poll does. + await asyncio.sleep(cooldown) + continue + # The run closed and its poll Update went with it, or the + # workflow does not exist; the caller describes to tell which. + if error.status != RPCStatusCode.NOT_FOUND: + raise + return + for item in result.items: + if item.topic == topic: + yield item.offset, Payload.FromString(base64.b64decode(item.data)) + offset = result.next_offset + if not result.more_ready and cooldown > 0: + await asyncio.sleep(cooldown) + + def _records( + self, decoder: RecordDecoder, run_id: str, offset: int, payload: Payload + ) -> list[StreamRecord[Any]]: + cursor = _cursor(run_id, offset) + wire = _unwrap(cursor, payload, logger.warning) + if wire is None: + return [] + return decoder.decode(cursor, wire) + + def _handle(self, run_id: str | None) -> WorkflowHandle[Any, Any]: + return self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + + async def _describe(self, handle: WorkflowHandle[Any, Any]) -> Any | None: + """The description of ``handle``'s run, or ``None`` when it is gone.""" + try: + return await handle.describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + return None + raise + + async def _status( + self, handle: WorkflowHandle[Any, Any] + ) -> WorkflowExecutionStatus | None: + description = await self._describe(handle) + return None if description is None else description.status + + async def _first_run(self) -> str: + """The oldest retained run of the chain, walking back from the latest.""" + try: + run_id = (await self._handle(None).describe()).run_id + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + assert run_id is not None + while True: + previous = await self._predecessor(run_id) + if previous is None: + return run_id + run_id = previous + + async def _predecessor(self, run_id: str) -> str | None: + try: + async for event in self._handle(run_id).fetch_history_events(page_size=1): + attributes = event.workflow_execution_started_event_attributes + if attributes.continued_execution_run_id: + return attributes.continued_execution_run_id + # A reset run's start event is the base run's, copied, and the + # original run id it carries is kept across resets, so it names + # the run the chain of resets began from. + original = attributes.original_execution_run_id + return original if original and original != run_id else None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + # The run's History is gone: the chain's retained part starts here. + return None + + async def _successor(self, handle: WorkflowHandle[Any, Any]) -> str | None: + events = handle.fetch_history_events( + event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT + ) + try: + async for event in events: + if event.HasField( + "workflow_execution_continued_as_new_event_attributes" + ): + attributes = ( + event.workflow_execution_continued_as_new_event_attributes + ) + return attributes.new_execution_run_id or None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {handle.run_id!r} was not found, " + "so its successor cannot be followed" + ) from error + return None + + async def _tail( + self, handle: WorkflowHandle[Any, Any], topic: str, from_offset: int + ) -> tuple[list[tuple[int, Payload]], int, bool]: + """One page of ``topic``'s tail as ``(items, next offset, more to come)``.""" + try: + wire = await handle.query( + _TAIL_QUERY, args=[from_offset, topic], result_type=dict + ) + except WorkflowQueryFailedError as error: + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the tail query on workflow {self._workflow_id!r} failed: {error}" + ) from error + # The workflow never opened a stream through this provider, so + # there is no tail to serve. + return [], from_offset, False + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + # The History is gone; nothing is left to serve. + return [], from_offset, False + if from_offset and wire.get("base_offset", from_offset) > from_offset: + # Restarting from the base would hand the caller records it + # already handled, and only the caller can decide to do that. + raise StreamCursorError( + f"offset {from_offset} of workflow {self._workflow_id!r} run " + f"{handle.run_id!r} is no longer retained" + ) + items = [ + (item["offset"], Payload.FromString(base64.b64decode(item["data"]))) + for item in wire["items"] + ] + return items, wire["next_offset"], bool(wire["more_ready"]) + + async def _start_on_current_run(self, topic: str, last_n: int) -> tuple[str, int]: + """The run a read at ``END`` or of the newest records starts on, and the offset. + + Raises: + StreamUnsupportedError: The workflow's worker runs a provider that + predates these reads and cannot answer where they start. + """ + handle = self._handle(self._run_id) + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") + run_id = description.run_id + assert run_id is not None + try: + offset = await self._handle(run_id).query( + _START_QUERY, args=[topic, last_n], result_type=int + ) + except WorkflowQueryFailedError as error: + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the start query on workflow {self._workflow_id!r} failed: {error}" + ) from error + raise StreamUnsupportedError( + f"workflow {self._workflow_id!r} does not answer where a read at END " + "or of the last records starts; its worker's provider predates them" + ) from error + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + return run_id, offset + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The newest position holding a record on ``topic``. + + The log is one per run, so the cursor names the run it was read from: + the pinned run, or the latest run of the chain. One log orders every + topic, so the answer comes from a Query that scans it for this topic + rather than from the head of the log, which usually names some other + topic's record. + """ + topic, _ = resolve_topic(topic) + wire_topic = self._wire_topic(topic) + handle = self._handle(self._run_id) + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") + try: + head = await handle.query(_LATEST_QUERY, wire_topic, result_type=int) + except WorkflowQueryFailedError as error: + if QUERY_HANDLER_NOT_FOUND not in str(error): + raise StreamError( + f"the latest query on workflow {self._workflow_id!r} failed: " + f"{error}" + ) from error + # The workflow never opened a stream through this provider, so + # the topic holds nothing. + head = -1 + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + run_id = description.run_id + assert run_id is not None + if head >= 0: + return _cursor(run_id, head) + # An empty topic on the chain's first run is the beginning of the + # stream; on a successor it is a position of its own, because + # BEGINNING would send a chain-following read back to the first run. + if await self._predecessor(run_id) is None: + return BEGINNING + return _cursor(run_id, -1) + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> WorkflowStreamsProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + wire_topic = self._wire_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return WorkflowStreamsProducer( + self._handle(self._run_id), + self._converter, + topic, + producer_id, + attempt, + item_topic=wire_topic, + transport=self._publish_transport, + retry_cooldown=self._poll_cooldown, + ) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this workflow's stream, pinned as this handle is.""" + return StreamRef.for_workflow( + self._workflow_id, run_id=self._run_id, topic=topic + ) + + async def close(self) -> None: + """Refuse: an owned stream ends with its owner, not by a caller.""" + raise ValueError( + "only a standalone stream can be closed; this handle is on an owned " + "stream, which ends when its workflow or activity does" + ) + + +class WorkflowStreamsActivityHandle(WorkflowStreamsHandle): + """The streams one activity of a workflow owns, kept in that workflow's log. + + Each topic is the reserved topic ``activity//`` of the + workflow's log, so an activity's ``tokens`` and its workflow's ``tokens`` + are two streams. The activity belongs to one run, so the handle pins the + run on first use and never follows a successor. A read ends when the + workflow closes, or once the activity has been seen pending, is pending + no longer and the read delivered a record: a stream the activity never + wrote has nothing to close, so a read on it waits for the workflow. + + A read here polls the tail Query and describes the workflow between + polls, every ``poll_cooldown``, rather than parking on the poll Update: + an Update parked in the workflow returns only when the log grows, so it + could not notice the activity ending, and each abandoned one would stay + parked against the run's Update caps. + """ + + def __init__( + self, + client: Client, + workflow_id: str, + run_id: str | None, + activity_id: str, + poll_cooldown: timedelta, + publish_transport: PublishTransport = "update", + ) -> None: + """Address ``activity_id``'s streams inside ``workflow_id``'s log.""" + super().__init__(client, workflow_id, run_id, poll_cooldown, publish_transport) + self._activity_id = activity_id + + def _wire_topic(self, topic: str) -> str: + return _activity_topic(self._activity_id, topic) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this activity's streams, through its workflow.""" + return StreamRef.for_activity( + self._activity_id, + workflow_id=self._workflow_id, + run_id=self._run_id, + topic=topic, + ) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + after: Cursor, + last: int | None, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + if named is not None: + run_id, offset = named[0], named[1] + 1 + elif last is not None or after == END: + run_id, offset = await self._start_on_current_run(topic, last or 0) + after = _cursor(run_id, offset - 1) + else: + run_id, offset = await self._current_run(), 0 + handle = self._handle(run_id) + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + cooldown = self._poll_cooldown.total_seconds() + seen_pending = False + delivered = False + ended = False + while True: + more = True + while more: + page, offset, more = await self._tail(handle, topic, offset) + for offset_, payload in page: + for record in self._records(decoder, run_id, offset_, payload): + delivered = True + yield record + if ended: + return + description = await self._describe(handle) + if description is None: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} run {run_id!r} was not found" + ) + pending = any( + info.activity_id == self._activity_id + for info in description.raw_description.pending_activities + ) + seen_pending = seen_pending or pending + # One more pass after learning the stream ended, so a record that + # landed between the tail and the describe is not lost. + ended = description.status != WorkflowExecutionStatus.RUNNING or ( + seen_pending and not pending and delivered + ) + if not ended: + await asyncio.sleep(cooldown) + + async def _current_run(self) -> str: + description = await self._describe(self._handle(self._run_id)) + if description is None: + raise StreamNotFoundError(f"workflow {self._workflow_id!r} was not found") + assert description.run_id is not None + return description.run_id + + +class WorkflowStreamsProvider(ProviderPlugin): + """The provider over the shipped Workflow Streams transport.""" + + def __init__( + self, + *, + poll_cooldown: timedelta = timedelta(milliseconds=100), + publish_transport: PublishTransport = "update", + ) -> None: + """Create the provider. + + Args: + poll_cooldown: How long an outside reader that is caught up waits + between polls. Backlogs drain at full speed regardless. Also + how long a producer waits before retrying a publish Update + the workflow's first task rejected. + publish_transport: How an outside producer's batches reach the + workflow. ``"update"``, the default, answers each batch with + its position and refuses a conflicting repeat; on a workflow + whose worker predates the publish Update the producer falls + back to the Signal by itself. ``"signal"`` is the transport + of the first release and skips the detection. Both cost one + Action per batch. + """ + self._poll_cooldown = poll_cooldown + self._publish_transport: PublishTransport = publish_transport + + def workflow_provider(self) -> _WSWorkflowProvider: + """The workflow half, over the running instance's shipped stream object.""" + return _WSWorkflowProvider() + + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> WorkflowStreamsHandle: + """A handle on ``workflow_id``'s log; without ``run_id`` it follows the chain.""" + return WorkflowStreamsHandle( + client, workflow_id, run_id, self._poll_cooldown, self._publish_transport + ) + + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> WorkflowStreamsActivityHandle: + """A handle on the streams an activity of ``workflow_id`` owns. + + They live in the workflow's log under the reserved topics + ``activity//``, and ``run_id`` pins the workflow's + run. Without ``workflow_id`` the activity is a standalone one, which + has no workflow to host a log, so this provider refuses it; such an + activity addresses a workflow's stream by ``workflow_id`` instead. + + Raises: + StreamUnsupportedError: ``workflow_id`` was not given. + """ + if workflow_id is None: + raise StreamUnsupportedError( + "the workflow_streams provider cannot hold a stream a standalone " + "activity owns: its log lives inside a running workflow, and no " + "workflow hosts this activity's" + ) + return WorkflowStreamsActivityHandle( + client, + workflow_id, + run_id, + activity_id, + self._poll_cooldown, + self._publish_transport, + ) + + async def create_standalone_stream( + self, + client: Client, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> WorkflowStreamsHandle: + """Refused: this provider has no store for a stream without an owner. + + Every log here is a running workflow's state, served by that + workflow's handlers, and a standalone stream has no workflow. A + housing workflow started for the stream id, with the retention policy + as its state and ``close()`` as a Signal, would be one way to offer + it on this transport; it is a design option for a later round, not + built here. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the workflow_streams provider does not host standalone streams: every " + "log is a running workflow's state, and a stream without an owner has " + "no workflow" + ) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> WorkflowStreamsHandle: + """Refused: this provider has no store for a stream without an owner. + + See :meth:`create_standalone_stream`. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the workflow_streams provider does not host standalone streams: every " + "log is a running workflow's state, and a stream without an owner has " + "no workflow" + ) + + async def close(self) -> None: + """Nothing to release: the provider holds no connection of its own.""" diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 13f2a1490..c53ea6f88 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -96,6 +96,13 @@ logger = logging.getLogger(__name__) +QUERY_HANDLER_NOT_FOUND = "expected but not found" +"""The phrase a query for an unregistered handler comes back with. + +A caller that has to recognise the condition has only the failure message to +go on, so it matches this constant rather than a copy of the sentence. +""" + # Set to true to log all cases where we're ignoring things during delete LOG_IGNORE_DURING_DELETE = False @@ -836,7 +843,8 @@ async def run_query() -> None: if not defn: known_queries = sorted([k for k in self._queries.keys() if k]) raise RuntimeError( - f"Query handler for '{job.query_type}' expected but not found, " + f"Query handler for '{job.query_type}' " + f"{QUERY_HANDLER_NOT_FOUND}, " f"known queries: [{' '.join(known_queries)}]" ) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 079cca975..793bf2acd 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -17,6 +17,16 @@ def pytest_configure(config: pytest.Config) -> None: "markers", "truncates: the case needs a way to drop a topic's oldest records", ) + config.addinivalue_line( + "markers", + "encodes_bodies: the case needs the outside path to run each body through " + "the client's data converter", + ) + config.addinivalue_line( + "markers", + "standalone_activities: the case needs the streams of an activity outside " + "any workflow", + ) config.addinivalue_line( "markers", "hosts_standalone_streams: the case needs a stream with an id of its own and " diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index bd0a675ac..61eaf7fea 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -7,10 +7,13 @@ activity execution, so a retry writes to the same stream and a reader sees the attempt change as ``SUPERSEDED``. -The memory provider always runs. A storage provider adds itself to -``SETUPS`` behind its own ``STREAMS_LIVE`` gate: its setup receives the -environment's client and hands back the provider and a client with it -registered, which the workers and the reads in these cases share. +The memory provider always runs, and so does Workflow Streams, whose store is +the workflow's own History. A storage provider adds itself to ``SETUPS`` +behind its own ``STREAMS_LIVE`` gate: its setup receives the environment's +client and hands back the provider and a client with it registered, which the +workers and the reads in these cases share, and says whether it holds the +streams of a standalone activity, so the cases marked +``standalone_activities`` are skipped with a reason where it does not. """ from __future__ import annotations @@ -34,6 +37,7 @@ topic, ) from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider from temporalio.testing import WorkflowEnvironment from tests.helpers import new_worker @@ -47,6 +51,8 @@ class ActivitySetup: name: str provider: StreamProvider client: Client + standalone_activities: bool = True + """The provider holds the streams of an activity outside any workflow.""" async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: @@ -57,8 +63,20 @@ async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider.reset() +async def _workflow_streams_setup(client: Client) -> AsyncIterator[ActivitySetup]: + provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + config = client.config() + config["plugins"] = [provider] + # An activity's streams live in its workflow's log, so a standalone + # activity has nowhere to put them. See the provider's module docstring. + yield ActivitySetup( + "workflow_streams", provider, Client(**config), standalone_activities=False + ) + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ActivitySetup]]] = { - "memory": _memory_setup + "memory": _memory_setup, + "workflow_streams": _workflow_streams_setup, } @@ -69,6 +87,13 @@ async def setup( if env.supports_time_skipping: pytest.skip("the time-skipping test server has no standalone activities") async for found in SETUPS[request.param](client): + if ( + request.node.get_closest_marker("standalone_activities") + and not found.standalone_activities + ): + pytest.skip( + f"the {found.name} provider does not hold a standalone activity's streams" + ) yield found @@ -167,6 +192,7 @@ async def test_scope_activity_gives_a_workflow_activity_its_own_streams( assert await workflow_stream.latest(topic=TOKENS) == BEGINNING +@pytest.mark.standalone_activities async def test_standalone_activity_defaults_to_its_own_stream(setup: ActivitySetup): client = setup.client activity_id = f"streams-saa-{uuid.uuid4().hex}" @@ -199,6 +225,7 @@ async def fail_once_after_writing() -> None: await producer.finish() +@pytest.mark.standalone_activities async def test_a_retry_inherits_the_stream_and_supersedes(setup: ActivitySetup): client = setup.client activity_id = f"streams-saa-retry-{uuid.uuid4().hex}" diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index f9c53bdcf..993cba69b 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -36,8 +36,9 @@ import pytest +from temporalio import workflow from temporalio.api.common.v1 import Payload -from temporalio.client import Client +from temporalio.client import Client, WorkflowHandle from temporalio.common import RawValue from temporalio.converter import ( DataConverter, @@ -68,7 +69,10 @@ topic, ) from temporalio.streams._ref import open_ref +from temporalio.streams.providers import workflow_streams from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider +from tests.helpers import new_worker # Defined once and shared by every case, the way an application shares them # between its workflow, its activities and its backend. @@ -90,8 +94,13 @@ class ProviderCase: """``append()`` returns where the records landed.""" detects_divergent_retries: bool = True """``append()`` compares a repeat's content with what it already holds.""" + encodes_bodies: bool = True + """The outside path runs each body through the client's data converter, so a + client with a converter of its own reads and writes another's records.""" host: Callable[[str], Awaitable[None]] | None = None """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" + task_queue: str | None = None + """Where the setup's worker runs the workflows below, when it has one.""" truncate: Callable[[str, str, int], Awaitable[None]] | None = None """Drops all but the newest records of a workflow's topic, standing in for retention, or ``None`` when the provider offers no way to.""" @@ -246,13 +255,98 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: provider.reset() +@workflow.defn +class TruncatingStreamHost: + """A stream host whose log an update can truncate, the way a workflow's retention would.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.update + def truncate(self, topic: str, keep: int) -> None: + stream = workflow_streams._instance().stream # pyright: ignore[reportPrivateUsage] + stream.truncate(workflow_streams.start_offset(stream, topic, keep)) + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +@workflow.defn +class DefaultTopicAnswer: + """Reads one value on its default topic and answers on the same topic.""" + + @workflow.run + async def run(self) -> None: + reader = workflow.stream_reader(result_type=dict) + async for value in reader.values(): + workflow.stream_writer().publish({"answer": value["n"] * 2}) + reader.close() + + +async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: + # No STREAMS_LIVE gate: the store is the workflow's own History, which the + # test environment's server provides. + provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + # Registered once, on the client: the host's worker inherits it and the + # cases open handles through client.get_stream_handle. + config = client.config() + config["plugins"] = [provider] + client = Client(**config) + hosts: dict[str, WorkflowHandle[Any, Any]] = {} + async with new_worker(client, TruncatingStreamHost, DefaultTopicAnswer) as worker: + + async def host(workflow_id: str) -> None: + if workflow_id not in hosts: + hosts[workflow_id] = await client.start_workflow( + TruncatingStreamHost.run, + id=workflow_id, + task_queue=worker.task_queue, + ) + + async def truncate(workflow_id: str, topic: str, keep: int) -> None: + await hosts[workflow_id].execute_update( + TruncatingStreamHost.truncate, args=[topic, keep] + ) + + # A publish is an Update, so the workflow answers with the position + # and refuses a divergent repeat: both capabilities hold here. Bodies + # meet the codec and external storage at the transport's envelope, + # which the worker's converter has to match, so a case whose client + # carries a converter of its own is skipped; see the provider's + # module docstring. + yield ProviderCase( + "workflow_streams", + provider, + client, + encodes_bodies=False, + host=host, + task_queue=worker.task_queue, + truncate=truncate, + # Every log is a running workflow's state; a stream with no owner + # has no workflow to live in, so neither of its policies exists. + hosts_standalone_streams=False, + bounds_standalone_bytes=False, + trims_open_stream_by_age=False, + refuses_appends_past_byte_cap=False, + ) + for handle in hosts.values(): + await handle.terminate() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { - "memory": _memory_case + "memory": _memory_case, + "workflow_streams": _workflow_streams_case, } _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, "detects_divergent_retries": lambda case: case.detects_divergent_retries, + "encodes_bodies": lambda case: case.encodes_bodies, "truncates": lambda case: case.truncate is not None, "hosts_standalone_streams": lambda case: case.hosts_standalone_streams, } @@ -482,6 +576,29 @@ async def test_naming_no_topic_addresses_the_default_topic(case: ProviderCase): assert await stream.latest(topic=DEFAULT_TOPIC) == records[0].cursor +async def test_a_workflow_answers_on_its_default_topic(case: ProviderCase): + if case.client is None or case.task_queue is None: + pytest.skip( + f"the {case.name} setup runs no worker; test_streams_workflow covers " + "its workflow half" + ) + workflow_id = new_workflow_id() + handle = await case.client.start_workflow( + DefaultTopicAnswer.run, id=workflow_id, task_queue=case.task_queue + ) + stream = case.client.get_stream_handle(workflow_id) + await stream.producer(producer_id="client", attempt=1).append({"n": 21}) + await handle.result() + + # The outside producer and the workflow, each naming no topic, meet on + # one topic that a reader naming none sees in order. + records = await take(stream.read(), 2, timeout=30.0) + assert [(r.topic, r.producer_id, r.value) for r in records] == [ + (DEFAULT_TOPIC, "client", {"n": 21}), + (DEFAULT_TOPIC, "", {"answer": 42}), + ] + + async def test_cursor_resumes_where_it_points(case: ProviderCase): workflow_id = new_workflow_id() stream = await case.open(workflow_id) @@ -643,8 +760,10 @@ async def test_beginning_starts_at_the_oldest_record_still_held(case: ProviderCa assert [r.value for r in records] == [{"n": 3}, {"n": 4}] newest = await take(stream.read(topic=OUT, last=3), 2) assert [r.value for r in newest] == [{"n": 3}, {"n": 4}] + # A cursor below the floor is refused. A provider that needs a round trip + # to know says so on the first step rather than on the call. with pytest.raises(StreamCursorError): - stream.read(topic=OUT, after=before[0].cursor) + await take(stream.read(topic=OUT, after=before[0].cursor), 1) async def test_a_read_start_names_one_place(case: ProviderCase): @@ -694,6 +813,7 @@ async def test_an_owned_stream_cannot_be_closed_by_a_handle(case: ProviderCase): await stream.close() +@pytest.mark.encodes_bodies async def test_a_body_above_the_threshold_is_offloaded_and_read_back( case: ProviderCase, client: Client ): @@ -721,6 +841,7 @@ async def test_a_body_above_the_threshold_is_offloaded_and_read_back( @pytest.mark.detects_divergent_retries +@pytest.mark.encodes_bodies async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( case: ProviderCase, client: Client ): diff --git a/tests/streams/test_workflow_streams_provider.py b/tests/streams/test_workflow_streams_provider.py new file mode 100644 index 000000000..218c55fdc --- /dev/null +++ b/tests/streams/test_workflow_streams_provider.py @@ -0,0 +1,1341 @@ +"""Conformance for the workflow_streams provider on the test environment's server. + +Runs the interface loop over the shipped Option 0 transport: an outside +producer appends through the publish Update, or the shipped publish Signal +where a workflow has no Update, the workflow reads and republishes through +its own state, and an outside reader follows the poll Update while the run is +open and the tail Query once it has closed. The outside-surface cases shared +by every provider run from ``test_streams_conformance``; this file covers +what the transport adds. +""" + +from __future__ import annotations + +import asyncio +import base64 +import dataclasses +import uuid +from collections.abc import AsyncIterator +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio import activity, workflow +from temporalio.api.common.v1 import Payload, WorkflowExecution +from temporalio.api.workflowservice.v1 import ResetWorkflowExecutionRequest +from temporalio.client import ( + Client, + WorkflowExecutionStatus, + WorkflowQueryFailedError, + WorkflowUpdateFailedError, +) +from temporalio.contrib.workflow_streams import PublishInput, WorkflowStream +from temporalio.converter import DataConverter, ExternalStorage +from temporalio.exceptions import ApplicationError +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import ( + BEGINNING, + CONTENT_HASH_KEY, + DEFAULT_TOPIC, + END, + Cursor, + RecordKind, + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, + StreamRef, + StreamUnsupportedError, + Supersession, + content_hash, +) +from temporalio.streams._wire import WireRecord +from temporalio.streams.providers import workflow_streams +from temporalio.streams.providers.workflow_streams import ( + WorkflowStreamsActivityHandle, + WorkflowStreamsHandle, + WorkflowStreamsProducer, + WorkflowStreamsProvider, + _InstanceStream, + _PublishResult, +) +from temporalio.worker._workflow_instance import QUERY_HANDLER_NOT_FOUND +from tests.helpers import new_worker + +INPUTS = "inputs" +DECISIONS = "decisions" + + +@pytest.fixture +def provider() -> WorkflowStreamsProvider: + return WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + + +@workflow.defn +class EchoLoop: + """Reads ``inputs``, echoes each value onto ``decisions``, ends on FINISH. + + Lingers until released so a reader can follow it while it runs; whoever + arrives after it closed is served by the tail Query instead. + """ + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> int: + inputs = workflow.stream_reader(INPUTS, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) + seen = 0 + async for record in inputs: + if record.kind is RecordKind.FINISH: + break + if record.kind is not RecordKind.DATA: + continue + assert record.value is not None + seen += 1 + decisions.publish({"echo": record.value["n"]}) + decisions.finish() + await workflow.wait_condition(lambda: self._released) + return seen + + +async def take(records: Any, count: int, timeout: float = 30.0) -> list: + out: list = [] + + async def _collect() -> None: + async for record in records: + out.append(record) + if len(out) >= count: + return + + await asyncio.wait_for(_collect(), timeout) + return out + + +async def _feed( + provider: WorkflowStreamsProvider, client: Client, workflow_id: str +) -> None: + stream = provider.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + await producer.append({"n": 3}) + await producer.finish() + + +ECHOED = [RecordKind.DATA, RecordKind.DATA, RecordKind.DATA, RecordKind.FINISH] + + +async def test_interface_loop_over_workflow_streams( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await _feed(provider, client, workflow_id) + + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS, result_type=dict), 4) + assert [r.kind for r in records] == ECHOED + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + assert all(r.topic == DECISIONS and r.producer_id == "" for r in records) + + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + +async def test_retried_producer_dedupes_and_new_attempt_supersedes( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + # The publish Update answers with where the batch landed. + landed = await first.append({"n": 1}) + assert landed is not None + # The retry of the same attempt re-sends its first batch and is + # answered with the same position. + retry = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + assert await retry.append({"n": 1}) == landed + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"n": 2}) + + records = await take(stream.read(topic=INPUTS, result_type=dict), 3) + assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 + assert records[0].cursor == landed + assert records[1].kind is RecordKind.SUPERSEDED + assert records[1].supersession == Supersession("model", 1, 2) + assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 + assert all(r.topic == INPUTS for r in records) + + # A sequence behind the producer's most recent one is stale, and a + # refusal is typed rather than a silent drop. + await second.append({"n": 3}) + stale = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + with pytest.raises(StreamProducerError, match="most recent"): + await stale.append({"n": "other"}) + assert stale.transport == "update" + + await second.finish() + await handle.signal(EchoLoop.release) + await handle.result() + + +async def test_cold_cache_serves_each_record_once( + client: Client, provider: WorkflowStreamsProvider +): + # Every task rebuilds the workflow from history, so the stream object + # has to belong to the instance that is running: a stale one would carry + # the previous instance's log and hand out every record twice. + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker( + client, EchoLoop, plugins=[provider], max_cached_workflows=0 + ) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await _feed(provider, client, workflow_id) + + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS, result_type=dict), 4) + assert [r.kind for r in records] == ECHOED + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + arrived = await take(stream.read(topic=INPUTS, result_type=dict), 4) + assert [r.kind for r in arrived] == ECHOED + assert [r.value["n"] for r in arrived[:3]] == [1, 2, 3] + + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + +async def test_a_closed_run_serves_its_tail_by_query_and_the_read_ends( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await _feed(provider, client, workflow_id) + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + # Nothing polled while the run was open. The poll Update is gone with + # the run, so everything below arrives through the tail Query, and + # the read ends by itself once the tail is delivered. + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS, result_type=dict)] + + records = await asyncio.wait_for(read_everything(), 30) + assert [r.kind for r in records] == ECHOED + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + + checkpoint = records[1].cursor + again = await take( + stream.read(topic=DECISIONS, result_type=dict, after=checkpoint), 2 + ) + assert [r.value["echo"] for r in again[:1]] == [3] + assert again[1].kind is RecordKind.FINISH + + +async def test_a_bounded_read_is_cancelled_within_its_timeout( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + async def read_forever() -> None: + async for _ in stream.read(topic=DECISIONS, result_type=dict): + pass + + # The run is open and nothing is published, so the read parks on the + # poll Update. The bound has to come out as a timeout rather than + # vanish into a resubscribe. + started = asyncio.get_running_loop().time() + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(read_forever(), timeout=1) + assert asyncio.get_running_loop().time() - started < 10 + # A reader task cancelled outright ends the same way. + task = asyncio.create_task(read_forever()) + await asyncio.sleep(0.2) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).finish() + await handle.signal(EchoLoop.release) + assert await handle.result() == 0 + + +@workflow.defn +class Truncating: + """Publishes on two topics and truncates its log when told to.""" + + def __init__(self) -> None: + # Constructed here so the shipped signal handler is registered before + # the provider looks for it, the way a migrating application holds it. + self._stream = WorkflowStream() + self._released = False + + @workflow.signal + def truncate_to(self, offset: int) -> None: + self._stream.truncate(offset) + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + other = workflow.stream_writer(INPUTS) + decisions.publish({"n": 0}) + other.publish({"side": "inputs"}) + decisions.publish({"n": 1}) + await workflow.wait_condition(lambda: self._released) + + +async def test_a_truncated_position_is_refused_rather_than_restarted( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Truncating, plugins=[provider]) as worker: + handle = await client.start_workflow( + Truncating.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + first = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert first[0].value == {"n": 0} + + # Everything the reader's cursor names is dropped from the log. + await handle.signal(Truncating.truncate_to, 3) + resumed = stream.read(topic=DECISIONS, result_type=dict, after=first[0].cursor) + # Starting over would hand back records the caller already handled, + # and only the caller can decide to do that. + with pytest.raises(StreamCursorError): + await take(resumed, 1, timeout=30) + + await handle.signal(Truncating.release) + await handle.result() + + +async def test_latest_names_the_newest_record_on_the_topic_asked_for( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Truncating, plugins=[provider]) as worker: + handle = await client.start_workflow( + Truncating.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + decisions = await take(stream.read(topic=DECISIONS, result_type=dict), 2) + + # One log orders both topics and the newest item on it belongs to + # `inputs`, so a log-global answer would be the wrong cursor here. + assert await stream.latest(topic=DECISIONS) == decisions[-1].cursor + inputs = await take(stream.read(topic=INPUTS, result_type=dict), 1) + assert await stream.latest(topic=INPUTS) == inputs[0].cursor + assert await stream.latest(topic="never-written") == BEGINNING + + await handle.signal(Truncating.release) + await handle.result() + + +async def test_the_tail_query_pages_instead_of_answering_in_one_blob(): + # The workflow-side half of the tail, driven directly: a Query response + # has to fit the server's blob limit, so a log larger than the cap comes + # back a page at a time with a position to resume from. + big = Payload(metadata={"encoding": b"binary/plain"}, data=b"x" * 400_000) + items = [(offset, DECISIONS, big) for offset in range(6)] + + class _Log: + next_offset = len(items) + + def items_from(self, offset: int) -> list[Any]: + return [item for item in items if item[0] >= offset] + + instance = _InstanceStream.__new__(_InstanceStream) + instance._stream = _Log() # type: ignore[assignment] # pyright: ignore[reportPrivateUsage] + + seen: list[int] = [] + offset, more = 0, True + while more: + page = instance._tail(offset, DECISIONS) + assert page["items"], "a page that fits nothing would never finish" + seen.extend(item["offset"] for item in page["items"]) + offset, more = page["next_offset"], page["more_ready"] + assert seen == list(range(6)) + + +@workflow.defn +class Relay: + """Publishes one record per run and continues as new once.""" + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"run": run}) + if run == 0: + workflow.continue_as_new(run + 1) + decisions.finish() + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[provider]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + first_run = handle.first_execution_run_id + assert first_run is not None + + chain = provider.get_stream_handle(client, workflow_id) + + async def read_everything(stream: Any) -> list[Any]: + return [ + (r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + # Each run keeps its own log, so following the chain means reading + # the first run to its close and then the successor from its start. + assert await asyncio.wait_for(read_everything(chain), 30) == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + pinned = provider.get_stream_handle(client, workflow_id, run_id=first_run) + assert await asyncio.wait_for(read_everything(pinned), 30) == [ + (RecordKind.DATA, {"run": 0}), + ] + # A cursor from the first run resumes into the successor. + records = await take(chain.read(topic=DECISIONS, result_type=dict), 1) + resumed = await asyncio.wait_for( + asyncio.ensure_future( + _values( + chain.read( + topic=DECISIONS, result_type=dict, after=records[0].cursor + ) + ) + ), + 30, + ) + assert resumed == [{"run": 1}, None] + + +@workflow.defn +class Rolling: + """Publishes what it is sent, continues as new onto ``successor_queue`` once.""" + + def __init__(self) -> None: + self._sent: list[int] = [] + self._roll_to = "" + self._released = False + + @workflow.signal + def emit(self, n: int) -> None: + self._sent.append(n) + + @workflow.signal + def roll(self, successor_queue: str) -> None: + self._roll_to = successor_queue + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + while True: + await workflow.wait_condition( + lambda: bool(self._sent or self._roll_to or self._released) + ) + while self._sent: + decisions.publish({"run": run, "n": self._sent.pop(0)}) + if self._roll_to: + workflow.continue_as_new(run + 1, task_queue=self._roll_to) + if self._released: + decisions.finish() + return + + +@workflow.defn +class Idle: + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def _collect_all(records: Any, into: list[Any]) -> None: + async for record in records: + into.append((record.kind, record.value)) + + +async def test_a_live_read_without_a_run_id_follows_continue_as_new( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + successor_queue = f"streams-ws-successor-{uuid.uuid4().hex}" + async with new_worker(client, Rolling, plugins=[provider]) as worker: + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_all( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + await handle.signal(Rolling.emit, 1) + await handle.signal(Rolling.emit, 2) + await assert_eventually_len(seen, 2, reader) + + # The successor runs on a queue nobody polls yet, so the reader's + # first poll on it waits for the successor's first task and lands in + # it, ahead of the hook that registers the handler. + await handle.signal(Rolling.roll, successor_queue) + successor = client.get_workflow_handle(workflow_id) + while (await successor.describe()).run_id == handle.result_run_id: + await asyncio.sleep(0.05) + await asyncio.sleep(1) + assert not reader.done() + async with new_worker( + client, Rolling, plugins=[provider], task_queue=successor_queue + ): + await successor.signal(Rolling.emit, 3) + await successor.signal(Rolling.emit, 4) + await assert_eventually_len(seen, 4, reader) + await successor.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await successor.result() + + assert seen == [ + (RecordKind.DATA, {"run": 0, "n": 1}), + (RecordKind.DATA, {"run": 0, "n": 2}), + (RecordKind.DATA, {"run": 1, "n": 3}), + (RecordKind.DATA, {"run": 1, "n": 4}), + (RecordKind.FINISH, None), + ] + + +async def test_a_poll_that_arrives_before_the_first_task_is_retried( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + task_queue = f"streams-ws-{uuid.uuid4().hex}" + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=task_queue + ) + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_all( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + # No worker yet, so the poll is delivered in the run's first task. + await asyncio.sleep(1) + async with new_worker(client, Rolling, plugins=[provider], task_queue=task_queue): + await handle.signal(Rolling.emit, 1) + await assert_eventually_len(seen, 1, reader) + await handle.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await handle.result() + assert seen == [(RecordKind.DATA, {"run": 0, "n": 1}), (RecordKind.FINISH, None)] + + +async def _reset_at_last_completed_task( + client: Client, workflow_id: str, run_id: str +) -> str: + """Reset ``run_id`` at its last completed task; the id of the run reset into.""" + completion_id = 0 + events = client.get_workflow_handle( + workflow_id, run_id=run_id + ).fetch_history_events() + async for event in events: + if event.HasField("workflow_task_completed_event_attributes"): + completion_id = event.event_id + assert completion_id + try: + answer = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=run_id + ), + reason="re-run from the last completed task", + workflow_task_finish_event_id=completion_id, + request_id=uuid.uuid4().hex, + ) + ) + except RPCError as error: + if error.status != RPCStatusCode.UNIMPLEMENTED: + raise + # The Java time-skipping test server has no reset; the real server + # does, and this never skips there. + pytest.skip("this test server does not implement ResetWorkflowExecution") + return answer.run_id + + +async def _collect_records(records: Any, into: list[Any]) -> None: + async for record in records: + into.append(record) + + +def _run_and_offset(record: Any) -> tuple[str, int]: + _, run_id, offset = record.cursor.token.split(":") + return run_id, int(offset) + + +async def test_a_live_read_without_a_run_id_follows_a_reset( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Rolling, plugins=[provider]) as worker: + handle = await client.start_workflow( + Rolling.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + base_run = handle.result_run_id + assert base_run is not None + seen: list[Any] = [] + reader = asyncio.create_task( + _collect_records( + provider.get_stream_handle(client, workflow_id).read( + topic=DECISIONS, result_type=dict + ), + seen, + ) + ) + await handle.signal(Rolling.emit, 1) + await handle.signal(Rolling.emit, 2) + await assert_eventually_len(seen, 2, reader) + + # The base run is closed by the reset with nothing in its own History + # to say so; the reader learns where it went from describe. The reset + # run replays the base run's History up to the last completed task, + # so its log holds the same two records at the same offsets, and the + # read carries on from the position it had reached. + reset_run = await _reset_at_last_completed_task(client, workflow_id, base_run) + assert reset_run != base_run + current = client.get_workflow_handle(workflow_id) + assert (await current.describe()).run_id == reset_run + await current.signal(Rolling.emit, 3) + await current.signal(Rolling.emit, 4) + await assert_eventually_len(seen, 4, reader) + await current.signal(Rolling.release) + await asyncio.wait_for(reader, 30) + await current.result() + + assert [(r.kind, r.value) for r in seen] == [ + (RecordKind.DATA, {"run": 0, "n": 1}), + (RecordKind.DATA, {"run": 0, "n": 2}), + (RecordKind.DATA, {"run": 0, "n": 3}), + (RecordKind.DATA, {"run": 0, "n": 4}), + (RecordKind.FINISH, None), + ] + # Two records from the base run, then the reset run's, whose offsets + # continue where the base run's log stood at the reset point. + assert [_run_and_offset(r) for r in seen] == [ + (base_run, 0), + (base_run, 1), + (reset_run, 2), + (reset_run, 3), + (reset_run, 4), + ] + # A handle pinned to the base run ends with it. Both closed runs are + # served by the tail Query, which needs the worker still up. + pinned = provider.get_stream_handle(client, workflow_id, run_id=base_run) + pinned_records: list[Any] = [] + await asyncio.wait_for( + _collect_records(pinned.read(topic=DECISIONS), pinned_records), 30 + ) + assert [_run_and_offset(r) for r in pinned_records] == [ + (base_run, 0), + (base_run, 1), + ] + # BEGINNING on the chain starts at the base run, whose start event the + # reset run copied, and a resume from a base run cursor crosses over. + chain = provider.get_stream_handle(client, workflow_id) + resumed = await take(chain.read(topic=DECISIONS, after=seen[1].cursor), 3) + assert [_run_and_offset(r) for r in resumed] == [ + (reset_run, 2), + (reset_run, 3), + (reset_run, 4), + ] + + +@workflow.defn +class ShippedOnly: + """Holds the shipped stream object alone, as a workflow on a worker without the provider.""" + + def __init__(self) -> None: + self._stream = WorkflowStream() + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def test_a_workflow_without_the_publish_update_is_appended_to_by_signal( + client: Client, provider: WorkflowStreamsProvider +): + # The worker has no provider, so the workflow serves the shipped Signal + # and nothing else. The producer learns that from two rejections and + # falls back, and the batches land in the shipped log all the same. + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, ShippedOnly) as worker: + handle = await client.start_workflow( + ShippedOnly.run, id=workflow_id, task_queue=worker.task_queue + ) + producer = provider.get_stream_handle(client, workflow_id).producer( + topic=INPUTS, producer_id="model", attempt=1 + ) + assert await producer.append({"n": 1}) is None + assert producer.transport == "signal" + await producer.append({"n": 2}, {"n": 3}) + assert ( + await handle.query("__temporal_workflow_stream_offset", result_type=int) + == 3 + ) + await handle.signal(ShippedOnly.release) + await handle.result() + + +async def test_a_running_workflow_without_the_provider_fails_the_read_clearly( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(client, Idle) as worker: + handle = await client.start_workflow( + Idle.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + with pytest.raises(StreamError, match="provider is not installed"): + await take(stream.read(topic=DECISIONS), 1) + await handle.signal(Idle.release) + await handle.result() + + +async def assert_eventually_len( + items: list[Any], count: int, reader: asyncio.Task[None] +) -> None: + async def _wait() -> None: + while len(items) < count: + if reader.done(): + # A read that failed surfaces its error instead of a timeout. + reader.result() + raise AssertionError(f"the read ended early with {items}") + await asyncio.sleep(0.05) + + await asyncio.wait_for(_wait(), 30) + + +async def _values(records: Any) -> list[Any]: + return [r.value async for r in records] + + +class _FlakyHandle: + """A workflow handle whose first Signal is accepted and then reported failed.""" + + id = "flaky" + + def __init__(self) -> None: + self.sent: list[PublishInput] = [] + self._fail_next = True + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.sent.append(arg) + if self._fail_next: + self._fail_next = False + raise ConnectionResetError( + "the server accepted the signal, the reply was lost" + ) + + +def _wires(publish: PublishInput) -> list[WireRecord]: + out = [] + for entry in publish.items: + payload = Payload.FromString(base64.b64decode(entry.data)) + out.append(WireRecord.FromString(payload.data)) + return out + + +def _sequences(sent: list[PublishInput]) -> list[tuple[int, list[int]]]: + return [ + (publish.sequence, [wire.sequence for wire in _wires(publish)]) + for publish in sent + ] + + +def _producer(handle: Any, transport: Any = "signal") -> WorkflowStreamsProducer: + return WorkflowStreamsProducer( + handle, + DataConverter.default.payload_converter, + INPUTS, + "model", + 1, + transport=transport, + retry_cooldown=timedelta(0), + ) + + +async def test_a_retried_append_after_an_ambiguous_failure_writes_once(): + # The Signal transport: what a producer that fell back to it does. + handle = _FlakyHandle() + producer = _producer(handle) + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + assert await producer.append({"n": 1}) is None + await producer.append({"n": 2}) + # The retry carries the same signal sequence and the same record + # sequence as the failed send, so the shipped dedupe drops the copy; the + # batch after it continues the numbering. + assert _sequences(handle.sent) == [(2, [1]), (2, [1]), (3, [2])] + assert all(publish.publisher_id == "model#1" for publish in handle.sent) + + +async def test_a_batch_whose_signal_failed_goes_out_before_the_next_one(): + handle = _FlakyHandle() + producer = _producer(handle) + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + await producer.append({"n": 2}, {"n": 3}) + await producer.finish() + assert _sequences(handle.sent) == [(2, [1]), (2, [1]), (4, [2, 3]), (5, [4])] + assert _wires(handle.sent[-1])[0].kind == int(RecordKind.FINISH) + + +class _UpdateHandle: + """A workflow handle serving the publish Update as the server and workflow would. + + Answers a repeated update id from the first outcome, the way the server + does, positions each new batch at the head of a log, the way the + workflow does, and loses the reply of the first call when told to. + Without a handler it rejects every Update the way a workflow whose + worker predates it does, and takes Signals. + """ + + id = "update" + + def __init__(self, *, lose_first_reply: bool = False, handler: bool = True): + self.sent: list[tuple[str, PublishInput]] = [] + self.signalled: list[PublishInput] = [] + self._outcomes: dict[str, _PublishResult] = {} + self._head = 0 + self._lose = lose_first_reply + self._handler = handler + + async def execute_update( + self, name: str, arg: PublishInput, *, id: str, result_type: Any + ) -> _PublishResult: + del name, result_type + self.sent.append((id, arg)) + if not self._handler: + raise WorkflowUpdateFailedError( + ApplicationError( + f"Update handler for 'x' {QUERY_HANDLER_NOT_FOUND}, known updates: []" + ) + ) + answer = self._outcomes.get(id) + if answer is None: + answer = _PublishResult("run", self._head + len(arg.items) - 1) + self._head += len(arg.items) + self._outcomes[id] = answer + if self._lose: + self._lose = False + raise ConnectionResetError("the server took the update, the reply was lost") + return answer + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.signalled.append(arg) + + +def _at(offset: int) -> Cursor: + return Cursor(f"workflow_streams:run:{offset}") + + +async def test_a_retried_update_after_a_lost_reply_is_answered_from_its_id(): + handle = _UpdateHandle(lose_first_reply=True) + producer = _producer(handle, "update") + assert await producer.append() is None + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + # The retry is the same Update: same producer, sequence and content make + # the same id, so the server answers with the outcome it already holds + # and the log takes the batch once. The batch after it is a new one. + assert await producer.append({"n": 1}) == _at(0) + assert await producer.append({"n": 2}) == _at(1) + assert await producer.append() == _at(1) + ids = [id for id, _ in handle.sent] + assert ids[0] == ids[1] != ids[2] + assert _sequences([arg for _, arg in handle.sent]) == [(2, [1]), (2, [1]), (3, [2])] + + +async def test_a_batch_whose_update_failed_goes_out_before_the_next_one(): + handle = _UpdateHandle(lose_first_reply=True) + producer = _producer(handle, "update") + with pytest.raises(ConnectionResetError): + await producer.append({"n": 1}) + # The pending batch lands first, at offset 0, so the new one follows it. + assert await producer.append({"n": 2}, {"n": 3}) == _at(2) + await producer.finish() + sent = [arg for _, arg in handle.sent] + assert _sequences(sent) == [(2, [1]), (2, [1]), (4, [2, 3]), (5, [4])] + assert _wires(sent[-1])[0].kind == int(RecordKind.FINISH) + + +async def test_a_producer_falls_back_to_the_signal_without_a_publish_update(): + handle = _UpdateHandle(handler=False) + producer = _producer(handle, "update") + # Rejected twice across a task boundary means the workflow's worker + # predates the Update; the batch goes by Signal, and so does every later + # one, without asking again. + assert await producer.append({"n": 1}) is None + assert producer.transport == "signal" + assert await producer.append({"n": 2}) is None + assert len(handle.sent) == 2 + assert _sequences(handle.signalled) == [(2, [1]), (3, [2])] + + +class _Description: + run_id = "the-only-run" + status = WorkflowExecutionStatus.COMPLETED + + +class _StubHandle: + """A workflow handle that answers describe and fails whatever the test names.""" + + id = "stub" + run_id = "the-only-run" + + def __init__( + self, + *, + query_error: BaseException | None = None, + events_error: BaseException | None = None, + ) -> None: + self._query_error = query_error + self._events_error = events_error + + async def describe(self) -> Any: + return _Description() + + async def start_update(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + # The run is closed, so its poll Update is gone with it and the read + # goes on to the tail Query, which is what these cases are about. + raise RPCError("no poll update", RPCStatusCode.NOT_FOUND, b"") + + async def query(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + assert self._query_error is not None + raise self._query_error + + def fetch_history_events(self, *args: Any, **kwargs: Any) -> Any: + del args, kwargs + error = self._events_error + + async def _events() -> AsyncIterator[Any]: + events: tuple[Any, ...] = () + for event in events: + yield event + if error is not None: + raise error + + return _events() + + +class _OneHandleClient: + """A client that answers every handle request with the same handle.""" + + data_converter = DataConverter.default + + def __init__(self, handle: Any) -> None: + self._handle = handle + + def get_workflow_handle( + self, workflow_id: str, *, run_id: str | None = None + ) -> Any: + del workflow_id, run_id + return self._handle + + +def _handle_over(stub: _StubHandle) -> WorkflowStreamsHandle: + return WorkflowStreamsHandle( + _OneHandleClient(stub), # type: ignore[arg-type] + "wf", + "the-only-run", + timedelta(0), + ) + + +async def test_a_failing_tail_query_comes_back_as_a_stream_error(): + # The handler being absent is the one benign case; anything else is a + # real failure, and the interface says a caller catches stream conditions + # by meaning rather than by the client's own exception types. + stub = _StubHandle(query_error=WorkflowQueryFailedError("the workflow rejected it")) + with pytest.raises(StreamError): + await take(_handle_over(stub).read(topic=DECISIONS, result_type=dict), 1) + + +async def test_a_missing_tail_handler_ends_the_read_instead_of_failing_it(): + stub = _StubHandle( + query_error=WorkflowQueryFailedError( + f"Query handler for 'x' {QUERY_HANDLER_NOT_FOUND}, known queries: []" + ) + ) + # The workflow never opened a stream through this provider, so the run + # holds no tail and the read is simply over. + assert [r async for r in _handle_over(stub).read(topic=DECISIONS)] == [] + + +async def test_a_failing_latest_query_comes_back_as_a_stream_error(): + stub = _StubHandle(query_error=WorkflowQueryFailedError("the workflow rejected it")) + with pytest.raises(StreamError): + await _handle_over(stub).latest(topic=DECISIONS) + + +async def test_a_missing_run_leaves_the_successor_lookup_as_a_stream_error(): + stub = _StubHandle(events_error=RPCError("gone", RPCStatusCode.NOT_FOUND, b"")) + with pytest.raises(StreamNotFoundError): + await _handle_over(stub)._successor(stub) # type: ignore[arg-type] + + +class _PlainHandle: + """A workflow handle that accepts every Signal and remembers it.""" + + id = "plain" + + def __init__(self) -> None: + self.sent: list[PublishInput] = [] + + async def signal(self, name: str, arg: PublishInput) -> None: + del name + self.sent.append(arg) + + +async def test_the_dedupe_sequence_names_where_the_records_end(): + # The shipped handler drops a batch whose sequence it has already passed, + # so the sequence has to say how far this producer's records reach. A + # count of signals does not: a retry that batches its records differently + # from the send it repeats then carries a sequence the workflow has not + # seen, and the records it already holds go in a second time. + first = _PlainHandle() + original = _producer(first) # type: ignore[arg-type] + await original.append({"n": 1}, {"n": 2}) + + second = _PlainHandle() + retry = _producer(second) # type: ignore[arg-type] + await retry.append({"n": 1}) + await retry.append({"n": 2}) + await retry.append({"n": 3}) + + # The original ended at record 2, so its sequence is 3. Neither half of + # the retry's re-split reaches past it, and only the new record does. + assert _sequences(first.sent) == [(3, [1, 2])] + assert _sequences(second.sent) == [(2, [1]), (3, [2]), (4, [3])] + + +@workflow.defn +class StartsWhenTold: + """Opens a reader on ``inputs`` at the start a signal names and returns what it read.""" + + def __init__(self) -> None: + self._start: str | None = None + + @workflow.signal + def begin(self, start: str) -> None: + self._start = start + + @workflow.run + async def run(self) -> list[Any]: + await workflow.wait_condition(lambda: self._start is not None) + if self._start == "end": + reader = workflow.stream_reader(INPUTS, result_type=dict, after=END) + want = 1 + else: + reader = workflow.stream_reader(INPUTS, result_type=dict, last=2) + want = 2 + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == want: + break + return values + + +async def test_a_workflow_reader_starts_at_end_or_the_newest_records( + client: Client, provider: WorkflowStreamsProvider +): + # The log is workflow state, so both starts resolve against it on the + # workflow thread; a cold cache replays every task and has to land the + # reader on the same offset each time. + async with new_worker( + client, StartsWhenTold, plugins=[provider], max_cached_workflows=0 + ) as worker: + for start, expected in (("last", [3, 4]), ("end", ["new"])): + handle = await client.start_workflow( + StartsWhenTold.run, + id=f"ws-start-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + stream = provider.get_stream_handle(client, handle.id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + await handle.signal(StartsWhenTold.begin, start) + result = asyncio.ensure_future(handle.result()) + if start == "end": + for _ in range(150): + try: + await producer.append({"n": "new"}) + except StreamNotFoundError: + # The Update returns once the workflow took the + # batch, and reading it is what closes the workflow, + # so the next append can find it gone before the + # result has been noticed. + break + done, _ = await asyncio.wait({result}, timeout=0.2) + if done: + break + assert await asyncio.wait_for(result, 30) == expected + + +async def test_a_standalone_activity_stream_is_refused( + client: Client, provider: WorkflowStreamsProvider +): + # The log lives inside a running workflow, so an activity outside any + # workflow has no place to put a stream of its own here. The refusal is + # the documented error, not an AttributeError or a stream silently put + # somewhere else. An activity a workflow scheduled is served. + with pytest.raises(StreamUnsupportedError, match="standalone"): + provider.get_activity_stream_handle(client, "act") + assert isinstance( + provider.get_activity_stream_handle(client, "act", workflow_id="wf"), + WorkflowStreamsActivityHandle, + ) + + +async def test_a_standalone_stream_is_refused( + client: Client, provider: WorkflowStreamsProvider +): + # Every log is a running workflow's state; a stream with no owner has no + # workflow to live in. Both the create and the lookup say so. + with pytest.raises(StreamUnsupportedError, match="standalone"): + await provider.create_standalone_stream(client, "shared") + with pytest.raises(StreamUnsupportedError, match="standalone"): + provider.get_standalone_stream_handle(client, "shared") + + +async def test_a_handle_names_its_stream_as_a_ref( + client: Client, provider: WorkflowStreamsProvider +): + workflow_stream = provider.get_stream_handle(client, "wf", run_id="run-1") + assert workflow_stream.ref(topic=INPUTS) == StreamRef.for_workflow( + "wf", run_id="run-1", topic=INPUTS + ) + own = provider.get_activity_stream_handle(client, "act", workflow_id="wf") + assert own.ref(topic=TOKENS) == StreamRef.for_activity( + "act", workflow_id="wf", topic=TOKENS + ) + assert own.ref().topic == DEFAULT_TOPIC + # An owned stream ends with its owner; only a standalone one closes. + with pytest.raises(ValueError, match="standalone"): + await own.close() + + +def test_a_record_carries_the_plaintext_hash_the_workflow_dedupes_by(): + handle = _UpdateHandle() + producer = _producer(handle, "update") + entries, _ = producer._entries( + [(RecordKind.DATA, {"n": 1}), (RecordKind.FINISH, None)] + ) # pyright: ignore[reportPrivateUsage] + data, finish = _wires(PublishInput(items=entries)) + # A DATA record is stamped with the hash of its converted body, where the + # workflow can read it without the body; FINISH has nothing to hash. + assert data.metadata[CONTENT_HASH_KEY].data.decode() == content_hash(data.body) + assert CONTENT_HASH_KEY not in finish.metadata + + # The workflow's identity for a batch is those hashes, so a body whose + # bytes a codec changed is still the same batch, and a different value + # is not. + def batch(*values: dict) -> PublishInput: + made, _ = _producer(handle, "update")._entries( # pyright: ignore[reportPrivateUsage] + [(RecordKind.DATA, value) for value in values] + ) + return PublishInput(items=made, publisher_id="model#1", sequence=3) + + same, recoded, other = batch({"n": 1}), batch({"n": 1}), batch({"n": 2}) + record = _wires(recoded)[0] + record.body.data = b"\x00" + record.body.data + recoded.items[0].data = base64.b64encode( + Payload( + metadata={"encoding": b"binary/plain"}, data=record.SerializeToString() + ).SerializeToString() + ).decode("ascii") + assert workflow_streams._content(same) == workflow_streams._content(recoded) # pyright: ignore[reportPrivateUsage] + assert workflow_streams._content(same) != workflow_streams._content(other) # pyright: ignore[reportPrivateUsage] + + +async def test_external_storage_applies_at_the_envelope_the_workflow_reads_through( + client: Client, provider: WorkflowStreamsProvider +): + # Worker and clients share one converter, as a deployment's do. A batch + # above the threshold leaves as a claim on the Update argument, the + # worker redeems it, the workflow reads the value, and the outside + # reader gets the poll response the same way. + from tests.streams.test_streams_conformance import RecordingDriver + + driver = RecordingDriver() + converter = dataclasses.replace( + DataConverter.default, + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=512), + ) + config = client.config() + config["data_converter"] = converter + shared = Client(**config) + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker(shared, EchoLoop, plugins=[provider]) as worker: + handle = await shared.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(shared, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + assert driver.stored == 0 + await producer.append({"n": 2, "blob": "x" * 4096}) + assert driver.stored == 1 + await producer.finish() + + records = await take(stream.read(topic=DECISIONS, result_type=dict), 3) + assert [r.value.get("echo") for r in records[:2]] == [1, 2] + assert records[2].kind is RecordKind.FINISH + assert driver.retrieved >= 1 + + await handle.signal(EchoLoop.release) + assert await handle.result() == 2 + + +TOKENS = "tokens" + + +@activity.defn +async def stream_then_finish(count: int) -> None: + producer = activity.stream_handle(scope="activity").producer(topic=TOKENS) + for n in range(count): + await producer.append({"n": n}) + # Long enough for a reader polling every few milliseconds to see this + # activity pending before it finishes. + await asyncio.sleep(1) + await producer.finish() + + +@workflow.defn +class RunsAnActivityThenLingers: + """Runs the streaming activity by a fixed id, then waits to be released.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self, count: int) -> None: + await workflow.execute_activity( + stream_then_finish, + count, + activity_id="streamer", + start_to_close_timeout=timedelta(seconds=30), + ) + await workflow.wait_condition(lambda: self._released) + + +async def test_an_activity_read_ends_with_the_activity_while_the_workflow_runs( + client: Client, provider: WorkflowStreamsProvider +): + workflow_id = f"streams-ws-{uuid.uuid4().hex}" + async with new_worker( + client, + RunsAnActivityThenLingers, + activities=[stream_then_finish], + plugins=[provider], + ) as worker: + handle = await client.start_workflow( + RunsAnActivityThenLingers.run, + 2, + id=workflow_id, + task_queue=worker.task_queue, + ) + own = provider.get_activity_stream_handle( + client, "streamer", workflow_id=workflow_id + ) + + async def read_everything() -> list[Any]: + return [r async for r in own.read(topic=TOKENS, result_type=dict)] + + records = await asyncio.wait_for(read_everything(), 30) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"n": 0}), + (RecordKind.DATA, {"n": 1}), + (RecordKind.FINISH, None), + ] + # The producer is the activity, and the record carries the plain + # topic name; the reserved name is the log's business. + assert all(r.producer_id == "streamer" and r.attempt == 1 for r in records) + assert all(r.topic == TOKENS for r in records) + # The activity's end ended the read: the workflow is still running. + assert (await handle.describe()).status == WorkflowExecutionStatus.RUNNING + + # The workflow's topic of the same name is another stream, and the + # reserved name cannot be reached as a workflow topic. + workflow_stream = provider.get_stream_handle(client, workflow_id) + assert await workflow_stream.latest(topic=TOKENS) == BEGINNING + with pytest.raises(ValueError, match="reserved"): + workflow_stream.read(topic="activity/streamer/tokens") + with pytest.raises(ValueError, match="reserved"): + workflow_stream.producer(topic="activity/x", producer_id="p", attempt=1) + + await handle.signal(RunsAnActivityThenLingers.release) + await handle.result()