diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index b5db0c4b7..999ad8bba 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -27,6 +27,7 @@ import temporalio.converter import temporalio.runtime import temporalio.service +import temporalio.streams import temporalio.workflow from temporalio.service import ( ConnectConfig, @@ -269,6 +270,7 @@ def __init__( default_workflow_query_reject_condition: None | (temporalio.common.QueryRejectCondition) = None, header_codec_behavior: HeaderCodecBehavior = HeaderCodecBehavior.NO_CODEC, + stream_provider: temporalio.streams.StreamProvider | None = None, ): """Create a Temporal client from a service client. @@ -283,6 +285,7 @@ def __init__( interceptors=interceptors, default_workflow_query_reject_condition=default_workflow_query_reject_condition, header_codec_behavior=header_codec_behavior, + stream_provider=stream_provider, ) self._initial_config = config.copy() @@ -3032,3 +3035,4 @@ class ClientConfig(TypedDict, total=False): temporalio.common.QueryRejectCondition | None ] header_codec_behavior: Required[HeaderCodecBehavior] + stream_provider: temporalio.streams.StreamProvider | None diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py new file mode 100644 index 000000000..9b7cb2df2 --- /dev/null +++ b/temporalio/streams/__init__.py @@ -0,0 +1,162 @@ +"""Streams: a channel a workflow reads, decides on, and writes. + +.. warning:: + This module is experimental and may change in future versions. The + design is meant to be the shape that goes GA; the label is the SDK's + release convention, not a licence to break it. + +The contract, in five statements: + +1. **A workflow publishes only to topics of its own stream, and it publishes + transactionally.** ``temporalio.workflow.StreamWriter.publish`` + returns at once. The record is visible when the Workflow Task is accepted, + and never if the task fails, so no reader can see a decision the workflow + did not commit. +2. **Reading is an observation, and the SDK records it.** What + ``temporalio.workflow.StreamReader`` handed to workflow code, + including the boundary where it found nothing, is committed with the + commands that reading produced. Recovery re-supplies the same records in + the same order. +3. **Anything that does I/O publishes on its own account.** An activity or an + outside process writes through a :class:`StreamProducer` with a producer + id, an attempt and a sequence, and its records are visible as soon as the + store accepts them. Those three let a reader tell a retry from a new + generation. A retry that carries the same content is written once; one that + carries different content at the same sequence is refused with + :class:`StreamProducerError`, so a divergent retry is never dropped in + silence. +4. **A cursor is opaque and belongs to its provider.** Hand it back to resume + strictly after the record it names; :meth:`StreamHandle.latest` positions a + follower. Do not compare two cursors or do arithmetic on one. A read with + no cursor yet starts at :data:`BEGINNING`, at :data:`END`, or at the last + ``N`` records with ``last=N``. +5. **A workflow addresses its streams relative to itself, by topic.** A topic + can be written by the workflow and by outside producers, and read by the + workflow and by outside consumers; which of those happen is the + application's business. A topic is defined once with :func:`topic`, with + the type its records decode to, and that definition is shared by the + workflow, its activities and the backend; a plain string names a topic + decided at runtime. A call that names no topic addresses the workflow's + default topic, :data:`DEFAULT_TOPIC`. + +A provider is an object, registered once as a plugin: +``Client.connect(plugins=[provider])``; workers built from that client inherit +it, and ``Worker(plugins=[provider])`` or ``Replayer(plugins=[provider])`` +registers it on a worker alone. Each context then asks for its stream the +same way. Workflow code uses ``temporalio.workflow.stream_reader`` and +``temporalio.workflow.stream_writer``. An activity uses +``temporalio.activity.stream_handle``, which is its own workflow pinned +to its run unless told otherwise. Any process holding a client uses +``temporalio.client.Client.get_stream_handle``, which mirrors +``get_workflow_handle``. The explicit form, +``provider.get_stream_handle(client, workflow_id)``, stays for a process that +talks to two stores. This module keeps the shared types, the errors and the +protocols a provider implements; nothing here that workflow code imports does +I/O. + +A handle is bound to its client and provider, so a stream is handed to another +process as a :class:`StreamRef`: the owner and the topic as plain data, with +no cursor and no provider name. :meth:`StreamHandle.ref` makes one, the +default data converter carries it as JSON, and the receiver opens it with +``client.get_stream_handle(ref)`` or ``activity.stream_handle(ref)`` on +whatever provider its client has. + +A stream can also stand alone, with an id of its own and no owner. +``client.create_stream(stream_id, retention=...)`` creates it with a retention +policy and returns its handle, ``client.get_stream_handle(stream_id=...)`` +reaches an existing one, and the handle's ``close()`` seals it, after which +appends are refused with :class:`StreamClosedError` and the retained records +stay readable. A provider whose store cannot hold an ownerless stream raises +:class:`StreamUnsupportedError` for both. + +What the contract does not promise: that a :attr:`RecordKind.FINISH` record +means the writing activity succeeded, that a superseded attempt's records can +be withdrawn, or that a stream outlives the retention its provider is +configured for. Reading somebody else's stream is out of scope for this +release. + +The record on the wire is ``temporal.api.stream.v1.StreamRecord`` on every +provider, with the user's value in ``body`` as an ordinary payload, so a +reader in any language decodes the same bytes and a payload codec applies. A +provider owes that body what the SDK gives every payload it sends: it encodes +it through the client's data converter, so the codec and the +:class:`temporalio.converter.ExternalStorage` drivers apply, it takes the +retry fingerprint over the converted bytes before either runs and leaves the +plaintext hash on the record under :data:`CONTENT_HASH_KEY`, and it offloads a +workflow's own publish off the workflow thread. :func:`encode_body`, +:func:`decode_body` and :func:`content_fingerprint` are the shared code for +that; :class:`StreamProvider` states the rule. +""" + +from __future__ import annotations + +from temporalio.streams._body import ( + CONTENT_HASH_KEY, + content_fingerprint, + content_hash, + decode_body, + encode_body, +) +from temporalio.streams._errors import ( + StreamClosedError, + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, + StreamUnsupportedError, +) +from temporalio.streams._provider import ( + ReadSource, + StreamHandle, + StreamProducer, + StreamProvider, + WorkflowStreamProvider, + WriteSink, +) +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + Supersession, +) +from temporalio.streams._ref import StreamOwnerKind, StreamRef +from temporalio.streams._topic import ( + DEFAULT_TOPIC, + StreamTopic, + resolve_topic, + topic, +) + +__all__ = [ + "BEGINNING", + "CONTENT_HASH_KEY", + "DEFAULT_TOPIC", + "END", + "Cursor", + "ReadSource", + "RecordKind", + "StreamClosedError", + "StreamCursorError", + "StreamError", + "StreamHandle", + "StreamNotFoundError", + "StreamOwnerKind", + "StreamProducer", + "StreamProducerError", + "StreamProvider", + "StreamRecord", + "StreamRef", + "StreamTopic", + "StreamUnsupportedError", + "Supersession", + "WorkflowStreamProvider", + "WriteSink", + "content_fingerprint", + "content_hash", + "decode_body", + "encode_body", + "resolve_topic", + "topic", +] diff --git a/temporalio/streams/_body.py b/temporalio/streams/_body.py new file mode 100644 index 000000000..2f10b6d4a --- /dev/null +++ b/temporalio/streams/_body.py @@ -0,0 +1,116 @@ +"""What a provider owes a record's body between the converter and its store. + +:func:`temporalio.streams._wire.to_wire` converts a value into the body with +the payload converter and stops there. Every other payload the SDK sends then +passes through the payload codec and external storage, and a stream body owes +the same, or a codec-protected deployment would leak plaintext through its +streams and a claim-check deployment would push oversized bodies at its store. +A provider runs the body through :func:`encode_body` before it stores or ships +a record and through :func:`decode_body` after it reads one back, off the +workflow thread in both directions. + +The order inside :func:`encode_body` is the point. The plaintext hash is taken +first and stamped on the record, and :func:`content_fingerprint` is taken over +the converted records too, before the codec runs, because a codec that +encrypts with a fresh nonce makes every retry's bytes differ, and a store that +fingerprinted those bytes would refuse the retry as a divergent write. The +store keeps the plaintext hash under :data:`CONTENT_HASH_KEY` and can compare +retries by it without ever seeing the plaintext. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Sequence + +from temporalio.api.common.v1 import Payload +from temporalio.api.stream.v1 import StreamRecord as WireRecord +from temporalio.converter import DataConverter + +__all__ = [ + "CONTENT_HASH_KEY", + "content_fingerprint", + "content_hash", + "decode_body", + "encode_body", +] + +CONTENT_HASH_KEY = "temporal.io/content-hash" +"""The record metadata key the plaintext hash of the body is stored under. + +Its value is a payload with ``encoding`` ``binary/plain`` whose data is the +hex digest :func:`content_hash` returns. A ``FINISH`` record carries no body +and no hash. +""" + +_HASH_ENCODING = b"binary/plain" + + +def content_hash(payload: Payload) -> str: + """The hex SHA-256 of ``payload`` as the converter produced it. + + Taken over the deterministic serialization of the whole payload, metadata + included, so two payloads that differ only in their encoding hash apart. + """ + return hashlib.sha256(payload.SerializeToString(deterministic=True)).hexdigest() + + +def content_fingerprint(records: Sequence[WireRecord]) -> bytes: + """The identity of one append, taken over its converted records. + + Length-delimited, so a batch split differently cannot collide with this + one. Take it before :func:`encode_body`, while the bodies are still what + the converter produced; that is what makes a retry through a + nondeterministic codec match its original. + """ + digest = hashlib.sha256() + for record in records: + body = record.SerializeToString(deterministic=True) + digest.update(len(body).to_bytes(8, "big")) + digest.update(body) + return digest.digest() + + +async def encode_body(converter: DataConverter, record: WireRecord) -> WireRecord: + """Stamp the plaintext hash on ``record`` and encode its body for the store. + + In place, and returned for convenience. The hash goes under + :data:`CONTENT_HASH_KEY` first; then the body passes through + ``converter``'s payload codec and external storage in the order + :meth:`temporalio.converter.DataConverter.encode` uses, so a body above + the external storage threshold is replaced by a claim and the claim is + what the store holds. A record without a body is returned untouched. + """ + if not record.HasField("body"): + return record + record.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload( + metadata={"encoding": _HASH_ENCODING}, + data=content_hash(record.body).encode(), + ) + ) + encoded = await converter._encode_payload_sequence([record.body]) + stored = await converter._external_store_payload_sequence(encoded) + record.body.CopyFrom(stored[0]) + return record + + +async def decode_body(converter: DataConverter, record: WireRecord) -> WireRecord: + """Undo :func:`encode_body` on a record read back from the store. + + In place, and returned for convenience. The body is retrieved from + external storage when it is a claim and then run through the payload + codec, in the order :meth:`temporalio.converter.DataConverter.decode` + uses, leaving the payload the converter can turn back into a value. The + hash stays on the record. + + Raises: + RuntimeError: The body is a claim and ``converter`` has no external + storage to redeem it with. + """ + if not record.HasField("body"): + return record + retrieved = await converter._external_retrieve_payload_sequence([record.body]) + decoded = await converter._decode_payload_sequence(retrieved) + record.body.CopyFrom(decoded[0]) + return record diff --git a/temporalio/streams/_errors.py b/temporalio/streams/_errors.py new file mode 100644 index 000000000..a2f4fecf5 --- /dev/null +++ b/temporalio/streams/_errors.py @@ -0,0 +1,48 @@ +"""The errors a stream call raises. + +Every stream condition is a :class:`StreamError`, so a caller can catch by +meaning the way it catches other :class:`temporalio.exceptions.TemporalError` +subclasses. Argument mistakes stay ``ValueError``. A provider's transport +failure surfaces as :class:`temporalio.service.RPCError`, never as the +transport's own exception type. +""" + +from __future__ import annotations + +import temporalio.exceptions + +__all__ = [ + "StreamClosedError", + "StreamCursorError", + "StreamError", + "StreamNotFoundError", + "StreamProducerError", + "StreamUnsupportedError", +] + + +class StreamError(temporalio.exceptions.TemporalError): + """Base for stream conditions.""" + + +class StreamNotFoundError(StreamError): + """The workflow, chain or topic does not exist or is past retention.""" + + +class StreamCursorError(StreamError): + """The cursor was minted by another provider or names a record no longer retained.""" + + +class StreamProducerError(StreamError): + """The producer attempt or sequence conflicts with what the store holds.""" + + +class StreamClosedError(StreamError): + """The standalone stream was sealed, so it takes no more records. + + Its retained records stay readable; only appends are refused. + """ + + +class StreamUnsupportedError(StreamError): + """This provider does not offer the requested capability.""" diff --git a/temporalio/streams/_ids.py b/temporalio/streams/_ids.py new file mode 100644 index 000000000..2a76dcf19 --- /dev/null +++ b/temporalio/streams/_ids.py @@ -0,0 +1,26 @@ +"""The store key a provider derives from a workflow id and a topic. + +A workflow id may contain any character, ``:`` included, so joining the pair +with a bare ``:`` is ambiguous: ``("a:b", "c")`` and ``("a", "b:c")`` would +land in one store. Every provider that keys a store by the pair goes through +:func:`topic_key`, so they all agree and none of them collides. +""" + +from __future__ import annotations + +__all__ = ["topic_key"] + + +def _escape(component: str) -> str: + # Percent first, so an escaped component cannot be mistaken for one that + # already contained the escape. + return component.replace("%", "%25").replace(":", "%3A") + + +def topic_key(workflow_id: str, topic: str) -> str: + """The store key for ``topic`` of ``workflow_id``'s stream. + + Both components are percent-encoded before joining, so the only bare ``:`` + in the result is the separator. + """ + return f"{_escape(workflow_id)}:{_escape(topic)}" diff --git a/temporalio/streams/_policy.py b/temporalio/streams/_policy.py new file mode 100644 index 000000000..81949e4fc --- /dev/null +++ b/temporalio/streams/_policy.py @@ -0,0 +1,66 @@ +"""Turning a producer's newer attempt into something a reader can act on. + +An activity that streams half an answer and then fails leaves those records in +the stream. Its retry calls the model again and writes different words. No +provider can undo the first half, and a workflow that already acted on it has +committed that decision, so the honest thing is to tell the reader that a new +attempt began and let the application decide. + +This runs in the reader over records it already observed, so it costs no round +trip and replays without the provider being involved. +""" + +from __future__ import annotations + +from typing import Any + +from temporalio.streams._record import ( + Cursor, + RecordKind, + StreamRecord, + Supersession, +) + +__all__ = ["AttemptTracker"] + + +class AttemptTracker: + """Watches producer attempts on one subscription.""" + + def __init__(self) -> None: + """Start with no producer seen.""" + self._attempts: dict[str, int] = {} + + def note( + self, producer_id: str, attempt: int, *, topic: str, previous: Cursor + ) -> StreamRecord[Any] | None: + """A supersession record when this record starts a newer attempt. + + ``previous`` is the cursor of the last record delivered before the one + being noted, or the cursor the read started from. The synthesized + record carries it, so a consumer that checkpoints the supersession and + resumes after it is handed the new attempt's first record next rather + than skipping it. + + A producer that declares no attempt supersedes nothing, because there + is no generation to compare. That is the same answer as an unnumbered + record: the interface reports what it was told and invents nothing. + """ + if not producer_id or attempt <= 0: + return None + seen = self._attempts.get(producer_id, 0) + if attempt <= seen: + return None + self._attempts[producer_id] = attempt + if seen == 0: + return None + return StreamRecord( + kind=RecordKind.SUPERSEDED, + cursor=previous, + topic=topic, + producer_id=producer_id, + attempt=attempt, + supersession=Supersession( + producer_id=producer_id, previous_attempt=seen, attempt=attempt + ), + ) diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py new file mode 100644 index 000000000..662bc7918 --- /dev/null +++ b/temporalio/streams/_provider.py @@ -0,0 +1,450 @@ +"""What a provider implements, in two halves. + +:class:`WorkflowStreamProvider` runs on the workflow thread and must keep the +contract's first two rules: publishes commit with the Workflow Task, and reads +are recorded observations. Nothing it needs may do I/O. :class:`StreamProvider` +is the half a process holds: it makes the workflow half for a worker and hands +out :class:`StreamHandle` objects to code outside a workflow. A Python provider +usually implements both on one class; the split is what lets a language whose +workflow code is bundled separately name the two halves in two packages. + +A provider only moves ``temporal.api.stream.v1.StreamRecord`` protos. The +handles around it convert values, synthesize supersession and mint cursors, +and turn a :class:`temporalio.streams.StreamTopic` into the plain name the +provider sees, through :func:`temporalio.streams.resolve_topic`. What it owes +a record's body on the way to and from its store, :class:`StreamProvider` +lists and :func:`temporalio.streams.encode_body` and +:func:`temporalio.streams.decode_body` do. +""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator +from datetime import timedelta +from typing import TYPE_CHECKING, Any, Protocol, TypeVar, overload + +from temporalio.api.stream.v1 import StreamRecord as WireRecord +from temporalio.streams._record import BEGINNING, Cursor, StreamRecord +from temporalio.streams._topic import StreamTopic + +if TYPE_CHECKING: + from temporalio.client import Client + from temporalio.streams._ref import StreamRef + +__all__ = [ + "ReadSource", + "StreamHandle", + "StreamProducer", + "StreamProvider", + "WorkflowStreamProvider", + "WriteSink", +] + +T = TypeVar("T") +T_contra = TypeVar("T_contra", contravariant=True) + + +class StreamProducer(Protocol[T_contra]): + """Appends to one topic from outside workflow code. + + Every append is visible as soon as the store accepts it, and carries the + producer id, attempt and sequence that let a reader tell a retried append + from a new generation. The type parameter is the topic definition's + value type; a producer on a string-named topic takes any value. + """ + + @property + def producer_id(self) -> str: + """Who this producer writes as.""" + ... + + @property + def attempt(self) -> int: + """The generation this producer is writing, or 0 when undeclared.""" + ... + + async def append(self, *values: T_contra) -> Cursor | None: + """Append ``values`` and return the cursor of the last record as the store holds it. + + A repeat of an earlier append (same producer, attempt and sequence) + carrying the same content is written once and returns the position the + original landed at. A repeat carrying different content is a conflict, + not a retry: it raises :class:`StreamProducerError` and writes nothing, + because the store cannot tell which of the two the reader was meant to + see. A provider that cannot compare content says so in its own + documentation rather than picking one silently. + + An empty call writes nothing and returns the same value a repeat would: + the position of this producer's last record, or ``BEGINNING`` when it + has written none. ``None`` means one thing only: this provider learns + positions at read time, and a caller that needs one positions itself + with :meth:`StreamHandle.latest`. + + Raises: + StreamProducerError: The attempt or sequence conflicts with what + the store holds. + """ + ... + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic. + + Says this producer has nothing more to send. It does not say the + activity behind it succeeded, and it does not end anyone's read. + """ + ... + + +class StreamHandle(Protocol): + """One owner's stream, addressed by topic, from outside workflow code. + + The owner is a workflow, an activity, or a standalone stream that has an + id of its own and no owner. A handle on a workflow follows its execution + chain unless it was opened with a ``run_id``, in which case it is pinned + to that run. A topic is a :class:`temporalio.streams.StreamTopic` + definition, which carries the record type, or a plain string with + ``result_type=`` for a name decided at runtime. A transport failure + surfaces as :class:`temporalio.service.RPCError`, never as the + transport's own exception type. + + A handle is bound to its client and provider. To hand a stream to another + process, :meth:`ref` names it as a :class:`temporalio.streams.StreamRef`, + which is plain data; the receiver opens it with + ``temporalio.client.Client.get_stream_handle`` or + ``temporalio.activity.stream_handle``, and calls that name no topic + on that handle address the ref's topic. + """ + + @overload + def read( + self, + *, + topic: StreamTopic[T], + after: Cursor = ..., + last: int | None = None, + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, + *, + topic: str | None = None, + after: Cursor = ..., + last: int | None = None, + result_type: type[T], + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, + *, + topic: str | None = None, + after: Cursor = ..., + last: int | None = None, + result_type: None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: ... + + 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 the records on ``topic`` after ``after`` as they arrive. + + Without ``topic`` it reads :data:`temporalio.streams.DEFAULT_TOPIC`. + ``BEGINNING`` yields everything the topic retains, starting at the + oldest record it still holds. ``END`` yields only what is appended + after the read starts. ``last=N`` starts at the newest ``N`` records, + or at all of them when there are fewer; it counts records of every + kind, so a ``FINISH`` among them leaves fewer than ``N`` values, and + it is exclusive with a cursor. Any other cursor came from a record a + reader saw, and reading resumes just past it, so a reader that stores + the last cursor it handled and hands it back sees every record + exactly once; that is the only way to resume. The read ends when the owning + execution, or its chain, is closed and every retained record after + ``after`` has been delivered; until then it waits. The result is a + generator, so a caller that stops early should ``aclose()`` it. How + much that releases is the provider's to say: one that holds only + local state lets go at once, and one that parked something on a + store it cannot un-park says in its own documentation what it + releases and when. Read the provider's ``read`` before relying on an + immediate release. + + Raises: + ValueError: ``result_type`` was passed with a topic definition, + the topic is empty, ``last`` is not positive, or ``last`` was + passed with a cursor. + StreamCursorError: ``after`` came from another provider or names + a record no longer retained. Raised by this call, not by the + first iteration. + StreamUnsupportedError: The provider cannot start a read where + ``END`` or ``last=`` asks. A provider that raises it says so + in its own documentation. + StreamNotFoundError: The workflow or topic does not exist or is + past retention. + """ + ... + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The cursor of the newest record on ``topic``, or ``BEGINNING`` when empty. + + Without ``topic`` it answers for :data:`temporalio.streams.DEFAULT_TOPIC`. + + For a reader that wants to follow from now: ``read(after=latest())`` + yields only what is published after this call returned, which is how + a client that is about to send a message positions itself before + sending, without the workflow having to report a position. + """ + ... + + @overload + def producer( + self, *, topic: StreamTopic[T], producer_id: str = ..., attempt: int = ... + ) -> StreamProducer[T]: ... + + @overload + def producer( + self, *, topic: str | None = None, producer_id: str = ..., attempt: int = ... + ) -> StreamProducer[Any]: ... + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> StreamProducer[Any]: + """A producer on ``topic``, or on the default topic without one. + + Inside an activity, leave ``producer_id`` and ``attempt`` unset: the + activity's own id and attempt are the right answer, and they are what + let a reader tell a retry from a new generation. Outside one, + ``producer_id`` is required and an empty one raises ``ValueError``. + """ + ... + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A :class:`temporalio.streams.StreamRef` to ``topic`` of this owner. + + Without ``topic`` it names :data:`temporalio.streams.DEFAULT_TOPIC`, + or the topic this handle was opened from a ref with. The ref carries + the owner exactly as this handle addresses it, a ``run_id`` included + when the handle is pinned, and no cursor or provider name, so it can + travel as a workflow argument, an activity result or a Nexus + operation input or result and be opened wherever a client is. + """ + ... + + async def close(self) -> None: + """Seal the standalone stream this handle is on. + + A sealed stream takes no more records: a later ``append`` raises + :class:`temporalio.streams.StreamClosedError`, while everything it + retains stays readable and a read on it ends once that tail has been + delivered. Idempotent. Only a standalone stream can be closed here, + because an owned stream ends with its owner. + + Raises: + ValueError: This handle is on a workflow's or an activity's + stream. + """ + ... + + +class ReadSource(Protocol): + """One subscription, as a provider supplies it to the workflow thread.""" + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + """The next records with their positions, waiting until there is at least one. + + A batch rather than a record because delivery boundaries are what a + provider actually records, and flattening them here keeps that out of + the contract. A record that cannot be parsed into a ``StreamRecord`` + proto is the provider's to skip. + + Raises: + StopAsyncIteration: This subscription has ended. + """ + ... + + def close(self) -> None: + """End the subscription. Idempotent.""" + ... + + +class WriteSink(Protocol): + """One topic of the running workflow's stream, as a provider binds it.""" + + def publish(self, record: WireRecord) -> None: + """Take one record into this Workflow Task's output. + + Synchronous: there is nothing to wait for inside a task, because the + task is the visibility boundary. The provider commits what it buffered + when the task completes and drops it when the task fails. A record + the provider cannot stage raises :class:`temporalio.streams.StreamError` + and fails the task, loudly. + """ + ... + + +class WorkflowStreamProvider(Protocol): + """The half of a provider that runs on the workflow thread. + + Imports nothing that does I/O. The worker creates one per workflow + instance through :meth:`StreamProvider.workflow_provider`, so state kept + here dies with the instance the way handlers do. It sees topics by name; + the definitions are resolved before it is called. + """ + + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + """Subscribe the running workflow to ``topic`` of its own stream. + + ``after`` and ``last`` mean what they mean on + :meth:`StreamHandle.read`, and arrive already checked. Where a start + is resolved has to be something replay reproduces, so a provider + resolves it in the store and records the result, never by reading + the store from the workflow thread. + + Raises: + StreamCursorError: ``after`` was minted by another provider. + StreamUnsupportedError: The provider cannot start where ``END`` + or ``last`` asks. + """ + ... + + def open_writer(self, topic: str) -> WriteSink: + """Bind ``topic`` of the running workflow's stream for publishing.""" + ... + + def on_workflow_start(self) -> None: + """Called before the workflow function runs. + + A provider that serves outside readers through handlers on the + workflow registers them here, before the first task completes. + """ + ... + + async def on_workflow_finish(self) -> None: + """Called after the workflow function returns, raises or continues as new. + + A provider that parked an outside reader against the run lets go + here, so the workflow can close. + """ + ... + + +class StreamProvider(Protocol): + """What a store ships. Also a :class:`temporalio.worker.Plugin` when it serves workers. + + Construct one, pass it to ``Client.connect(plugins=[provider])`` so the + client and the workers built from it carry it, or to + ``Worker(plugins=[provider])`` and ``Replayer(plugins=[provider])`` for a + worker alone, and open handles from it anywhere else. Nothing is global: + two workers in one process may hold two providers. + + **What a provider owes a record's body.** The handles convert a value + into the body with the payload converter and no more; what the SDK does + to every other payload it sends, the codec and external storage, the + provider owes the body too, through the client's data converter, so the + :class:`temporalio.converter.ExternalStorage` drivers an application + configured apply to stream bodies as well. It does that in one order. + First it takes the retry fingerprint, the identity a repeated append is + matched by, over the converted bytes, before the codec and before any + offload, so a codec that encrypts with a fresh nonce cannot turn a retry + into a divergent write; the plaintext hash also rides the record under + :data:`temporalio.streams.CONTENT_HASH_KEY`, where the store can read it. + Then it encodes the body and offloads it, and on a read it does the + reverse before the record reaches a reader. A workflow's own publish is + converted on the workflow thread and no further: the codec and the offload + run when the provider commits the task's batch, off that thread. + :func:`temporalio.streams.encode_body`, + :func:`temporalio.streams.decode_body` and + :func:`temporalio.streams.content_fingerprint` are that rule in code. + + **Standalone streams.** A stream can have an id of its own and no owner. + It is created on purpose, with :meth:`create_standalone_stream` and a + retention policy, and sealed on purpose, with the handle's ``close``. It + is addressed by topic like an owner's streams; how a provider lays its + topics out in the store is its own. A provider whose store cannot hold a + stream without an owner raises + :class:`temporalio.streams.StreamUnsupportedError` from both standalone + calls. + """ + + def workflow_provider(self) -> WorkflowStreamProvider: + """The half that serves one workflow instance on its thread.""" + ... + + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> StreamHandle: + """A handle on ``workflow_id``'s stream. + + Without ``run_id`` it follows the execution chain, so a consumer keeps + reading across continue-as-new; with one it is pinned to that run. + """ + ... + + 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, + ) -> StreamHandle: + """Create the standalone stream ``stream_id`` and return a handle on it. + + The three policy arguments bound what the stream retains: records + older than ``retention``, beyond the newest ``max_records``, or past + ``max_bytes`` of stored records are dropped, and ``None`` leaves that + bound to the provider's default. Creating a stream that exists with + the same policy returns a handle on it, so a retried create is + harmless. + + Raises: + ValueError: ``stream_id`` is empty, a bound is not positive, or + the stream exists with a different policy. + StreamUnsupportedError: The provider's store cannot hold a stream + without an owner. + """ + ... + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> StreamHandle: + """A handle on the standalone stream ``stream_id``, which must exist. + + Nothing here creates the stream: the first ``read``, ``latest`` or + ``producer`` on a stream that does not exist raises + :class:`temporalio.streams.StreamNotFoundError`, unless the provider + can wait for the stream to be created, in which case a ``read`` parks + until the first write and says so in its own documentation. + + Raises: + StreamUnsupportedError: The provider's store cannot hold a stream + without an owner. + """ + ... + + async def close(self) -> None: + """Release what this provider holds for the process. + + A provider that keeps a connection pool or an HTTP session open needs + a moment where the process says it is done; this is it. A provider + that holds nothing returns at once. + + The application calls this, not the worker and not the client. One + provider serves the workers built from a client and every handle + opened outside them, so no single one of those owns its lifetime and + a worker shutting down would close a connection its siblings are + still reading through. A provider that outlives the process it was + made in is the application's to close. + """ + ... diff --git a/temporalio/streams/_record.py b/temporalio/streams/_record.py new file mode 100644 index 000000000..c31050ede --- /dev/null +++ b/temporalio/streams/_record.py @@ -0,0 +1,146 @@ +"""The value types the stream contract is expressed in. + +Nothing here touches Temporal or a provider, so every provider shares it +unchanged. +""" + +from __future__ import annotations + +import enum +from dataclasses import dataclass +from typing import Generic, TypeVar + +__all__ = [ + "BEGINNING", + "END", + "Cursor", + "RecordKind", + "StreamRecord", + "Supersession", +] + +T = TypeVar("T") + + +@enum.unique +class RecordKind(enum.IntEnum): + """What a record is. + + Mirrors ``temporal.api.stream.v1.StreamRecordKind`` value for value, so a + record's kind crosses the wire as the integer the proto holds. + """ + + UNSPECIFIED = 0 + """The proto's zero value. + + A stored record whose writer set no kind is read as :attr:`DATA`, as the + proto defines it, so a reader never sees this kind on a record. + """ + + DATA = 1 + """Carries a value published by a workflow or a producer.""" + + FINISH = 2 + """The producer named in ``producer_id`` will write nothing more on this topic. + + An empty ``producer_id`` names the owning workflow. It does not end a + read, which ends when the owning execution or its chain is closed and the + retained tail has been delivered, and it says nothing about the producer's + outcome: an activity can still time out after writing it. + """ + + SUPERSEDED = 3 + """A later attempt of the same producer started writing. + + Synthesized by the reader from what it observed, never stored, so every + provider delivers it identically and replay reproduces it without the + provider's help. Its cursor is the position before the new attempt's + first record, so resuming after it delivers that record next. + """ + + +@dataclass(frozen=True) +class Cursor: + """A position in a stream, ordered by its provider rather than by value. + + Opaque on purpose. One provider numbers records with integers and another + with a millisecond-and-sequence pair, so comparing tokens here would be + right for one and wrong for the other. Hand a cursor back to resume after + the record it names; nothing here advances one. The token starts with the + name of the provider that minted it, and a provider refuses a token from + another with :class:`temporalio.streams.StreamCursorError`. + """ + + token: str + + def __str__(self) -> str: + """The provider's position token.""" + return self.token + + +BEGINNING = Cursor("") +"""Read from the oldest record the stream still retains.""" + +END = Cursor("$end") +"""Read only what is appended after the read starts. + +Provider-neutral, like :data:`BEGINNING`. It is resolved when the read +starts, not when it is called, so it cannot position a client before it +sends something; :meth:`temporalio.streams.StreamHandle.latest` does that. +""" + + +def check_read_start(after: Cursor, last: int | None) -> None: + """Refuse a read start that names two places, or a count that names none. + + ``after=`` resumes a read and ``last=`` starts one, so a call gives one or + the other. ``BEGINNING`` is the default for ``after=``, and passing it + alongside ``last=`` is the same as passing ``last=`` alone. + + Raises: + ValueError: ``last`` is not a positive int, or it was given together + with a cursor. + """ + if last is None: + return + if isinstance(last, bool) or not isinstance(last, int) or last <= 0: + raise ValueError(f"last must be a positive int, got {last!r}") + if after != BEGINNING: + raise ValueError( + "pass either after= or last=, not both: after= resumes a read and " + "last= starts one" + ) + + +@dataclass(frozen=True) +class Supersession: + """What a :attr:`RecordKind.SUPERSEDED` record reports.""" + + producer_id: str + previous_attempt: int + attempt: int + + +@dataclass(frozen=True) +class StreamRecord(Generic[T]): + """One record as a reader sees it. + + ``value`` is set on a :attr:`RecordKind.DATA` record and ``supersession`` + on a :attr:`RecordKind.SUPERSEDED` one; every other kind carries neither. + Each field means one thing, so a consumer narrows on ``kind`` and reads + the field that kind promises. + """ + + kind: RecordKind + cursor: Cursor + topic: str + producer_id: str = "" + """Who wrote it, or empty when the owning workflow wrote it itself.""" + attempt: int = 0 + """The producer's attempt, or 0 when it did not declare one.""" + sequence: int = 0 + """The producer's position within its attempt, or 0 when it does not number.""" + value: T | None = None + """The published value. Set on ``DATA`` only.""" + supersession: Supersession | None = None + """The attempt change being reported. Set on ``SUPERSEDED`` only.""" diff --git a/temporalio/streams/_ref.py b/temporalio/streams/_ref.py new file mode 100644 index 000000000..3d0eb0a5d --- /dev/null +++ b/temporalio/streams/_ref.py @@ -0,0 +1,142 @@ +"""A reference to one stream that crosses a process boundary as data. + +A handle is bound to a client and a provider, so it cannot be a workflow +argument, an activity result or a Nexus operation result. A +:class:`StreamRef` can: it names the owner and the topic, nothing more, and +whoever receives it opens the stream on the provider its own client carries. +It carries no cursor, because a position belongs to a reader, and no provider +name, because the same owner and topic name the same stream on every provider +a deployment runs. +""" + +from __future__ import annotations + +import dataclasses +from dataclasses import dataclass +from typing import Any, Literal + +from temporalio.streams._topic import DEFAULT_TOPIC, StreamTopic, resolve_topic + +__all__ = ["StreamOwnerKind", "StreamRef"] + +StreamOwnerKind = Literal["workflow", "activity", "standalone"] +"""What owns a stream: a workflow, an activity, or the stream itself.""" + + +@dataclass(frozen=True) +class StreamRef: + """One stream, named by its owner and its topic. + + ``kind`` says what owns the stream. A ``"workflow"`` ref carries + ``workflow_id`` and, when pinned to one run, ``run_id``. An ``"activity"`` + ref carries ``activity_id``, plus ``workflow_id`` (and its ``run_id``) + when a workflow scheduled the activity; without ``workflow_id`` it is a + standalone activity, and ``run_id`` then pins one run of it. A + ``"standalone"`` ref carries ``stream_id`` and nothing else. ``topic`` is + the topic's name, and a handle opened from a ref addresses it whenever a + call names no topic. + + A ref comes from :meth:`temporalio.streams.StreamHandle.ref`, or from + :meth:`for_workflow`, :meth:`for_activity` and :meth:`for_standalone` + when only the ids are at hand. The default data converter carries it as + JSON, so it can be a workflow argument, an activity result, or a Nexus + operation input or result, and + ``temporalio.client.Client.get_stream_handle`` and + ``temporalio.activity.stream_handle`` open one directly. + """ + + kind: StreamOwnerKind + topic: str = DEFAULT_TOPIC + workflow_id: str | None = None + run_id: str | None = None + activity_id: str | None = None + stream_id: str | None = None + + def __post_init__(self) -> None: + """Refuse a ref that names an owner its kind does not have.""" + if not self.topic: + raise ValueError("a StreamRef needs a topic name") + if self.kind == "workflow": + if not self.workflow_id: + raise ValueError("a workflow StreamRef needs a workflow_id") + if self.activity_id is not None or self.stream_id is not None: + raise ValueError( + "a workflow StreamRef carries no activity_id and no stream_id" + ) + elif self.kind == "activity": + if not self.activity_id: + raise ValueError("an activity StreamRef needs an activity_id") + if self.stream_id is not None: + raise ValueError("an activity StreamRef carries no stream_id") + elif self.kind == "standalone": + if not self.stream_id: + raise ValueError("a standalone StreamRef needs a stream_id") + if ( + self.workflow_id is not None + or self.run_id is not None + or self.activity_id is not None + ): + raise ValueError("a standalone StreamRef carries only its stream_id") + else: + raise ValueError( + f"unknown StreamRef kind {self.kind!r}; expected 'workflow', " + "'activity' or 'standalone'" + ) + + @classmethod + def for_workflow( + cls, + workflow_id: str, + *, + run_id: str | None = None, + topic: str | StreamTopic[Any] | None = None, + ) -> StreamRef: + """A ref to ``topic`` of ``workflow_id``'s stream. + + Without ``run_id`` a handle opened from it follows the execution + chain; with one it is pinned to that run. Without ``topic`` it names + :data:`temporalio.streams.DEFAULT_TOPIC`. + """ + name, _ = resolve_topic(topic) + return cls("workflow", name, workflow_id=workflow_id, run_id=run_id) + + @classmethod + def for_activity( + cls, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + topic: str | StreamTopic[Any] | None = None, + ) -> StreamRef: + """A ref to ``topic`` of the streams ``activity_id`` owns. + + With ``workflow_id`` the activity is one that workflow scheduled and + ``run_id`` is the workflow's run; without it the activity is a + standalone one and ``run_id`` pins one run of it. + """ + name, _ = resolve_topic(topic) + return cls( + "activity", + name, + workflow_id=workflow_id, + run_id=run_id, + activity_id=activity_id, + ) + + @classmethod + def for_standalone( + cls, stream_id: str, *, topic: str | StreamTopic[Any] | None = None + ) -> StreamRef: + """A ref to ``topic`` of the standalone stream ``stream_id``.""" + name, _ = resolve_topic(topic) + return cls("standalone", name, stream_id=stream_id) + + def with_topic(self, topic: str | StreamTopic[Any] | None) -> StreamRef: + """The same owner, naming ``topic`` instead. + + ``None`` names :data:`temporalio.streams.DEFAULT_TOPIC`, as a call + that passes no topic does. + """ + name, _ = resolve_topic(topic) + return dataclasses.replace(self, topic=name) diff --git a/temporalio/streams/_topic.py b/temporalio/streams/_topic.py new file mode 100644 index 000000000..f887b6b57 --- /dev/null +++ b/temporalio/streams/_topic.py @@ -0,0 +1,109 @@ +"""Typed topic definitions. + +Temporal's idiom is to define once and refer by reference: signals, queries +and updates are decorated methods, activities and workflows are functions, +Nexus operations are typed definitions. A topic follows the same rule. It is +defined once, at module level, with the type its records decode to, and the +workflow, its activities and the backend all refer to that one definition. A +plain string names a topic too, the way a string names a signal chosen at +runtime; then the decode hint travels as ``result_type=`` on each call. + +The wire does not change: a topic is a string on the proto and in every +store, and :attr:`temporalio.streams.StreamRecord.topic` is that string. + +Every workflow also has a default topic, :data:`DEFAULT_TOPIC`, which a call +addresses by naming no topic at all. It is an ordinary name, so naming it +explicitly is the same topic, not an error. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Generic, TypeVar, overload + +__all__ = ["DEFAULT_TOPIC", "StreamTopic", "resolve_topic", "topic"] + +T = TypeVar("T") + +DEFAULT_TOPIC = "output" +"""The topic a call addresses when it names none. + +The server resolves an unnamed stream of a workflow to this same name, so on +the native provider the default topic is the server's default stream, and a +store with no default of its own holds it under this name. It stays an +ordinary name rather than a reserved one for the same reason the server does +not reserve it: refusing it would make one stream reachable under two rules. +""" + + +@dataclass(frozen=True) +class StreamTopic(Generic[T]): + """A topic of a workflow's stream, with the type its records decode to. + + Made with :func:`topic`. Hand it to ``temporalio.workflow.stream_reader``, + ``temporalio.workflow.stream_writer``, and to a handle's ``read``, + ``latest`` and ``producer``, and the record and value types follow from + it; ``result_type=`` is not passed alongside a definition. + """ + + name: str + result_type: type[T] | None = None + """The type records decode to, or ``None`` for the converter's default.""" + + +@overload +def topic(name: str, result_type: type[T]) -> StreamTopic[T]: ... + + +@overload +def topic(name: str, result_type: None = None) -> StreamTopic[Any]: ... + + +def topic(name: str, result_type: type | None = None) -> StreamTopic[Any]: + """Define a topic of a workflow's stream. + + Define it once, at module level, and share it: the workflow reads or + publishes it, an activity or a backend produces onto it or reads it, and + the type it carries is inferred wherever it is used. Use a plain string + instead only when the name is decided at runtime. + + Args: + name: The topic's name, as it appears on every record. + result_type: The type records on this topic decode to. Without one, + the payload converter's default applies. + + Raises: + ValueError: ``name`` is empty. + """ + if not name: + raise ValueError("topic name must not be empty") + return StreamTopic(name, result_type) + + +def resolve_topic( + topic: str | StreamTopic[Any] | None = None, result_type: type | None = None +) -> tuple[str, type | None]: + """The name and decode hint a call means, from either form of topic. + + Providers call this once at the top of ``read``, ``latest`` and + ``producer``, so a definition and a string are the same to the store. + ``None`` is a call that named no topic, and means :data:`DEFAULT_TOPIC`. + + Raises: + ValueError: A definition was given together with ``result_type``, + which would name two types for one topic, or the name is empty. + """ + if isinstance(topic, StreamTopic): + if result_type is not None: + raise ValueError( + f"topic {topic.name!r} already carries its type; do not pass " + "result_type= with a definition" + ) + name, result_type = topic.name, topic.result_type + elif topic is None: + name = DEFAULT_TOPIC + else: + name = topic + if not name: + raise ValueError("topic must not be empty") + return name, result_type diff --git a/temporalio/streams/_wire.py b/temporalio/streams/_wire.py new file mode 100644 index 000000000..4913cc308 --- /dev/null +++ b/temporalio/streams/_wire.py @@ -0,0 +1,181 @@ +"""How a record crosses a provider: the proto is the record. + +``temporal.api.stream.v1.StreamRecord`` is the wire format on every provider. +A store that keeps bytes keeps ``SerializeToString()`` of it, the native +server stores the proto it is handed, and a reader in any language parses the +same bytes. ``body`` is the user's payload, produced and consumed through the +payload converter, so a codec applies to it like any other payload and a +pre-encoded :class:`temporalio.common.RawValue` passes through untouched. + +Cursors are self-describing: a token starts with the name of the provider that +minted it, so a provider can refuse a foreign one at the call rather than +misreading it deep in a generator. +""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import Any + +import temporalio.converter +from temporalio.api.stream.v1 import StreamRecord as WireRecord +from temporalio.api.stream.v1 import StreamRecordKind +from temporalio.streams._errors import StreamCursorError +from temporalio.streams._policy import AttemptTracker +from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord + +__all__ = [ + "RecordDecoder", + "WireRecord", + "cursor_position", + "from_wire", + "mint_cursor", + "producer_identity", + "to_wire", +] + + +def to_wire( + converter: temporalio.converter.PayloadConverter, + *, + topic: str, + kind: RecordKind, + value: Any = None, + producer_id: str = "", + attempt: int = 0, + sequence: int = 0, +) -> WireRecord: + """Build the record a provider stores or ships. + + Only a ``DATA`` record carries a body; the converter encodes ``value`` + into it, which is where a pre-encoded ``RawValue`` passes through. + """ + record = WireRecord( + topic=topic, + kind=StreamRecordKind.ValueType(int(kind)), + producer_id=producer_id, + attempt=attempt, + sequence=sequence, + ) + if kind is RecordKind.DATA: + record.body.CopyFrom(converter.to_payloads([value])[0]) + return record + + +def from_wire( + converter: temporalio.converter.PayloadConverter, + cursor: Cursor, + wire: WireRecord, + result_type: type | None, +) -> StreamRecord[Any]: + """Turn a stored record into the record a reader yields. + + Raises: + ValueError: The kind is one no store may hold, such as a synthesized + ``SUPERSEDED`` or a value this SDK does not know. + """ + kind = RecordKind(wire.kind) + if kind is RecordKind.SUPERSEDED: + raise ValueError("a SUPERSEDED record is synthesized by readers, never stored") + if kind is RecordKind.UNSPECIFIED: + # The proto defines an unset kind as DATA, so every reader agrees. + kind = RecordKind.DATA + value: Any = None + if kind is RecordKind.DATA and wire.HasField("body"): + hints = [result_type] if result_type is not None else None + value = converter.from_payloads([wire.body], hints)[0] + return StreamRecord( + kind=kind, + cursor=cursor, + topic=wire.topic, + producer_id=wire.producer_id, + attempt=wire.attempt, + sequence=wire.sequence, + value=value, + ) + + +class RecordDecoder: + """Turns the records a provider hands over into the records a reader yields. + + One per read. It synthesizes supersession from the attempts it observes, + positions each synthesized record at the cursor before the record that + triggered it, and skips a record it cannot decode with a warning rather + than raising, so a poisoned record cannot pin a workflow on every retry + while an outside reader of the same stream moves past it. + """ + + def __init__( + self, + converter: temporalio.converter.PayloadConverter, + result_type: type | None, + *, + after: Cursor, + warn: Callable[[str], None], + ) -> None: + """Decode with ``converter`` into ``result_type``, resuming after ``after``.""" + self._converter = converter + self._result_type = result_type + self._previous = after + self._warn = warn + self._attempts = AttemptTracker() + + def decode(self, cursor: Cursor, wire: WireRecord) -> list[StreamRecord[Any]]: + """The records to yield for one stored record, in order.""" + try: + record = from_wire(self._converter, cursor, wire, self._result_type) + except Exception as error: + self._warn(f"skipping stream record at {cursor}: {error}") + # The skipped record still holds its position, so a resume after + # it moves on rather than tripping over it again. + self._previous = cursor + return [] + out: list[StreamRecord[Any]] = [] + superseded = self._attempts.note( + wire.producer_id, wire.attempt, topic=wire.topic, previous=self._previous + ) + if superseded is not None: + out.append(superseded) + out.append(record) + self._previous = cursor + return out + + +def mint_cursor(provider: str, position: str) -> Cursor: + """A cursor that names ``position`` and the provider that understands it.""" + return Cursor(f"{provider}:{position}") + + +def cursor_position(cursor: Cursor, *, provider: str) -> str | None: + """The position inside a cursor ``provider`` minted, or ``None`` for BEGINNING. + + Raises: + StreamCursorError: The cursor came from another provider. + """ + if cursor == BEGINNING: + return None + prefix = f"{provider}:" + if not cursor.token.startswith(prefix): + raise StreamCursorError( + f"cursor {cursor.token!r} was not minted by the {provider} stream provider" + ) + return cursor.token[len(prefix) :] + + +def producer_identity(producer_id: str, attempt: int) -> tuple[str, int]: + """Resolve who a producer is, defaulting to the running activity. + + Imported lazily so the module workflow code imports carries no activity + machinery; the default only means anything inside an activity anyway. + """ + if producer_id: + return producer_id, attempt + import temporalio.activity + + if not temporalio.activity.in_activity(): + raise ValueError( + "producer_id is required outside an activity; inside one it defaults " + "to the activity's id and attempt" + ) + info = temporalio.activity.info() + return info.activity_id, attempt or info.attempt diff --git a/temporalio/streams/providers/__init__.py b/temporalio/streams/providers/__init__.py new file mode 100644 index 000000000..fcefdaa1b --- /dev/null +++ b/temporalio/streams/providers/__init__.py @@ -0,0 +1,118 @@ +"""Stream providers, one module each. + +A provider that serves workers is a :class:`temporalio.worker.Plugin`, and +one registered on a client is a :class:`temporalio.client.Plugin` too. +:class:`ProviderPlugin` is the plugin half the providers in this tree share: +it hands the provider to the client, the worker and the replayer as their +``stream_provider`` and leaves their execution alone, so +``Client.connect(plugins=[provider])`` registers it once and every worker +built from that client inherits it. A provider that holds connections closes +them through its own ``close()``, not with the worker, because the same +provider serves handles outside any worker. +""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import AbstractAsyncContextManager + +import temporalio.client +import temporalio.worker +from temporalio.client import ClientConfig, WorkflowHistory +from temporalio.service import ConnectConfig, ServiceClient +from temporalio.streams._provider import StreamProvider +from temporalio.worker import ( + Replayer, + ReplayerConfig, + Worker, + WorkerConfig, + WorkflowReplayResult, +) + +__all__ = ["ProviderPlugin"] + + +class ProviderPlugin( + StreamProvider, temporalio.client.Plugin, temporalio.worker.Plugin +): + """The plugin every provider in this tree is built on. + + Subclasses implement :class:`temporalio.streams.StreamProvider`; this + class supplies the plugin hooks, so ``Client.connect(plugins=[provider])``, + ``Worker(plugins=[provider])`` and ``Replayer(plugins=[provider])`` reach + the provider through their ``stream_provider`` option. A worker built from + a client that carries the plugin inherits it, and the worker installs the + interceptor that calls the workflow half's lifecycle hooks. + """ + + def _claim(self, held: StreamProvider | None) -> StreamProvider: + """This provider, unless another one already holds the slot. + + There is one slot and it decides where every workflow on the thing + being configured reads and publishes. A last-write-wins here would + silently drop a provider the user passed by hand or a second provider + plugin, so say it instead; a process that talks to two stores opens + the other one's handles from the provider object. + + Raises: + ValueError: Another provider already holds the slot. + """ + if held is not None and held is not self: + raise ValueError( + f"stream provider {held!r} is already registered; pass one provider " + f"and open the other's handles from the object itself" + ) + return self + + def configure_client(self, config: ClientConfig) -> ClientConfig: + """Set this provider as the client's ``stream_provider``. + + Raises: + ValueError: Another provider already holds the slot. + """ + config["stream_provider"] = self._claim(config.get("stream_provider")) + return config + + async def connect_service_client( + self, + config: ConnectConfig, + next: Callable[[ConnectConfig], Awaitable[ServiceClient]], + ) -> ServiceClient: + """Connect unchanged.""" + return await next(config) + + def configure_worker(self, config: WorkerConfig) -> WorkerConfig: + """Set this provider as the worker's ``stream_provider``. + + Raises: + ValueError: Another provider already holds the slot. + """ + config["stream_provider"] = self._claim(config.get("stream_provider")) + return config + + def configure_replayer(self, config: ReplayerConfig) -> ReplayerConfig: + """Set this provider as the replayer's ``stream_provider``. + + Raises: + ValueError: Another provider already holds the slot. + """ + config["stream_provider"] = self._claim(config.get("stream_provider")) + return config + + async def run_worker( + self, worker: Worker, next: Callable[[Worker], Awaitable[None]] + ) -> None: + """Run the worker unchanged.""" + await next(worker) + + def run_replayer( + self, + replayer: Replayer, + histories: AsyncIterator[WorkflowHistory], + next: Callable[ + [Replayer, AsyncIterator[WorkflowHistory]], + AbstractAsyncContextManager[AsyncIterator[WorkflowReplayResult]], + ], + ) -> AbstractAsyncContextManager[AsyncIterator[WorkflowReplayResult]]: + """Run the replayer unchanged.""" + return next(replayer, histories) 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 diff --git a/temporalio/worker/_replayer.py b/temporalio/worker/_replayer.py index b3eb1a4d1..1291f5066 100644 --- a/temporalio/worker/_replayer.py +++ b/temporalio/worker/_replayer.py @@ -17,6 +17,7 @@ import temporalio.client import temporalio.converter import temporalio.runtime +import temporalio.streams import temporalio.worker import temporalio.workflow @@ -56,13 +57,14 @@ def __init__( runtime: temporalio.runtime.Runtime | None = None, disable_safe_workflow_eviction: bool = False, header_codec_behavior: HeaderCodecBehavior = HeaderCodecBehavior.NO_CODEC, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: """Create a replayer to replay workflows from history. See :py:meth:`temporalio.worker.Worker.__init__` for a description of most of the arguments. Most of the same arguments need to be passed to the replayer that were passed to the worker when the workflow originally - ran. + ran, ``stream_provider`` included when the workflow used streams. Note, unlike the worker, for the replayer the workflow_task_executor will default to a new thread pool executor with no max_workers set that @@ -86,6 +88,7 @@ def __init__( runtime=runtime, disable_safe_workflow_eviction=disable_safe_workflow_eviction, header_codec_behavior=header_codec_behavior, + stream_provider=stream_provider, ) self._initial_config = self._config.copy() self._default_workflow_logic_flags = set(_DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS) @@ -433,6 +436,7 @@ class ReplayerConfig(TypedDict, total=False): runtime: temporalio.runtime.Runtime | None disable_safe_workflow_eviction: bool header_codec_behavior: HeaderCodecBehavior + stream_provider: temporalio.streams.StreamProvider | None @dataclass(frozen=True) diff --git a/temporalio/worker/_worker.py b/temporalio/worker/_worker.py index 60f824c4d..25dcf697d 100644 --- a/temporalio/worker/_worker.py +++ b/temporalio/worker/_worker.py @@ -24,6 +24,7 @@ import temporalio.common import temporalio.runtime import temporalio.service +import temporalio.streams from temporalio.common import ( HeaderCodecBehavior, VersioningBehavior, @@ -152,6 +153,7 @@ def __init__( ), disable_payload_error_limit: bool = False, max_workflow_task_external_storage_concurrency: int = _DEFAULT_WORKFLOW_TASK_EXTERNAL_STORAGE_CONCURRENCY, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: """Create a worker to process workflows and/or activities. @@ -343,6 +345,10 @@ def __init__( Defaults to 3. Adjust this value based on your workload's needs. Please report any issues you encounter with this setting or if you feel the default should be changed. + stream_provider: Experimental. The stream provider that workflows + on this worker read and publish through, see + :py:mod:`temporalio.streams`. A provider that is also a + :py:class:`Plugin` sets this itself when passed in ``plugins``. WARNING: This setting is experimental. """ @@ -392,6 +398,7 @@ def __init__( nexus_task_poller_behavior=nexus_task_poller_behavior, disable_payload_error_limit=disable_payload_error_limit, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, + stream_provider=stream_provider, ) plugins_from_client = cast( @@ -1020,6 +1027,7 @@ class WorkerConfig(TypedDict, total=False): nexus_task_poller_behavior: PollerBehavior disable_payload_error_limit: bool max_workflow_task_external_storage_concurrency: int + stream_provider: temporalio.streams.StreamProvider | None def _warn_if_activity_executor_max_workers_is_inconsistent( diff --git a/tests/streams/__init__.py b/tests/streams/__init__.py new file mode 100644 index 000000000..e69de29bb 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 new file mode 100644 index 000000000..8e3d09feb --- /dev/null +++ b/tests/streams/test_streams_internals.py @@ -0,0 +1,278 @@ +"""Unit tests for the pieces under ``temporalio.streams`` that no provider owns. + +The wire format, the supersession policy, the store key, the cursor prefix and +the plugin registration are shared by every provider and implemented once, so +they are tested once, here, against the private modules. What a provider owes +is in ``test_streams_conformance``; keeping the two apart is what makes that +file answerable by a new provider. +""" + +from __future__ import annotations + +import dataclasses +import os +from collections.abc import Sequence + +import pytest + +from temporalio.api.common.v1 import Payload +from temporalio.client import ClientConfig +from temporalio.converter import ( + DataConverter, + ExternalStorage, + PayloadCodec, + StorageDriver, + StorageDriverClaim, + StorageDriverRetrieveContext, + StorageDriverStoreContext, +) +from temporalio.streams import ( + BEGINNING, + CONTENT_HASH_KEY, + Cursor, + RecordKind, + StreamCursorError, + StreamRef, + Supersession, + _ids, + _wire, + content_fingerprint, + content_hash, + decode_body, + 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(): + converter = DataConverter.default.payload_converter + wire = _wire.to_wire( + converter, + topic="decisions", + kind=RecordKind.DATA, + value={"n": 1}, + producer_id="model", + attempt=3, + sequence=7, + ) + parsed = _wire.WireRecord.FromString(wire.SerializeToString()) + record = _wire.from_wire(converter, Cursor("memory:0"), parsed, dict) + assert ( + record.kind, + record.topic, + record.producer_id, + record.attempt, + record.sequence, + record.value, + ) == (RecordKind.DATA, "decisions", "model", 3, 7, {"n": 1}) + assert record.supersession is None + finish = _wire.to_wire(converter, topic="decisions", kind=RecordKind.FINISH) + assert not finish.HasField("body") + assert _wire.from_wire(converter, Cursor("memory:1"), finish, dict).value is None + + +def test_a_stored_supersession_is_not_a_record(): + converter = DataConverter.default.payload_converter + wire = _wire.WireRecord(topic="t", kind=int(RecordKind.SUPERSEDED)) # type: ignore[arg-type] + with pytest.raises(ValueError, match="synthesized"): + _wire.from_wire(converter, Cursor("memory:0"), wire, None) + + +def test_an_unset_kind_is_read_as_data(): + converter = DataConverter.default.payload_converter + wire = _wire.WireRecord(topic="t", body=converter.to_payloads([{"n": 1}])[0]) + record = _wire.from_wire(converter, Cursor("memory:0"), wire, dict) + assert record.kind is RecordKind.DATA + assert record.value == {"n": 1} + + +def test_supersession_is_synthesized_from_observations(): + attempts = AttemptTracker() + assert attempts.note("model", 1, topic="t", previous=BEGINNING) is None + superseded = attempts.note("model", 2, topic="t", previous=Cursor("memory:0")) + assert superseded is not None + assert superseded.kind is RecordKind.SUPERSEDED + assert superseded.supersession == Supersession("model", 1, 2) + assert superseded.value is None + # Positioned before the triggering record, so a resume after it delivers + # that record next. + assert superseded.cursor == Cursor("memory:0") + # The same attempt again is not a new generation. + assert attempts.note("model", 2, topic="t", previous=Cursor("memory:1")) is None + + +def test_topic_keys_cannot_collide(): + # A colon in a workflow id must not make two addresses one key. + assert _ids.topic_key("a:b", "c") != _ids.topic_key("a", "b:c") + assert _ids.topic_key("a%3Ab", "c") != _ids.topic_key("a:b", "c") + assert _ids.topic_key("wf", "inputs") == "wf:inputs" + + +def test_cursors_name_their_provider(): + assert _wire.cursor_position(BEGINNING, provider="memory") is None + assert _wire.cursor_position(Cursor("memory:42"), provider="memory") == "42" + with pytest.raises(StreamCursorError): + _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.""" + + def __init__(self) -> None: + self.held: list[bytes] = [] + + def name(self) -> str: + return "hold" + + async def store( + self, context: StorageDriverStoreContext, payloads: Sequence[Payload] + ) -> list[StorageDriverClaim]: + claims = [] + for payload in payloads: + claims.append(StorageDriverClaim(claim_data={"i": str(len(self.held))})) + self.held.append(payload.SerializeToString()) + return claims + + async def retrieve( + self, + context: StorageDriverRetrieveContext, + claims: Sequence[StorageDriverClaim], + ) -> list[Payload]: + return [Payload.FromString(self.held[int(c.claim_data["i"])]) for c in claims] + + +class _NonceCodec(PayloadCodec): + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=os.urandom(8) + p.SerializeToString(), + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [Payload.FromString(p.data[8:]) for p in payloads] + + +async def test_encode_body_stamps_the_plaintext_hash_and_offloads_the_body(): + driver = _HoldEverything() + converter = dataclasses.replace( + DataConverter.default, + payload_codec=_NonceCodec(), + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=0), + ) + wire = _wire.to_wire( + converter.payload_converter, topic="t", kind=RecordKind.DATA, value={"n": 1} + ) + plaintext = Payload() + plaintext.CopyFrom(wire.body) + + await encode_body(converter, wire) + # The hash is over what the converter produced, not over what the codec + # or the driver made of it, and it rides the record where the store can + # read it without the plaintext. + stamped = wire.metadata[CONTENT_HASH_KEY] + assert stamped.metadata["encoding"] == b"binary/plain" + assert stamped.data.decode() == content_hash(plaintext) + assert len(stamped.data) == 64 + # With a threshold of zero the body was offloaded: the record holds the + # claim and the driver holds the coded payload. + assert wire.body != plaintext + assert len(wire.body.external_payloads) == 1 + assert len(driver.held) == 1 + + await decode_body(converter, wire) + assert wire.body == plaintext + assert wire.metadata[CONTENT_HASH_KEY] == stamped + + # A record without a body has nothing to hash or offload. + finish = _wire.to_wire( + converter.payload_converter, topic="t", kind=RecordKind.FINISH + ) + await encode_body(converter, finish) + assert CONTENT_HASH_KEY not in finish.metadata + assert len(driver.held) == 1 + + +async def test_content_fingerprint_is_taken_before_the_codec(): + converter = dataclasses.replace(DataConverter.default, payload_codec=_NonceCodec()) + plain = converter.payload_converter + + def batch(*values: dict) -> list[_wire.WireRecord]: + return [ + _wire.to_wire(plain, topic="t", kind=RecordKind.DATA, value=v, sequence=i) + for i, v in enumerate(values, 1) + ] + + first, retry = batch({"n": 1}, {"n": 2}), batch({"n": 1}, {"n": 2}) + before = content_fingerprint(first) + assert before == content_fingerprint(retry) + # Different content, and the same content split differently, both differ. + assert before != content_fingerprint(batch({"n": 1}, {"n": 3})) + assert before != content_fingerprint(batch({"n": 1}) + batch({"n": 2})) + + for record in first + retry: + await encode_body(converter, record) + # The codec made the two batches' bytes differ; the identity taken first + # is what lets a store still recognise the retry. + assert first[0].body != retry[0].body + assert content_fingerprint(first) != content_fingerprint(retry) + # Decoding gives the converted bodies back; the hash stays stamped on the + # record, which is why the identity is taken before encoding, not after. + for record in first: + await decode_body(converter, record) + assert [r.body for r in first] == [r.body for r in batch({"n": 1}, {"n": 2})] + assert all(CONTENT_HASH_KEY in r.metadata for r in first) + + +async def test_a_stream_ref_names_one_owner_and_travels_as_json(): + workflow = StreamRef.for_workflow("wf", run_id="r", topic="out") + activity = StreamRef.for_activity("act", workflow_id="wf", topic="progress") + standalone = StreamRef.for_standalone("shared") + assert workflow == StreamRef("workflow", "out", workflow_id="wf", run_id="r") + assert activity.kind == "activity" and activity.activity_id == "act" + assert standalone == StreamRef("standalone", "output", stream_id="shared") + assert standalone.with_topic("x").topic == "x" + + for bad in ( + dict(kind="workflow"), + dict(kind="workflow", workflow_id="wf", stream_id="s"), + dict(kind="activity", workflow_id="wf"), + dict(kind="standalone", stream_id="s", workflow_id="wf"), + dict(kind="standalone"), + dict(kind="nexus", stream_id="s"), + dict(kind="workflow", workflow_id="wf", topic=""), + ): + with pytest.raises(ValueError): + StreamRef(**bad) # type: ignore[arg-type] + + converter = DataConverter.default + for ref in (workflow, activity, standalone): + [carried] = await converter.decode(await converter.encode([ref]), [StreamRef]) + assert carried == ref + payload = (await converter.encode([standalone]))[0] + assert payload.metadata["encoding"] == b"json/plain"