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 a1e7b5f40..f142a4cc5 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/contrib/external_workflow_streams/__init__.py b/temporalio/contrib/external_workflow_streams/__init__.py index 80acb17ae..a04387f21 100644 --- a/temporalio/contrib/external_workflow_streams/__init__.py +++ b/temporalio/contrib/external_workflow_streams/__init__.py @@ -100,6 +100,7 @@ Offset, OffsetComparator, RecordKind, + StartAtTail, StreamRecord, ) from temporalio.contrib.external_workflow_streams._wake import WakeRequest @@ -128,6 +129,7 @@ "ExternalStreamProducerTopic", "ExternalStreamSubscription", "ExternalStreamTopic", + "StartAtTail", "IdempotencyKey", "Offset", "OffsetComparator", diff --git a/temporalio/contrib/external_workflow_streams/_api.py b/temporalio/contrib/external_workflow_streams/_api.py index f69674c25..90bfa4d85 100644 --- a/temporalio/contrib/external_workflow_streams/_api.py +++ b/temporalio/contrib/external_workflow_streams/_api.py @@ -35,6 +35,9 @@ classify_read_failure, ) from temporalio.contrib.external_workflow_streams._record import ( + Cursor, + Offset, + StartAtTail, StreamRecord, ) from temporalio.contrib.external_workflow_streams._wake import channel_for @@ -95,6 +98,8 @@ def register( wait_id: int, stream_key: StreamKey, idle_timeout: timedelta, + start_cursor: Cursor | None = None, + start_at_tail: StartAtTail | None = None, ) -> None: """Registers a wait with the Worker's subscription manager. @@ -275,7 +280,12 @@ class ExternalStreamTopic(Generic[AnyType]): value_type: type[AnyType] | None options: ExternalStreamOptions - def subscribe(self) -> ExternalStreamSubscription[AnyType]: + def subscribe( + self, + *, + start_cursor: Cursor | None = None, + start_at_tail: StartAtTail | 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 +296,18 @@ 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. + start_at_tail: Start at the stream's tail instead, or at its newest + ``last`` records. The Worker resolves the boundary against the + store after this Workflow Task and records it with the + subscription, so replay starts where the live run did rather + than asking the store again. Exclusive with ``start_cursor``. """ state = _run_state() if state.runtime is None: @@ -297,6 +319,13 @@ 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} + ) + if start_at_tail is not None: + start["start_at_tail"] = start_at_tail state.runtime.register( wait_id=wait_id, stream_key=stream_key, @@ -307,6 +336,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 +641,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 +680,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/temporalio/contrib/external_workflow_streams/_backend.py b/temporalio/contrib/external_workflow_streams/_backend.py index 90f0be3e0..6c05adaa5 100644 --- a/temporalio/contrib/external_workflow_streams/_backend.py +++ b/temporalio/contrib/external_workflow_streams/_backend.py @@ -275,6 +275,23 @@ def compare_offsets(self, left: Offset, right: Offset) -> int: inside a validation loop. """ + async def tail_cursor(self, key: StreamKey, *, before_last: int = 0) -> Cursor: + """The boundary the newest ``before_last`` records begin after. + + With ``before_last=0`` it is the boundary after the newest record, so a + read from it sees only what is appended after this call. ``BEGINNING`` + when the stream holds fewer records than asked for. Resolved on the + Worker, never on the Workflow thread, and recorded with the + subscription so replay reads it from the marker instead of asking + again. A provider that cannot answer leaves this as it is, and a + subscription that asks for a tail start fails when the Worker resolves + it. + """ + raise NotImplementedError( + f"{type(self).__name__} cannot resolve a start {before_last} records " + f"before the tail of {key}; subscribe with a cursor instead" + ) + # --- waking ------------------------------------------------------------- wake_transport: WakeTransport = "auto" diff --git a/temporalio/contrib/external_workflow_streams/_record.py b/temporalio/contrib/external_workflow_streams/_record.py index 549c5c0b3..8b4d96b46 100644 --- a/temporalio/contrib/external_workflow_streams/_record.py +++ b/temporalio/contrib/external_workflow_streams/_record.py @@ -11,7 +11,29 @@ from dataclasses import dataclass, field from typing import Final, Protocol + +@dataclass(frozen=True) +class StartAtTail: + """A subscription start the Worker resolves against the store: the tail. + + ``last`` is how many of the newest records the read begins with; zero + means the boundary after the newest record, so the read sees only what + is appended after it was opened. The Workflow thread cannot ask the store + where the tail is, so the runtime records the request and the Worker + resolves it before the watcher starts, writing the boundary into the + marker beside the subscription as it does a cursor the Workflow named. + """ + + last: int = 0 + + def __post_init__(self) -> None: + """Refuse a negative count.""" + if self.last < 0: + raise ValueError(f"last must not be negative, got {self.last}") + + __all__ = [ + "StartAtTail", "AFTER", "BEGINNING", "Cursor", diff --git a/temporalio/contrib/external_workflow_streams/_replay.py b/temporalio/contrib/external_workflow_streams/_replay.py index f136c9cbb..ac293c172 100644 --- a/temporalio/contrib/external_workflow_streams/_replay.py +++ b/temporalio/contrib/external_workflow_streams/_replay.py @@ -252,7 +252,9 @@ async def _read_range( """ try: return await backend.read_range(key, run.first_offset, run.last_offset) - except StreamStorageError: + except (StreamStorageError, StreamIntegrityError): + # A backend that can tell the range is gone reports the loss itself; + # wrapping it would file a permanent loss under a transient failure. raise except Exception as err: # Not integrity loss: nothing has been shown to be missing, only diff --git a/temporalio/contrib/external_workflow_streams/_runtime.py b/temporalio/contrib/external_workflow_streams/_runtime.py index 9e7c5e53a..d7b631758 100644 --- a/temporalio/contrib/external_workflow_streams/_runtime.py +++ b/temporalio/contrib/external_workflow_streams/_runtime.py @@ -88,6 +88,7 @@ BEGINNING, Cursor, RecordKind, + StartAtTail, StreamRecord, ) from temporalio.contrib.external_workflow_streams._replay import ReplayPlan @@ -200,6 +201,14 @@ class _SubscriptionState: generation: int = 0 """Increments each time this wait re-enters the blocked state.""" + start_pending: bool = False + """The start is a tail the Worker has yet to resolve against the store. + + Until it does, the cursors hold ``BEGINNING`` as a placeholder, no watcher + runs, and the binding stays out of every header and bindings frame, so the + marker never records the placeholder as where this wait began. + """ + announced: bool = False """Whether the *current* annotation has carried this subscription's binding. @@ -291,6 +300,9 @@ def __init__( #: would give replay whatever the stream holds now (ADR-022). self._continuation = continuation self._subscriptions: dict[int, _SubscriptionState] = {} + #: Tail starts registered in the activation under way, resolved by the + #: Worker once it returns. See :meth:`resolve_pending_starts`. + self._pending_starts: dict[int, StartAtTail] = {} #: Where the *current* annotation begins, per wait. Captured when the #: annotation begins rather than when its header is first needed -- #: lazily reading the delivery cursor would let a record delivered @@ -604,7 +616,7 @@ def _reserve_bytes(self) -> int: late = { wait_id: self._binding(state) for wait_id, state in sorted(self._subscriptions.items()) - if not state.announced + if not state.announced and not state.start_pending } if late and self._accumulator is not None: # Only with an accumulator: without one the header is about to carry @@ -1134,13 +1146,16 @@ def register( stream_key: StreamKey, idle_timeout: timedelta | None = None, start_cursor: Cursor | None = None, + start_at_tail: StartAtTail | None = None, ) -> None: """Registers a wait with the Worker's manager. Non-blocking, no I/O. The start cursor defaults to what the predecessor Run committed for this ``wait_id``, so a chain resumes where it left off without the Workflow code saying anything about it -- and a first execution gets ``BEGINNING`` - from the same path. + from the same path. A ``start_at_tail`` is resolved by the Worker after + this activation and recorded with the subscription; on replay the + recorded boundary is taken from the marker and the store is not asked. Raises: ExternalStreamCapacityError: This subscription set cannot be recorded @@ -1161,7 +1176,20 @@ def register( "external streams are not configured on this Worker; pass " "external_stream_backend=... to the Worker" ) - if start_cursor is None: + if start_at_tail is not None and start_cursor is not None: + raise ValueError( + "a subscription starts at a cursor or at the tail, not at both" + ) + pending = False + if start_at_tail is not None: + recorded = (self._replay_bindings or {}).get(wait_id) + if recorded is not None: + # The live run resolved it and the marker carries the answer. + start_cursor = recorded.start_cursor + else: + pending = True + start_cursor = BEGINNING + elif start_cursor is None: start_cursor = self.restored_start(wait_id, stream_key.stream_name) if self._replay_bindings is not None: # A subscription made while a marker is being replayed -- which is @@ -1179,6 +1207,7 @@ def register( delivery_cursor=start_cursor, consumption_cursor=start_cursor, idle_timeout=idle_timeout or self._default_idle_timeout, + start_pending=pending, ) self._subscriptions[wait_id] = state # A subscription created part-way through an annotation begins at its @@ -1194,16 +1223,57 @@ def register( del self._annotation_start[wait_id] raise self._update_reserve() - self._manager.register( - run_id=self._run_id, - wait_id=wait_id, - stream_key=stream_key, - start_cursor=start_cursor, - ) + if pending: + assert start_at_tail is not None + self._pending_starts[wait_id] = start_at_tail + else: + self._manager.register( + run_id=self._run_id, + wait_id=wait_id, + stream_key=stream_key, + start_cursor=start_cursor, + ) # A registration alone is replay-visible: it is what puts the stream in # the annotation header, without which replay cannot start. self._observed_this_activation = True + async def resolve_pending_starts(self) -> None: + """Resolves every tail start the activation just finished registered. + + On the Worker's loop, between the activation and its completion. The + Workflow thread can do no I/O, and the boundary has to be fixed before + anything is delivered on the wait, so the marker records it beside the + subscription and replay reads it there instead of asking the store, + which by then holds something else. + """ + if not self._pending_starts: + return + backend = self._backend + if backend is None: + raise RuntimeError( + "external streams are not configured on this Worker; pass " + "external_stream_backend=... to the Worker" + ) + pending, self._pending_starts = self._pending_starts, {} + for wait_id, tail in pending.items(): + state = self._subscriptions.get(wait_id) + if state is None or not state.start_pending: + continue + cursor = await backend.tail_cursor(state.stream_key, before_last=tail.last) + state.start_cursor = cursor + state.delivery_cursor = cursor + state.consumption_cursor = cursor + state.start_pending = False + self._annotation_start[wait_id] = cursor + self._check_annotation_capacity(wait_id) + self._update_reserve() + self._manager.register( + run_id=self._run_id, + wait_id=wait_id, + stream_key=state.stream_key, + start_cursor=cursor, + ) + def drain(self, wait_id: int, max_records: int | None = None) -> list[StreamRecord]: """Pops buffered records. Performs no I/O and never blocks. @@ -2026,7 +2096,7 @@ def _announce_late_subscriptions(self) -> None: late = { wait_id: state for wait_id, state in sorted(self._subscriptions.items()) - if not state.announced + if not state.announced and not state.start_pending } if not late: return @@ -2053,11 +2123,13 @@ def _header(self) -> AnnotationHeader: wrong store for all the others. """ for state in self._subscriptions.values(): - state.announced = True + if not state.start_pending: + state.announced = True return AnnotationHeader( streams={ wait_id: self._binding(state) for wait_id, state in sorted(self._subscriptions.items()) + if not state.start_pending }, ) @@ -2072,6 +2144,7 @@ def _header_preview(self) -> AnnotationHeader: streams={ wait_id: self._binding(state) for wait_id, state in sorted(self._subscriptions.items()) + if not state.start_pending }, ) diff --git a/temporalio/streams/providers/redis.py b/temporalio/streams/providers/redis.py new file mode 100644 index 000000000..4171bac6d --- /dev/null +++ b/temporalio/streams/providers/redis.py @@ -0,0 +1,2117 @@ +"""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. ``END`` and + ``last=N`` are positioned against the log when the read starts: outside, + on the first step of the generator, since the call itself cannot reach the + store; inside a workflow, by the worker right after the Workflow Task that + opened the subscription, which records the entry it resolved in the marker + beside the subscription, so replay and a cold start read it from History + and never ask the log again. +- 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. +- An activity's own streams are one Redis stream per topic, keyed by the + namespace, the workflow id, empty for a standalone activity, the run the + activity execution belongs to, and the activity id. The run is the + workflow's for a workflow's activity and the activity's own for a + standalone one, so a retry writes to the same stream, which a reader sees + as ``SUPERSEDED``, and an id started again, in a new run, starts a new + one, as it does on the server-side provider. Inside the activity the run + is in its info; a handle opened outside without one describes the owner + and takes its current run. The owner is one key component joined with + ``/``, a character the chain keys percent-encode out of every id, so no + chain key and no key derived from one can name an activity's stream. + There is no input stream, because no workflow reads these, and no + staging, because an activity's append is visible as soon as the store + accepts it. A read ends when the owner is terminal and the retained tail + is delivered: a standalone activity is described directly, and a + workflow's activity through its workflow, whose close ends the read, as + does the activity leaving the pending set once its stream exists. An + activity that never wrote has no stream, so a read on it waits for the + workflow. Nothing gates an append after that, so a late attempt still + lands, and a read that has ended does not see it. +- 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. +- Retention is trimming, with no consumer floor. By default a record older + than :data:`DEFAULT_RETENTION`, seven days, is trimmed by the next append + the provider makes to its key, an activity's stream included, whatever any + reader has reached; ``retention=None`` turns the age trim off, and + ``max_len`` adds a count cap that is off by default. A replay that reaches + a recorded range the trim removed fails its Workflow Task, an outside + cursor below the trim is refused, and a fully trimmed topic reads as + empty. A run that has to replay cold after seven days of consuming fails, + so a long-lived consumer continues as new inside the window, or is + configured with a longer one. + +- 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. +- A standalone stream is one more key scheme: ``standalone/`` as the + owner component, a hash under it that holds the policy and the seal, and one + log per topic beside the hash. ``create_stream`` writes the hash once and + refuses a different policy for an id that has one; a handle on an id with no + hash raises ``StreamNotFoundError`` at its first use. Every append applies + the policy to the topic it wrote, ``max_records`` and ``retention`` as on the + other logs and ``max_bytes`` by a byte total the hash keeps per topic, + dropping the oldest entries until the topic fits. ``close()`` sets the seal: + a later append is refused with ``StreamClosedError``, the retained records + stay readable and a read ends once it has delivered them. Nothing in + Temporal knows these keys, so no tooling lists them. + +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, Sequence +from dataclasses import dataclass, replace +from datetime import timedelta +from typing import Any, Final, Generic, TypeVar, get_args +from urllib.parse import quote + +from google.protobuf.message import DecodeError + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.client import ActivityExecutionStatus, Client, WorkflowExecutionStatus +from temporalio.contrib.external_workflow_streams import ( + AFTER, + AppendConflictError, + AppendNotAcknowledgedError, + ChainKeyMismatchError, + ExternalStreamProducer, + Offset, + StartAtTail, + 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, + StagedOutputRecord, +) +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 ( + ActivitySerializationContext, + SerializationContext, + WorkflowSerializationContext, +) +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import CONTENT_HASH_KEY, content_hash +from temporalio.streams._errors import ( + StreamClosedError, + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, +) +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__ = ["DEFAULT_RETENTION", "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 record is kept when the constructor is not told otherwise. +#: +#: An age rather than a count, because the count a topic can afford depends on +#: its record size and the age does not, and because a count cap refuses a task +#: whose batch does not fit under it. Seven days is long enough to replay a +#: consumer that was evicted over a weekend and short enough that a chain +#: nobody reads any more does not keep its records for good. +DEFAULT_RETENTION: Final = timedelta(days=7) + +#: 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 _trim_floor(backend: Any) -> str: + retention = getattr(backend, "_retention", None) + if retention is None: + return "" + # The worker's clock names the floor, so a skewed worker shifts the window by + # its skew, the same as the backend's own trim. + floor = int((time.time() - retention.total_seconds()) * 1000) + return f"{max(floor, 0)}-0" + + +def _trim_maxlen(backend: Any) -> str: + max_len = getattr(backend, "_max_len", None) + return "" if max_len is None else str(max_len) + + +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(backend: Any, record: TransportRecord, digest: str) -> list[Any]: + """What the log append script takes: the trims, the identity and digest, 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 [ + _trim_floor(backend).encode(), + _trim_maxlen(backend).encode(), + 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 the retention trims +#: ride along instead of costing their own round trips; they are exact for the +#: reason the backend's own trims are. 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 minid = ARGV[1] +local maxlen = ARGV[2] +local existing = redis.call('HGET', KEYS[2], ARGV[3]) +if existing then + local sep = string.find(existing, '|') + if string.sub(existing, sep + 1) ~= ARGV[4] then + return {'conflict', ''} + end + return {'ok', string.sub(existing, 1, sep - 1)} +end +local id = redis.call('XADD', KEYS[1], '*', unpack(ARGV, 5)) +redis.call('HSET', KEYS[2], ARGV[3], id .. '|' .. ARGV[4]) +if minid ~= '' then + redis.call('XTRIM', KEYS[1], 'MINID', minid) +end +if maxlen ~= '' then + redis.call('XTRIM', KEYS[1], 'MAXLEN', maxlen) +end +return {'ok', id} +""" + + +class _LogAppend: + """One record on one log, a topic's or an activity's, 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(self._backend, record, digest), + ) + if _text(outcome) == "conflict": + raise AppendConflictError(record.idempotency_key) + return Offset(_text(placed)) + + +@dataclass(frozen=True) +class _StandaloneOwner: + """A stream with an id of its own and no owner, and where Redis keeps it.""" + + namespace: str + stream_id: str + + def meta(self, key_prefix: str) -> str: + """The hash holding the stream's policy, seal and byte totals.""" + namespace = quote(self.namespace, safe="") + return f"{key_prefix}:{namespace}:standalone/{quote(self.stream_id, safe='')}" + + def key(self, key_prefix: str, topic: str) -> str: + """The log holding ``topic`` of this stream; one more component than the hash.""" + return f"{self.meta(key_prefix)}:{quote(topic, safe='')}" + + def __str__(self) -> str: + """The stream, for messages.""" + return f"standalone stream {self.stream_id!r}" + + +def _policy_fields( + retention: timedelta | None, max_records: int | None, max_bytes: int | None +) -> list[bytes]: + """The policy as the hash stores it: milliseconds, counts, and blanks for none.""" + return [ + b"" + if retention is None + else str(int(retention.total_seconds() * 1000)).encode(), + b"" if max_records is None else str(max_records).encode(), + b"" if max_bytes is None else str(max_bytes).encode(), + ] + + +#: Create a standalone stream's hash, or say whether the one there agrees. +_CREATE_STANDALONE_LUA: Final = """ +if redis.call('EXISTS', KEYS[1]) == 1 then + local held = redis.call('HMGET', KEYS[1], 'retention_ms', 'max_records', 'max_bytes') + if held[1] == ARGV[1] and held[2] == ARGV[2] and held[3] == ARGV[3] then + return 'same' + end + return 'conflict' +end +redis.call('HSET', KEYS[1], 'retention_ms', ARGV[1], 'max_records', ARGV[2], + 'max_bytes', ARGV[3], 'sealed', '0') +return 'created' +""" + + +#: Append one record to a standalone stream's topic under the stream's policy. +#: +#: The policy lives in the hash, so the script reads it rather than being told: +#: a sealed stream refuses, a missing one says so, and the trims apply to the +#: topic just written. ``max_bytes`` has no Redis trim of its own, so each entry +#: carries the size of the record it holds, the hash keeps a byte total per +#: topic, and the oldest entries are dropped one at a time until the topic fits. +_STANDALONE_APPEND_LUA: Final = """ +local meta = redis.call('HMGET', KEYS[3], 'sealed', 'retention_ms', 'max_records', + 'max_bytes', ARGV[6]) +if not meta[1] then + return {'missing', ''} +end +if meta[1] == '1' then + return {'closed', ''} +end +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, 7)) +redis.call('HSET', KEYS[2], ARGV[1], id .. '|' .. ARGV[2]) +local total = tonumber(meta[5] or '0') + tonumber(ARGV[4]) +local floor = nil +if meta[2] ~= '' then + floor = tonumber(ARGV[3]) - tonumber(meta[2]) +end +while true do + local entries = redis.call('XRANGE', KEYS[1], '-', '+', 'COUNT', 1) + if #entries == 0 then break end + local entry = entries[1] + local drop = false + if meta[3] ~= '' and redis.call('XLEN', KEYS[1]) > tonumber(meta[3]) then drop = true end + if meta[4] ~= '' and total > tonumber(meta[4]) then drop = true end + if floor and tonumber(string.match(entry[1], '^(%d+)')) < floor then drop = true end + if not drop then break end + local size = 0 + local fields = entry[2] + for i = 1, #fields, 2 do + if fields[i] == ARGV[5] then size = tonumber(fields[i + 1]) end + end + redis.call('XDEL', KEYS[1], entry[1]) + total = total - size +end +redis.call('HSET', KEYS[3], ARGV[6], total) +return {'ok', id} +""" + + +#: The entry field holding the size a standalone stream's policy counts: the +#: serialized record, as the memory provider measures it, not the stored payload +#: with its envelope. +_SIZE_FIELD: Final = b"__size" + + +class _StandaloneAppend: + """One record on a standalone stream's topic, under its policy, in one call.""" + + def __init__(self, backend: RedisStreamBackend, owner: _StandaloneOwner) -> None: + """Bind to ``backend``'s client and ``owner``'s hash.""" + self._owner = owner + self._meta = owner.meta(_prefix(backend)) + self._script = backend._client.register_script(_STANDALONE_APPEND_LUA) + + async def write( + self, *, name: str, topic: str, record: TransportRecord, digest: str, size: int + ) -> Offset: + """Append ``record`` to the log ``name`` and return where it landed. + + ``size`` is what the stream's ``max_bytes`` counts for this record. + + Raises: + StreamNotFoundError: The stream was never created. + StreamClosedError: The stream is sealed. + AppendConflictError: The identity is held with a different digest. + """ + outcome, placed = await self._script( + keys=[name, f"{name}:idem", self._meta], + args=[ + str(record.idempotency_key).encode(), + digest.encode(), + str(int(time.time() * 1000)).encode(), + str(size).encode(), + _SIZE_FIELD, + f"bytes:{quote(topic, safe='')}".encode(), + *_fields_args(record), + _SIZE_FIELD, + str(size).encode(), + ], + ) + result = _text(outcome) + if result == "missing": + raise StreamNotFoundError(f"{self._owner} was not found") + if result == "closed": + raise StreamClosedError( + f"{self._owner} is closed and takes no more records" + ) + if result == "conflict": + raise AppendConflictError(record.idempotency_key) + return Offset(_text(placed)) + + +def _prefix(backend: Any) -> str: + """The key prefix ``backend`` writes under, for keys the provider derives itself.""" + return backend._key_prefix + + +@dataclass(frozen=True) +class _ActivityOwner: + """The activity whose streams a handle addresses, and where Redis keeps them. + + ``workflow_id`` is ``None`` for a standalone activity. ``run_id`` is the + workflow's run for a workflow's activity and the activity's own run for a + standalone one. It is part of the key, so an activity execution's streams + are its own and an id started again in a new run starts new ones, and it + decides whose close ends a read. ``None`` means the run is not known yet: + a handle opened outside without one asks the server for the current run + before it touches a key. + """ + + namespace: str + workflow_id: str | None + activity_id: str + run_id: str | None + + def key(self, key_prefix: str, topic: str) -> str: + """The Redis stream holding ``topic`` of this activity's streams.""" + if self.run_id is None: + raise RuntimeError(f"the run of {self} was not resolved before its key") + # The chain keys percent-encode every id, so none of their components + # holds a "/", and a component built around one can never equal a + # chain key or a key derived from one. + namespace = quote(self.namespace, safe="") + owner = f"activity/{quote(self.workflow_id or '', safe='')}/" + owner += f"{quote(self.run_id, safe='')}/{quote(self.activity_id, safe='')}" + return f"{key_prefix}:{namespace}:{owner}:{quote(topic, safe='')}" + + def context(self) -> SerializationContext: + """What the owner's payloads are coded under. + + A workflow's activity writes the workflow's data, as it does on the + workflow's own topics; a standalone activity has no workflow, so its + own identity is the context. + """ + if self.workflow_id is not None: + return WorkflowSerializationContext( + namespace=self.namespace, workflow_id=self.workflow_id + ) + return ActivitySerializationContext( + namespace=self.namespace, + activity_id=self.activity_id, + activity_type=None, + activity_task_queue=None, + workflow_id=None, + workflow_type=None, + is_local=False, + ) + + def __str__(self) -> str: + """The owner, for messages.""" + if self.workflow_id is None: + return f"activity {self.activity_id!r}" + return f"activity {self.activity_id!r} of workflow {self.workflow_id!r}" + + +async def _resolve_owner(client: Client, owner: _ActivityOwner) -> _ActivityOwner: + """``owner`` with its run known, or ``StreamNotFoundError`` when the server does not know it. + + One describe: it says whether the owner exists and, for a handle opened + without a run, which run is current. A standalone activity is described + itself; a workflow's activity is described through its workflow, because + the server does not describe it on its own. + """ + try: + if owner.workflow_id is None: + described = await client.get_activity_handle( + owner.activity_id, run_id=owner.run_id + ).describe() + run_id = owner.run_id or described.activity_run_id + else: + description = await client.get_workflow_handle( + owner.workflow_id, run_id=owner.run_id + ).describe() + run_id = owner.run_id or description.run_id + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError(f"{owner} was not found") from error + raise + if not run_id: + raise StreamNotFoundError(f"the server reports no run for {owner}") + return replace(owner, run_id=run_id) + + +async def _retained(client: Any, name: str, offset: Offset) -> bool: + """Whether the record at ``offset`` on the stream ``name`` survived trimming. + + A record at or after the first retained entry is there. On an emptied + stream the last id Redis generated says whether the record ever was. + """ + if not await client.exists(name): + # Nothing was ever written under this key, so nothing was trimmed from it. + return True + info = await client.xinfo_stream(name) + wanted = _entry_id(offset.token) + first = info.get("first-entry") + if first: + return _entry_id(first[0]) <= wanted + return wanted > _entry_id(info["last-generated-id"]) + + +async def _require_standalone(store: Any, meta: str, owner: _StandaloneOwner) -> None: + """Raise ``StreamNotFoundError`` unless ``owner``'s hash exists.""" + if not await store.exists(meta): + raise StreamNotFoundError(f"{owner} was not found") + + +async def _sealed(store: Any, meta: str) -> bool: + """Whether the standalone stream behind ``meta`` is sealed.""" + return _text(await store.hget(meta, "sealed") or b"0") == "1" + + +async def _newest(store: Any, name: str) -> Cursor: + """The cursor of the newest entry of the log ``name``, or ``BEGINNING``.""" + newest: Any = await store.xrevrange(name, "+", "-", count=1) + if not newest: + return BEGINNING + return mint_cursor(_PROVIDER, _text(newest[0][0])) + + +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 and trims. + + 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. + + Trims are exact rather than approximate: Redis's approximate trim drops + whole macro nodes only, so a stream shorter than one node, a hundred + entries by default, would never trim and the window would not mean what + it says. Only the logs are trimmed; the idempotency and stage hashes + beside them keep one entry per record and stage. + """ + + def __init__( + self, + *, + client: Any, + key_prefix: str, + retention: timedelta | None, + max_len: int | None, + wake_transport: WakeTransport = "auto", + ) -> None: + super().__init__(client=client, key_prefix=key_prefix) + self._retention = retention + self._max_len = max_len + 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 + + def describe_window(self) -> str: + """The configured window, for messages.""" + parts = [] + if self._retention is not None: + parts.append(f"retention={self._retention}") + if self._max_len is not None: + parts.append(f"max_len={self._max_len}") + return ", ".join(parts) or "no retention" + + async def append(self, key: StreamKey, record: Any) -> Any: + placed = await super().append(key, record) + await self._trim(key) + return placed + + async def stage_output( + self, manifest: OutputStageManifest, records: Sequence[StagedOutputRecord] + ) -> OutputStage: + # A stage is invisible until its task commits, and the trim has no consumer + # floor to hold it: a window at or below the batch takes entries out of the + # stage that was just written, and the commit then fails on a missing record + # for as long as the task retries. Flooring the trim instead would need the + # floor this provider deliberately does not keep, and would not hold anyway, + # because the trim that removes the stage is not the one that staged it. + if self._max_len is not None and manifest.record_count >= self._max_len: + raise ValueError( + f"this task publishes {manifest.record_count} records and max_len is " + f"{self._max_len}: the window has to exceed the largest batch a task " + "publishes, or the batch is trimmed before it commits" + ) + stage = await super().stage_output(manifest, records) + await self._trim(manifest.stream_key) + return stage + + 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]: + # The replay read. Said here, where the trim is known, rather than left + # to the range checks, which can only report the record as missing. + if not await self.retains(key, first): + raise StreamIntegrityError( + f"the recorded range [{first}, {last}] on topic " + f"{key.stream_name!r} is past the redis provider's retention " + f"({self.describe_window()}): the records were trimmed, so this " + "run cannot be replayed" + ) + 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 + + async def retains(self, key: StreamKey, offset: Offset) -> bool: + """Whether the record at ``offset`` survived trimming.""" + return await _retained(self._client, self.stream_key(key), offset) + + async def _trim(self, key: StreamKey) -> None: + name = self.stream_key(key) + if self._retention is not None: + # The worker's clock names the floor, so a skewed worker shifts + # the window by its skew. + floor = int((time.time() - self._retention.total_seconds()) * 1000) + await self._client.xtrim( + name, minid=f"{max(floor, 0)}-0", approximate=False + ) + if self._max_len is not None: + await self._client.xtrim(name, maxlen=self._max_len, approximate=False) + + +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) + subscribe = self._input.topic(topic, type=bytes).subscribe + 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: the transport has the worker resolve + # it after this task and record the entry with the subscription. + return _RedisReadSource(subscribe(start_at_tail=StartAtTail(last or 0))) + 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 trimmed 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, + owner: _ActivityOwner | None = None, + standalone: _StandaloneOwner | None = None, + ) -> None: + """Bind this producer to ``topic`` of the chain, of ``owner`` or of ``standalone``.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._owner = owner + self._standalone = standalone + 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 | _StandaloneAppend | 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() + if self._standalone is not None: + # No owner to code under and no one to wake: a standalone stream + # is read from outside only, and its hash says whether it exists. + self._name = self._standalone.key(_prefix(backend), self._topic) + self._append = _StandaloneAppend(backend, self._standalone) + self._codec = StreamPayloadCodec(self._client.data_converter, bytes) + return + self._append = _LogAppend(backend) + if self._owner is not None: + # No transport producer to bind and no one to wake: an activity's + # stream has no workflow reader and no chain to check the key against. + # Inside the activity the run is known and nothing is described. + if self._owner.run_id is None: + self._owner = await _resolve_owner(self._client, self._owner) + self._name = self._owner.key(_prefix(backend), self._topic) + self._codec = StreamPayloadCodec( + self._client.data_converter.with_context(self._owner.context()), bytes + ) + return + 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 the trims ride along with it. + 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 + # trims ride 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, + ) + if isinstance(self._append, _StandaloneAppend): + last = await self._append.write( + name=self._name, + topic=self._topic, + record=staged, + digest=digest, + size=record.ByteSize(), + ) + else: + last = await self._append.write( + name=self._name, record=staged, digest=digest + ) + if self._owner is None and self._standalone is None: + 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 owner's topics from outside. + + The owner is ``workflow_id``'s workflow, whose topic logs are read through + the transport's output read so a staged batch is a barrier until its task + settles; with ``activity_id`` an activity, read from the stream the + provider keeps for it, a standalone activity without ``workflow_id`` or an + activity that workflow scheduled; or with ``stream_id`` a standalone + stream, read from the logs under its hash until it is sealed. + """ + + def __init__( + self, + streams: RedisStreams, + client: Client, + workflow_id: str | None, + run_id: str | None, + activity_id: str | None = None, + *, + stream_id: str | None = None, + ) -> None: + """Address the owner'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 + self._owner: _ActivityOwner | None = None + self._standalone: _StandaloneOwner | None = None + converter = client.data_converter + if stream_id is not None: + # No owner, so nothing to code the bodies under. + self._standalone = _StandaloneOwner(client.namespace, stream_id) + elif activity_id is not None: + self._owner = _ActivityOwner( + client.namespace, workflow_id, activity_id, run_id + ) + converter = converter.with_context(self._owner.context()) + else: + if workflow_id is None: + raise ValueError( + "a stream handle needs a workflow_id, an activity_id or a stream_id" + ) + converter = 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 owner closes. + + For a workflow that is the chain, or the pinned run; for an activity it is + the activity reaching a terminal status, learned as the module docstring + says. ``END`` and ``last=`` are positioned against the log on the first + step of the generator, since this call cannot reach the store. + + Two refusals and they do not land together. 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. A well-formed cursor the retention has + trimmed is refused on the first step of the generator, because answering that + needs a round trip and this call is not a coroutine. Neither yields a record + first. + + 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) + if self._standalone is not None: + return self._read_standalone( + self._standalone, topic, position, after, result_type, tail=tail + ) + if self._owner is not None: + return self._read_owned( + self._owner, topic, position, after, result_type, tail=tail + ) + return self._read(topic, position, after, result_type, tail=tail) + + async def _read_standalone( + self, + owner: _StandaloneOwner, + topic: str, + position: Offset | None, + after: Cursor, + result_type: type | None, + *, + tail: int | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + backend = self._streams._require_backend() + store = backend._client + meta = owner.meta(_prefix(backend)) + await _require_standalone(store, meta, owner) + name = owner.key(_prefix(backend), topic) + if tail is not None: + position = await _tail_after(store, name, tail) + after = ( + BEGINNING + if position is None + else mint_cursor(_PROVIDER, position.token) + ) + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + if position is not None and not await _retained(store, name, position): + raise StreamCursorError( + f"cursor {after.token!r} names a record on {topic!r} that the " + f"stream's policy has dropped" + ) + start = _BEGINNING_SENTINEL if position is None else position.token + block = int(self._streams._poll.total_seconds() * 1000) or None + closed = False + while True: + found: Any = await store.xread( + {name: start}, count=_READ_BATCH, block=block + ) + entries = found[0][1] if found else [] + for entry_id, fields in entries: + placed = _to_record(entry_id, fields) + assert placed.offset is not None + start = placed.offset.token + 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 entries: + continue + if closed: + return + # One more pass after learning of the seal, so a record that landed + # between the read and the seal is not lost. + closed = await _sealed(store, meta) + + async def _read_owned( + self, + owner: _ActivityOwner, + topic: str, + position: Offset | None, + after: Cursor, + result_type: type | None, + *, + tail: int | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + backend = self._streams._require_backend() + # Described on every read, so a handle without a run reads the run that + # is current when the read starts, and ends with it. + owner = await _resolve_owner(self._client, owner) + store = backend._client + name = owner.key(_prefix(backend), topic) + if tail is not None: + position = await _tail_after(store, name, tail) + after = ( + BEGINNING + if position is None + else mint_cursor(_PROVIDER, position.token) + ) + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + if ( + position is not None + and isinstance(backend, _TopicLogBackend) + and not await _retained(store, name, position) + ): + # Refused rather than resumed from the first retained record, + # which would skip whatever the trim took in between. + raise StreamCursorError( + f"cursor {after.token!r} names a record on {topic!r} that the " + f"provider's retention has trimmed ({backend.describe_window()})" + ) + start = _BEGINNING_SENTINEL if position is None else position.token + # XREAD BLOCK 0 waits forever, which is not what a zero poll asks for. + block = int(self._streams._poll.total_seconds() * 1000) or None + closed = False + while True: + found: Any = await store.xread( + {name: start}, count=_READ_BATCH, block=block + ) + entries = found[0][1] if found else [] + for entry_id, fields in entries: + placed = _to_record(entry_id, fields) + assert placed.offset is not None + start = placed.offset.token + 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 entries: + continue + if closed: + return + # One more pass after learning the owner is terminal, so a record + # that landed between the read and the describe is not lost. + closed = await self._owner_closed(owner, store, name) + + async def _owner_closed(self, owner: _ActivityOwner, store: Any, name: str) -> bool: + """Whether the owning activity is terminal, as far as this store can tell. + + A standalone activity says so itself. The server does not describe a + workflow's activity, so its workflow is asked: the workflow closing ends + the read, and so does the activity leaving the pending set once its + stream exists. Before the first write there is no stream, so an + activity that never wrote is read until its workflow closes. + """ + try: + if owner.workflow_id is None: + described = await self._client.get_activity_handle( + owner.activity_id, run_id=owner.run_id + ).describe() + return described.status != ActivityExecutionStatus.RUNNING + description = await self._client.get_workflow_handle( + owner.workflow_id, run_id=owner.run_id + ).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + # The owner's History is gone, so there is nothing left to follow. + return True + raise + status = description.status + if status is not None and status != WorkflowExecutionStatus.RUNNING: + # An activity's streams are not the chain's: a run that continued + # as new took its activities with it. + return True + pending = any( + info.activity_id == owner.activity_id + for info in description.raw_description.pending_activities + ) + if pending: + return False + return bool(await store.exists(name)) + + 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 + ) + if ( + position is not None + and isinstance(backend, _TopicLogBackend) + and not await backend.retains(key, position) + ): + # Refused rather than resumed from the first retained record, + # which would skip whatever the trim took in between. + raise StreamCursorError( + f"cursor {after.token!r} names a record on {topic!r} that the " + f"provider's retention has trimmed ({backend.describe_window()})" + ) + 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, which a topic whose + records retention has all trimmed answers too: the two are the same state to + a reader, and a read from it starts at the first record retained after it + rather than at the tail the caller asked to follow from. + """ + topic, _ = resolve_topic(topic) + backend = self._streams._require_backend() + if self._standalone is not None: + await _require_standalone( + backend._client, + self._standalone.meta(_prefix(backend)), + self._standalone, + ) + return await _newest( + backend._client, self._standalone.key(_prefix(backend), topic) + ) + if self._owner is not None: + owner = await _resolve_owner(self._client, self._owner) + return await _newest(backend._client, owner.key(_prefix(backend), topic)) + 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.""" + if self._standalone is not None: + return StreamRef.for_standalone(self._standalone.stream_id, topic=topic) + if self._owner is not None: + return StreamRef.for_activity( + self._owner.activity_id, + workflow_id=self._owner.workflow_id, + run_id=self._owner.run_id, + topic=topic, + ) + 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: + """Seal a standalone stream. An owned stream ends with its owner, not by a caller. + + The seal is a flag in the stream's hash that every append reads, so a + later append is refused with ``StreamClosedError`` and a read ends once + it has delivered the retained tail. Idempotent. + + Raises: + ValueError: This handle is on a workflow's or an activity's stream. + StreamNotFoundError: The stream was never created. + """ + if self._standalone is None: + raise ValueError( + "only a standalone stream can be closed; this handle is on an owned " + "stream, which ends when its workflow or activity does" + ) + backend = self._streams._require_backend() + meta = self._standalone.meta(_prefix(backend)) + await _require_standalone(backend._client, meta, self._standalone) + await backend._client.hset(meta, "sealed", "1") + + 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, + owner=self._owner, + standalone=self._standalone, + ) + + +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), + retention: timedelta | None = DEFAULT_RETENTION, + max_len: int | None = None, + 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 and trims 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. + retention: Trim records older than this from a topic's log on + every append the provider makes to it. + :data:`DEFAULT_RETENTION`, seven days, unless the caller says + otherwise; ``None`` keeps every record until ``max_len`` + trims it, or for good when that is unset too. This is + retention without a consumer floor: nothing holds a record + for a reader that has not reached it. A workflow whose replay + reaches a recorded range past the window fails its Workflow + Task with the transport's ``StreamIntegrityError`` until the + window is raised, an outside ``read(after=)`` below the window + raises ``StreamCursorError``, and a live reader that falls + behind the window misses records. The floor the server-side + provider keeps would need a consumer registry in Redis. + max_len: Keep at most this many entries per key, trimmed on the + same appends and with the same consequences. Off unless set. + It must exceed the largest batch a task publishes, or a stage + is trimmed before its commit; a batch at or above it is + refused where it is staged. + 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}") + if retention is not None and retention <= timedelta(0): + raise ValueError("retention must be positive") + if max_len is not None and max_len < 1: + raise ValueError("max_len must be positive") + 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._retention = retention + self._max_len = max_len + 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, + retention=self._retention, + max_len=self._max_len, + 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: + """A handle on the topics ``activity_id`` owns, apart from any workflow's. + + Without ``workflow_id`` the activity is a standalone one and ``run_id`` + names its run; with one it is that workflow's activity and ``run_id`` + names the workflow's run. The streams are keyed by that run, so a + retry writes to the same ones and a reader sees the attempt change as + ``SUPERSEDED``, while an id started again in a new run starts new + ones. Without ``run_id`` each read, ``latest()`` and the first append + of a producer describe the owner and take the run current at that + moment. A read ends when the activity is terminal and the retained + tail is delivered; the store has no gate on an append after that, so + a late attempt still lands, and a read that has ended does not see it. + """ + return RedisStreamHandle(self, client, workflow_id, run_id, activity_id) + + 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: + """Create the standalone stream ``stream_id``, or find it with the same policy. + + The policy is written once into the stream's hash and applied on every + append to any of its topics. ``retention`` left unset takes this + provider's default window; the other two bounds are off unless set. + + Raises: + ValueError: ``stream_id`` is empty, a bound is not positive, or + the stream exists with a different policy. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + if retention is not None and retention <= timedelta(0): + raise ValueError("retention must be positive") + if max_records is not None and max_records <= 0: + raise ValueError("max_records must be positive") + if max_bytes is not None and max_bytes <= 0: + raise ValueError("max_bytes must be positive") + backend = self._require_backend() + owner = _StandaloneOwner(client.namespace, stream_id) + policy = _policy_fields( + self._retention if retention is None else retention, max_records, max_bytes + ) + outcome = await backend._client.register_script(_CREATE_STANDALONE_LUA)( + keys=[owner.meta(_prefix(backend))], args=policy + ) + if _text(outcome) == "conflict": + raise ValueError( + f"{owner} exists with another policy; a policy is set when the " + "stream is created and does not change" + ) + return RedisStreamHandle(self, client, None, None, stream_id=stream_id) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> RedisStreamHandle: + """A handle on the standalone stream ``stream_id``. + + Nothing is checked here: a ``read``, ``latest``, ``producer`` or + ``close`` on a stream that was never created raises + :class:`temporalio.streams.StreamNotFoundError` when it is used. + """ + return RedisStreamHandle(self, client, None, None, stream_id=stream_id) + + 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() diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 2ea360f58..d72d210a3 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -650,6 +650,11 @@ async def _handle_activation( raise deadlock_exc from None output_runtime = self._external_stream_runtimes.get(act.run_id) + if output_runtime is not None and completion.HasField("successful"): + # A subscription opened at a stream's tail is positioned here, + # off the Workflow thread and before the completion goes out, so + # the marker records the boundary the watcher starts from. + await output_runtime.resolve_pending_starts() if ( output_runtime is not None and completion.HasField("successful") @@ -913,6 +918,7 @@ async def _handle_cache_eviction( # skips the instance's eviction job entirely, and a Run whose # watchers outlived it would keep a backend connection open forever. self._external_stream_runtimes.pop(act.run_id, None) + await self._abort_dead_external_output(run_id=act.run_id) self._pending_external_output_stages.pop(act.run_id, None) if self._external_stream_manager is not None: await self._external_stream_manager.evict_run(act.run_id) @@ -1301,6 +1307,56 @@ async def _promote_external_output( else: self._pending_external_output_stages.pop(run_id, None) + async def _abort_dead_external_output(self, *, run_id: str) -> None: + """Abort the staged batches of an evicted Run that History has rejected. + + A Run evicted because its completion was not accepted may have staged + a batch that no marker will ever name. History already says so when + the task was failed or timed out, so such a stage is aborted here + instead of waiting for a reader to find the barrier. A stage whose + marker is authoritative is never re-promoted from here: promotion is + attempted once after the completion, and a stage it left pending is + owed to a cold client's repair, as the contract tests pin down. + """ + backend = self._external_stream_backend + client = self._client + pending = self._pending_external_output_stages.get(run_id) + if backend is None or client is None or not pending: + return + + from temporalio.contrib.external_workflow_streams._output_client import ( + _apply_output_stage_decision, + _history_decision, + _HistoryDecision, + ) + + for manifest in pending: + workflow_id = manifest.stream_key.workflow_id + try: + decision = await _history_decision( + client=client, + workflow_id=workflow_id, + manifest=manifest, + ) + if decision is not _HistoryDecision.ABORT: + continue + await _apply_output_stage_decision( + backend=backend, + manifest=manifest, + decision=decision, + ) + except asyncio.CancelledError: + raise + except Exception as err: + self._stream_metrics.record(err) + logger.warning( + "Could not abort the dead external output stage %s for " + "Workflow %s at eviction; it remains pending for a client", + manifest.stage_token, + workflow_id, + exc_info=True, + ) + def _replay_stream_converters( self, plan: Any ) -> dict[int, temporalio.converter.DataConverter]: diff --git a/tests/contrib/external_workflow_streams/test_api.py b/tests/contrib/external_workflow_streams/test_api.py index 0557ddb25..cb257cb0f 100644 --- a/tests/contrib/external_workflow_streams/test_api.py +++ b/tests/contrib/external_workflow_streams/test_api.py @@ -29,6 +29,8 @@ ) from temporalio.contrib.external_workflow_streams._manager import PreparedRecord from temporalio.contrib.external_workflow_streams._record import ( + AFTER, + Cursor, Offset, RecordKind, StreamRecord, @@ -48,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]] = [] @@ -72,9 +75,15 @@ def register( wait_id: int, stream_key: StreamKey, idle_timeout: timedelta, + start_cursor: Cursor | None = None, + start_at_tail: Any = None, ) -> None: self.registrations.append((wait_id, stream_key)) self.idle_timeouts[wait_id] = idle_timeout + # A tail start stands in for the cursor the Worker would resolve. + self.start_cursors[wait_id] = ( + start_cursor if start_at_tail is None else start_at_tail + ) def drain(self, wait_id: int, max_records: int | None = None) -> list[StreamRecord]: buffered = self.buffers.get(wait_id, []) @@ -195,6 +204,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_registry.py b/tests/contrib/external_workflow_streams/test_registry.py index c368be89c..d18d658b4 100644 --- a/tests/contrib/external_workflow_streams/test_registry.py +++ b/tests/contrib/external_workflow_streams/test_registry.py @@ -104,6 +104,7 @@ def test_supported_entry_points_are_public() -> None: "PrecedingWriteFailedError", "RecordKind", "RedisStreamBackend", + "StartAtTail", "StreamBackend", "StreamDecodeError", "StreamError", diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index bd0a675ac..8d4f7c59c 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -16,6 +16,7 @@ from __future__ import annotations import asyncio +import os import uuid from collections.abc import AsyncIterator, Callable from dataclasses import dataclass @@ -34,6 +35,7 @@ topic, ) from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.redis import RedisStreams from temporalio.testing import WorkflowEnvironment from tests.helpers import new_worker @@ -57,9 +59,32 @@ async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider.reset() +async def _redis_setup(client: Client) -> AsyncIterator[ActivitySetup]: + # 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-activity-{uuid.uuid4().hex}", + poll_interval=timedelta(milliseconds=100), + ) + config = client.config() + config["plugins"] = [provider] + yield ActivitySetup("redis", provider, Client(**config)) + await provider.close() + + SETUPS: dict[str, Callable[[Client], AsyncIterator[ActivitySetup]]] = { "memory": _memory_setup } +if os.environ.get("STREAMS_LIVE") == "redis": + SETUPS["redis"] = _redis_setup @pytest.fixture(params=sorted(SETUPS)) diff --git a/tests/streams/test_redis_activity_streams.py b/tests/streams/test_redis_activity_streams.py new file mode 100644 index 000000000..9208fb54f --- /dev/null +++ b/tests/streams/test_redis_activity_streams.py @@ -0,0 +1,288 @@ +"""Live checks for what the Redis provider decides about an activity's streams. + +``test_activity_streams`` runs the shared cases on this provider behind +``STREAMS_LIVE=redis``. This module covers what only this store does: how a +read learns that a workflow's activity is terminal without the server saying +so, that an activity that never wrote is read until its workflow closes, +that retention trims an activity's stream like any other, that nothing +gates an append once the owner is terminal, and that an activity's streams +belong to the run its execution is in, so an id started again in a new run +starts new ones. All need a dev server (``TEMPORAL_ADDRESS`` or the test +environment's own) and a Redis (``TEMPORAL_TEST_REDIS_URL`` or +``AI198_REDIS_URL``). +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from collections.abc import AsyncIterator +from datetime import timedelta + +import pytest + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.streams import BEGINNING, RecordKind, StreamCursorError +from temporalio.streams.providers.redis import RedisStreams +from temporalio.testing import WorkflowEnvironment +from tests.helpers import new_worker +from tests.streams.test_activity_streams import ( + TOKENS, + RunsOneActivity, + read_all, + summary, + write_by_default, + write_to_own_streams, +) +from tests.streams.test_streams_conformance import StreamHost, take + +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 live(client: Client, env: WorkflowEnvironment) -> AsyncIterator[Client]: + """A client with a fresh provider registered, on a server with standalone activities.""" + if env.supports_time_skipping: + pytest.skip("the time-skipping test server has no standalone activities") + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = RedisStreams( + url=redis_url(), + # A prefix per case, because the store keeps what earlier cases wrote. + key_prefix=f"streams-redis-activity-{uuid.uuid4().hex}", + poll_interval=timedelta(milliseconds=100), + ) + config = client.config() + config["plugins"] = [provider] + try: + yield Client(**config) + finally: + await provider.close() + + +@activity.defn +async def write_nothing(label: str) -> str: + return label + + +@workflow.defn +class RunsOneActivityThenWaits: + """Runs one activity by name, then stays open until released.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self, name: str, label: str) -> str: + result = await workflow.execute_activity( + name, + label, + activity_id="streamer", + start_to_close_timeout=timedelta(seconds=30), + ) + await workflow.wait_condition(lambda: self._released) + return result + + +async def test_a_read_opened_after_the_activity_finished_ends_by_itself( + live: Client, +): + # The server does not describe a workflow's activity, so the provider ends + # the read once the activity is no longer pending and its stream exists. + workflow_id = f"streams-redis-wfa-{uuid.uuid4().hex}" + async with new_worker( + live, RunsOneActivityThenWaits, activities=[write_to_own_streams] + ) as worker: + handle = await live.start_workflow( + RunsOneActivityThenWaits.run, + args=["write_to_own_streams", "to the activity"], + id=workflow_id, + task_queue=worker.task_queue, + ) + own = live.get_stream_handle(workflow_id, activity_id="streamer") + # The first read follows the activity to its FINISH while it runs. + await take(own.read(topic=TOKENS), 2, 30) + # A read opened afterwards ends on its own, with the workflow still open. + records = await read_all(own.read(topic=TOKENS), timeout=10) + assert summary(records) == [ + (RecordKind.DATA, 1, "to the activity"), + (RecordKind.FINISH, 1, None), + ] + assert (await handle.describe()).status is not None + await handle.signal(RunsOneActivityThenWaits.release) + await handle.result() + + +async def test_an_activity_that_never_wrote_is_read_until_its_workflow_closes( + live: Client, +): + workflow_id = f"streams-redis-wfa-silent-{uuid.uuid4().hex}" + async with new_worker( + live, RunsOneActivityThenWaits, activities=[write_nothing] + ) as worker: + handle = await live.start_workflow( + RunsOneActivityThenWaits.run, + args=["write_nothing", "quiet"], + id=workflow_id, + task_queue=worker.task_queue, + ) + assert await handle.query("__temporal_workflow_metadata") is not None + own = live.get_stream_handle(workflow_id, activity_id="streamer") + reading = asyncio.ensure_future(read_all(own.read(topic=TOKENS), timeout=30)) + # No stream exists for an activity that wrote nothing, so the read has + # nothing to say the activity is over and waits for the workflow. + await asyncio.sleep(1.0) + assert not reading.done() + await handle.signal(RunsOneActivityThenWaits.release) + await handle.result() + assert await reading == [] + + +async def test_retention_trims_an_activity_stream_too(client: Client): + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect(address) + provider = RedisStreams( + url=redis_url(), + key_prefix=f"streams-redis-activity-{uuid.uuid4().hex}", + poll_interval=timedelta(milliseconds=100), + max_len=2, + ) + config = client.config() + config["plugins"] = [provider] + live = Client(**config) + workflow_id = f"streams-redis-wfa-trim-{uuid.uuid4().hex}" + try: + async with new_worker(live, StreamHost) as worker: + host = await live.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = live.get_stream_handle(workflow_id, activity_id="tool") + producer = stream.producer(topic=TOKENS, producer_id="tool", attempt=1) + first = await producer.append({"token": "one"}) + assert first is not None + await producer.append({"token": "two"}, {"token": "three"}) + # The window holds two entries, so the first record is gone. + records = await take(stream.read(topic=TOKENS), 2, 10) + assert [r.value["token"] for r in records] == ["two", "three"] + with pytest.raises(StreamCursorError, match="retention has trimmed"): + await take(stream.read(topic=TOKENS, after=first), 1, 10) + await host.terminate() + finally: + await provider.close() + + +async def test_an_append_after_the_owner_is_terminal_still_lands(live: Client): + # The store has no gate the server would have: a late attempt writes, and + # a reader that already ended is not told. The handle names no run, so it + # describes the finished activity and lands on that run's stream. + activity_id = f"streams-redis-saa-late-{uuid.uuid4().hex}" + async with new_worker(live, activities=[write_by_default]) as worker: + handle = await live.start_activity( + write_by_default, + "standalone", + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=30), + ) + assert await handle.result() == activity_id + stream = live.get_stream_handle(activity_id=activity_id) + before = await stream.latest(topic=TOKENS) + assert before != BEGINNING + late = stream.producer(topic=TOKENS, producer_id=activity_id, attempt=2) + after = await late.append({"token": "late"}) + assert after != before + records = await read_all(stream.read(topic=TOKENS), timeout=10) + assert summary(records) == [ + (RecordKind.DATA, 1, "standalone"), + (RecordKind.FINISH, 1, None), + (RecordKind.SUPERSEDED, 2, None), + (RecordKind.DATA, 2, "late"), + ] + # Pinned to the run, the same stream reads the same way. + pinned = live.get_stream_handle(activity_id=activity_id, run_id=handle.run_id) + assert summary(await read_all(pinned.read(topic=TOKENS), timeout=10)) == summary( + records + ) + + +async def test_a_standalone_activity_id_started_again_starts_a_new_stream( + live: Client, +): + # The stream belongs to the activity execution, which is its run, not to + # the id: the second execution under the same id writes to a new stream, + # a handle without a run reads the current one, and a run pins the other. + activity_id = f"streams-redis-saa-again-{uuid.uuid4().hex}" + async with new_worker(live, activities=[write_by_default]) as worker: + runs = [] + for label in ("first", "second"): + handle = await live.start_activity( + write_by_default, + label, + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=30), + ) + assert await handle.result() == activity_id + runs.append(handle.run_id) + assert runs[0] != runs[1] + current = live.get_stream_handle(activity_id=activity_id) + assert summary(await read_all(current.read(topic=TOKENS), timeout=10)) == [ + (RecordKind.DATA, 1, "second"), + (RecordKind.FINISH, 1, None), + ] + earlier = live.get_stream_handle(activity_id=activity_id, run_id=runs[0]) + assert summary(await read_all(earlier.read(topic=TOKENS), timeout=10)) == [ + (RecordKind.DATA, 1, "first"), + (RecordKind.FINISH, 1, None), + ] + + +async def test_a_workflow_activity_in_a_new_run_starts_a_new_stream(live: Client): + # The same workflow id run again schedules the same activity id; keyed by + # the workflow's run, the two executions keep their streams apart. + workflow_id = f"streams-redis-wfa-again-{uuid.uuid4().hex}" + async with new_worker( + live, RunsOneActivity, activities=[write_to_own_streams] + ) as worker: + runs = [] + for label in ("first", "second"): + handle = await live.start_workflow( + RunsOneActivity.run, + args=["write_to_own_streams", label], + id=workflow_id, + task_queue=worker.task_queue, + ) + await handle.result() + runs.append(handle.result_run_id) + assert runs[0] != runs[1] + current = live.get_stream_handle(workflow_id, activity_id="streamer") + assert summary(await read_all(current.read(topic=TOKENS), timeout=10)) == [ + (RecordKind.DATA, 1, "second"), + (RecordKind.FINISH, 1, None), + ] + earlier = live.get_stream_handle( + workflow_id, activity_id="streamer", run_id=runs[0] + ) + assert summary(await read_all(earlier.read(topic=TOKENS), timeout=10)) == [ + (RecordKind.DATA, 1, "first"), + (RecordKind.FINISH, 1, None), + ] diff --git a/tests/streams/test_redis_provider.py b/tests/streams/test_redis_provider.py new file mode 100644 index 000000000..3a3ee0870 --- /dev/null +++ b/tests/streams/test_redis_provider.py @@ -0,0 +1,379 @@ +"""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 ( + StreamDirection, + WakeNotAcknowledgedError, +) +from temporalio.contrib.external_workflow_streams import ( + StreamError as TransportStreamError, +) +from temporalio.contrib.external_workflow_streams._backend import StreamKey +from temporalio.contrib.external_workflow_streams._record import Offset +from temporalio.contrib.external_workflow_streams._redis import RedisStreamBackend +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 ( + DEFAULT_RETENTION, + RedisProducer, + RedisStreams, + _ActivityOwner, + _drive, + _position, + _StandaloneOwner, + _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_activity_keys_encode_their_ids_and_never_meet_chain_keys(): + # The run is part of the key: a workflow's activity is keyed by the + # workflow's run and a standalone one by its own, so an id started again + # in a new run starts a new stream. + assert ( + _ActivityOwner("ns", "wf", "act", "run").key("p", "t") + == "p:ns:activity/wf/run/act:t" + ) + assert ( + _ActivityOwner("ns", None, "act", "run").key("p", "t") + == "p:ns:activity//run/act:t" + ) + # An id holding a separator is encoded, so it cannot move a boundary. + assert ( + _ActivityOwner("n:s", "w/f", "a:c", "r/1").key("p", "t/u") + == "p:n%3As:activity/w%2Ff/r%2F1/a%3Ac:t%2Fu" + ) + # A key needs the run; a handle opened without one resolves it first. + with pytest.raises(RuntimeError, match="not resolved"): + _ActivityOwner("ns", "wf", "act", None).key("p", "t") + # A chain key percent-encodes every id, so none of its components holds a + # "/" however the ids are chosen, and the owner component here always does. + backend = RedisStreamBackend(client=_NoRedis(), key_prefix="p") + forged = StreamKey( + namespace="ns", + workflow_id="activity/wf/act", + first_execution_run_id="t", + stream_name="t", + direction=StreamDirection.OUTPUT, + ) + assert "/" not in backend.stream_key(forged) + assert backend.stream_key(forged) != _ActivityOwner("ns", "wf", "act", "t").key( + "p", "t" + ) + + +def test_an_activity_owner_names_itself_for_messages(): + assert str(_ActivityOwner("ns", None, "act", None)) == "activity 'act'" + assert ( + str(_ActivityOwner("ns", "wf", "act", "run")) + == "activity 'act' of workflow 'wf'" + ) + + +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()) + + +def test_retention_options_are_checked_at_construction(): + with pytest.raises(ValueError, match="retention"): + RedisStreams(retention=timedelta(0)) + with pytest.raises(ValueError, match="max_len"): + RedisStreams(max_len=0) + RedisStreams(retention=timedelta(hours=1), max_len=10) + + +async def test_a_client_the_caller_opened_gets_the_providers_layout_and_stays_open(): + # The layout and the trims are 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, max_len=10) + backend = provider._require_backend() + assert backend._client is client + assert backend.describe_window() == f"retention={DEFAULT_RETENTION}, max_len=10" + await provider.close() + assert provider._require_backend() is not backend + assert provider._require_backend()._client is client + + +async def test_the_default_window_is_an_age_and_can_be_turned_off(): + # Nothing is trimmed on a topic nobody appends to, so the default has to + # be a window that every append applies; a count cap would refuse a task + # whose batch does not fit under it, so that one stays off. + provider = RedisStreams() + try: + backend = provider._require_backend() + assert backend._retention == DEFAULT_RETENTION == timedelta(days=7) + assert backend._max_len is None + assert backend.describe_window() == f"retention={DEFAULT_RETENTION}" + finally: + await provider.close() + unbounded = RedisStreams(retention=None) + try: + assert unbounded._require_backend().describe_window() == "no retention" + finally: + await unbounded.close() + + +def test_standalone_keys_have_their_own_owner_component(): + owner = _StandaloneOwner("ns", "shared") + assert owner.meta("p") == "p:ns:standalone/shared" + assert owner.key("p", "t") == "p:ns:standalone/shared:t" + # A topic log always has one more component than the hash, and every id and + # topic is encoded, so no name can land on the hash or on another topic. + assert owner.key("p", "meta") != owner.meta("p") + tricky = _StandaloneOwner("n:s", "a/b:c") + assert tricky.meta("p") == "p:n%3As:standalone/a%2Fb%3Ac" + assert tricky.key("p", "t/u") == "p:n%3As:standalone/a%2Fb%3Ac:t%2Fu" + assert str(owner) == "standalone stream 'shared'" diff --git a/tests/streams/test_redis_replay.py b/tests/streams/test_redis_replay.py new file mode 100644 index 000000000..9d078a712 --- /dev/null +++ b/tests/streams/test_redis_replay.py @@ -0,0 +1,1396 @@ +"""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, trims by +retention and shows what a replay and a read past the trim do, 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 datetime import timedelta +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, + StreamCursorError, + 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, _TopicLogBackend +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() + + +async def _trim_everything( + streams: RedisStreams, client: Client, workflow_id: str, topic: str +) -> None: + """What retention elsewhere, or an operator, does to a topic's output key.""" + import redis.asyncio + + backend = streams._require_backend() + chain = await _chain(client, workflow_id) + store = redis.asyncio.from_url(redis_url()) + try: + await store.xtrim( + backend.stream_key( + chain.stream_key(topic, direction=StreamDirection.OUTPUT) + ), + maxlen=0, + approximate=False, + ) + finally: + await store.aclose() + + +async def test_a_replay_past_the_retention_window_fails_loudly(live_client: Client): + # Six entries per key: the loop's four input records stay while it runs, + # and six more appends afterwards push them out. + streams = RedisStreams( + url=redis_url(), key_prefix=f"streams-redis-{uuid.uuid4().hex}", max_len=6 + ) + workflow_id = f"streams-redis-retention-{uuid.uuid4().hex}" + try: + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[ContractLoop], + plugins=[streams], + max_cached_workflows=100, + ): + handle = await live_client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = streams.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 len(await handle.result()) == 4 + before = await take(stream.read(topic=INPUTS), 4, 60) + history = await handle.fetch_history() + assert await _log_length(streams, live_client, workflow_id, INPUTS.name) == 4 + + # Inside the window the recorded ranges read back and the replay passes. + await Replayer(workflows=[ContractLoop], plugins=[streams]).replay_workflow( + history + ) + + late = stream.producer(topic=INPUTS, producer_id="late", attempt=1) + cursors = [await late.append({"n": n}) for n in range(10, 16)] + assert await _log_length(streams, live_client, workflow_id, INPUTS.name) == 6 + + # The recorded input ranges are gone, and the replay says so rather + # than delivering fewer records. + with pytest.raises(Exception) as failure: + await Replayer(workflows=[ContractLoop], plugins=[streams]).replay_workflow( + history + ) + # The task fails under the transport's integrity row: its type, the + # external storage cause, and a message that names the window. + message = str(failure.value) + assert "StreamIntegrityError" in message and "ExternalStorageFailure" in message + assert "past the redis provider's retention (" in message + assert "max_len=6)" in message + + # An outside cursor below the trim is refused, not resumed from the + # first retained record. + with pytest.raises(StreamCursorError, match="retention has trimmed"): + await take(stream.read(topic=INPUTS, after=before[0].cursor), 1, 10) + + # Inside the window a read still works and lands where append said. + records = [r async for r in stream.read(topic=INPUTS, after=cursors[1])] + assert [r.value for r in records] == [{"n": n} for n in range(12, 16)] + assert [r.cursor for r in records] == cursors[2:] + finally: + await streams.close() + + +async def test_retention_by_age_trims_older_entries(live_client: Client): + streams = RedisStreams( + url=redis_url(), + key_prefix=f"streams-redis-{uuid.uuid4().hex}", + retention=timedelta(milliseconds=300), + ) + workflow_id = f"streams-redis-age-{uuid.uuid4().hex}" + try: + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[streams], + ): + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + stream = streams.get_stream_handle(live_client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + first = await producer.append({"n": 1}) + await producer.append({"n": 2}) + await asyncio.sleep(0.6) + # The append that crosses the window is what trims the two before it. + third = await producer.append({"n": 3}) + assert ( + await _log_length(streams, live_client, workflow_id, INPUTS.name) == 1 + ) + assert await stream.latest(topic=INPUTS) == third + with pytest.raises(StreamCursorError, match="retention has trimmed"): + await take(stream.read(topic=INPUTS, after=first), 1, 10) + + # A fully trimmed topic answers the way an empty one does. + await _trim_everything(streams, live_client, workflow_id, INPUTS.name) + assert await stream.latest(topic=INPUTS) == BEGINNING + await handle.signal(StreamHost.release) + await handle.result() + assert [r async for r in stream.read(topic=INPUTS)] == [] + finally: + await streams.close() + + +@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_batch_the_window_cannot_hold_is_refused_at_the_stage( + live_client: Client, +): + # max_len at or below a task's batch trims the stage before its commit, and + # the commit then fails on a missing record for as long as the task retries. + streams = RedisStreams( + url=redis_url(), key_prefix=f"streams-redis-{uuid.uuid4().hex}", max_len=2 + ) + workflow_id = f"streams-redis-window-{uuid.uuid4().hex}" + try: + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[StreamHost], + plugins=[streams], + ) as worker: + handle = await live_client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + run_id = handle.first_execution_run_id or "" + with pytest.raises(ValueError, match="has to exceed the largest batch"): + await _stage_one( + streams, + live_client, + workflow_id, + run_id, + INPUTS.name, + [{"n": 1}, {"n": 2}], + ) + # One below the window still stages. + await _stage_one( + streams, live_client, workflow_id, run_id, INPUTS.name, [{"n": 1}] + ) + await handle.signal(StreamHost.release) + await handle.result() + finally: + await streams.close() + + +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 test_a_workflow_reader_starts_at_the_last_n_records_and_replays_there( + live_client: Client, provider: RedisStreams +): + # The worker positions the subscription against the log after the task that + # opened it and records the entry with it; replay takes the entry from + # History, so what lands in the log later does not move the start. The + # producer needs the run to exist, so the reader opens on a signal. + workflow_id = f"streams-redis-last-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[NewestTwoOnGo], + plugins=[provider], + ): + handle = await live_client.start_workflow( + NewestTwoOnGo.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="tool", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + await handle.signal(NewestTwoOnGo.go) + assert await asyncio.wait_for(handle.result(), 60) == [3, 4] + history = await handle.fetch_history() + await producer.append({"n": 5}, {"n": 6}) + # Offline, with two more records in the log than the run ever saw. + await Replayer(workflows=[NewestTwoOnGo], plugins=[provider]).replay_workflow( + history + ) + + +async def test_a_workflow_reader_at_end_skips_what_was_there_and_replays( + live_client: Client, provider: RedisStreams +): + workflow_id = f"streams-redis-end-{uuid.uuid4().hex}" + async with Worker( + live_client, + task_queue=f"tq-{workflow_id}", + workflows=[FromNowOnGo], + plugins=[provider], + ): + handle = await live_client.start_workflow( + FromNowOnGo.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="tool", attempt=1) + await producer.append({"n": "old"}) + await handle.signal(FromNowOnGo.go) + result = asyncio.ensure_future(handle.result()) + # The subscription is positioned when the worker gets to it, which the + # test does not observe, so appends keep coming until the run takes one. + for _ in range(150): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({result}, timeout=0.2) + if done: + break + assert await asyncio.wait_for(result, 60) == "new" + history = await handle.fetch_history() + await Replayer(workflows=[FromNowOnGo], plugins=[provider]).replay_workflow(history) + + +async def test_a_batch_staged_for_a_rejected_task_does_not_stay_in_the_log( + live_client: Client, monkeypatch: pytest.MonkeyPatch +): + # The first stage takes longer than the Workflow Task timeout, so the server + # times the task out and the worker's completion is rejected; the run is + # evicted and run again, and the second attempt stages the same batch under + # a new token. The first stage can never be named by a marker. It is + # settled against History when the run is evicted and its entries leave the + # log, so the log holds the batch once and nothing counts the dead one. + class SlowFirstStage(_TopicLogBackend): + delayed = False + + async def stage_output(self, manifest: Any, records: Any) -> Any: + if not SlowFirstStage.delayed: + SlowFirstStage.delayed = True + await asyncio.sleep(3) + return await super().stage_output(manifest, records) + + monkeypatch.setattr(redis_provider, "_TopicLogBackend", SlowFirstStage) + provider = RedisStreams( + url=redis_url(), key_prefix=f"streams-redis-{uuid.uuid4().hex}" + ) + workflow_id = f"streams-redis-rejected-{uuid.uuid4().hex}" + try: + 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}", + task_timeout=timedelta(seconds=1), + ) + await asyncio.wait_for(handle.result(), 60) + assert SlowFirstStage.delayed + # Counted before any reader could settle a barrier: the worker did. + assert ( + await _log_length(provider, live_client, workflow_id, DECISIONS.name) + == 2 + ) + stream = provider.get_stream_handle(live_client, workflow_id) + records = [r async for r in stream.read(topic=DECISIONS)] + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"from": "workflow"}), + (RecordKind.FINISH, None), + ] + finally: + await provider.close() + + +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 a5c32adcb..c9febb4dc 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 ( @@ -68,6 +68,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 @@ -250,9 +251,82 @@ 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, + # The refusal of a trimmed cursor lands on the first step on this + # provider, where the case wants it at the call; its own live module + # covers the trimmed floor. + truncate=None, + # A standalone stream's append script keeps a byte total per topic + # and trims by age on every append, so both bounds hold while the + # stream is open. + bounds_standalone_bytes=True, + trims_open_stream_by_age=True, + # 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,