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/streams/providers/redis.py b/temporalio/streams/providers/redis.py index d39167b5b..1dc105a6b 100644 --- a/temporalio/streams/providers/redis.py +++ b/temporalio/streams/providers/redis.py @@ -58,6 +58,15 @@ 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, 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 @@ -86,7 +95,7 @@ import logging import re import time -from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine +from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Sequence from dataclasses import replace from datetime import timedelta from typing import Any, Final, Generic, TypeVar @@ -135,6 +144,7 @@ OutputStageResolutionError, OutputStageStatus, OutputStreamRecord, + StagedOutputRecord, ) from temporalio.contrib.external_workflow_streams._output_client import ( _reconcile_output_stage, @@ -180,7 +190,7 @@ from temporalio.streams.providers import ProviderPlugin from temporalio.worker import ReplayerConfig, WorkerConfig -__all__ = ["RedisProducer", "RedisStreamHandle", "RedisStreams"] +__all__ = ["DEFAULT_RETENTION", "RedisProducer", "RedisStreamHandle", "RedisStreams"] T = TypeVar("T") @@ -192,6 +202,15 @@ _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. @@ -278,6 +297,21 @@ async def _tail_after( 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] = [] @@ -287,14 +321,16 @@ def _fields_args(record: TransportRecord) -> list[Any]: return args -def _append_args(record: TransportRecord, digest: str) -> list[Any]: - """What the log append script takes: the identity and digest, then the fields. +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), @@ -321,23 +357,32 @@ def _stamp_hash(wire: WireRecord) -> None: #: Append one record to a log, or answer where it already is. #: -#: The provider's own write rather than the transport's, so a retry is matched -#: by its plaintext digest rather than by the encoded bytes. The idempotency -#: hash is read before the log is touched, so an identity already used with a -#: different plaintext digest refuses the record rather than writing it, and one -#: used with the same digest answers with the original position, which is what -#: settles a call whose answer was lost. +#: 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 existing = redis.call('HGET', KEYS[2], ARGV[1]) +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[2] then + 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, 3)) -redis.call('HSET', KEYS[2], ARGV[1], id .. '|' .. ARGV[2]) +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} """ @@ -358,20 +403,37 @@ async def write(self, *, name: str, record: TransportRecord, digest: str) -> Off """ outcome, placed = await self._script( keys=[name, f"{name}:idem"], - args=_append_args(record, digest), + args=_append_args(self._backend, record, digest), ) if _text(outcome) == "conflict": raise AppendConflictError(record.idempotency_key) return Offset(_text(placed)) +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"]) + + 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. + """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, @@ -381,8 +443,26 @@ class _TopicLogBackend(RedisStreamBackend): 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, + ) -> None: + super().__init__(client=client, key_prefix=key_prefix) + self._retention = retention + self._max_len = max_len + 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)) @@ -428,6 +508,39 @@ async def read_after( 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. @@ -475,6 +588,15 @@ async def _output_stage_from_offsets( 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() ) @@ -532,6 +654,22 @@ async def aborted(page: list[Any]) -> set[str]: 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. @@ -678,7 +816,7 @@ def _storage_error(error: Exception, what: str) -> StreamError: def _integrity_error(error: Exception, what: str) -> StreamError: # Named as a loss rather than a transient read failure, because no retry brings - # a lost record back and the caller's next move is different. + # a trimmed record back and the caller's next move is different. return StreamNotFoundError(f"{what}: {error}") @@ -791,7 +929,7 @@ async def _connect(self) -> None: chain = await _chain(self._client, self._workflow_id) try: # The transport's producer is bound for its wake and its key: the - # append itself is the provider's, so a retry is matched by digest. + # append itself is the provider's, so the trims ride along with it. input_ = await ExternalStreamProducer.connect( backend=backend, workflow=chain, @@ -843,7 +981,7 @@ async def _place(self, records: list[WireRecord]) -> Offset | None: digest = _plaintext_digest(record) _stamp_hash(record) # Built here rather than handed to the transport's publish, so the - # digest rides along. The identity is the one the transport would + # 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, @@ -963,8 +1101,12 @@ def read( positioned against the log on the first step of the generator, since this call cannot reach the store. - A cursor another provider minted, or one that is not a Redis entry id, - is refused by this call: reading the token needs nothing from the store. + 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. @@ -1003,6 +1145,17 @@ async def _read( 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: @@ -1093,7 +1246,10 @@ async def _closed(self) -> bool: async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: """The cursor of the newest committed record on ``topic``, for following from now. - ``BEGINNING`` when the topic holds no committed record. + ``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() @@ -1161,6 +1317,8 @@ def __init__( 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, ) -> None: """Create the provider. @@ -1172,10 +1330,32 @@ def __init__( holds its Workflow Task open before the worker parks it. client: A ``redis.asyncio.Redis`` the caller opened, with ``decode_responses=False``, and closes itself; the provider - puts its own key layout on top of it. + 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. """ + 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 @@ -1183,6 +1363,8 @@ def __init__( self._backend: _TopicLogBackend | None = None self._owned_client: Any = None self._poll = poll_interval + self._retention = retention + self._max_len = max_len def _require_backend(self) -> _TopicLogBackend: if self._backend is None: @@ -1201,6 +1383,8 @@ def _require_backend(self) -> _TopicLogBackend: self._backend = _TopicLogBackend( client=client, key_prefix=self._key_prefix, + retention=self._retention, + max_len=self._max_len, ) return self._backend diff --git a/tests/streams/test_redis_provider.py b/tests/streams/test_redis_provider.py index c02c250c2..d4e7e1914 100644 --- a/tests/streams/test_redis_provider.py +++ b/tests/streams/test_redis_provider.py @@ -21,6 +21,7 @@ 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, _drive, @@ -218,13 +219,41 @@ async def run() -> None: 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 is the provider's whichever connection it runs on, and - # closing the provider does not close a caller's client. + # 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) + 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() diff --git a/tests/streams/test_redis_replay.py b/tests/streams/test_redis_replay.py index 7d8b016fd..e97756583 100644 --- a/tests/streams/test_redis_replay.py +++ b/tests/streams/test_redis_replay.py @@ -3,7 +3,8 @@ The conformance suite covers the outside surface when ``STREAMS_LIVE=redis``. This module runs the interface loop inside a workflow over the staged commit, lets a read end with the workflow, shares a topic between an outside producer -and the workflow, queries a completed run, which replays it, and seeds a +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. @@ -15,6 +16,7 @@ import os import uuid from collections.abc import AsyncIterator +from datetime import timedelta from typing import Any import pytest @@ -37,6 +39,7 @@ END, Cursor, RecordKind, + StreamCursorError, StreamProducerError, ) from temporalio.streams._wire import to_wire @@ -250,6 +253,130 @@ async def _log_length( 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.""" @@ -560,6 +687,45 @@ async def test_a_committed_stage_is_released_in_the_order_it_was_staged( ] +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 ): diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 1eb10e656..1f6d902fc 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -395,6 +395,10 @@ async def host(workflow_id: str) -> None: client, host=host, task_queue=worker.task_queue, + # 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, # Every stream this provider keeps belongs to a workflow. hosts_standalone_streams=False, )