From 6250dfe687846c7a43d8c39534480e7d52f4ebb7 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 04:30:39 -0700 Subject: [PATCH 1/5] Let an external stream subscription name its start and yield offsets. A provider over the transport needs both: a reader seeded from a cursor starts after it, and records() hands back the offset each value was read from. The test fakes now place records the way a provider does, since records() needs an offset on every one. --- .../contrib/external_workflow_streams/_api.py | 34 +++++++++++++++- .../external_workflow_streams/test_api.py | 39 ++++++++++++++++++- .../test_output_runtime.py | 4 +- 3 files changed, 73 insertions(+), 4 deletions(-) diff --git a/temporalio/contrib/external_workflow_streams/_api.py b/temporalio/contrib/external_workflow_streams/_api.py index f69674c25..38986819b 100644 --- a/temporalio/contrib/external_workflow_streams/_api.py +++ b/temporalio/contrib/external_workflow_streams/_api.py @@ -35,6 +35,8 @@ classify_read_failure, ) from temporalio.contrib.external_workflow_streams._record import ( + Cursor, + Offset, StreamRecord, ) from temporalio.contrib.external_workflow_streams._wake import channel_for @@ -95,6 +97,7 @@ def register( wait_id: int, stream_key: StreamKey, idle_timeout: timedelta, + start_cursor: Cursor | None = None, ) -> None: """Registers a wait with the Worker's subscription manager. @@ -275,7 +278,9 @@ class ExternalStreamTopic(Generic[AnyType]): value_type: type[AnyType] | None options: ExternalStreamOptions - def subscribe(self) -> ExternalStreamSubscription[AnyType]: + def subscribe( + self, *, start_cursor: Cursor | None = None + ) -> ExternalStreamSubscription[AnyType]: """Starts a new subscription and returns its async iterator. Each call is an **independent** subscription with its own ``wait_id``, @@ -286,6 +291,13 @@ def subscribe(self) -> ExternalStreamSubscription[AnyType]: the same hazard class as timers and activities: inserting, removing, or reordering a ``subscribe()`` call renumbers every later wait in the Run and must be gated behind ``workflow.patched()``. + + Args: + start_cursor: The boundary the subscription begins after. ``None`` + resumes where the predecessor Run committed this wait, or at + ``BEGINNING`` on a first execution. A boundary the Workflow + names is recorded in the marker's header like the restored one, + so it must be derived deterministically: replay names it again. """ state = _run_state() if state.runtime is None: @@ -297,6 +309,11 @@ def subscribe(self) -> ExternalStreamSubscription[AnyType]: state.next_wait_id += 1 stream_key = state.runtime.stream_key(self.name) + # Only a named boundary travels; the runtime derives the default itself + # from the predecessor Run's continuation. + start: dict[str, Any] = ( + {} if start_cursor is None else {"start_cursor": start_cursor} + ) state.runtime.register( wait_id=wait_id, stream_key=stream_key, @@ -307,6 +324,7 @@ def subscribe(self) -> ExternalStreamSubscription[AnyType]: # `with_options` was given, so no configured value can ever reach # the reduction and every set parks after one second. idle_timeout=self.options.idle_timeout, + **start, ) # The channel the stream's writers notify. Asked of the SDK's object for # the Run rather than of the stream runtime, because the answer is part @@ -611,6 +629,16 @@ def __aiter__(self) -> AsyncIterator[AnyType]: return self._iterate() async def _iterate(self) -> AsyncIterator[AnyType]: + async for _offset, value in self.records(): + yield value + + async def records(self) -> AsyncIterator[tuple[Offset, AnyType]]: + """Each value with the provider offset it was read from. + + Same delivery, same consumption, same commit order as iterating values. + A reader that has to name where it got to, or hand a position to + something outside the Workflow, cannot do it from the values alone. + """ while not self._finished: # Re-filled before *every* record rather than once per batch. A # record buffered while Workflow code was doing something else -- a @@ -640,8 +668,10 @@ async def _iterate(self) -> AsyncIterator[AnyType]: # becomes a value or raises, and neither outcome can leave the # ready list half-consumed. value = self._decode(record) + offset = record.offset self._commit(record) - yield value + assert offset is not None + yield offset, value continue await self._await_readiness() diff --git a/tests/contrib/external_workflow_streams/test_api.py b/tests/contrib/external_workflow_streams/test_api.py index 207bc3b4e..4ab024e0a 100644 --- a/tests/contrib/external_workflow_streams/test_api.py +++ b/tests/contrib/external_workflow_streams/test_api.py @@ -29,6 +29,9 @@ ) from temporalio.contrib.external_workflow_streams._manager import PreparedRecord from temporalio.contrib.external_workflow_streams._record import ( + AFTER, + Cursor, + Offset, RecordKind, StreamRecord, ) @@ -47,6 +50,7 @@ def __init__(self) -> None: #: `wait_id -> configured idle timeout`, so a test can see that the #: value `with_options` was given actually reached the Worker. self.idle_timeouts: dict[int, timedelta] = {} + self.start_cursors: dict[int, Cursor | None] = {} self.buffers: dict[int, list[StreamRecord]] = {} self.deliveries: list[tuple[int, StreamRecord]] = [] self.consumed: list[tuple[int, StreamRecord]] = [] @@ -58,6 +62,9 @@ def __init__(self) -> None: #: Overrides what `codec_for` hands back, so a test can control when -- #: and whether -- decoding a record succeeds. self.codec: Any = None + #: Offsets handed out by `drain`, so buffered records come back in the + #: read-path shape whatever a test put in. + self._placed = 0 def stream_key(self, stream_name: str) -> StreamKey: return StreamKey("ns", "wf", "first-run", stream_name) @@ -68,9 +75,11 @@ def register( wait_id: int, stream_key: StreamKey, idle_timeout: timedelta, + start_cursor: Cursor | None = None, ) -> None: self.registrations.append((wait_id, stream_key)) self.idle_timeouts[wait_id] = idle_timeout + self.start_cursors[wait_id] = start_cursor def drain(self, wait_id: int, max_records: int | None = None) -> list[StreamRecord]: buffered = self.buffers.get(wait_id, []) @@ -81,7 +90,22 @@ def drain(self, wait_id: int, max_records: int | None = None) -> list[StreamReco ) else: self.buffers[wait_id] = [] - return buffered + # A provider places a record when it appends it, so nothing without an + # offset ever reaches a reader. Place what a test left bare rather than + # let the fake deliver a shape the real path cannot produce. + placed = [] + for record in buffered: + if record.offset is not None: + placed.append(record) + continue + self._placed += 1 + at = record.placed_at(Offset(f"{self._placed:020d}")) + if isinstance(record, PreparedRecord): + at = PreparedRecord.of( + at, record.prepared_payload, record.prepare_error + ) + placed.append(at) + return placed def delivery_budget_remaining(self) -> int: return self.budget @@ -176,6 +200,19 @@ def test_topics_inherit_their_options( assert subscription.idle_timeout == timedelta(seconds=3) +def test_a_subscription_may_name_the_boundary_it_starts_after( + runtime: FakeRuntime, +) -> None: + """A named boundary reaches the runtime; an unnamed one leaves the default to it.""" + boundary = AFTER(Offset("1700000000000-3")) + + seeded = external_stream.topic("tokens").subscribe(start_cursor=boundary) + default = external_stream.topic("tokens").subscribe() + + assert runtime.start_cursors[seeded.wait_id] == boundary + assert runtime.start_cursors[default.wait_id] is None + + # --- topics ------------------------------------------------------------------- diff --git a/tests/contrib/external_workflow_streams/test_output_runtime.py b/tests/contrib/external_workflow_streams/test_output_runtime.py index c194bcf0d..51c785733 100644 --- a/tests/contrib/external_workflow_streams/test_output_runtime.py +++ b/tests/contrib/external_workflow_streams/test_output_runtime.py @@ -168,7 +168,9 @@ def backend() -> _OutputMemoryBackend: @pytest.fixture -def runtime(backend: _OutputMemoryBackend) -> WorkflowStreamRuntime: +async def runtime(backend: _OutputMemoryBackend) -> WorkflowStreamRuntime: + # The manager captures the loop it is built on, which in the Worker is the + # Worker's own. A sync fixture has none to capture. manager = StreamSubscriptionManager( backend=backend, notify_ready=_notify, From 640e594dca1863d757a408e878de8ed8fbf6ea4b Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 04:30:39 -0700 Subject: [PATCH 2/5] Added the Redis provider over External Workflow Streams. RedisStreams serves the stream interface from a store the customer runs: one log per topic, a workflow's publishes staged with its task and promoted by the completion marker, outside appends that wake the reader over the channel or the Signal, and a content-hash match for retries. CI gets a Redis service. --- .github/workflows/ci.yml | 50 + CHANGELOG.md | 6 +- streams_demo/provider_setup.py | 26 +- temporalio/streams/providers/redis.py | 1343 +++++++++++++++++++++++++ 4 files changed, 1415 insertions(+), 10 deletions(-) create mode 100644 temporalio/streams/providers/redis.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a53035a7d..f55413c2a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -100,6 +100,56 @@ jobs: npx doctoc README.md [[ -z $(git status --porcelain README.md) ]] || (git diff README.md; echo "README changed"; exit 1) + # The client-side (Redis) stream provider's own evidence. Its tests need a store + # this repo does not otherwise stand up, so without this job nothing that proves + # retention, cursor ownership, the staged commit or the paired producer write ever + # runs anywhere but a developer's machine. + streams-redis: + timeout-minutes: 30 + runs-on: ubuntu-latest + services: + redis: + image: redis:8-alpine + ports: + - 6379:6379 + options: >- + --health-cmd "redis-cli ping" + --health-interval 5s + --health-timeout 3s + --health-retries 10 + steps: + - uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + submodules: recursive + - uses: dtolnay/rust-toolchain@29eef336d9b2848a0b548edc03f92a220660cdb8 # stable + - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5 + with: + python-version: "3.13" + - uses: Swatinem/rust-cache@e18b497796c12c097a38f9edb9d0641fb99eee32 # v2 + with: + workspaces: temporalio/bridge -> target + key: streams-redis-${{ env.pythonLocation }} + - uses: arduino/setup-protoc@c65c819552d16ad3c9b72d9dfd5ba5237b9c906b # v3 + with: + version: "23.x" + repo-token: ${{ secrets.GITHUB_TOKEN }} + - uses: astral-sh/setup-uv@cec208311dfd045dd5311c1add060b2062131d57 # v8 + - run: uv tool install poethepoet + - run: uv sync --all-extras + - run: poe build-develop + # The dev server comes from the test environment, the store from the service + # above. Run serially: the cases measure real timing and share one Redis. + - run: uv run pytest tests/streams -p no:randomly -s + timeout-minutes: 20 + env: + STREAMS_LIVE: redis + TEMPORAL_TEST_REDIS_URL: redis://127.0.0.1:6379 + AI198_REDIS_URL: redis://127.0.0.1:6379 + # Also without the store, so the gate itself keeps working and the memory + # provider's conformance run stays honest. + - run: uv run pytest tests/streams -p no:randomly -s + timeout-minutes: 10 + # Verify the optional FIPS build: the Rust core must link aws-lc-fips-sys # (aws-lc-rs FIPS mode) and must NOT link `ring` (the cargo-tree guard, ported # from sdk-ruby PR #466's `fips_tree` guard); then run the test suite against the diff --git a/CHANGELOG.md b/CHANGELOG.md index 9d0e47e37..6ccb8fc4d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -52,7 +52,11 @@ to include examples, links to docs, or any other relevant information. `close()` seals it. A provider runs record bodies through the client's data 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. + reference provider the conformance tests run against, and + `temporalio.streams.providers.redis.RedisStreams` serves the same interface + over External Workflow Streams, one topic as an input and an output stream. +- `ExternalStreamSubscription.records()` yields each value with the provider + offset it was read from, for a reader that has to name where it got to. - Added experimental External Workflow Streams in `temporalio.contrib.external_workflow_streams`. Workflow stream payloads are diff --git a/streams_demo/provider_setup.py b/streams_demo/provider_setup.py index 2e4bcc350..7fbf9316c 100644 --- a/streams_demo/provider_setup.py +++ b/streams_demo/provider_setup.py @@ -1,8 +1,9 @@ """Pick the provider for a demo run from the environment. -``STREAMS_PROVIDER`` names a provider; this base tree carries only ``memory``, -and each provider branch adds its own name here. The demo needs a Temporal -server to run the workflow either way; ``TEMPORAL_ADDRESS`` points at it. +``STREAMS_PROVIDER`` names a provider; this tree carries ``memory`` and +``redis``. The demo needs a Temporal server to run the workflow either way; +``TEMPORAL_ADDRESS`` points at it, and the Redis demo reads its store from +``AI198_REDIS_URL`` and ``AI198_REDIS_PREFIX``. """ from __future__ import annotations @@ -21,16 +22,23 @@ async def open() -> tuple[str, ProviderPlugin]: """The server to connect to and the provider the worker and the client share.""" - if NAME != "memory": - raise SystemExit(f"this tree carries no stream provider named {NAME!r}") - return os.environ.get("TEMPORAL_ADDRESS", "localhost:7233"), MemoryStreams() + address = os.environ.get("TEMPORAL_ADDRESS", "localhost:7233") + if NAME == "memory": + return address, MemoryStreams() + if NAME == "redis": + from temporalio.streams.providers.redis import RedisStreams + + return address, RedisStreams( + url=os.environ.get("AI198_REDIS_URL", "redis://127.0.0.1:6379"), + key_prefix=os.environ.get("AI198_REDIS_PREFIX", "ai198-contract"), + ) + raise SystemExit(f"this tree carries no stream provider named {NAME!r}") async def close(provider: ProviderPlugin) -> None: """Let go of whatever :func:`open` acquired. - The memory provider holds no connection, so this is its ``close()`` and - nothing more. A provider branch that opens one closes it the same way, so - the demo's teardown reads the same on every provider. + The memory provider holds no connection and the Redis provider closes + the client it opened, so the demo's teardown reads the same on both. """ await provider.close() diff --git a/temporalio/streams/providers/redis.py b/temporalio/streams/providers/redis.py new file mode 100644 index 000000000..7a5445f97 --- /dev/null +++ b/temporalio/streams/providers/redis.py @@ -0,0 +1,1343 @@ +"""The client-side (Redis) provider. + +Streams live in a store the customer runs, Redis here, behind the external +workflow streams transport. A workflow's publish is buffered by the worker, +staged invisibly under a token when the Workflow Task completes, and promoted +only once a marker in History proves the task was accepted. Consumption is +recorded the same way, as ranges and boundaries in History. + +The mapping, in one place: + +- One Redis stream per topic, the topic's log, read by the workflow and by + outside readers alike. The transport derives an input key and an output key + for a topic, because the direction is part of every key it renders; this + provider's backend renders both onto the log. An outside producer appends a + record once, and it is where the workflow's subscription reads and where an + outside read follows; the workflow's staged batches land in the same log + when its task completes and are promoted there, so an outside reader sees + the producer's records and the workflow's in one order. The workflow does + not read its own records: the backend drops the entries the workflow staged + from every read it serves the transport, the live read and the replay read + alike, so a recorded range and its replay agree. The wake path is the + transport's own: a producer appends, then wakes the subscribed run through + the input key. A wake the server refuses because the consuming run is + closing is sent again after a short wait: a run that continues as new hands + the log to its successor, the records are already where the successor + reads them, and only the wake has to follow. The retries stop when the + chain's current run takes the wake or the chain proves terminal. A run + still refusing when the window passes is inside a Workflow Task that tried + to close it while the wake sat buffered; the wake is dropped, because the + record is in the log and a run that does not close after all rechecks the + log at its next park. +- A reset run inherits the base run's History up to the reset point, and + with it the ranges the base run's task completions recorded. It re-reads + them from the log and replays them against the inherited markers, so the + batches those tasks published are not published again, and live reading + continues from the last inherited boundary. The reset-point task itself is + run again, and what the base run consumed after the reset point is still + in the log, so the reset run reads it again and publishes again: an + outside reader sees that task's batch twice, once from each run. Nothing + in the log marks the reset. +- A record rides as the transport's payload: the serialized ``StreamRecord`` + proto as a ``binary/plain`` value, which the worker's codec encodes and + decodes like any other payload. +- Producer identity dedupes through the transport's idempotency: the session + is ``producer#attempt`` and every record carries a sequence, so a retried + batch is dropped with the original position and a new attempt passes. +- Cursors are ``redis:-`` and name an entry of the topic's log, so a + cursor a workflow reader returned seeds an outside read and the other way + round. A reader opened without a cursor starts where the chain's + predecessor run committed, which is the transport's own rule. An outside + read positions ``END`` and ``last=N`` against the log on the first step of + the generator, since the call itself cannot reach the store. A workflow + reader starts at a cursor or at the beginning. +- A workflow's streams are keyed by the chain's first run, so a handle + follows continue-as-new by construction and ``run_id`` only decides whose + close ends a read. +- A task's publishes are staged as one batch. The transport's own per-task + batch limits are lifted for this provider, because a synchronous publish + cannot wait for the worker to stage a full batch; a batch it cannot stage + fails the task. + +- Every record carries the SHA-256 of its converted body under + ``temporal.io/content-hash``, stamped before the payload codec runs, and the + append script compares that hash rather than the stored bytes, so a retry + through a codec that differs on every call is still the same append while a + divergent one is refused. A record without a body is matched by its + plaintext record instead. + +Keys written by the earlier layout. Before the log, a topic was two keys: the +transport's input key, ``::::`` +with every id percent-encoded, and its output key, the same with an ``output`` +component before the topic. The log is the input key, so the records outside +producers wrote, the ranges a workflow recorded and the cursors its readers +returned, ``redis:in:-``, read as they did. The output key is not read +any more: batches a workflow promoted under the old layout are not delivered to +an outside reader, and an outside cursor minted then names an entry of that +key, so a reader holding one starts over from ``BEGINNING`` or positions itself +with ``latest()``. Nothing moves or deletes an old output key; it ages out under +the operator's own retention. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import logging +import re +import time +from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine +from dataclasses import replace +from datetime import timedelta +from typing import Any, Final, Generic, TypeVar, get_args + +from google.protobuf.message import DecodeError + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.client import Client, WorkflowExecutionStatus +from temporalio.contrib.external_workflow_streams import ( + AFTER, + AppendConflictError, + AppendNotAcknowledgedError, + ChainKeyMismatchError, + ExternalStreamProducer, + Offset, + StreamDirection, + WakeNotAcknowledgedError, + WorkflowChainKey, + external_output_stream, + external_stream, +) +from temporalio.contrib.external_workflow_streams import ( + BEGINNING as TRANSPORT_BEGINNING, +) +from temporalio.contrib.external_workflow_streams import ( + RecordKind as TransportRecordKind, +) +from temporalio.contrib.external_workflow_streams import ( + StreamError as TransportStreamError, +) +from temporalio.contrib.external_workflow_streams import ( + StreamRecord as TransportRecord, +) +from temporalio.contrib.external_workflow_streams._backend import ( + DEFAULT_WATCH_BLOCK, + StreamKey, +) +from temporalio.contrib.external_workflow_streams._codec import StreamPayloadCodec +from temporalio.contrib.external_workflow_streams._errors import StreamIntegrityError +from temporalio.contrib.external_workflow_streams._output_backend import ( + OutputStage, + OutputStageConflictError, + OutputStageManifest, + OutputStageNotFoundError, + OutputStageResolutionError, + OutputStageStatus, + OutputStreamRecord, +) +from temporalio.contrib.external_workflow_streams._output_client import ( + _reconcile_output_stage, +) +from temporalio.contrib.external_workflow_streams._redis import ( + _BEGINNING_SENTINEL, + _OUTPUT_STAGE_FIELD, + RedisStreamBackend, + _text, + _to_record, +) +from temporalio.contrib.external_workflow_streams._redis import ( + _parse as _parse_entry_id, +) +from temporalio.contrib.external_workflow_streams._wake import WakeTransport +from temporalio.converter import ( + WorkflowSerializationContext, +) +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 import ReplayerConfig, WorkerConfig + +__all__ = ["RedisProducer", "RedisStreamHandle", "RedisStreams"] + +T = TypeVar("T") + +_PROVIDER = "redis" +#: What a workflow reader's cursor carried when a topic was two keys. It named +#: an entry of the input key, which is the log now, so the form is still read. +_LEGACY_INPUT_PREFIX = "in:" +_REDIS_ID = re.compile(r"\d+-\d+") +_READ_BATCH = 256 +_STAGE_FIELD: Final = _OUTPUT_STAGE_FIELD.encode() + +#: How long a wake the server refused is sent again before it is given up, and +#: the pause between attempts. The pause is there because in the instant between +#: two runs of a chain the successor is not yet the run a Signal resolves to. +_WAKE_RETRY_WINDOW: Final = timedelta(seconds=5) +_WAKE_RETRY_BACKOFF: Final = timedelta(milliseconds=200) + +#: Bits of a wake counter given to an entry id's sequence part. Redis assigns +#: sequence numbers from 0 within one millisecond, so a million appends to one +#: log inside the same millisecond would be needed to reach the cap. +_WAKE_SEQUENCE_BITS: Final = 20 + + +def _wake_counter(offset: Offset) -> int: + """A wake counter from a Redis entry id ``-``. + + ``ms * 2**20 + min(seq, 2**20 - 1)``, which increases with the id's own + order and fits the server's signed 64-bit counter for any millisecond + timestamp before the year 2248. Ids past the cap within one millisecond + share a counter, which only lets a later wake fold into a pending one that + reports an earlier id of that same millisecond. + """ + ms, seq = _parse_entry_id(offset) + cap = (1 << _WAKE_SEQUENCE_BITS) - 1 + return (ms << _WAKE_SEQUENCE_BITS) + min(seq, cap) + + +# The transport's per-task batch limits are a backpressure point that awaits +# inside the workflow, and a synchronous publish has nowhere to await. Lifted +# to where they cannot fire, so a task's publishes stage as one batch. +_UNBOUNDED_OUTPUT = external_output_stream.with_options( + max_records=1 << 30, max_logical_bytes=1 << 40 +) + +logger = logging.getLogger(__name__) + + +def _require_topic(topic: str) -> None: + if not topic: + raise ValueError("topic must not be empty") + + +def _position(after: Cursor) -> Offset | None: + """The log entry a cursor names, or ``None`` for BEGINNING. + + Raises: + StreamCursorError: The cursor is another provider's, or does not name + a Redis entry id. + """ + token = cursor_position(after, provider=_PROVIDER) + if token is None: + return None + if token.startswith(_LEGACY_INPUT_PREFIX): + token = token[len(_LEGACY_INPUT_PREFIX) :] + if not _REDIS_ID.fullmatch(token): + raise StreamCursorError( + f"cursor {after.token!r} does not name a Redis stream position" + ) + return Offset(token) + + +def _entry_id(token: str | bytes) -> tuple[int, int]: + """A Redis entry id as the ``(ms, seq)`` pair it orders by.""" + text = token.decode() if isinstance(token, bytes) else token + ms, _, seq = text.partition("-") + return int(ms), int(seq or 0) + + +#: Given one page of log entries, newest first, the ids of the ones that are +#: not records to the reader asking. +_SkipIn = Callable[[list[Any]], Awaitable[set[str]]] + + +async def _tail_after( + store: Any, name: str, before_last: int, *, skip_in: _SkipIn | None = None +) -> Offset | None: + """The entry the newest ``before_last`` records of the log ``name`` come after. + + ``None`` when the log holds no more than ``before_last`` records, which + means a read from the beginning. With ``before_last=0`` it is the newest + record itself, the boundary a read at ``END`` starts after. Walks the log + from its newest entry, leaving out what ``skip_in`` names. + """ + needed = before_last + 1 + end = "+" + seen = 0 + while True: + page: Any = await store.xrevrange(name, end, "-", count=max(needed - seen, 16)) + if not page: + return None + skipped = await skip_in(page) if skip_in is not None else set() + for entry_id, _fields in page: + token = _text(entry_id) + if token in skipped: + continue + seen += 1 + if seen == needed: + return Offset(token) + ms, seq = _entry_id(page[-1][0]) + if seq: + end = f"{ms}-{seq - 1}" + elif ms: + end = f"{ms - 1}-18446744073709551615" + else: + return None + + +def _fields_args(record: TransportRecord) -> list[Any]: + """The record's stored fields, name then value, as the append scripts take them.""" + args: list[Any] = [] + for name, value in sorted(record.to_fields().items()): + args.append(name.encode()) + args.append(value) + return args + + +def _append_args(record: TransportRecord, digest: str) -> list[Any]: + """What the log append script takes: the identity and digest, then the fields. + + ``digest`` is the plaintext hash of what the record carries, taken before + the payload codec ran, so a retry the codec encoded differently is still + recognized as the same append. + """ + return [ + str(record.idempotency_key).encode(), + digest.encode(), + *_fields_args(record), + ] + + +def _plaintext_digest(wire: WireRecord) -> str: + """The digest an append is matched by: the body's hash, or the record's without one.""" + if wire.HasField("body"): + return content_hash(wire.body) + return hashlib.sha256(wire.SerializeToString(deterministic=True)).hexdigest() + + +def _stamp_hash(wire: WireRecord) -> None: + """Put the body's plaintext hash on ``wire`` under the shared metadata key.""" + if wire.HasField("body"): + wire.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload( + metadata={"encoding": b"binary/plain"}, + data=content_hash(wire.body).encode(), + ) + ) + + +#: Append one record to a log, or answer where it already is. +#: +#: The provider's own write rather than the transport's, so a retry is matched +#: by its plaintext digest rather than by the encoded bytes. The idempotency +#: hash is read before the log is touched, so an identity already used with a +#: different plaintext digest refuses the record rather than writing it, and one +#: used with the same digest answers with the original position, which is what +#: settles a call whose answer was lost. +_LOG_APPEND_LUA: Final = """ +local existing = redis.call('HGET', KEYS[2], ARGV[1]) +if existing then + local sep = string.find(existing, '|') + if string.sub(existing, sep + 1) ~= ARGV[2] then + return {'conflict', ''} + end + return {'ok', string.sub(existing, 1, sep - 1)} +end +local id = redis.call('XADD', KEYS[1], '*', unpack(ARGV, 3)) +redis.call('HSET', KEYS[2], ARGV[1], id .. '|' .. ARGV[2]) +return {'ok', id} +""" + + +class _LogAppend: + """One record on a topic's log, in one Redis call.""" + + def __init__(self, backend: RedisStreamBackend) -> None: + """Bind to ``backend``'s client.""" + self._backend = backend + self._script = backend._client.register_script(_LOG_APPEND_LUA) + + async def write(self, *, name: str, record: TransportRecord, digest: str) -> Offset: + """Append ``record`` to the log ``name`` and return where it landed. + + A repeat of the same ``(session, sequence)`` with the same ``digest`` + returns the original position and writes nothing. + """ + outcome, placed = await self._script( + keys=[name, f"{name}:idem"], + args=_append_args(record, digest), + ) + if _text(outcome) == "conflict": + raise AppendConflictError(record.idempotency_key) + return Offset(_text(placed)) + + +def _is_staged(fields: Any) -> bool: + """Whether a log entry is one the workflow staged, rather than a producer's.""" + return _STAGE_FIELD in fields + + +class _TopicLogBackend(RedisStreamBackend): + """The transport's Redis backend with this provider's layout. + + One key per topic: the transport renders a topic's input key and its + output key apart, because the direction is part of every key it derives, + and this backend renders both onto the log, the input key's rendering. + Every read the transport makes through an input key, the live watch and + the replay range alike, drops the entries the workflow staged itself, so + a workflow never reads its own records and a recorded range replays to + what was delivered. Reads through an output key are the outside reader's + and see the whole log, with the stage protocol deciding what is visible. + """ + + def __init__( + self, + *, + client: Any, + key_prefix: str, + wake_transport: WakeTransport = "auto", + ) -> None: + super().__init__(client=client, key_prefix=key_prefix) + self.wake_transport = wake_transport + + def wake_counter_for(self, offset: Offset) -> int: + """The entry id's own order, so producers and workers rank wakes alike.""" + return _wake_counter(offset) + + def wake_counter_now(self) -> int: + """The entry id rule applied to the clock, with the largest sequence. + + Above every entry appended before now, so a wake that reports no + position, the worker's shutdown sweep, is not folded under the + channel's latest notification. + """ + now_ms = int(time.time() * 1000) + return (now_ms << _WAKE_SEQUENCE_BITS) + (1 << _WAKE_SEQUENCE_BITS) - 1 + + def stream_key(self, key: StreamKey) -> str: + """The topic's log, whichever direction the transport asks for.""" + return super().stream_key(replace(key, direction=StreamDirection.INPUT)) + + async def read_after( + self, + key: StreamKey, + after: Any, + *, + max_records: int, + block: timedelta | None = DEFAULT_WATCH_BLOCK, + ) -> list[TransportRecord]: + """The live read, without the entries the workflow staged itself. + + A batch that held nothing but the workflow's own entries is read past + without blocking again, so the transport's recheck, which asks for one + record, is not answered "nothing" while a producer's record sits behind + a staged batch. + """ + if key.direction is StreamDirection.OUTPUT: + return await super().read_after( + key, after, max_records=max_records, block=block + ) + name = self.stream_key(key) + start = _BEGINNING_SENTINEL if after.is_beginning else after.offset.serialize() + block_ms = None if block is None else int(block.total_seconds() * 1000) + if block_ms is not None and block_ms <= 0: + block_ms = None + while True: + found: Any = await self._client.xread( + {name: start}, count=max_records, block=block_ms + ) + entries = found[0][1] if found else [] + if not entries: + return [] + records = [ + _to_record(entry_id, fields) + for entry_id, fields in entries + if not _is_staged(fields) + ] + if records: + return records + start = _text(entries[-1][0]) + block_ms = None + + async def abort_output(self, manifest: OutputStageManifest) -> OutputStage: + """Resolve the stage as aborted and take its entries out of the log. + + A batch whose task was not accepted yields nothing to any reader, and + on this provider the workflow's own entries share the log with the + producers' records, so leaving them would count them against every + window and bound the log holds. The stage's hashes stay, status and + offsets included, so a reader that captured the barrier settles it the + same way. + """ + stage = await super().abort_output(manifest) + if stage.records: + await self._client.xdel( + self.stream_key(manifest.stream_key), + *(record.offset.serialize() for record in stage.records), + ) + return stage + + async def _output_stage_from_offsets( + self, + manifest: OutputStageManifest, + status: OutputStageStatus, + encoded_offsets: str, + ) -> OutputStage: + # An aborted stage's entries are gone from the log here, so they are + # stood in for by their positions: the stage still names what it held. + if status is not OutputStageStatus.ABORTED: + return await super()._output_stage_from_offsets( + manifest, status, encoded_offsets + ) + offsets = [Offset(value) for value in encoded_offsets.split(",") if value] + if len(offsets) != manifest.record_count: + raise StreamIntegrityError( + "an output stage's stored offsets do not match its manifest" + ) + return OutputStage( + manifest, + tuple( + OutputStreamRecord(TransportRecordKind.DATA, b"", offset) + for offset in offsets + ), + status, + ) + + async def read_range( + self, key: StreamKey, first: Offset, last: Offset + ) -> list[TransportRecord]: + entries: Any = await self._client.xrange( + self.stream_key(key), first.serialize(), last.serialize() + ) + # The same entries the live read dropped, so the range holds the count + # the marker recorded. + return [ + _to_record(entry_id, fields) + for entry_id, fields in entries + if key.direction is StreamDirection.OUTPUT or not _is_staged(fields) + ] + + async def tail_cursor(self, key: StreamKey, *, before_last: int = 0) -> Any: + """The boundary the newest ``before_last`` records begin after. + + Through an input key the workflow's own staged entries are not + records, as on every read the transport makes; through an output key + only an aborted batch's entries are left out, since a pending one may + still commit. + """ + name = self.stream_key(key) + skip_in = ( + self._staged_in + if key.direction is StreamDirection.INPUT + else self._aborted_in(key) + ) + after = await _tail_after(self._client, name, before_last, skip_in=skip_in) + return TRANSPORT_BEGINNING if after is None else AFTER(after) + + @staticmethod + async def _staged_in(page: list[Any]) -> set[str]: + return {_text(entry_id) for entry_id, fields in page if _is_staged(fields)} + + def _aborted_in(self, key: StreamKey) -> _SkipIn: + async def aborted(page: list[Any]) -> set[str]: + staged = { + _text(entry_id): _text(fields[_STAGE_FIELD]) + for entry_id, fields in page + if _is_staged(fields) + } + if not staged: + return set() + stage_ids = sorted(set(staged.values())) + statuses = dict( + zip( + stage_ids, + await self._client.hmget(self._output_status_key(key), stage_ids), + ) + ) + return { + entry_id + for entry_id, stage_id in staged.items() + if _text(statuses.get(stage_id) or "") + == OutputStageStatus.ABORTED.value + } + + return aborted + + +def _drive(coroutine: Coroutine[Any, Any, None]) -> None: + """Run a transport publish to completion without yielding to the loop. + + The transport's publish only awaits when the task's batch is full. A + synchronous publish has nowhere to wait, so a publish that would have + waited fails the task instead, loudly. + """ + try: + coroutine.send(None) + except StopIteration: + return + coroutine.close() + raise StreamError( + "this Workflow Task's output batch is full, and a synchronous publish " + "cannot wait for the worker to stage it" + ) + + +def _parse(cursor: Cursor, raw: bytes, warn: Any) -> WireRecord | None: + try: + return WireRecord.FromString(raw) + except DecodeError as error: + # Same answer as an undecodable body: skip and say so, so one bad + # record cannot pin a reader. + warn("skipping stream record at %s: %s", cursor, error) + return None + + +class _RedisReadSource: + """One subscription of the running workflow, over the transport's input key.""" + + def __init__(self, subscription: Any) -> None: + self._subscription = subscription + self._records = subscription.records() + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + while True: + # One record per batch: the transport reports readiness per record, + # and a batch here would invent a boundary replay never observed. + offset, body = await self._records.__anext__() + cursor = mint_cursor(_PROVIDER, offset.token) + wire = _parse(cursor, body, workflow.logger.warning) + if wire is not None: + return [(cursor, wire)] + + def close(self) -> None: + self._subscription.close() + + +class _RedisWriteSink: + def __init__(self, topic: str) -> None: + self._topic = _UNBOUNDED_OUTPUT.topic(topic, type=bytes) + + def publish(self, record: WireRecord) -> None: + # Buffered in the transport's per-task batch, staged by the worker when + # the task completes and promoted once History proves the task was + # accepted: rule 1 through the transport's own commit. The hash is a + # digest over bytes already in hand, so it costs no I/O here. + _stamp_hash(record) + _drive(self._topic.publish(record.SerializeToString())) + + +class _RedisWorkflowProvider: + """The workflow half: the transport's subscriptions and staged output.""" + + def __init__(self, idle_timeout: timedelta) -> None: + self._input = external_stream.with_options(idle_timeout=idle_timeout) + + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + _require_topic(topic) + check_read_start(after, last) + if last is not None or after == END: + # The tail is where the log is when the worker looks, which the + # workflow thread cannot see. + raise StreamUnsupportedError( + "a workflow reader on the redis provider starts at a cursor or at " + "the beginning" + ) + subscribe = self._input.topic(topic, type=bytes).subscribe + position = _position(after) + # Without a position the transport resumes where the chain's + # predecessor run committed; with one, that is where the wait starts + # and what the marker's header records. + start = None if position is None else AFTER(position) + return _RedisReadSource(subscribe(start_cursor=start)) + + def open_writer(self, topic: str) -> WriteSink: + _require_topic(topic) + return _RedisWriteSink(topic) + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +async def _chain(client: Client, workflow_id: str) -> WorkflowChainKey: + try: + description = await client.get_workflow_handle(workflow_id).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {workflow_id!r} was not found" + ) from error + raise + return WorkflowChainKey( + client.namespace, + workflow_id, + description.raw_description.workflow_execution_info.first_run_id, + ) + + +#: A chain still has a consumer while a run is in either of these: a run that +#: continued as new hands the stream to its successor rather than ending it. +_STILL_CONSUMING: Final = ( + WorkflowExecutionStatus.RUNNING, + WorkflowExecutionStatus.CONTINUED_AS_NEW, +) + +#: What the server says when a Signal reaches a run whose Workflow Task is +#: closing it: on a continue-as-new handover for an instant, and after a task +#: that tried to close while the Signal sat buffered, until the next one settles. +_CLOSING_REFUSAL: Final = "workflow is closing" + + +def _chain_ended(error: WakeNotAcknowledgedError) -> bool: + """Whether the server refused a wake because the chain has ended.""" + cause = error.__cause__ + return isinstance(cause, RPCError) and cause.status == RPCStatusCode.NOT_FOUND + + +def _refused_as_closing(error: WakeNotAcknowledgedError) -> bool: + """Whether the server refused a wake because the run is closing.""" + return _CLOSING_REFUSAL in str(error) + + +def _storage_error(error: Exception, what: str) -> StreamError: + return StreamError(f"{what}: {error}") + + +def _integrity_error(error: Exception, what: str) -> StreamError: + # Named as a loss rather than a transient read failure, because no retry brings + # a lost record back and the caller's next move is different. + return StreamNotFoundError(f"{what}: {error}") + + +class RedisProducer(Generic[T]): + """Appends to a topic from outside workflow code. + + Every append is visible as soon as the store accepts it. Each record goes + to the topic's log, where the workflow's subscription reads it and where + outside readers see it beside the workflow's own records, and a workflow + subscribed to the topic is woken; the cursor returned names the entry. + """ + + def __init__( + self, + streams: RedisStreams, + client: Client, + workflow_id: str | None, + topic: str, + producer_id: str, + attempt: int, + ) -> None: + """Bind this producer to ``topic`` of the chain.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._topic = topic + self._producer_id = producer_id + self._attempt = attempt + self._converter = client.data_converter.payload_converter + self._sequence = 0 + self._last = BEGINNING + self._input: Any = None + self._append: _LogAppend | None = None + self._name: str | None = None + self._codec: StreamPayloadCodec[bytes] | None = None + + @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 _session(self) -> str: + # The transport dedupes on this and the sequence. The attempt is part + # of it so a retried append is dropped while a new generation writing + # different words at the same sequence is not. + return ( + f"{self._producer_id}#{self._attempt}" + if self._attempt + else self._producer_id + ) + + async def append(self, *values: T) -> Cursor: + """Append ``values`` and return the cursor of the last record as stored. + + A repeat returns where the original landed, because the transport + answers a byte-identical repeat with the original offset; an empty + call returns the position of this producer's last record. + """ + if not values: + return self._last + return await self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.DATA, + value=value, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence + index, + ) + for index, value in enumerate(values) + ] + ) + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic. + + Written as an ordinary record rather than the transport's own + terminal, which would end every outside read of the topic; a + ``FINISH`` here only says this producer is done. + """ + await self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.FINISH, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + ) + ] + ) + + async def _connect(self) -> None: + if self._codec is not None: + return + backend = self._streams._require_backend() + self._append = _LogAppend(backend) + assert self._workflow_id is not None + chain = await _chain(self._client, self._workflow_id) + try: + # The transport's producer is bound for its wake and its key: the + # append itself is the provider's, so a retry is matched by digest. + input_ = await ExternalStreamProducer.connect( + backend=backend, + workflow=chain, + client=self._client, + session_id=self._session, + ) + except ChainKeyMismatchError as error: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} is not the chain this producer " + f"was opened on: {error}" + ) from error + self._input = input_.topic(self._topic, type=bytes) + self._name = backend.stream_key(self._input.stream_key) + self._codec = StreamPayloadCodec( + self._client.data_converter.with_context( + WorkflowSerializationContext( + namespace=chain.namespace, workflow_id=chain.workflow_id + ) + ), + bytes, + ) + + async def _write(self, records: list[WireRecord]) -> Cursor: + await self._connect() + try: + last = await self._place(records) + except AppendConflictError as error: + raise StreamProducerError( + f"producer {self._session!r} already wrote a different record at " + f"sequence {error.key.sequence}" + ) from error + except AppendNotAcknowledgedError as error: + raise _storage_error(error, "an append was not acknowledged") from error + except TransportStreamError as error: + raise _storage_error(error, "the store refused an append") from error + self._sequence += len(records) + assert last is not None + self._last = mint_cursor(_PROVIDER, last.token) + return self._last + + async def _place(self, records: list[WireRecord]) -> Offset | None: + """Append ``records`` to the owner's log and return the last position.""" + assert self._codec is not None and self._append is not None + assert self._name is not None + last: Offset | None = None + for index, record in enumerate(records): + # The digest is taken and stamped before the codec runs, so a retry + # the codec encodes differently still matches its original. + digest = _plaintext_digest(record) + _stamp_hash(record) + # Built here rather than handed to the transport's publish, so the + # digest rides along. The identity is the one the transport would + # derive, so a record the log already holds is reused. + staged = TransportRecord( + kind=TransportRecordKind.DATA, + payload=await self._codec.encode(record.SerializeToString()), + producer_session_id=self._session, + sequence=self._sequence + index, + ) + last = await self._append.write( + name=self._name, record=staged, digest=digest + ) + await self._wake(last) + return last + + async def _wake(self, position: Offset | None) -> None: + """Wake the consuming workflow, following the chain past a closing run. + + The records are appended before this is called; what can fail is + telling the consumer. The server refuses a Signal while the run it + resolves to is closing, and a consumer that reads a terminal record + straight from the store and continues as new on it closes in exactly + that way, ahead of the wake for that record. The streams are keyed by + the chain, so the records are already where the successor reads them + and only the wake has to follow: it is sent again, addressed to the + workflow id as every wake is, until the chain's current run takes it. + A chain that has ended instead is the ordinary ending of a terminal + record racing the consumer acting on it, not an error. + + A run that still refuses when the window passes is inside a Workflow + Task that tried to close it while this wake sat buffered, which the + server answers by failing that task and holding the run closed to + Signals until the next one settles. The wake is dropped then: if the + run closes, nothing is owed; if it does not, its next park rechecks + the log and finds the record, the transport's own rule for a record + appended before a park. Any other refusal that outlasts the window is + raised. + + ``NOT_FOUND`` answers the chain question without a describe: the wake + call names the chain's first run, and the server refuses it that way + only once the chain has ended. A Signal, which names the Workflow ID + alone, refuses an ended run with it too. + """ + deadline = time.monotonic() + _WAKE_RETRY_WINDOW.total_seconds() + while True: + try: + await self._input.wake(position=position) + return + except WakeNotAcknowledgedError as error: + if _chain_ended(error) or await self._chain_is_terminal(): + return + if time.monotonic() >= deadline: + if _refused_as_closing(error): + logger.info( + "dropping the wake for %r on topic %r: the run is closing a " + "Workflow Task, and its next park rechecks the log", + self._workflow_id, + self._topic, + ) + return + raise + except TransportStreamError as error: + raise _storage_error(error, "the wake could not be sent") from error + await asyncio.sleep(_WAKE_RETRY_BACKOFF.total_seconds()) + + async def _chain_is_terminal(self) -> bool: + """Whether the chain has ended for good, rather than handing over. + + A run that continued as new is not the end: its successor is the + consumer now, and a closing run still describes as running. A chain + whose History is gone has ended. + """ + assert self._workflow_id is not None + handle = self._client.get_workflow_handle(self._workflow_id) + try: + status = (await handle.describe()).status + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + return True + raise + return status is not None and status not in _STILL_CONSUMING + + +class RedisStreamHandle: + """One workflow's topics from outside, over the transport's output streams.""" + + def __init__( + self, + streams: RedisStreams, + client: Client, + workflow_id: str | None, + run_id: str | None, + ) -> None: + """Address the workflow's topics; ``run_id`` decides whose close ends a read.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + if workflow_id is None: + raise ValueError("a stream handle needs a workflow_id") + converter = client.data_converter.with_context( + WorkflowSerializationContext( + namespace=client.namespace, workflow_id=workflow_id + ) + ) + self._converter = client.data_converter.payload_converter + self._codec: StreamPayloadCodec[bytes] = StreamPayloadCodec(converter, bytes) + + def read( + self, + *, + topic: str | StreamTopic[Any], + 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 workflow closes. + + That is the chain, or the pinned run. ``END`` and ``last=`` are + positioned against the log on the first step of the generator, since + this call cannot reach the store. + + A cursor another provider minted, or one that is not a Redis entry id, + is refused by this call: reading the token needs nothing from the store. + + Raises: + ValueError: ``last`` is not positive or came with a cursor. + StreamCursorError: The cursor is another provider's, or does not name + a Redis entry. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + # A tail start is resolved in the generator; a cursor is parsed here so + # a foreign one fails this call, not the first iteration. + tail = last if last is not None else (0 if after == END else None) + position = None if tail is not None else _position(after) + return self._read(topic, position, after, result_type, tail=tail) + + async def _read( + self, + topic: str, + position: Offset | None, + after: Cursor, + result_type: type | None, + *, + tail: int | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + backend = self._streams._require_backend() + assert self._workflow_id is not None + chain = await _chain(self._client, self._workflow_id) + key = chain.stream_key(topic, direction=StreamDirection.OUTPUT) + if tail is not None: + resolved = await backend.tail_cursor(key, before_last=tail) + position = None if resolved.is_beginning else resolved.offset + after = ( + BEGINNING + if position is None + else mint_cursor(_PROVIDER, position.token) + ) + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + cursor = TRANSPORT_BEGINNING if position is None else AFTER(position) + closed = False + while True: + try: + result = await backend.read_output_after( + key, + cursor, + max_records=_READ_BATCH, + block=self._streams._poll, + ) + except StreamIntegrityError as error: + # Integrity loss is permanent, and the taxonomy the transport builds + # on it is the difference between a retry that clears and one that + # never will. Filed as a store read, an operator retries forever. + raise _integrity_error(error, "the store lost records") from error + except TransportStreamError as error: + raise _storage_error(error, "the store could not be read") from error + for placed in result.records: + cursor = AFTER(placed.offset) + if placed.kind is not TransportRecordKind.DATA: + continue + minted = mint_cursor(_PROVIDER, placed.offset.token) + wire = _parse( + minted, await self._codec.decode(placed.payload), logger.warning + ) + if wire is None: + continue + for record in decoder.decode(minted, wire): + yield record + if result.pending is not None: + # A staged batch whose task History has not settled yet is a + # barrier: nothing past it is read until History says whether + # the task was accepted. + try: + resolved = await _reconcile_output_stage( + backend=backend, + client=self._client, + workflow_id=self._workflow_id, + manifest=result.pending.manifest, + ) + except StreamIntegrityError as error: + raise _integrity_error( + error, "a staged batch lost records" + ) from error + except ( + OutputStageConflictError, + OutputStageNotFoundError, + OutputStageResolutionError, + ) as error: + # The transport's own stage vocabulary does not cross the public + # surface: to a reader this is a batch that cannot be settled. + raise _storage_error( + error, "a staged batch could not be settled" + ) from error + except TransportStreamError as error: + raise _storage_error( + error, "a staged batch could not be settled" + ) from error + if not resolved: + await asyncio.sleep(self._streams._poll.total_seconds()) + continue + if result.records: + continue + if closed: + return + # One more pass after learning the workflow closed, so a batch + # promoted between the read and the describe is not lost. + closed = await self._closed() + + async def _closed(self) -> bool: + assert self._workflow_id is not None + handle = self._client.get_workflow_handle( + self._workflow_id, run_id=self._run_id + ) + try: + status = (await handle.describe()).status + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + # The run's History is gone; nothing more can be promoted. + return True + raise + if status is None or status == WorkflowExecutionStatus.RUNNING: + return False + # The streams are keyed by the chain, so a run that continued as new + # is not the end unless the handle was pinned to it. + return not (self._run_id is None and status in _STILL_CONSUMING) + + async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: + """The cursor of the newest committed record on ``topic``, for following from now. + + ``BEGINNING`` when the topic holds no committed record. + """ + topic, _ = resolve_topic(topic) + backend = self._streams._require_backend() + assert self._workflow_id is not None + chain = await _chain(self._client, self._workflow_id) + try: + tail = await backend.output_tail( + chain.stream_key(topic, direction=StreamDirection.OUTPUT) + ) + except TransportStreamError as error: + raise _storage_error(error, "the store's tail could not be read") from error + if tail.is_beginning: + return BEGINNING + assert tail.offset is not None + return mint_cursor(_PROVIDER, tail.offset.token) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this owner's stream, pinned as this handle is.""" + assert self._workflow_id is not None + 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" + ) + + def producer( + self, + *, + topic: str | StreamTopic[Any], + producer_id: str = "", + attempt: int = 0, + ) -> RedisProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return RedisProducer( + self._streams, + self._client, + self._workflow_id, + topic, + producer_id, + attempt, + ) + + +class RedisStreams(ProviderPlugin): + """The client-side provider over Redis streams. + + Construct one, pass it to the worker as a plugin and open handles from + it anywhere else. The Redis client is opened on first use, so it belongs + to the loop that uses it, and released by :meth:`close`; a ``backend`` + the caller hands in stays the caller's to close. + """ + + def __init__( + self, + *, + url: str = "redis://127.0.0.1:6379", + key_prefix: str = "temporal-streams", + idle_timeout: timedelta = timedelta(seconds=1), + client: Any | None = None, + poll_interval: timedelta = timedelta(milliseconds=500), + wake_transport: WakeTransport = "auto", + ) -> None: + """Create the provider. + + Args: + url: The Redis to connect to when no ``client`` is given. + key_prefix: Prepended to every key, so one Redis serves several + deployments. + idle_timeout: How long a workflow reader with nothing to read + holds its Workflow Task open before the worker parks it. + client: A ``redis.asyncio.Redis`` the caller opened, with + ``decode_responses=False``, and closes itself; the provider + puts its own key layout on top of it. + poll_interval: How long an outside reader that is caught up waits + for a record before asking whether the workflow closed. + wake_transport: How a producer and a worker wake a workflow after + an append. ``"channel"`` notifies the stream's channel with + the entry id as the position: on a server with channels + linked to a workflow the channel lives in the reading + workflow's own state and nothing subscribes, on one with + independent channels every subscribed reader is woken. + ``"signal"`` uses the reserved Signal, which writes a + History event. ``"auto"`` notifies the channel and steps + down to the Signal on a server without channels. + """ + # Checked against the alias so an untyped caller still gets a ValueError. + if wake_transport not in get_args(WakeTransport): + raise ValueError(f"unknown wake transport {wake_transport!r}") + self._url = url + self._key_prefix = key_prefix + self._idle_timeout = idle_timeout + self._client = client + self._backend: _TopicLogBackend | None = None + self._owned_client: Any = None + self._poll = poll_interval + self._wake_transport: WakeTransport = wake_transport + + def _require_backend(self) -> _TopicLogBackend: + if self._backend is None: + client = self._client + if client is None: + import redis.asyncio + + # A dead peer would otherwise hold a blocking read open forever. + # Several block periods, so a healthy socket that is merely idle + # inside one read window is never abandoned. + client = self._owned_client = redis.asyncio.from_url( + self._url, + decode_responses=False, + socket_timeout=DEFAULT_WATCH_BLOCK.total_seconds() * 6, + ) + self._backend = _TopicLogBackend( + client=client, + key_prefix=self._key_prefix, + wake_transport=self._wake_transport, + ) + return self._backend + + def configure_worker(self, config: WorkerConfig) -> WorkerConfig: + """Set this provider and its transport backend on the worker.""" + config = super().configure_worker(config) + config["external_stream_backend"] = self._require_backend() + return config + + def configure_replayer(self, config: ReplayerConfig) -> ReplayerConfig: + """Set this provider and its transport backend on the replayer.""" + config = super().configure_replayer(config) + config["external_stream_backend"] = self._require_backend() + return config + + def workflow_provider(self) -> _RedisWorkflowProvider: + """The workflow half, over the transport's subscriptions and staged output.""" + return _RedisWorkflowProvider(self._idle_timeout) + + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> RedisStreamHandle: + """A handle on ``workflow_id``'s topics; it follows the chain by construction.""" + return RedisStreamHandle(self, client, workflow_id, run_id) + + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> RedisStreamHandle: + """Refuse: this provider keeps no stream an activity owns. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the redis provider does not host activity streams" + ) + + 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, + ) -> RedisStreamHandle: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the redis provider does not host standalone streams" + ) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> RedisStreamHandle: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the redis provider does not host standalone streams" + ) + + async def close(self) -> None: + """Release the Redis connection this provider opened; a caller's stays open.""" + self._backend = None + client, self._owned_client = self._owned_client, None + if client is not None: + await client.aclose() From 69213bcc5797d47604f3a3cebe05f27abfe9483e Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 04:30:39 -0700 Subject: [PATCH 3/5] Covered the Redis provider. The conformance suite runs the Redis case when STREAMS_LIVE=redis, and the live module checks the staged commit, replay, resets and the channel wake against a server and a Redis. --- tests/streams/test_redis_provider.py | 288 ++++++ tests/streams/test_redis_replay.py | 1117 +++++++++++++++++++++ tests/streams/test_streams_conformance.py | 69 +- 3 files changed, 1473 insertions(+), 1 deletion(-) create mode 100644 tests/streams/test_redis_provider.py create mode 100644 tests/streams/test_redis_replay.py diff --git a/tests/streams/test_redis_provider.py b/tests/streams/test_redis_provider.py new file mode 100644 index 000000000..0f2151234 --- /dev/null +++ b/tests/streams/test_redis_provider.py @@ -0,0 +1,288 @@ +"""What the Redis provider decides without a store: cursors, the sync publish, the wake.""" + +from __future__ import annotations + +import asyncio +import time +from datetime import timedelta +from types import SimpleNamespace +from typing import Any, cast + +import pytest + +from temporalio.client import WorkflowExecutionStatus +from temporalio.contrib.external_workflow_streams import ( + StreamError as TransportStreamError, +) +from temporalio.contrib.external_workflow_streams import ( + WakeNotAcknowledgedError, +) +from temporalio.contrib.external_workflow_streams._record import Offset +from temporalio.converter import DataConverter +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import BEGINNING, Cursor, StreamCursorError, StreamError +from temporalio.streams.providers import redis as redis_provider +from temporalio.streams.providers.redis import ( + RedisProducer, + RedisStreams, + _drive, + _position, + _wake_counter, +) + + +class _ChainClient: + """A client whose describes of the chain answer from a script of statuses. + + The last entry repeats, so a chain that stays running keeps saying so. + """ + + def __init__(self, *answers: WorkflowExecutionStatus | Exception) -> None: + self.data_converter = DataConverter.default + self.namespace = "default" + self._answers = list(answers) + self.describes = 0 + + def get_workflow_handle(self, _workflow_id: str, **_: Any) -> Any: + return self + + async def describe(self) -> Any: + self.describes += 1 + answer = self._answers.pop(0) if len(self._answers) > 1 else self._answers[0] + if isinstance(answer, Exception): + raise answer + return SimpleNamespace(status=answer) + + +class _RefusingInput: + """The transport's input topic, refusing the first ``refusals`` wakes.""" + + def __init__(self, refusals: int, error: Exception | None = None) -> None: + self._refusals = refusals + self._error = error + self.calls = 0 + self.positions: list[Any] = [] + + async def wake(self, *, position: Any = None) -> list[str]: + self.calls += 1 + self.positions.append(position) + if self.calls <= self._refusals: + raise self._error or WakeNotAcknowledgedError( + "workflow operation can not be applied because workflow is closing", + pending=[], + ) + return ["sent"] + + +def _waking_producer(client: _ChainClient, wakes: _RefusingInput) -> RedisProducer[Any]: + producer: RedisProducer[Any] = RedisProducer( + RedisStreams(), cast(Any, client), "wf", "inputs", "console", 1 + ) + producer._input = wakes + return producer + + +@pytest.fixture +def quick_wake_retries(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + redis_provider, "_WAKE_RETRY_BACKOFF", timedelta(milliseconds=10) + ) + monkeypatch.setattr( + redis_provider, "_WAKE_RETRY_WINDOW", timedelta(milliseconds=300) + ) + + +@pytest.mark.usefixtures("quick_wake_retries") +async def test_a_wake_refused_by_a_closing_run_is_sent_again_to_its_successor(): + # The chain describes as running: the run that refused the wake handed over + # to a successor, which is where the records already are. + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + wakes = _RefusingInput(refusals=1) + await _waking_producer(client, wakes)._wake(None) + assert (wakes.calls, client.describes) == (2, 1) + + +@pytest.mark.usefixtures("quick_wake_retries") +async def test_a_refused_wake_waits_for_the_successor_to_become_current(): + # In the instant between two runs the chain still describes as the + # predecessor that continued as new, and the next Signal is refused too. + client = _ChainClient( + WorkflowExecutionStatus.CONTINUED_AS_NEW, WorkflowExecutionStatus.RUNNING + ) + wakes = _RefusingInput(refusals=2) + await _waking_producer(client, wakes)._wake(None) + assert (wakes.calls, client.describes) == (3, 2) + + +@pytest.mark.usefixtures("quick_wake_retries") +@pytest.mark.parametrize( + "ending", + [ + WorkflowExecutionStatus.COMPLETED, + WorkflowExecutionStatus.FAILED, + WorkflowExecutionStatus.CANCELED, + WorkflowExecutionStatus.TERMINATED, + WorkflowExecutionStatus.TIMED_OUT, + RPCError("gone", RPCStatusCode.NOT_FOUND, b""), + ], + ids=lambda ending: getattr(ending, "name", "not found"), +) +async def test_a_wake_refused_by_a_finished_chain_is_the_ordinary_ending( + ending: WorkflowExecutionStatus | Exception, +): + client = _ChainClient(ending) + wakes = _RefusingInput(refusals=99) + await _waking_producer(client, wakes)._wake(None) + assert wakes.calls == 1 + + +@pytest.mark.usefixtures("quick_wake_retries") +async def test_a_wake_refused_as_closing_for_the_whole_window_is_dropped(): + # The run is inside a Workflow Task that tried to close it while the wake + # sat buffered. The record is in the log, so if the run stays open its + # next park rechecks the log and finds it; the producer is not failed. + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + wakes = _RefusingInput(refusals=99) + await _waking_producer(client, wakes)._wake(None) + assert wakes.calls > 1 + + +@pytest.mark.usefixtures("quick_wake_retries") +async def test_a_wake_refused_for_another_reason_for_the_whole_window_is_raised(): + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + wakes = _RefusingInput( + refusals=99, error=WakeNotAcknowledgedError("signal rate limited", pending=[]) + ) + with pytest.raises(WakeNotAcknowledgedError, match="rate limited"): + await _waking_producer(client, wakes)._wake(None) + assert wakes.calls > 1 + + +@pytest.mark.usefixtures("quick_wake_retries") +async def test_a_wake_refused_as_not_found_ends_without_a_describe(): + # The server answers NOT_FOUND for a chain that has ended, whether the + # wake went as a Signal or to the owner's linked channel, so nothing is + # left to describe. + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + refusal = WakeNotAcknowledgedError("gone", pending=[]) + refusal.__cause__ = RPCError("gone", RPCStatusCode.NOT_FOUND, b"") + wakes = _RefusingInput(refusals=99, error=refusal) + await _waking_producer(client, wakes)._wake(None) + assert (wakes.calls, client.describes) == (1, 0) + + +async def test_a_wake_reports_the_entry_it_follows(): + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + wakes = _RefusingInput(refusals=0) + placed = Offset("1700000000000-4") + await _waking_producer(client, wakes)._wake(placed) + assert wakes.positions == [placed] + + +def test_a_wake_counter_follows_the_entry_id_order(): + ids = ["1700000000000-0", "1700000000000-1", "1700000000001-0", "1800000000000-7"] + counters = [_wake_counter(Offset(entry)) for entry in ids] + assert counters == sorted(counters) + assert len(set(counters)) == len(counters) + assert _wake_counter(Offset("1700000000000-3")) == 1700000000000 * 2**20 + 3 + + +def test_a_wake_counter_caps_the_sequence_inside_its_millisecond(): + capped = _wake_counter(Offset(f"5-{2**20 + 9}")) + assert capped == 5 * 2**20 + 2**20 - 1 + assert capped < _wake_counter(Offset("6-0")) + + +def test_a_wake_counter_fits_the_servers_signed_64_bit_field(): + year_2200_ms = 7258118400000 + assert _wake_counter(Offset(f"{year_2200_ms}-{2**20}")) < 2**63 + + +def test_a_positionless_wake_counter_follows_the_entry_id_rule(): + backend = RedisStreams(client=_NoRedis())._require_backend() + before = _wake_counter(Offset(f"{int(time.time() * 1000)}-{2**20 - 1}")) + now = backend.wake_counter_now() + after = _wake_counter(Offset(f"{int(time.time() * 1000) + 1}-0")) + assert before <= now < after + # Any entry appended before the call orders below it. + assert now > _wake_counter(Offset(f"{int(time.time() * 1000) - 1}-7")) + + +def test_the_wake_transport_reaches_the_backend_both_sides_share(): + streams = RedisStreams(client=_NoRedis(), wake_transport="signal") + backend = streams._require_backend() + assert backend.wake_transport == "signal" + assert backend.wake_counter_for(Offset("2-1")) == _wake_counter(Offset("2-1")) + assert RedisStreams(client=_NoRedis())._require_backend().wake_transport == "auto" + channel = RedisStreams(client=_NoRedis(), wake_transport="channel") + assert channel._require_backend().wake_transport == "channel" + + +def test_an_unknown_wake_transport_is_refused_at_construction(): + with pytest.raises(ValueError, match="wake transport"): + RedisStreams(wake_transport=cast(Any, "carrier-pigeon")) + + +async def test_a_wake_the_store_could_not_send_is_a_storage_error(): + client = _ChainClient(WorkflowExecutionStatus.RUNNING) + wakes = _RefusingInput(refusals=1, error=TransportStreamError("no connection")) + with pytest.raises(StreamError, match="could not be sent"): + await _waking_producer(client, wakes)._wake(None) + assert (wakes.calls, client.describes) == (1, 0) + + +class _NoRedis: + """Enough of a Redis client to construct a backend and render its keys.""" + + def register_script(self, _script: str) -> None: + return None + + +def test_cursors_name_entries_of_the_topics_one_log(): + assert _position(BEGINNING) is None + position = _position(Cursor("redis:1700000000000-3")) + assert position is not None and position.token == "1700000000000-3" + # A workflow reader and an outside reader name the same log, so either + # side's cursor seeds the other. The form the two-key layout minted for + # workflow readers named an entry of the input key, which is the log now. + legacy = _position(Cursor("redis:in:1700000000000-3")) + assert legacy is not None and legacy.token == "1700000000000-3" + with pytest.raises(StreamCursorError): + _position(Cursor("memory:3")) + with pytest.raises(StreamCursorError): + _position(Cursor("redis:not-an-id")) + with pytest.raises(StreamCursorError): + _position(Cursor("redis:in:not-an-id")) + + +def test_a_publish_that_completes_at_once_is_driven_to_the_end(): + done = [] + + async def publish() -> None: + done.append(True) + + _drive(publish()) + assert done == [True] + + +def test_a_publish_that_would_wait_fails_loudly(): + async def publish() -> None: + await asyncio.get_running_loop().create_future() + + async def run() -> None: + with pytest.raises(StreamError, match="batch is full"): + _drive(publish()) + + asyncio.run(run()) + + +async def test_a_client_the_caller_opened_gets_the_providers_layout_and_stays_open(): + # The layout is the provider's whichever connection it runs on, and + # closing the provider does not close a caller's client. + client = _NoRedis() + provider = RedisStreams(client=client) + backend = provider._require_backend() + assert backend._client is client + await provider.close() + assert provider._require_backend() is not backend + assert provider._require_backend()._client is client diff --git a/tests/streams/test_redis_replay.py b/tests/streams/test_redis_replay.py new file mode 100644 index 000000000..81c6f4781 --- /dev/null +++ b/tests/streams/test_redis_replay.py @@ -0,0 +1,1117 @@ +"""Live checks for the client-side (Redis) provider inside a workflow. + +The conformance suite covers the outside surface when ``STREAMS_LIVE=redis``. +This module runs the interface loop inside a workflow over the staged commit, +lets a read end with the workflow, shares a topic between an outside producer +and the workflow, queries a completed run, which replays it, and seeds a +workflow reader from a cursor. All need a dev server (``TEMPORAL_ADDRESS``) +and a Redis (``TEMPORAL_TEST_REDIS_URL`` or ``AI198_REDIS_URL``). The worker +keeps a warm cache because the transport holds the task open between records. +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from collections.abc import AsyncIterator +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.client import Client +from temporalio.common import Execution, ExecutionType +from temporalio.contrib.external_workflow_streams import ( + RecordKind as TransportRecordKind, +) +from temporalio.contrib.external_workflow_streams import StreamDirection +from temporalio.contrib.external_workflow_streams._codec import StreamPayloadCodec +from temporalio.contrib.external_workflow_streams._output_backend import ( + OutputStageManifest, + OutputStageStatus, + StagedOutputRecord, +) +from temporalio.converter import WorkflowSerializationContext +from temporalio.streams import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamProducerError, +) +from temporalio.streams._wire import to_wire +from temporalio.streams.providers import redis as redis_provider +from temporalio.streams.providers.redis import RedisStreams, _chain +from temporalio.worker import Replayer, Worker +from tests.streams.test_streams_conformance import StreamHost, take +from tests.streams.test_streams_workflow import ( + DECISIONS, + INPUTS, + ContractLoop, + OneLine, +) + +pytestmark = pytest.mark.skipif( + os.environ.get("STREAMS_LIVE") != "redis", + reason="needs a live server and redis; run with STREAMS_LIVE=redis", +) + + +def redis_url() -> str: + return os.environ.get("TEMPORAL_TEST_REDIS_URL") or os.environ.get( + "AI198_REDIS_URL", "redis://127.0.0.1:6379" + ) + + +@pytest.fixture +async def provider() -> AsyncIterator[RedisStreams]: + # A prefix per case, because the store keeps what earlier cases wrote. + streams = RedisStreams( + url=redis_url(), key_prefix=f"streams-redis-{uuid.uuid4().hex}" + ) + try: + yield streams + finally: + await streams.close() + + +@pytest.fixture +async def live_client(client: Client) -> Client: + # The test environment's own server, unless TEMPORAL_ADDRESS names another. + address = os.environ.get("TEMPORAL_ADDRESS") + return await Client.connect(address) if address else client + + +async def test_interface_loop_over_redis(live_client: Client, provider: RedisStreams): + workflow_id = f"streams-redis-live-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[ContractLoop], + plugins=[provider], + max_cached_workflows=100, + ): + handle = await live_client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_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() + + records = await take(stream.read(topic=DECISIONS), 4, 60) + assert [r.kind for r in records] == [ + RecordKind.DATA, + RecordKind.DATA, + RecordKind.DATA, + RecordKind.FINISH, + ] + assert [r.value["decided"] for r in records[:3]] == [1, 2, 3] + assert all(r.producer_id == "" and r.topic == DECISIONS.name for r in records) + assert await handle.result() == [ + {"kind": "decision", "n": 1, "attempt": 1}, + {"kind": "decision", "n": 2, "attempt": 1}, + {"kind": "decision", "n": 3, "attempt": 1}, + {"kind": "finish", "producer": "model"}, + ] + + # The read ends by itself once the workflow is closed and every + # promoted record has been handed over. + async def read_everything() -> list[Any]: + return [r.value async for r in stream.read(topic=DECISIONS)] + + assert await asyncio.wait_for(read_everything(), 60) == [ + {"decided": 1}, + {"decided": 2}, + {"decided": 3}, + None, + ] + # The producer's own records are readable from outside as well, on + # the topic it wrote. + inputs = await take(stream.read(topic=INPUTS), 4, 60) + assert [r.value for r in inputs[:3]] == [{"n": 1}, {"n": 2}, {"n": 3}] + assert inputs[3].kind is RecordKind.FINISH + assert all(r.producer_id == "model" and r.attempt == 1 for r in inputs) + + +async def test_an_outside_producer_and_the_workflow_share_a_topic( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-shared-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[OneLine], + plugins=[provider], + ): + handle = await live_client.start_workflow( + OneLine.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_client, workflow_id) + await stream.producer(topic=DECISIONS, producer_id="tool", attempt=1).append( + {"from": "producer"} + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 60) + # Both writers land on one topic, each under its own identity. The order + # between them is whatever the store took first. + assert sorted( + ((r.producer_id, r.kind, r.value) for r in records), key=str + ) == sorted( + [ + ("tool", RecordKind.DATA, {"from": "producer"}), + ("", RecordKind.DATA, {"from": "workflow"}), + ("", RecordKind.FINISH, None), + ], + key=str, + ) + + +@workflow.defn +class SignalWokenPublish: + """Publishes in the task a signal wakes, then completes in that same task.""" + + def __init__(self) -> None: + self._closed = False + + @workflow.signal + def close(self) -> None: + self._closed = True + + @workflow.run + async def run(self) -> int: + out = workflow.stream_writer("out") + out.publish({"n": 0}) + await workflow.wait_condition(lambda: self._closed) + for n in range(1, 4): + out.publish({"n": n}) + return 4 + + @workflow.query + def probe(self) -> int: + return 1 + + +async def test_query_after_completion_replays_the_final_task( + live_client: Client, provider: RedisStreams +): + # A query against a completed run replays it. The final task's publishes + # were woken by a signal, which the activation applies before the replay + # marker, so the marker's expectations have to be installed before that + # task's code runs or the replay records nothing against a manifest of + # three. + workflow_id = f"streams-redis-replay-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[SignalWokenPublish], + plugins=[provider], + ): + handle = await live_client.start_workflow( + SignalWokenPublish.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_client, workflow_id) + await take(stream.read(topic="out", result_type=dict), 1, timeout=60) + await handle.signal(SignalWokenPublish.close) + assert await handle.result() == 4 + assert await handle.query(SignalWokenPublish.probe) == 1 + values = [r.value async for r in stream.read(topic="out", result_type=dict)] + assert values == [{"n": 0}, {"n": 1}, {"n": 2}, {"n": 3}] + + +async def _log_length( + streams: RedisStreams, client: Client, workflow_id: str, topic: str +) -> int: + """How many entries the topic's log holds. + + Both of the transport's keys for the topic render onto the log, which is + what makes it one; asked through both so a test would notice if they came + apart. + """ + import redis.asyncio + + backend = streams._require_backend() + chain = await _chain(client, workflow_id) + store = redis.asyncio.from_url(redis_url()) + try: + by_input = backend.stream_key( + chain.stream_key(topic, direction=StreamDirection.INPUT) + ) + by_output = backend.stream_key( + chain.stream_key(topic, direction=StreamDirection.OUTPUT) + ) + assert by_input == by_output + return await store.xlen(by_input) + finally: + await store.aclose() + + +@workflow.defn +class ResumeAfter: + """Reads ``inputs``; the first run hands its first record's cursor to the next.""" + + def __init__(self) -> None: + self._seen: list[int] = [] + + @workflow.query + def seen(self) -> list[int]: + return self._seen + + @workflow.run + async def run(self, after: str | None) -> list[int]: + reader = workflow.stream_reader( + INPUTS, after=Cursor(after) if after else BEGINNING + ) + first: str | None = None + async for record in reader: + if record.kind is RecordKind.FINISH: + break + assert record.value is not None + self._seen.append(record.value["n"]) + first = first or record.cursor.token + if after is None and len(self._seen) == 2: + workflow.continue_as_new(first) + return self._seen + + +async def test_a_workflow_reader_started_from_a_cursor_skips_the_earlier_records( + live_client: Client, provider: RedisStreams +): + # The first run reads two records and continues as new with the cursor of + # the first. Without the cursor the successor would resume after the + # second, which is where the chain left off; with it, it reads the second + # again and then the rest. + workflow_id = f"streams-redis-resume-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[ResumeAfter], + plugins=[provider], + max_cached_workflows=100, + ): + handle = await live_client.start_workflow( + ResumeAfter.run, None, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + await producer.finish() + assert await handle.result() == [2, 3] + # The completed run is evicted, so the query replays it cold with the + # seeded position and the recorded ranges. + assert await handle.query(ResumeAfter.seen) == [2, 3] + + # Both runs replay offline against the store, the seeded one included. + replayer = Replayer(workflows=[ResumeAfter], plugins=[provider]) + await replayer.replay_workflow( + await live_client.get_workflow_handle( + workflow_id, run_id=handle.first_execution_run_id + ).fetch_history() + ) + await replayer.replay_workflow(await handle.fetch_history()) + + +@workflow.defn +class BatchConsumer: + """Applies one batch of ``inputs`` per run, handing over at the sender's FINISH.""" + + @workflow.run + async def run(self, applied: list[str]) -> list[str]: + stop = False + async for record in workflow.stream_reader(INPUTS): + if record.kind is RecordKind.FINISH: + break + if record.kind is not RecordKind.DATA: + continue + assert record.value is not None + if record.value["op"] == "stop": + stop = True + continue + applied.append(record.value["op"]) + if stop: + return applied + workflow.continue_as_new(applied) + + +async def _current_open_run( + client: Client, workflow_id: str, previous: str | None +) -> str: + """Wait until the chain's newest run is a new one and still open.""" + while True: + description = await client.get_workflow_handle(workflow_id).describe() + assert description.run_id is not None + if description.run_id != previous and description.close_time is None: + return description.run_id + await asyncio.sleep(0.1) + + +async def test_a_wake_refused_by_a_run_continuing_as_new_reaches_its_successor( + live_client: Client, provider: RedisStreams +): + # The consumer reads the sender's FINISH straight from the store while it + # holds its task open and continues as new on it, so the wake Signal for + # that record resolves to a run that is closing and the server refuses it. + # The records are keyed by the chain and already where the successor reads + # them; the wake has to follow, and finish() must not fail the sender. + workflow_id = f"streams-redis-batches-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[BatchConsumer], + plugins=[provider], + max_cached_workflows=100, + ): + handle = await live_client.start_workflow( + BatchConsumer.run, [], id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_client, workflow_id) + run_id: str | None = None + for number, batch in enumerate([["a", "b"], ["c"], ["d", "stop"]], start=1): + run_id = await asyncio.wait_for( + _current_open_run(live_client, workflow_id, run_id), 30 + ) + # A sender per batch: the stream spans the chain, so one identity + # numbering its records from the start again would collide with + # the batch before. + sender = stream.producer( + topic=INPUTS, producer_id=f"console-{number}", attempt=1 + ) + for op in batch: + await sender.append({"op": op}) + await sender.finish() + assert await asyncio.wait_for(handle.result(), 30) == ["a", "b", "c", "d"] + + +async def _stage_one( + streams: RedisStreams, + client: Client, + workflow_id: str, + run_id: str, + topic: str, + values: list[Any], +) -> tuple[Any, Any]: + """Stage ``values`` on ``topic``'s output key without committing them. + + Staged the way the worker stages a task's publishes, so the records read back as + records rather than as bytes the reader has to skip. + """ + backend = streams._require_backend() + chain = await _chain(client, workflow_id) + key = chain.stream_key(topic, direction=StreamDirection.OUTPUT) + codec: StreamPayloadCodec[bytes] = StreamPayloadCodec( + client.data_converter.with_context( + WorkflowSerializationContext( + namespace=client.namespace, workflow_id=workflow_id + ) + ), + bytes, + ) + records = [] + for index, value in enumerate(values): + wire = to_wire( + client.data_converter.payload_converter, + topic=topic, + kind=RecordKind.DATA, + value=value, + ) + records.append( + StagedOutputRecord( + kind=TransportRecordKind.DATA, + payload=await codec.encode(wire.SerializeToString()), + publish_index=index, + ) + ) + manifest = OutputStageManifest( + stream_key=key, + provider_id=backend.provider_id, + provider_format_version=backend.provider_format_version, + stage_token=uuid.uuid4().hex, + run_id=run_id, + history_floor_event_id=3, + sub_batch_id=0, + fingerprint_version=1, + fingerprint=b"\x01" * 32, + record_count=len(records), + logical_byte_count=sum(len(r.payload) for r in records), + ) + await backend.stage_output(manifest, records) + return backend, manifest + + +async def test_a_pending_stage_is_settled_by_the_reader_rather_than_wedging_it( + live_client: Client, provider: RedisStreams +): + # The commit protocol's own failure modes, which nothing else here reaches: a + # stage whose task never landed in History is a barrier until the reader + # reconciles it, the reconcile aborts it, and the records it held are never + # handed over while everything behind it flows. + workflow_id = f"streams-redis-stage-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + + backend, manifest = await _stage_one( + provider, + live_client, + workflow_id, + handle.first_execution_run_id or "", + INPUTS.name, + [{"n": 10}, {"n": 11}], + ) + assert ( + await backend.output_stage(manifest) + ).status is OutputStageStatus.PENDING + + # Behind the stage so the read has to get past it to reach this. + await producer.append({"n": 2}) + + assert [r.value for r in await take(stream.read(topic=INPUTS), 2, 30)] == [ + {"n": 1}, + {"n": 2}, + ] + # The task never reached History, so the reader settled the stage as + # aborted and its records are gone rather than pending forever. + assert ( + await backend.output_stage(manifest) + ).status is OutputStageStatus.ABORTED + + await handle.signal(StreamHost.release) + await handle.result() + + +async def test_an_aborted_stage_yields_nothing_and_stops_being_a_barrier( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-abort-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + backend, manifest = await _stage_one( + provider, + live_client, + workflow_id, + handle.first_execution_run_id or "", + INPUTS.name, + [{"n": 10}], + ) + await backend.abort_output(manifest) + await producer.append({"n": 2}) + + await handle.signal(StreamHost.release) + await handle.result() + assert [r.value async for r in stream.read(topic=INPUTS)] == [ + {"n": 1}, + {"n": 2}, + ] + + +async def test_a_committed_stage_is_released_in_the_order_it_was_staged( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-commit-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(live_client, workflow_id) + backend, manifest = await _stage_one( + provider, + live_client, + workflow_id, + handle.first_execution_run_id or "", + INPUTS.name, + [{"n": 10}, {"n": 11}], + ) + await backend.commit_output(manifest) + + assert await stream.latest(topic=INPUTS) != BEGINNING + await handle.signal(StreamHost.release) + await handle.result() + assert [r.value async for r in stream.read(topic=INPUTS)] == [ + {"n": 10}, + {"n": 11}, + ] + + +async def test_a_producer_record_lands_once_however_often_it_is_sent( + live_client: Client, provider: RedisStreams +): + # One log per topic: an outside record is written once, where the + # workflow's subscription and outside readers both find it. + workflow_id = f"streams-redis-once-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + producer = provider.get_stream_handle(live_client, workflow_id).producer( + topic=INPUTS, producer_id="model", attempt=1 + ) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + assert await _log_length(provider, live_client, workflow_id, INPUTS.name) == 3 + + # A producer that comes back with a fresh sequence re-appends the same + # identities, which the script reuses rather than doubling. + again = provider.get_stream_handle(live_client, workflow_id).producer( + topic=INPUTS, producer_id="model", attempt=1 + ) + await again.append({"n": 1}, {"n": 2}, {"n": 3}) + assert await _log_length(provider, live_client, workflow_id, INPUTS.name) == 3 + + await handle.signal(StreamHost.release) + await handle.result() + + +async def test_a_repeat_under_one_identity_with_other_bytes_writes_nothing( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-conflict-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(live_client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).append( + {"n": 1} + ) + other = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + with pytest.raises(StreamProducerError, match="sequence 0"): + await other.append({"n": 99}) + # The refusal wrote nothing. + assert await _log_length(provider, live_client, workflow_id, INPUTS.name) == 1 + await handle.signal(StreamHost.release) + await handle.result() + + +async def test_an_identity_claimed_with_other_bytes_refuses_the_first_append( + live_client: Client, provider: RedisStreams +): + # The idempotency hash is read before the log is touched, so a claim that + # disagrees with the record refuses it without writing. + workflow_id = f"streams-redis-claimed-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + backend = provider._require_backend() + chain = await _chain(live_client, workflow_id) + key = chain.stream_key(INPUTS.name, direction=StreamDirection.OUTPUT) + await backend._client.hset( + backend._idempotency_key(key), "model#1/0", "1-0|" + "0" * 64 + ) + + producer = provider.get_stream_handle(live_client, workflow_id).producer( + topic=INPUTS, producer_id="model", attempt=1 + ) + with pytest.raises(StreamProducerError, match="sequence 0"): + await producer.append({"n": 1}) + + assert await _log_length(provider, live_client, workflow_id, INPUTS.name) == 0 + await handle.signal(StreamHost.release) + await handle.result() + + +@workflow.defn +class NudgedLoop: + """The contract loop with a signal that does nothing but complete a task.""" + + def __init__(self) -> None: + self._nudges = 0 + + @workflow.signal + def nudge(self) -> None: + self._nudges += 1 + + @workflow.run + async def run(self) -> list[int]: + decisions = workflow.stream_writer(DECISIONS) + seen: list[int] = [] + async for record in workflow.stream_reader(INPUTS): + if record.kind is RecordKind.FINISH: + break + if record.kind is not RecordKind.DATA: + continue + assert record.value is not None + decisions.publish({"decided": record.value["n"]}) + seen.append(record.value["n"]) + decisions.finish() + return seen + + +async def _completed_tasks(client: Client, workflow_id: str, run_id: str) -> list[int]: + return [ + event.event_id + async for event in client.get_workflow_handle( + workflow_id, run_id=run_id + ).fetch_history_events() + if event.HasField("workflow_task_completed_event_attributes") + ] + + +_MARKER_DETAILS_KEY = "external_stream" + + +async def _consuming_tasks(client: Client, workflow_id: str, run_id: str) -> list[int]: + """The completed Workflow Tasks whose marker delivered a record, in order. + + Read from the markers rather than counted by position: which task opens + the reader, and whether that task can stay open for the first record, + depends on the server, so a task index names a different task from one + server to the next while the marker says what each task consumed. + """ + from temporalio.bridge.proto.external_data import ExternalStreamMarkerData + from temporalio.contrib.external_workflow_streams._annotation import ( + decode_annotation, + ) + + handle = client.get_workflow_handle(workflow_id, run_id=run_id) + completed: int | None = None + consuming: list[int] = [] + async for event in handle.fetch_history_events(): + if event.HasField("workflow_task_completed_event_attributes"): + completed = event.event_id + continue + if not event.HasField("marker_recorded_event_attributes"): + continue + details = event.marker_recorded_event_attributes.details + if _MARKER_DETAILS_KEY not in details or completed is None: + continue + data = ExternalStreamMarkerData() + data.ParseFromString(details[_MARKER_DETAILS_KEY].payloads[0].data) + annotation = decode_annotation(data.replay_annotation) + if any(segment.runs for segment in annotation.segments): + consuming.append(completed) + return consuming + + +async def _reset_at( + client: Client, workflow_id: str, run_id: str, *, finish_event_id: int +) -> str: + """Reset ``run_id`` at the Workflow Task completed by ``finish_event_id``. + + The server keeps History up to that task's completion and runs the task + again, so the tasks before it are inherited and the task itself is not. + Returns the new run id. + """ + from temporalio.api.common.v1 import WorkflowExecution + from temporalio.api.workflowservice.v1 import ResetWorkflowExecutionRequest + + response = await client.workflow_service.reset_workflow_execution( + ResetWorkflowExecutionRequest( + namespace=client.namespace, + workflow_execution=WorkflowExecution( + workflow_id=workflow_id, run_id=run_id + ), + reason="streams: reset mid-stream", + workflow_task_finish_event_id=finish_event_id, + request_id=uuid.uuid4().hex, + ) + ) + return response.run_id + + +async def _consume_two_then_nudge( + live_client: Client, provider: RedisStreams, workflow_id: str, task_queue: str +) -> tuple[str, Any, Any]: + """Two records, each consumed by a task of its own, then a nudged task; the base run.""" + handle = await live_client.start_workflow( + NudgedLoop.run, id=workflow_id, task_queue=task_queue + ) + base_run = handle.result_run_id + assert base_run is not None + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + await take(stream.read(topic=DECISIONS), 1, 60) + await producer.append({"n": 2}) + await take(stream.read(topic=DECISIONS), 2, 60) + before = len(await _completed_tasks(live_client, workflow_id, base_run)) + await handle.signal(NudgedLoop.nudge) + for _ in range(300): + if len(await _completed_tasks(live_client, workflow_id, base_run)) > before: + break + await asyncio.sleep(0.1) + return base_run, stream, producer + + +async def test_a_reset_run_replays_the_inherited_ranges_and_continues( + live_client: Client, provider: RedisStreams +): + # Reset at the completion of the task the nudge woke, which consumed and + # published nothing: both consuming tasks are inherited. Their ranges are + # re-read from the log and replayed against the inherited markers, so + # their decisions are not published again, and reading continues from the + # last inherited boundary. + workflow_id = f"streams-redis-reset-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[NudgedLoop], + plugins=[provider], + max_cached_workflows=100, + ): + base_run, stream, producer = await _consume_two_then_nudge( + live_client, provider, workflow_id, f"tq-{workflow_id}" + ) + last_task = (await _completed_tasks(live_client, workflow_id, base_run))[-1] + reset_run = await _reset_at( + live_client, workflow_id, base_run, finish_event_id=last_task + ) + assert reset_run != base_run + + await producer.append({"n": 3}) + await producer.finish() + continued = live_client.get_workflow_handle(workflow_id, run_id=reset_run) + assert await asyncio.wait_for(continued.result(), 90) == [1, 2, 3] + decisions = [r async for r in stream.read(topic=DECISIONS)] + assert [(r.kind, r.value) for r in decisions] == [ + (RecordKind.DATA, {"decided": 1}), + (RecordKind.DATA, {"decided": 2}), + (RecordKind.DATA, {"decided": 3}), + (RecordKind.FINISH, None), + ] + + # Offline, the reset run's History reads the inherited ranges from the log. + await Replayer(workflows=[NudgedLoop], plugins=[provider]).replay_workflow( + await continued.fetch_history() + ) + + +async def test_a_reset_point_task_is_run_again_from_the_log( + live_client: Client, provider: RedisStreams +): + # Reset at the completion of the task that consumed the second record: the + # tasks before it are inherited, the one that consumed the first record + # among them, and this one is run again. The record it consumed is still + # in the log, so the reset run reads it again and publishes again, and an + # outside reader sees that decision from both runs. + workflow_id = f"streams-redis-reset-rerun-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[NudgedLoop], + plugins=[provider], + max_cached_workflows=100, + ): + base_run, stream, producer = await _consume_two_then_nudge( + live_client, provider, workflow_id, f"tq-{workflow_id}" + ) + second = (await _consuming_tasks(live_client, workflow_id, base_run))[1] + reset_run = await _reset_at( + live_client, workflow_id, base_run, finish_event_id=second + ) + + await producer.append({"n": 3}) + await producer.finish() + continued = live_client.get_workflow_handle(workflow_id, run_id=reset_run) + assert await asyncio.wait_for(continued.result(), 90) == [1, 2, 3] + decisions = [r.value async for r in stream.read(topic=DECISIONS)] + assert decisions == [ + {"decided": 1}, + {"decided": 2}, + {"decided": 2}, + {"decided": 3}, + None, + ] + + await Replayer(workflows=[NudgedLoop], plugins=[provider]).replay_workflow( + await continued.fetch_history() + ) + + +@workflow.defn +class EchoOnOneTopic: + """Reads ``inputs`` and answers each value on the same topic. + + Its own records land in the log it reads, so what it returns says whether + it read them back. + """ + + @workflow.run + async def run(self) -> list[Any]: + seen: list[Any] = [] + out = workflow.stream_writer(INPUTS) + async for record in workflow.stream_reader(INPUTS): + if record.kind is RecordKind.FINISH: + break + if record.kind is not RecordKind.DATA: + continue + seen.append(record.value) + out.publish({"echo": record.value}) + out.finish() + return seen + + @workflow.query + def probe(self) -> int: + return 1 + + +async def test_a_workflow_does_not_read_its_own_records_from_the_shared_log( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-echo-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[EchoOnOneTopic], + plugins=[provider], + ) as worker: + handle = await live_client.start_workflow( + EchoOnOneTopic.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + # The echoes are promoted into the same log the producer wrote to. + first_four = await take(stream.read(topic=INPUTS), 4, 60) + assert sorted( + ((r.producer_id, r.value) for r in first_four), key=str + ) == sorted( + [ + ("model", {"n": 1}), + ("model", {"n": 2}), + ("", {"echo": {"n": 1}}), + ("", {"echo": {"n": 2}}), + ], + key=str, + ) + await producer.finish() + # The workflow saw the producer's records and none of its echoes. + assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}, {"n": 2}] + assert await _log_length(provider, live_client, workflow_id, INPUTS.name) == 6 + + # A cold query replays the run against the log with its own entries + # in every recorded range; the read filters them the way the live one did. + assert await handle.query(EchoOnOneTopic.probe) == 1 + everything = [r async for r in stream.read(topic=INPUTS)] + assert [(r.producer_id, r.kind) for r in everything if r.producer_id == ""] == [ + ("", RecordKind.DATA), + ("", RecordKind.DATA), + ("", RecordKind.FINISH), + ] + + await Replayer(workflows=[EchoOnOneTopic], plugins=[provider]).replay_workflow( + await handle.fetch_history() + ) + + +@workflow.defn +class NewestTwoOnGo: + """Opens ``inputs`` at its newest two records once told to, and returns them.""" + + def __init__(self) -> None: + self._go = False + + @workflow.signal + def go(self) -> None: + self._go = True + + @workflow.run + async def run(self) -> list[Any]: + await workflow.wait_condition(lambda: self._go) + reader = workflow.stream_reader(INPUTS, last=2) + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == 2: + reader.close() + return values + + +@workflow.defn +class FromNowOnGo: + """Opens ``inputs`` at its end once told to, and returns the first value.""" + + def __init__(self) -> None: + self._go = False + + @workflow.signal + def go(self) -> None: + self._go = True + + @workflow.run + async def run(self) -> Any: + await workflow.wait_condition(lambda: self._go) + reader = workflow.stream_reader(INPUTS, after=END) + async for value in reader.values(): + reader.close() + return value["n"] + return None + + +async def _channel_support(client: Client) -> Any: + """What the server offers the readers: no channels, independent ones or linked ones.""" + from tests.contrib.external_workflow_streams.conftest import server_channel_support + + return await server_channel_support(client) + + +@pytest.mark.needs_channel_server +async def test_an_outside_producer_wakes_the_reader_through_the_channel( + live_client: Client, +): + from temporalio.contrib.external_workflow_streams._wake import ChannelSupport + + support = await _channel_support(live_client) + if support is ChannelSupport.NONE: + pytest.skip("the server does not implement notification channels") + if support is ChannelSupport.LINKED: + pytest.skip("a workflow-owned stream listens on its linked channel there") + events = await _consume_three_through_the_channel(live_client) + subscribed = _subscribed(events) + assert len(subscribed) == 1, "the run subscribes once per channel" + [index] = [ + i + for i, e in enumerate(events) + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + # Core issues the subscription after the marker of the completion that + # ends the task, so the task that opened the reader stayed retained. + assert events[index - 1].HasField("marker_recorded_event_attributes"), ( + "the subscription did not wait for the completion that ends the task" + ) + notified = _notified(events) + assert {n.channel for n in notified} == set(subscribed) + assert not any(n.HasField("linked_to") for n in notified) + + +@pytest.mark.needs_linked_server +async def test_an_outside_producer_wakes_the_reader_through_its_linked_channel( + live_client: Client, +): + """The stream's channel is the reading workflow's own, so nothing subscribes.""" + from temporalio.contrib.external_workflow_streams._wake import ChannelSupport + + if await _channel_support(live_client) is not ChannelSupport.LINKED: + pytest.skip("the server does not serve channels linked to a workflow") + events = await _consume_three_through_the_channel(live_client) + assert _subscribed(events) == [], "the owner is the listener by construction" + notified = _notified(events) + assert len({n.channel for n in notified}) == 1 + [workflow_id] = { + e.workflow_execution_started_event_attributes.workflow_id + for e in events + if e.HasField("workflow_execution_started_event_attributes") + } or {""} + owners = {Execution.from_proto(n.linked_to) for n in notified} + assert {(owner.type, owner.business_id) for owner in owners} == { + (ExecutionType.WORKFLOW, workflow_id) + } + + +async def _consume_three_through_the_channel(live_client: Client) -> list[Any]: + """Runs the contract loop with three spaced appends and returns its History. + + The channel outright rather than ``"auto"``, so a step down to the Signal + fails the checks here instead of passing quietly. Common to both channel + kinds: no Signal, a notification per wake, and the entry id as the + position with the counter derived from it, so producers and workers order + the same wakes alike. + """ + from temporalio.contrib.external_workflow_streams._record import Offset + + provider = RedisStreams( + url=redis_url(), + key_prefix=f"streams-redis-{uuid.uuid4().hex}", + wake_transport="channel", + ) + workflow_id = f"streams-redis-channel-{uuid.uuid4().hex}" + try: + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[ContractLoop], + plugins=[provider], + max_cached_workflows=100, + ): + handle = await live_client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = provider.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + # Spaced past the reader's idle timeout, so the reader parks between + # appends and only a notification from outside can move it. + for n in (1, 2, 3): + await producer.append({"n": n}) + await asyncio.sleep(2) + await producer.finish() + + result = await asyncio.wait_for(handle.result(), 60) + assert [entry.get("n") for entry in result[:3]] == [1, 2, 3] + assert result[3] == {"kind": "finish", "producer": "model"} + + events = [e async for e in handle.fetch_history_events()] + finally: + await provider.close() + signalled = [ + e for e in events if e.HasField("workflow_execution_signaled_event_attributes") + ] + assert signalled == [], "a Signal woke the reader" + notified = _notified(events) + assert notified, "no Workflow Task was scheduled with a notification" + positioned = [note for note in notified if note.position] + assert positioned, "no notification carried the store's position" + for note in positioned: + offset = Offset(note.position.decode()) + assert note.counter == redis_provider._wake_counter(offset) + return events + + +def _subscribed(events: list[Any]) -> list[str]: + return [ + e.workflow_notification_channel_subscribed_event_attributes.channel + for e in events + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + + +def _notified(events: list[Any]) -> list[Any]: + return [ + n + for e in events + if e.HasField("workflow_task_scheduled_event_attributes") + for n in e.workflow_task_scheduled_event_attributes.notifications + ] diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 1dec4d6f8..59337ec06 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -38,7 +38,7 @@ 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 Execution, ExecutionType, RawValue from temporalio.contrib.external_workflow_streams._wake import ChannelSupport from temporalio.converter import ( @@ -70,6 +70,7 @@ ) from temporalio.streams._ref import RefHandle, open_ref from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.redis import RedisStreams from temporalio.testing import WorkflowEnvironment from tests.contrib.external_workflow_streams.conftest import server_channel_support from tests.helpers import new_worker @@ -256,9 +257,75 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: provider.reset() +@workflow.defn +class StreamHost: + """Owns a stream and lingers, so outside code has a running workflow to address.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + +async def _redis_case(client: Client) -> AsyncIterator[ProviderCase]: + # The store is a Redis the test environment does not start; the server + # is the environment's own unless TEMPORAL_ADDRESS names another. + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = RedisStreams( + url=os.environ.get("TEMPORAL_TEST_REDIS_URL") + or os.environ.get("AI198_REDIS_URL", "redis://127.0.0.1:6379"), + # A prefix per setup, because the store keeps what earlier runs wrote. + key_prefix=f"streams-conformance-{uuid.uuid4().hex}", + ) + # Registered once, on the client: the host's worker inherits it and the + # cases open handles through client.get_stream_handle. + config = client.config() + config["plugins"] = [provider] + client = Client(**config) + hosts: dict[str, WorkflowHandle[Any, Any]] = {} + async with new_worker(client, StreamHost) as worker: + + async def host(workflow_id: str) -> None: + if workflow_id not in hosts: + hosts[workflow_id] = await client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + + yield ProviderCase( + "redis", + provider, + client, + host=host, + # Every stream this provider keeps belongs to a workflow. + hosts_standalone_streams=False, + # An outside append notifies the stream's channel, addressed to + # the workflow that owns the stream; the reader's run is + # subscribed to it when the task that opened the read ends where + # the server has no linked kind, and listens by construction + # where it has. + wakes_by_notification=True, + wakes_by_linked_notification=True, + ) + for handle in hosts.values(): + await handle.terminate() + await provider.close() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { "memory": _memory_case } +if os.environ.get("STREAMS_LIVE") == "redis": + SETUPS["redis"] = _redis_case _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, From 13703e75621103120df599a32860b48f72a3e441 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 04:57:15 -0700 Subject: [PATCH 4/5] Said which provider carries a reader across continue-as-new. The stream_reader docstring said nothing crosses continue-as-new, which the Redis provider contradicts: its transport resumes a successor's reader where the predecessor committed. --- temporalio/workflow/_streams.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index 550398869..5531855a4 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -242,8 +242,11 @@ def stream_reader( to whichever loop pulls first; such a call may pass no ``after``, no ``last`` and no different type. Adding a reader on a new topic is a new command, so gate it with :func:`temporalio.workflow.patched` as you would a timer. A - reader in a successor run starts a new subscription: nothing crosses - continue-as-new implicitly. + reader in a successor run starts a new subscription. Whether it picks up + where the predecessor left off is the provider's: the Redis provider + resumes an ``after=BEGINNING`` reader where the predecessor committed, + because its transport carries that boundary across continue-as-new, and + the memory provider starts it at ``after``. Args: topic: The topic, relative to this workflow's stream. From f1b65f9e1280e5b08ac51508461c52f8b6013e8c Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 04:59:29 -0700 Subject: [PATCH 5/5] Described the Redis provider's one log per topic in the changelog. The entry still said each topic was an input and an output stream, which the provider stopped doing when it moved to one shared log per topic. --- CHANGELOG.md | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 6ccb8fc4d..0d7cfb9d9 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -54,7 +54,8 @@ to include examples, links to docs, or any other relevant information. `temporalio.streams.providers.memory.MemoryStreams` is the in-memory reference provider the conformance tests run against, and `temporalio.streams.providers.redis.RedisStreams` serves the same interface - over External Workflow Streams, one topic as an input and an output stream. + over External Workflow Streams, with one Redis log per topic that the workflow + and outside readers share. - `ExternalStreamSubscription.records()` yields each value with the provider offset it was read from, for a reader that has to name where it got to.