From c73a3bfe8632db033602d5933e7448c32f2a42ab Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:49:22 -0700 Subject: [PATCH 1/2] Added MemoryStreams, the in-memory reference provider. It holds streams in process memory, so the interface has a provider the conformance suite and unit tests can run without a server. --- temporalio/streams/providers/memory.py | 613 +++++++++++++++++++++++++ 1 file changed, 613 insertions(+) create mode 100644 temporalio/streams/providers/memory.py diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py new file mode 100644 index 000000000..87f4d5116 --- /dev/null +++ b/temporalio/streams/providers/memory.py @@ -0,0 +1,613 @@ +"""The in-process reference provider. + +Exists so the conformance suite can exercise the whole surface without a +store, and to document in one file what a provider owes. Its limits, stated +so nobody mistakes it for evidence: + +- It is not replay-safe. Workflow-side state lives in plain process memory, + so run it with a warm workflow cache and do not use it to demonstrate + recovery. +- A workflow's publish becomes visible at ``publish`` time rather than at + task acceptance, and a failed task's records stay, so it only approximates + rule 1 of the contract. +- Topics are keyed by workflow id rather than by run, so a successor run's + reader from ``BEGINNING`` sees the chain's records. A ``run_id`` on a + handle only decides which run's close ends a read. +- It learns that a workflow closed by describing it, so a handle opened + without a client reads until the caller closes it. +- It keeps every record until :meth:`MemoryStreams.truncate` drops the + oldest ones, which stands in for a store's retention in tests. +- It does not host standalone streams; both standalone calls raise + :class:`temporalio.streams.StreamUnsupportedError`. +- The outside path encodes and decodes bodies through the client's data + converter, codec and external storage included, and fingerprints a retry + over the converted bytes first. The workflow half has no client, so a + workflow's own publish is stored as the payload converter produced it and + a workflow-side read hands records over as stored. + +The outside surface (producer identity, retry deduplication, positions, +supersession, cursors) is faithful, which is what the conformance tests lean +on. One list per topic; a topic is written by the workflow and by outside +producers alike and read from either side. +""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import AsyncGenerator +from datetime import timedelta +from typing import Any, Generic, TypeVar + +from google.protobuf.message import DecodeError + +import temporalio.converter +from temporalio import workflow +from temporalio.client import Client, WorkflowExecutionStatus +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import content_fingerprint, decode_body, encode_body +from temporalio.streams._errors import ( + StreamCursorError, + StreamProducerError, + StreamUnsupportedError, +) +from temporalio.streams._ids import topic_key +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 + +__all__ = ["MemoryProducer", "MemoryStreamHandle", "MemoryStreams"] + +_PROVIDER = "memory" + +T = TypeVar("T") + +logger = logging.getLogger(__name__) + + +def _wake(future: asyncio.Future[None]) -> None: + if not future.done(): + future.set_result(None) + + +class _Topic: + """One topic's records, and the waiters parked on its tail.""" + + def __init__(self) -> None: + # The retained records, the first of which sits at offset ``base``. + # Offsets are never reused, so a cursor keeps naming the same record + # after truncation drops the ones before it. + self.base = 0 + self.records: list[bytes] = [] + # Dedupe identity is (producer#attempt, first sequence of the append), + # the same pair the storage providers use, mapped to where the batch + # landed and a digest of what it held, so a repeat answers with the + # original position and a divergent one is told apart from it. + self.seen: dict[tuple[str, int], tuple[int, int, bytes]] = {} + # Each waiter is parked with the loop it belongs to. A workflow's + # publish runs on the workflow thread, and waking a foreign loop's + # future from there needs call_soon_threadsafe or the loop stays + # blocked in select until unrelated I/O happens to wake it. + self._waiters: list[tuple[asyncio.AbstractEventLoop, asyncio.Future[None]]] = [] + + def append( + self, + wires: list[WireRecord], + *, + writer: str | None = None, + sequence: int = 0, + content: bytes | None = None, + ) -> tuple[int, int]: + """Store ``wires`` and return where they landed as ``(first offset, count)``. + + With a ``writer``, a repeat of ``(writer, sequence)`` carrying the same + content stores nothing and returns where the original landed. + ``content`` is the fingerprint the repeat is matched by; a producer + takes it over the records before their bodies are encoded, and + without one it is taken over ``wires`` as they are. + + Raises: + StreamProducerError: ``(writer, sequence)`` is held with different + content. + """ + key = (writer or "", sequence) + bodies = [wire.SerializeToString(deterministic=True) for wire in wires] + if content is None: + content = content_fingerprint(wires) + if writer is not None: + held = self.seen.get(key) + if held is not None: + first, count, seen_content = held + if seen_content != content: + raise StreamProducerError( + f"producer sequence {sequence} already used with different " + f"content by {writer!r}" + ) + return first, count + first = self.head + self.records.extend(bodies) + if writer is not None: + self.seen[key] = (first, len(wires), content) + waiters, self._waiters = self._waiters, [] + for loop, future in waiters: + loop.call_soon_threadsafe(_wake, future) + return first, len(wires) + + @property + def head(self) -> int: + """The offset the next record lands at.""" + return self.base + len(self.records) + + def at(self, offset: int) -> bytes: + """The retained record at ``offset``.""" + return self.records[offset - self.base] + + def truncate(self, keep: int) -> None: + """Drop all but the newest ``keep`` records.""" + drop = max(0, len(self.records) - keep) + self.base += drop + del self.records[:drop] + + async def wait_past(self, offset: int, timeout: float | None) -> None: + """Wait until a record exists at ``offset``, or ``timeout`` passes.""" + if self.head > offset: + return + loop = asyncio.get_running_loop() + future: asyncio.Future[None] = loop.create_future() + self._waiters.append((loop, future)) + try: + await asyncio.wait_for(future, timeout) + except asyncio.TimeoutError: + pass + finally: + # Dropped on every exit, cancellation included, so a reader that + # aclose()s while parked here leaves nothing behind on the topic. + self._waiters = [w for w in self._waiters if w[1] is not future] + + +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 _MemReadSource: + """Workflow-side read that wakes by polling a timer. + + A real provider wakes the workflow by delivering; polling is the price of + having no delivery path, and it is why this provider is for tests. + """ + + def __init__(self, store: _Topic, start: int, poll: timedelta) -> None: + self._store = store + self._offset = start + self._poll = poll + self._closed = False + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + while not self._closed: + head = self._store.head + if head > self._offset: + batch: list[tuple[Cursor, WireRecord]] = [] + for offset in range(max(self._offset, self._store.base), head): + cursor = mint_cursor(_PROVIDER, str(offset)) + wire = _parse( + cursor, self._store.at(offset), workflow.logger.warning + ) + if wire is not None: + batch.append((cursor, wire)) + self._offset = head + if batch: + return batch + continue + await workflow.sleep(self._poll) + raise StopAsyncIteration + + def close(self) -> None: + self._closed = True + + +class _MemWriteSink: + def __init__(self, store: _Topic) -> None: + self._store = store + + def publish(self, record: WireRecord) -> None: + # Visible at once rather than at task acceptance: the documented gap + # between this provider and rule 1. + self._store.append([record]) + + +class _MemoryWorkflowProvider: + """The workflow half. Nothing to install and nothing to release.""" + + def __init__(self, streams: MemoryStreams) -> None: + self._streams = streams + + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + check_read_start(after, last) + store = self._streams._topic(workflow.info().workflow_id, topic) + start = self._streams._start(store, after, last) + return _MemReadSource(store, start, self._streams._poll) + + def open_writer(self, topic: str) -> WriteSink: + return _MemWriteSink(self._streams._topic(workflow.info().workflow_id, topic)) + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +class MemoryProducer(Generic[T]): + """The outside producer, faithful to the contract.""" + + def __init__( + self, + store: _Topic, + converter: temporalio.converter.DataConverter, + topic: str, + producer_id: str, + attempt: int, + ) -> None: + """Bind this producer to ``topic``'s ``store``.""" + self._store = store + self._converter = converter + self._topic = topic + self._producer_id = producer_id + self._attempt = attempt + # One-based, because zero on the wire says the producer does not + # number its records and this one does. + self._sequence = 1 + self._last = BEGINNING + + @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 _writer(self) -> str: + 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 of the same content returns where the original landed; an + empty call returns the position of this producer's last record. + + Raises: + StreamProducerError: This sequence is held with different content. + """ + if not values: + return self._last + return await self._write( + [ + to_wire( + self._converter.payload_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.""" + await self._write( + [ + to_wire( + self._converter.payload_converter, + topic=self._topic, + kind=RecordKind.FINISH, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + ) + ] + ) + + async def _write(self, wires: list[WireRecord]) -> Cursor: + # The fingerprint comes first, over the converted records, so a codec + # that encrypts with a fresh nonce cannot make a retry look divergent. + content = content_fingerprint(wires) + for wire in wires: + await encode_body(self._converter, wire) + first, count = self._store.append( + wires, writer=self._writer, sequence=self._sequence, content=content + ) + self._sequence += len(wires) + self._last = mint_cursor(_PROVIDER, str(first + count - 1)) + return self._last + + +class MemoryStreamHandle: + """One workflow's stream from outside, with the shared reader rules.""" + + def __init__( + self, + streams: MemoryStreams, + client: Client | None, + workflow_id: str, + run_id: str | None, + ) -> None: + """Address ``workflow_id``'s topics in ``streams``.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._converter = ( + client.data_converter + if client is not None + else temporalio.converter.DataConverter.default + ) + + def read( + self, + *, + topic: str | StreamTopic[Any] | None = None, + after: Cursor = BEGINNING, + last: int | None = None, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` from where the read starts until the workflow closes. + + ``END`` and ``last=`` are resolved by this call, against what the + topic holds when it is made. + """ + check_read_start(after, last) + name, result_type = resolve_topic(topic, result_type) + store = self._streams._topic(self._workflow_id, name) + # Parsed here so a foreign cursor fails this call, not the first + # iteration of the generator. + start = self._streams._start(store, after, last) + # The decoder positions a synthesized record at the one before it, so + # it is told the position before the first record this read yields. + previous = mint_cursor(_PROVIDER, str(start - 1)) if start else BEGINNING + return self._read(store, start, previous, result_type) + + async def _read( + self, + store: _Topic, + offset: int, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + decoder = RecordDecoder( + self._converter.payload_converter, + result_type, + after=after, + warn=logger.warning, + ) + closed = False + while True: + while offset < store.head: + if offset < store.base: + raise StreamCursorError( + f"offset {offset} was truncated while this read was behind; " + f"the topic now starts at {store.base}" + ) + cursor = mint_cursor(_PROVIDER, str(offset)) + wire = _parse(cursor, store.at(offset), logger.warning) + offset += 1 + if wire is None: + continue + await decode_body(self._converter, wire) + for record in decoder.decode(cursor, wire): + yield record + if closed: + return + # One more pass after learning the workflow closed, so a record + # that landed between the scan and the describe is not lost. + closed = await self._closed() + if not closed: + await store.wait_past( + offset, + None + if self._client is None + else self._streams._poll.total_seconds(), + ) + + async def _closed(self) -> bool: + if self._client is None: + return False + handle = self._client.get_workflow_handle( + self._workflow_id, run_id=self._run_id + ) + try: + description = await handle.describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + # A producer may write before the workflow exists; there is + # nothing to follow yet, so keep waiting. + return False + raise + status = description.status + if status is None or status == WorkflowExecutionStatus.RUNNING: + return False + # Following the chain, a run that continued as new is not the end: + # the next describe without a run id finds its successor. + return not ( + self._run_id is None and status == WorkflowExecutionStatus.CONTINUED_AS_NEW + ) + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The cursor of the newest record on ``topic``, for following from now.""" + name, _ = resolve_topic(topic) + head = self._streams._topic(self._workflow_id, name).head + return mint_cursor(_PROVIDER, str(head - 1)) if head else BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> MemoryProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + name, _ = resolve_topic(topic) + store = self._streams._topic(self._workflow_id, name) + producer_id, attempt = producer_identity(producer_id, attempt) + return MemoryProducer(store, self._converter, name, producer_id, attempt) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this workflow's stream, pinned as this handle is.""" + return StreamRef.for_workflow( + self._workflow_id, run_id=self._run_id, topic=topic + ) + + async def close(self) -> None: + """Refuse: a workflow's stream ends with the workflow, not by a caller.""" + raise ValueError( + "only a standalone stream can be closed; this handle is on a workflow's " + "stream, which ends when the workflow does" + ) + + +class MemoryStreams(ProviderPlugin): + """The in-memory provider, one list per topic. + + Construct one and pass the same instance to the worker and to the code + that opens handles; two instances share nothing. + """ + + def __init__( + self, *, poll_interval: timedelta = timedelta(milliseconds=100) + ) -> None: + """Create an empty provider. + + Args: + poll_interval: How often a workflow-side reader with nothing to + read checks again, and how often an outside reader asks + whether the workflow closed. + """ + self._poll = poll_interval + self._topics: dict[str, _Topic] = {} + + def reset(self) -> None: + """Drop every topic. For tests.""" + self._topics.clear() + + def truncate(self, workflow_id: str, topic: str, *, keep: int) -> None: + """Drop all but the newest ``keep`` records of a topic. For tests. + + Stands in for a store's retention: offsets are kept, so a cursor from + before still names its record, and a read from ``BEGINNING`` starts + at the oldest one left. + """ + self._topic(workflow_id, topic).truncate(keep) + + def workflow_provider(self) -> _MemoryWorkflowProvider: + """The workflow half, over this provider's topics.""" + return _MemoryWorkflowProvider(self) + + def get_stream_handle( + self, client: Client | None, workflow_id: str, *, run_id: str | None = None + ) -> MemoryStreamHandle: + """A handle on ``workflow_id``'s topics. + + ``client`` may be ``None`` here, unlike on a storage provider; then + the handle cannot see the workflow close and a read waits until the + caller closes it. + """ + return MemoryStreamHandle(self, client, workflow_id, run_id) + + async def create_standalone_stream( + self, + client: Client | None, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> MemoryStreamHandle: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the memory provider does not host standalone streams" + ) + + def get_standalone_stream_handle( + self, client: Client | None, stream_id: str + ) -> MemoryStreamHandle: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the memory provider does not host standalone streams" + ) + + async def close(self) -> None: + """Nothing to release: the provider holds no connection.""" + + def _topic(self, workflow_id: str, topic: str) -> _Topic: + if not topic: + raise ValueError("topic must not be empty") + key = topic_key(workflow_id, topic) + found = self._topics.get(key) + if found is None: + found = self._topics[key] = _Topic() + return found + + def _start(self, store: _Topic, after: Cursor, last: int | None) -> int: + """The offset a read starts at, resolved against what ``store`` holds now.""" + if last is not None: + return max(store.base, store.head - last) + if after == END: + return store.head + position = cursor_position(after, provider=_PROVIDER) + if position is None: + return store.base + try: + start = int(position) + 1 + except ValueError: + raise StreamCursorError( + f"cursor {after.token!r} does not name a position on the memory provider" + ) from None + if start < store.base: + raise StreamCursorError( + f"cursor {after.token!r} names a record no longer retained; the " + f"topic starts at offset {store.base}" + ) + return start From c8f45d1d7e168303861eb25f72b9de47d09b36c2 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 01:49:23 -0700 Subject: [PATCH 2/2] Added the provider conformance suite and ran it on memory. The suite is parametrised by provider, and capability markers skip what a provider does not claim. Every later provider registers here. --- tests/streams/conftest.py | 17 + tests/streams/test_memory_provider.py | 69 +++ tests/streams/test_streams_conformance.py | 674 ++++++++++++++++++++++ tests/streams/test_streams_internals.py | 24 + 4 files changed, 784 insertions(+) create mode 100644 tests/streams/conftest.py create mode 100644 tests/streams/test_memory_provider.py create mode 100644 tests/streams/test_streams_conformance.py diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py new file mode 100644 index 000000000..b8a152f20 --- /dev/null +++ b/tests/streams/conftest.py @@ -0,0 +1,17 @@ +import pytest + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "reports_positions: the case needs append() to return where records landed", + ) + config.addinivalue_line( + "markers", + "detects_divergent_retries: the case needs append() to compare a repeat's " + "content with what the store holds", + ) + config.addinivalue_line( + "markers", + "truncates: the case needs a way to drop a topic's oldest records", + ) diff --git a/tests/streams/test_memory_provider.py b/tests/streams/test_memory_provider.py new file mode 100644 index 000000000..3fd83a9fb --- /dev/null +++ b/tests/streams/test_memory_provider.py @@ -0,0 +1,69 @@ +"""What the reference provider does that the conformance suite cannot see. + +The conformance suite is the public surface, so it can say that closing a read +returns and that the topic still works afterwards, but not that the provider +let go of what the read parked on. That is this file: a few assertions against +``MemoryStreams`` internals, where holding on would leak quietly. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from temporalio.api.stream.v1 import StreamRecord +from temporalio.streams import CONTENT_HASH_KEY, StreamProducerError, topic +from temporalio.streams.providers.memory import MemoryStreams + +OUT = topic("out", dict) + + +async def test_closing_a_parked_read_drops_its_waiter(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-parked") # type: ignore[arg-type] + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}) + store = provider._topic("wf-parked", OUT.name) + + records = stream.read(topic=OUT) + await asyncio.wait_for(records.__anext__(), 5.0) + + pending = asyncio.ensure_future(records.__anext__()) + await asyncio.sleep(0.2) + assert store._waiters, "the read should be parked on the topic by now" + + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + # Nothing left behind: a reader that comes and goes must not grow this + # list for the life of the topic. + assert store._waiters == [] + await asyncio.wait_for(records.aclose(), 5.0) + + +async def test_a_divergent_retry_leaves_the_store_alone(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-divergent") # type: ignore[arg-type] + store = provider._topic("wf-divergent", OUT.name) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}) + assert len(store.records) == 1 + + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + with pytest.raises( + StreamProducerError, match="already used with different content" + ): + await retry.append({"n": 2}) + assert len(store.records) == 1 + + +async def test_a_stored_record_carries_the_plaintext_hash(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-hash") # type: ignore[arg-type] + await stream.producer(topic=OUT, producer_id="model", attempt=1).append({"n": 1}) + stored = StreamRecord.FromString(provider._topic("wf-hash", OUT.name).records[0]) + # What the store holds is the record after encode_body: the hash the + # server-side dedupe reads is on it, under the shared key. + assert stored.metadata[CONTENT_HASH_KEY].data.decode().isalnum() + assert len(stored.metadata[CONTENT_HASH_KEY].data) == 64 diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py new file mode 100644 index 000000000..fb819bcb4 --- /dev/null +++ b/tests/streams/test_streams_conformance.py @@ -0,0 +1,674 @@ +"""Conformance tests for the stream contract's outside surface. + +Written against the public surface, parametrised over the providers this +tree can stand up. The memory provider always runs, with no server and no +store. A storage provider adds itself to ``SETUPS``, behind its own +``STREAMS_LIVE`` gate when it needs a store the test environment does not +start: its setup receives the environment's client and hands back a provider +instance and which capabilities it lacks, so the cases marked +``reports_positions`` are skipped with a reason on a provider whose +``append()`` learns positions at read time. + +What this file pins down is what a provider owes: producer identity, retry +deduplication, positions, supersession, topic addressing, cursor resumption, +cursor ownership, releasing a read the caller stopped early, naming a stream +as a ``StreamRef``, and running bodies through the client's data converter so +external storage applies and a retry through a nondeterministic codec still +matches its original. Every case here goes through the public surface, so a +new provider answers this file and nothing else. The shared pieces no provider +implements are unit-tested in ``test_streams_internals``; the workflow-side +handles and the two rules about Workflow Tasks live in +``test_streams_workflow``. +""" + +from __future__ import annotations + +import asyncio +import dataclasses +import os +import uuid +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from dataclasses import dataclass +from typing import Any + +import pytest + +from temporalio.api.common.v1 import Payload +from temporalio.client import Client +from temporalio.common import RawValue +from temporalio.converter import ( + DataConverter, + ExternalStorage, + PayloadCodec, + StorageDriver, + StorageDriverClaim, + StorageDriverRetrieveContext, + StorageDriverStoreContext, +) +from temporalio.streams import ( + BEGINNING, + DEFAULT_TOPIC, + END, + Cursor, + RecordKind, + StreamCursorError, + StreamHandle, + StreamProducerError, + StreamProvider, + StreamRef, + Supersession, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams + +# Defined once and shared by every case, the way an application shares them +# between its workflow, its activities and its backend. +OUT = topic("out", dict) +A = topic("a", dict) +B = topic("b", dict) +Y = topic("y", dict) +XY = topic("x:y", dict) + + +@dataclass +class ProviderCase: + """One provider under test, and what the cases may ask of it.""" + + name: str + provider: StreamProvider + reports_positions: bool = True + """``append()`` returns where the records landed.""" + detects_divergent_retries: bool = True + """``append()`` compares a repeat's content with what it already holds.""" + truncate: Callable[[str, str, int], Awaitable[None]] | None = None + """Drops all but the newest records of a workflow's topic, standing in + for retention, or ``None`` when the provider offers no way to.""" + bounds_standalone_bytes: bool = True + """A standalone stream's policy can bound the bytes it keeps.""" + refuses_appends_past_byte_cap: bool = False + """The byte bound refuses an append that would cross it, instead of + dropping the oldest records to make room.""" + trims_open_stream_by_age: bool = True + """A standalone stream drops records older than ``retention`` while it is + open, rather than keeping them that long after it closes.""" + + async def open( + self, + workflow_id: str, + *, + run_id: str | None = None, + client: Client | None = None, + ) -> StreamHandle: + if client is not None: + # The explicit form, for a case that needs the handle to encode + # bodies through this client's data converter. + return self.provider.get_stream_handle(client, workflow_id, run_id=run_id) + # Only the memory provider gets here, and it takes no client. + return self.provider.get_stream_handle( + None, # type: ignore[arg-type] + workflow_id, + run_id=run_id, + ) + + +class RecordingDriver(StorageDriver): + """An in-memory external storage driver that counts what it was asked to hold.""" + + def __init__(self) -> None: + self.held: dict[str, bytes] = {} + self.stored = 0 + self.retrieved = 0 + + def name(self) -> str: + return "recording" + + async def store( + self, context: StorageDriverStoreContext, payloads: Sequence[Payload] + ) -> list[StorageDriverClaim]: + claims: list[StorageDriverClaim] = [] + for payload in payloads: + key = f"payload-{len(self.held)}" + self.held[key] = payload.SerializeToString() + self.stored += 1 + claims.append(StorageDriverClaim(claim_data={"key": key})) + return claims + + async def retrieve( + self, + context: StorageDriverRetrieveContext, + claims: Sequence[StorageDriverClaim], + ) -> list[Payload]: + self.retrieved += len(claims) + return [Payload.FromString(self.held[c.claim_data["key"]]) for c in claims] + + +class NonceCodec(PayloadCodec): + """A codec whose output differs on every call, as one that encrypts with a fresh nonce does.""" + + def __init__(self) -> None: + self.encoded = 0 + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + self.encoded += len(payloads) + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=os.urandom(16) + p.SerializeToString(), + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [Payload.FromString(p.data[16:]) for p in payloads] + + +def _client_with(client: Client, converter: DataConverter) -> Client: + # The same connection, carrying the converter the case wants bodies to + # pass through. + config = client.config() + config["data_converter"] = converter + return Client(**config) + + +async def _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: + provider = MemoryStreams() + + async def truncate(workflow_id: str, topic: str, keep: int) -> None: + provider.truncate(workflow_id, topic, keep=keep) + + yield ProviderCase( + "memory", + provider, + truncate=truncate, + bounds_standalone_bytes=True, + refuses_appends_past_byte_cap=False, + trims_open_stream_by_age=True, + ) + provider.reset() + + +SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { + "memory": _memory_case +} + +_CAPABILITIES = { + "reports_positions": lambda case: case.reports_positions, + "detects_divergent_retries": lambda case: case.detects_divergent_retries, + "truncates": lambda case: case.truncate is not None, +} + + +@pytest.fixture(params=sorted(SETUPS)) +async def case( + request: pytest.FixtureRequest, client: Client +) -> AsyncIterator[ProviderCase]: + async for provider_case in SETUPS[request.param](client): + for marker, supported in _CAPABILITIES.items(): + if request.node.get_closest_marker(marker) and not supported(provider_case): + pytest.skip(f"the {provider_case.name} provider does not {marker}") + yield provider_case + + +def new_workflow_id() -> str: + # Unique per case, because a storage provider keeps what earlier cases + # wrote and the memory provider only happens to forget. + return f"wf-{uuid.uuid4().hex}" + + +async def take(records: Any, count: int, timeout: float = 5.0) -> list: + out: list = [] + + async def _collect() -> None: + async for record in records: + out.append(record) + if len(out) >= count: + return + + await asyncio.wait_for(_collect(), timeout) + return out + + +async def test_append_read_roundtrip(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + assert (producer.producer_id, producer.attempt) == ("model", 1) + await producer.append({"id": "r1"}, {"id": "r2"}) + await producer.finish() + + records = await take(stream.read(topic=OUT), 3) + assert [r.kind for r in records] == [ + RecordKind.DATA, + RecordKind.DATA, + RecordKind.FINISH, + ] + assert [r.value for r in records[:2]] == [{"id": "r1"}, {"id": "r2"}] + assert records[2].value is None + assert all(r.producer_id == "model" and r.attempt == 1 for r in records) + assert [r.sequence for r in records] == [1, 2, 3] + assert all(r.topic == OUT.name for r in records) + + +async def test_raw_values_pass_through_untouched(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + payload = Payload(metadata={"encoding": b"binary/plain"}, data=b"\x00\x01raw") + producer = stream.producer(topic=OUT.name, producer_id="model", attempt=1) + await producer.append(RawValue(payload)) + + records = await take(stream.read(topic="out", result_type=RawValue), 1) + assert isinstance(records[0].value, RawValue) + assert records[0].value.payload == payload + + +@pytest.mark.reports_positions +async def test_retried_append_returns_the_original_position(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + landed = await first.append({"id": "r1"}) + assert landed is not None + # The retry of the same attempt starts its sequence over and appends the + # same record. The provider stores it once and answers with where the + # original landed, so the retry can checkpoint the same position. + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + assert await retry.append({"id": "r1"}) == landed + # An empty call writes nothing and answers the same way. + assert await retry.append() == landed + + records = await take(stream.read(topic=OUT), 1) + assert records[0].value == {"id": "r1"} + assert records[0].cursor == landed + # The store holds exactly the one record: the newest position is its cursor. + assert await stream.latest(topic=OUT) == landed + + +async def test_retried_append_is_stored_once(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + await first.append({"id": "r1"}) + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + await retry.append({"id": "r1"}) + await retry.append({"id": "r2"}) + + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [{"id": "r1"}, {"id": "r2"}] + + +@pytest.mark.detects_divergent_retries +async def test_a_divergent_retry_is_refused(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + await first.append({"id": "r1"}) + + # Same producer, attempt and sequence, different content. The store has no + # way to know which of the two the reader was meant to see, so it says so + # rather than answering with the position of the one it kept. + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + with pytest.raises(StreamProducerError): + await retry.append({"id": "other"}) + + # And it wrote nothing: the producer that owns the sequence carries on + # past the original, with no second record wedged in front of it. + await first.append({"id": "r2"}) + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [{"id": "r1"}, {"id": "r2"}] + + +async def test_closing_a_read_early_releases_it(case: ProviderCase): + # A read with nothing left to hand over waits against the store. Closing + # the generator is how a caller that stops early says so, and it has to + # let go of whatever it parked instead of hanging on it. + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}) + + records = stream.read(topic=OUT) + assert (await asyncio.wait_for(records.__anext__(), 5.0)).value == {"n": 1} + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(records.__anext__(), 0.5) + await asyncio.wait_for(records.aclose(), 5.0) + + # The topic is untouched by the close: a new read still sees everything. + await producer.append({"n": 2}) + again = await take(stream.read(topic=OUT), 2) + assert [r.value for r in again] == [{"n": 1}, {"n": 2}] + + +async def test_new_attempt_supersedes_the_old_one(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + await first.append({"text": "The capital of"}) + second = stream.producer(topic=OUT, producer_id="model", attempt=2) + await second.append({"text": "Paris is the capital"}) + + records = await take(stream.read(topic=OUT), 3) + assert records[0].kind is RecordKind.DATA and records[0].attempt == 1 + assert records[1].kind is RecordKind.SUPERSEDED + assert records[1].supersession == Supersession("model", 1, 2) + assert records[1].value is None + assert records[2].kind is RecordKind.DATA and records[2].attempt == 2 + + +async def test_a_superseded_record_resumes_to_the_triggering_record( + case: ProviderCase, +): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + await first.append({"n": 1}) + second = stream.producer(topic=OUT, producer_id="model", attempt=2) + await second.append({"n": 2}) + + records = await take(stream.read(topic=OUT), 3) + superseded = records[1] + assert superseded.kind is RecordKind.SUPERSEDED + # The synthesized record sits at the position before the new attempt's + # first record, so a consumer that checkpoints it and restarts is handed + # that record rather than skipping it. + assert superseded.cursor == records[0].cursor + resumed = await take(stream.read(topic=OUT, after=superseded.cursor), 1) + assert resumed[0].kind is RecordKind.DATA + assert resumed[0].value == {"n": 2} + + +async def test_topics_are_addressed_by_name(case: ProviderCase): + # Two producers on two topics of the same workflow's stream: each read + # names its topic and sees only that topic's records. + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + on_a = stream.producer(topic=A, producer_id="tool-a", attempt=1) + await on_a.append({"n": 1}) + on_b = stream.producer(topic=B, producer_id="tool-b", attempt=1) + await on_b.append({"n": 2}) + + only_a = await take(stream.read(topic=A), 1) + assert [(r.topic, r.value) for r in only_a] == [("a", {"n": 1})] + only_b = await take(stream.read(topic=B), 1) + assert [(r.topic, r.value) for r in only_b] == [("b", {"n": 2})] + + +async def test_naming_no_topic_addresses_the_default_topic(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + assert await stream.latest() == BEGINNING + producer = stream.producer(producer_id="model", attempt=1) + await producer.append({"n": 1}) + await stream.producer(topic=OUT, producer_id="model", attempt=1).append({"n": 2}) + + records = await take(stream.read(), 1) + assert [(r.topic, r.value) for r in records] == [(DEFAULT_TOPIC, {"n": 1})] + assert await stream.latest() == records[0].cursor + # The default is an ordinary name, so naming it is the same topic. + named = await take(stream.read(topic=DEFAULT_TOPIC, result_type=dict), 1) + assert [r.value for r in named] == [{"n": 1}] + assert await stream.latest(topic=DEFAULT_TOPIC) == records[0].cursor + + +async def test_cursor_resumes_where_it_points(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + + records = await take(stream.read(topic=OUT), 3) + checkpoint = records[0].cursor + + # Resuming after a record hands back everything past it and nothing + # twice, without the reader ever advancing a cursor itself. + again = await take(stream.read(topic=OUT, after=checkpoint), 2) + assert [r.value for r in again] == [{"n": 2}, {"n": 3}] + + +@pytest.mark.reports_positions +async def test_append_cursor_names_the_last_record_of_the_batch(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + appended = await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + assert appended is not None + then = await producer.append({"n": 4}) + + # A producer that resumes a reader after its own append must see only + # what came later, not the tail of the batch it just wrote. + records = await take(stream.read(topic=OUT, after=appended), 1) + assert [r.value for r in records] == [{"n": 4}] + assert records[0].cursor == then + assert await producer.append() == then + + +async def test_latest_positions_a_reader_at_the_end(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + assert await stream.latest(topic=OUT) == BEGINNING + + await producer.append({"n": 1}, {"n": 2}) + since = await stream.latest(topic=OUT) + await producer.append({"n": 3}) + + # A reader that positioned itself before the last append sees only what + # came after, which is how a client follows a turn it is about to start. + records = await take(stream.read(topic=OUT, after=since), 1) + assert [r.value for r in records] == [{"n": 3}] + + +async def test_topic_addresses_with_colons_do_not_share_a_store(case: ProviderCase): + # ("wf:x", "y") and ("wf", "x:y") differ only in where the colon sits. + base = new_workflow_id() + left = await case.open(f"{base}:x") + right = await case.open(base) + await left.producer(topic=Y, producer_id="l", attempt=1).append({"side": "left"}) + await right.producer(topic=XY, producer_id="r", attempt=1).append({"side": "right"}) + + only_left = await take(left.read(topic=Y), 1) + assert [r.value for r in only_left] == [{"side": "left"}] + assert await left.latest(topic=Y) == only_left[0].cursor + only_right = await take(right.read(topic=XY), 1) + assert [r.value for r in only_right] == [{"side": "right"}] + assert await right.latest(topic=XY) == only_right[0].cursor + + +async def test_a_foreign_cursor_is_refused_at_the_call(case: ProviderCase): + stream = await case.open(new_workflow_id()) + # Refused by read() itself, not by the first iteration of its generator, + # so the caller's except clause is where the mistake surfaces. + with pytest.raises(StreamCursorError): + stream.read(topic=OUT, after=Cursor("elsewhere:42")) + + +async def test_argument_mistakes_are_value_errors(case: ProviderCase): + stream = await case.open(new_workflow_id()) + with pytest.raises(ValueError): + stream.read(topic="") + with pytest.raises(ValueError): + stream.producer(topic="", producer_id="model", attempt=1) + # Outside an activity there is no identity to fall back on. + with pytest.raises(ValueError, match="producer_id is required"): + stream.producer(topic=OUT) + + +async def test_a_definition_carries_its_type_once(case: ProviderCase): + stream = await case.open(new_workflow_id()) + with pytest.raises(ValueError, match="already carries its type"): + stream.read(topic=OUT, result_type=dict) # type: ignore[call-overload] + with pytest.raises(ValueError): + topic("", dict) + # A string names a topic decided at runtime, and the hint rides the call. + assert await stream.latest(topic=OUT.name) == BEGINNING + + +async def test_last_n_starts_at_the_newest_records(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + + newest = await take(stream.read(topic=OUT, last=2), 2) + assert [r.value for r in newest] == [{"n": 3}, {"n": 4}] + # Fewer records than asked for is all of them, not an error. + everything = await take(stream.read(topic=OUT, last=100), 4) + assert [r.value for r in everything] == [{"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}] + # The cursors it yields are ordinary cursors, so a resume after one works. + again = await take(stream.read(topic=OUT, after=newest[0].cursor), 1) + assert [r.value for r in again] == [{"n": 4}] + + +async def test_last_n_counts_finish_records(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + await producer.finish() + + records = await take(stream.read(topic=OUT, last=2), 2) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"n": 2}), + (RecordKind.FINISH, None), + ] + + +async def test_end_reads_only_what_arrives_after_the_read_starts(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": "old"}, {"n": "old"}) + + records = stream.read(topic=OUT, after=END) + first = asyncio.ensure_future(records.__anext__()) + # END resolves when the read starts, and nothing says when that was, so + # appends keep coming until the reader takes one. + try: + for _ in range(100): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({first}, timeout=0.1) + if done: + break + record = await asyncio.wait_for(first, 5) + finally: + await records.aclose() + assert record.value == {"n": "new"} + + +@pytest.mark.truncates +async def test_beginning_starts_at_the_oldest_record_still_held(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + before = await take(stream.read(topic=OUT), 1) + assert case.truncate is not None + await case.truncate(workflow_id, OUT.name, 2) + + # BEGINNING is the oldest record retained, not offset zero, which a + # truncated stream no longer holds. + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [{"n": 3}, {"n": 4}] + newest = await take(stream.read(topic=OUT, last=3), 2) + assert [r.value for r in newest] == [{"n": 3}, {"n": 4}] + with pytest.raises(StreamCursorError): + stream.read(topic=OUT, after=before[0].cursor) + + +async def test_a_read_start_names_one_place(case: ProviderCase): + stream = await case.open(new_workflow_id()) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + appended = await producer.append({"n": 1}) + for last in (0, -1, True): + with pytest.raises(ValueError, match="positive"): + stream.read(topic=OUT, last=last) + if appended is not None: + with pytest.raises(ValueError, match="either after= or last="): + stream.read(topic=OUT, after=appended, last=1) + with pytest.raises(ValueError, match="either after= or last="): + stream.read(topic=OUT, after=END, last=1) + + +async def test_a_ref_names_the_stream_and_round_trips_as_data(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + ref = stream.ref(topic=OUT) + assert ref == StreamRef.for_workflow(workflow_id, topic="out") + assert (ref.kind, ref.run_id, ref.activity_id, ref.stream_id) == ( + "workflow", + None, + None, + None, + ) + # Without a topic the ref names the default topic, like every other call. + assert stream.ref().topic == DEFAULT_TOPIC + assert stream.ref().with_topic(A) == stream.ref(topic=A) + # A pinned handle hands out a pinned ref. + pinned = await case.open(workflow_id, run_id="run-1") + assert pinned.ref(topic=OUT).run_id == "run-1" + + # Plain data through the default converter, so it can be a workflow + # argument, an activity result or a Nexus operation input or result. + converter = DataConverter.default + [carried] = await converter.decode(await converter.encode([ref]), [StreamRef]) + assert carried == ref + + +async def test_an_owned_stream_cannot_be_closed_by_a_handle(case: ProviderCase): + # A workflow's stream ends with the workflow; close() is for a stream + # that stands alone. + stream = await case.open(new_workflow_id()) + with pytest.raises(ValueError, match="standalone"): + await stream.close() + + +async def test_a_body_above_the_threshold_is_offloaded_and_read_back( + case: ProviderCase, client: Client +): + driver = RecordingDriver() + converter = dataclasses.replace( + DataConverter.default, + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=256), + ) + workflow_id = new_workflow_id() + stream = await case.open(workflow_id, client=_client_with(client, converter)) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + small = {"n": 1} + large = {"blob": "x" * 1024} + await producer.append(small) + await producer.append(large) + # Only the body over the threshold left the record; the small one stayed + # inline, as it would on any other payload the SDK sends. + assert driver.stored == 1 + + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [small, large] + assert driver.retrieved == 1 + + +@pytest.mark.detects_divergent_retries +async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( + case: ProviderCase, client: Client +): + codec = NonceCodec() + converter = dataclasses.replace(DataConverter.default, payload_codec=codec) + workflow_id = new_workflow_id() + stream = await case.open(workflow_id, client=_client_with(client, converter)) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + landed = await first.append({"id": "r1"}) + assert codec.encoded == 1 + + # The codec produced different bytes for the retry. The provider matched + # it by the plaintext it converted, so it is the same append: stored + # once, answered with the original position. + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + again = await retry.append({"id": "r1"}) + if landed is not None: + assert again == landed + # And a retry that really does differ is still told apart. + divergent = stream.producer(topic=OUT, producer_id="model", attempt=1) + with pytest.raises(StreamProducerError): + await divergent.append({"id": "other"}) + + await first.append({"id": "r2"}) + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [{"id": "r1"}, {"id": "r2"}] diff --git a/tests/streams/test_streams_internals.py b/tests/streams/test_streams_internals.py index 598877eaa..8e3d09feb 100644 --- a/tests/streams/test_streams_internals.py +++ b/tests/streams/test_streams_internals.py @@ -16,6 +16,7 @@ import pytest from temporalio.api.common.v1 import Payload +from temporalio.client import ClientConfig from temporalio.converter import ( DataConverter, ExternalStorage, @@ -41,6 +42,8 @@ encode_body, ) from temporalio.streams._policy import AttemptTracker +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.worker import ReplayerConfig, WorkerConfig def test_record_roundtrips_through_the_wire(): @@ -114,6 +117,27 @@ def test_cursors_name_their_provider(): _wire.cursor_position(Cursor("redis:1700000000000-0"), provider="memory") +def test_registering_a_provider_twice_is_refused(): + # There is one slot on each of the three, and a user who passes a provider + # by hand and a provider plugin, or two provider plugins, meant both. + first, second = MemoryStreams(), MemoryStreams() + with pytest.raises(ValueError, match="already registered"): + second.configure_client(ClientConfig(stream_provider=first)) # type: ignore[typeddict-item] + with pytest.raises(ValueError, match="already registered"): + second.configure_worker(WorkerConfig(stream_provider=first)) # type: ignore[typeddict-item] + with pytest.raises(ValueError, match="already registered"): + second.configure_replayer(ReplayerConfig(stream_provider=first)) # type: ignore[typeddict-item] + + +def test_registering_the_same_provider_twice_is_fine(): + # A worker built from a client that already carries the plugin configures + # it again with the same object, which is not a conflict. + provider = MemoryStreams() + config = provider.configure_client(ClientConfig(stream_provider=provider)) # type: ignore[typeddict-item] + assert config.get("stream_provider") is provider + assert provider.configure_client(ClientConfig()).get("stream_provider") is provider # type: ignore[typeddict-item] + + class _HoldEverything(StorageDriver): """A driver that keeps every payload it is handed, in memory."""