diff --git a/.gitignore b/.gitignore index 8cd439e05..42baf649d 100644 --- a/.gitignore +++ b/.gitignore @@ -15,3 +15,6 @@ temporalio/bridge/temporal_sdk_bridge* tags /.claude tmpclaude-* + +# Demo run output, written per provider; the numbers live in the doc. +streams_demo/results-*/ diff --git a/CHANGELOG.md b/CHANGELOG.md index 4b2f8f941..a1e7b5f40 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -35,6 +35,25 @@ to include examples, links to docs, or any other relevant information. ### Added +- **Experimental**: `temporalio.streams` defines one stream interface a workflow + can read, decide on, and write. A provider is registered once as a plugin, + `Client.connect(plugins=[provider])`, and workers built from that client + inherit it; each context then asks for its stream the same way: + `workflow.stream_reader()` and `workflow.stream_writer()` in workflow code, + `activity.stream_handle()` in an activity, and `client.get_stream_handle()` + anywhere a client is held. A topic is a typed definition, + `streams.topic("inputs", Token)`, shared by workflow, activity and client + code; a plain string names a topic decided at runtime. The record on the wire + is `temporal.api.stream.v1.StreamRecord` on every provider. A stream is + handed to another process as a `streams.StreamRef`, plain data naming the + owner and, when it has one, the topic, which `client.get_stream_handle(ref)` + and `activity.stream_handle(ref)` open; `client.create_stream(stream_id, ...)` + creates a standalone stream with a retention policy, and its handle's + `close()` seals it. A provider runs record bodies through the client's data + converter, so a payload codec and external storage apply to them. + `temporalio.streams.providers.memory.MemoryStreams` is the in-memory + reference provider the conformance tests run against. + - Added experimental External Workflow Streams in `temporalio.contrib.external_workflow_streams`. Workflow stream payloads are stored in a configured external backend instead of Temporal History, with a diff --git a/streams_demo/agent_loop.py b/streams_demo/agent_loop.py new file mode 100644 index 000000000..df5eac4d6 --- /dev/null +++ b/streams_demo/agent_loop.py @@ -0,0 +1,110 @@ +"""One agent loop, written once, run on every stream provider. + +Reads input, decides on it, publishes the decision, and runs an ordinary +Activity in the same workflow task; the Activity reports on its own +workflow's stream in turn. Also handles the two control records the contract +defines, so a retried producer and a finished topic are exercised rather than +described. + +This file is byte-identical in the server-side tree and the client-side tree. +Nothing in it names a provider: the workflow asks its runtime, the Activity +asks its context, and the process that runs them registered the provider once. +""" + +from __future__ import annotations + +from datetime import timedelta +from typing import Any + +from temporalio import activity, streams, workflow +from temporalio.streams import RecordKind + +# Defined once and shared by the workflow, the Activity and the demo's +# reader, so the type each topic carries is stated in one place. +DECISIONS = streams.topic("decisions", dict[str, Any]) +INPUTS = streams.topic("inputs", dict[str, Any]) +RECEIPTS = streams.topic("receipts", dict[str, Any]) + + +@activity.defn(name="RecordDecision") +async def record_decision(decision: dict[str, Any]) -> str: + """An ordinary command in the same task as the publish. + + Appends its receipt onto its own workflow's stream too, under the + Activity's own identity, so a reader outside sees the decision and the + record of it side by side. + """ + receipt = f"recorded:{decision['source']}:{decision['branch']}" + await activity.stream_handle().producer(topic=RECEIPTS).append({"receipt": receipt}) + return receipt + + +def decide(token: dict[str, Any]) -> dict[str, Any]: + """The decision the workflow is here to make.""" + if token["value"] % 2 == 0: + return { + "source": token["id"], + "branch": "even", + "computed": token["value"] * 10, + } + return {"source": token["id"], "branch": "odd", "computed": token["value"] + 100} + + +@workflow.defn(name="StreamContractDemo", sandboxed=False) +class AgentLoop: + """Read, decide, write, until the producer says it has finished.""" + + @workflow.run + async def run(self, limit: int) -> list[dict[str, Any]]: + """Decide on at most ``limit`` inputs, then return the trace.""" + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + accepted = 0 + try: + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + # A newer attempt of the same producer started writing. The + # decisions already published stand, so the workflow says so + # rather than pretending they can be withdrawn. + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "producer": record.producer_id, + "replaced": record.supersession.previous_attempt, + "attempt": record.supersession.attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + # The producer says it is done, which is what ends the loop. + # Counting decisions instead would leave the terminal record + # unread and let the workflow finish while its producer is + # still writing. + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decision = decide(record.value) + decisions.publish(decision) + receipt = await workflow.execute_activity( + record_decision, + decision, + activity_id=f"decision-{decision['source']}", + start_to_close_timeout=timedelta(seconds=10), + ) + trace.append( + {"kind": "decision", "value": decision, "receipt": receipt} + ) + accepted += 1 + if accepted >= limit: + # A bound so a stuck producer cannot run this forever. The + # terminal record above is the ordinary way out. + break + finally: + inputs.close() + decisions.finish() + return trace diff --git a/streams_demo/provider_setup.py b/streams_demo/provider_setup.py new file mode 100644 index 000000000..2e4bcc350 --- /dev/null +++ b/streams_demo/provider_setup.py @@ -0,0 +1,36 @@ +"""Pick the provider for a demo run from the environment. + +``STREAMS_PROVIDER`` names a provider; this base tree carries only ``memory``, +and each provider branch adds its own name here. The demo needs a Temporal +server to run the workflow either way; ``TEMPORAL_ADDRESS`` points at it. +""" + +from __future__ import annotations + +import os + +from temporalio.streams.providers import ProviderPlugin +from temporalio.streams.providers.memory import MemoryStreams + +NAME = os.environ.get("STREAMS_PROVIDER", "memory") + +# The memory provider is not replay-safe, so its demo keeps the cache warm. +# Storage providers run with the smallest cache they support instead. +WORKFLOW_CACHE = int(os.environ.get("STREAMS_WORKFLOW_CACHE", "512")) + + +async def open() -> tuple[str, ProviderPlugin]: + """The server to connect to and the provider the worker and the client share.""" + if NAME != "memory": + raise SystemExit(f"this tree carries no stream provider named {NAME!r}") + return os.environ.get("TEMPORAL_ADDRESS", "localhost:7233"), MemoryStreams() + + +async def close(provider: ProviderPlugin) -> None: + """Let go of whatever :func:`open` acquired. + + The memory provider holds no connection, so this is its ``close()`` and + nothing more. A provider branch that opens one closes it the same way, so + the demo's teardown reads the same on every provider. + """ + await provider.close() diff --git a/streams_demo/run_demo.py b/streams_demo/run_demo.py new file mode 100644 index 000000000..e91bf2e31 --- /dev/null +++ b/streams_demo/run_demo.py @@ -0,0 +1,172 @@ +"""Run the shared agent loop against whichever provider is configured. + +Three cases, the same on every provider: + +- read, decide, publish and an ordinary Activity in the same workflow task, + with the smallest workflow cache the provider supports, so that as much of + the run as it allows is rebuilt rather than remembered; +- a producer whose second attempt supersedes its first, which the reader has + to report and the workflow has to act on; +- an outside reader following what the workflow published, and the receipts + the Activity appended on the workflow's stream from inside its own context. + +The provider is registered once, on the client; the worker inherits it and +every context asks for its stream without naming it. Byte-identical in every +tree. ``provider_setup`` is what differs, and it is the only import here that +names a provider. +""" + +from __future__ import annotations + +import asyncio +import json +import sys +import time +import uuid +from pathlib import Path +from typing import Any + +from temporalio.api.enums.v1 import EventType +from temporalio.client import Client +from temporalio.streams import RecordKind, StreamHandle +from temporalio.worker import Worker + +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import provider_setup # noqa: E402 +from agent_loop import ( # noqa: E402 + DECISIONS, + INPUTS, + RECEIPTS, + AgentLoop, + record_decision, +) + +DECISION_LIMIT = 8 +EXPECTED_OUTPUT = 6 + + +async def collect_output(stream: StreamHandle, want: int) -> list[dict]: + """Read ``want`` decisions off the workflow's stream from outside it.""" + seen: list[dict] = [] + async for record in stream.read(topic=DECISIONS): + seen.append({"kind": record.kind.name, "value": record.value}) + if len(seen) >= want: + break + return seen + + +async def main() -> int: + """Run the demo once and write what happened next to this file.""" + out = Path(__file__).resolve().parent / f"results-{provider_setup.NAME}" + out.mkdir(exist_ok=True) + target, provider = await provider_setup.open() + client = await Client.connect(target, namespace="default", plugins=[provider]) + + uid = f"ai198-contract-{provider_setup.NAME}-" + uuid.uuid4().hex + record: dict[str, Any] = { + "provider": provider_setup.NAME, + "workflow_id": uid, + "target": target, + "max_cached_workflows": provider_setup.WORKFLOW_CACHE, + } + + async with Worker( + client, + task_queue=uid, + workflows=[AgentLoop], + activities=[record_decision], + max_cached_workflows=provider_setup.WORKFLOW_CACHE, + ): + handle = await client.start_workflow( + AgentLoop.run, DECISION_LIMIT, id=uid, task_queue=uid + ) + stream = client.get_stream_handle(uid) + output = asyncio.create_task(collect_output(stream, EXPECTED_OUTPUT)) + + # The first attempt writes two records and then stops, as a failed + # activity would. The second writes different inputs under the same + # logical producer, which is what the reader has to report. + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"id": "r1", "value": 1}, {"id": "r2", "value": 2}) + await asyncio.sleep(0.5) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"id": "r3", "value": 3}, {"id": "r4", "value": 4}) + await second.finish() + + # A failed Workflow Task is not an outcome. The server rejects a + # completion that raced newly buffered events, and the retry usually + # gets through, so reading the first failure as the result reports a + # working run as a broken one. Wait for a terminal event, then report + # the retries separately so they are neither the headline nor hidden. + deadline = time.monotonic() + 120 + # Taken from the enum rather than written out, because guessing these + # numbers is how a run that completed gets reported as terminated. + terminal = { + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_COMPLETED: "completed", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_FAILED: "workflow_failed", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_TIMED_OUT: "execution_timed_out", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_CANCELED: "canceled", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_TERMINATED: "terminated", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_CONTINUED_AS_NEW: "continued_as_new", + } + while time.monotonic() < deadline: + history = await handle.fetch_history() + reached = [ + terminal[e.event_type] + for e in history.events + if e.event_type in terminal + ] + if reached: + record["outcome"] = reached[-1] + if record["outcome"] == "completed": + record["trace"] = await handle.result() + break + await asyncio.sleep(0.1) + else: + record["outcome"] = "timed_out_waiting" + history = await handle.fetch_history() + + retried = [ + e + for e in history.events + if e.event_type == EventType.EVENT_TYPE_WORKFLOW_TASK_FAILED + ] + record["task_failures"] = [ + { + "event_id": e.event_id, + "cause": int(e.workflow_task_failed_event_attributes.cause), + "message": e.workflow_task_failed_event_attributes.failure.message, + } + for e in retried + ] + + try: + record["observed_output"] = await asyncio.wait_for(output, timeout=20) + except asyncio.TimeoutError: + output.cancel() + record["observed_output"] = "timed_out" + + async def receipts() -> list[Any]: + # Ends by itself once the workflow is closed and the tail served. + return [ + r.value + async for r in stream.read(topic=RECEIPTS) + if r.kind is RecordKind.DATA + ] + + try: + record["receipts"] = await asyncio.wait_for(receipts(), timeout=20) + except asyncio.TimeoutError: + record["receipts"] = "timed_out" + + history = await handle.fetch_history() + record["history_events"] = len(history.events) + (out / "history.json").write_text(history.to_json()) + await provider_setup.close(provider) + (out / "results.json").write_text(json.dumps(record, indent=2) + "\n") + print(json.dumps(record, indent=2)) + return 0 if record["outcome"] == "completed" else 1 + + +if __name__ == "__main__": + sys.exit(asyncio.run(main())) diff --git a/temporalio/activity.py b/temporalio/activity.py index 3f69bc17f..c906c5f05 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -20,6 +20,7 @@ from typing import ( TYPE_CHECKING, Any, + Literal, NoReturn, overload, ) @@ -29,9 +30,11 @@ import temporalio.bridge.proto.activity_task import temporalio.common import temporalio.converter +import temporalio.streams from temporalio.converter._payload_converter import ( _TemporalTransferTypePayloadConverter, ) +from temporalio.streams._ref import open_ref from .types import CallableType @@ -209,6 +212,7 @@ class _Context: runtime_metric_meter: temporalio.common.MetricMeter | None client: Client | None cancellation_details: _ActivityCancellationDetailsHolder + stream_provider: temporalio.streams.StreamProvider | None = None _logger_details: Mapping[str, Any] | None = None _payload_converter: temporalio.converter.PayloadConverter | None = None _metric_meter: temporalio.common.MetricMeter | None = None @@ -298,6 +302,132 @@ def client() -> Client: return client +@overload +def stream_handle( + workflow_id: temporalio.streams.StreamRef, + *, + run_id: str | None = None, + scope: Literal["workflow", "activity"] | None = None, +) -> temporalio.streams.RefHandle: ... + + +@overload +def stream_handle( + workflow_id: str | None = None, + *, + run_id: str | None = None, + scope: Literal["workflow", "activity"] | None = None, +) -> temporalio.streams.StreamHandle: ... + + +def stream_handle( + workflow_id: str | temporalio.streams.StreamRef | None = None, + *, + run_id: str | None = None, + scope: Literal["workflow", "activity"] | None = None, +) -> temporalio.streams.StreamHandle: + """Return a stream handle from the provider the worker was given. + + A :py:class:`temporalio.streams.StreamRef` in place of ``workflow_id``, + such as one this activity received as an argument, opens the stream it + names, whatever owns it, and takes no other argument; the handle's calls + that name no topic then address a topic the ref names. + + Which stream a call with no ``workflow_id`` reaches is decided by where + the activity runs, never by what exists: + + - In an activity a workflow scheduled, it is that workflow's stream, + pinned to the run the activity belongs to. + - In a standalone activity, it is the activity's own stream. + - ``scope="activity"`` gives an activity a workflow scheduled its own + streams instead, apart from the workflow's. ``scope="workflow"`` asks + for the workflow explicitly, and a standalone activity has none. + + The rule is static because a stream is created by its first write, so a + rule that looked for one would send attempt 1 to the workflow and a retry + to the stream attempt 1 created. An activity's own streams are one per + activity execution, not per attempt: a retry writes to the same stream, + under a new attempt, and they end when the activity reaches a terminal + status. + + Name a ``workflow_id`` to address another workflow; ``run_id`` then pins + the handle to one run and its absence follows the execution chain. A + ``read`` starts at :py:data:`temporalio.streams.BEGINNING`, at + :py:data:`temporalio.streams.END` or at the last ``N`` records with + ``last=N``. See :py:mod:`temporalio.streams`. + + Like :py:func:`client`, this is only available in ``async def`` + activities. + + Args: + workflow_id: Another workflow whose stream to address, or a + :py:class:`temporalio.streams.StreamRef` naming the stream. + run_id: The run of ``workflow_id`` to pin to. + scope: ``"activity"`` for this activity's own streams, + ``"workflow"`` for its workflow's. Without it the rule above + decides. + + Returns: + :py:class:`temporalio.streams.StreamHandle` for use in the current + activity. + + Raises: + RuntimeError: When the client is not available, which is what a + ``def`` activity gets, or when ``scope="workflow"`` is asked of + an activity that belongs to no workflow. + temporalio.streams.StreamUnsupportedError: The worker has no stream + provider, or its provider cannot hold a stream an activity owns. + Register one with ``Client.connect(plugins=[provider])`` or + ``Worker(plugins=[provider])``. + ValueError: ``run_id`` was given without ``workflow_id``, + ``scope="activity"`` with one, or a ref with either. + """ + context = _Context.current() + provider = context.stream_provider + if provider is None: + raise temporalio.streams.StreamUnsupportedError( + "no stream provider is configured on this worker; register one with " + "Client.connect(plugins=[provider]) or Worker(plugins=[provider])" + ) + if isinstance(workflow_id, temporalio.streams.StreamRef): + if run_id is not None or scope is not None: + raise ValueError( + "a StreamRef names the stream in full, so it takes no run_id or scope" + ) + return open_ref(provider, client(), workflow_id) + if workflow_id is not None: + if scope == "activity": + raise ValueError( + "scope='activity' addresses this activity's own streams, so it takes " + "no workflow_id" + ) + return provider.get_stream_handle(client(), workflow_id, run_id=run_id) + if run_id is not None: + raise ValueError("run_id needs a workflow_id") + info = context.info() + if scope is None: + scope = "workflow" if info.in_workflow else "activity" + if scope == "workflow": + if info.workflow_id is None: + raise RuntimeError( + "this activity belongs to no workflow, so name the workflow_id to " + "address, or leave scope unset for the activity's own streams" + ) + return provider.get_stream_handle( + client(), info.workflow_id, run_id=info.workflow_run_id + ) + if info.workflow_id is not None: + return provider.get_activity_stream_handle( + client(), + info.activity_id, + workflow_id=info.workflow_id, + run_id=info.workflow_run_id, + ) + return provider.get_activity_stream_handle( + client(), info.activity_id, run_id=info.activity_run_id + ) + + def in_activity() -> bool: """Whether the current code is inside an activity. diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 368cee9ed..ab6e820a8 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -28,6 +28,7 @@ import temporalio.converter import temporalio.runtime import temporalio.service +import temporalio.streams import temporalio.workflow from temporalio.service import ( ConnectConfig, @@ -40,6 +41,7 @@ ServiceClient, TLSConfig, ) +from temporalio.streams._ref import open_ref from ..common import HeaderCodecBehavior from ..types import ( @@ -188,6 +190,7 @@ async def connect( grpc_compression: GrpcCompression = GrpcCompression.GZIP, payload_limits: PayloadLimitsConfig = PayloadLimitsConfig(), header_codec_behavior: HeaderCodecBehavior = HeaderCodecBehavior.NO_CODEC, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> Self: """Connect to a Temporal server. @@ -253,6 +256,11 @@ async def connect( payload_limits: Warning thresholds for outbound payload/memo sizes. Over-threshold fields are logged but still sent. Set a threshold to 0 to disable it. header_codec_behavior: Encoding behavior for headers sent by the client. + stream_provider: Experimental. The stream provider + :py:meth:`get_stream_handle` opens handles from, see + :py:mod:`temporalio.streams`. A provider that is also a + :py:class:`Plugin` sets this itself when passed in ``plugins``, + and workers built from this client inherit it. """ connect_config = temporalio.service.ConnectConfig( target_host=target_host, @@ -289,6 +297,7 @@ def make_lambda( default_workflow_query_reject_condition=default_workflow_query_reject_condition, header_codec_behavior=header_codec_behavior, plugins=plugins, + stream_provider=stream_provider, ) def __init__( @@ -302,6 +311,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. @@ -316,6 +326,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() @@ -928,6 +939,157 @@ def get_workflow_handle( result_type=result_type, ) + @overload + def get_stream_handle( + self, + workflow_id: temporalio.streams.StreamRef, + *, + run_id: str | None = None, + activity_id: str | None = None, + stream_id: str | None = None, + ) -> temporalio.streams.RefHandle: ... + + @overload + def get_stream_handle( + self, + workflow_id: str | None = None, + *, + run_id: str | None = None, + activity_id: str | None = None, + stream_id: str | None = None, + ) -> temporalio.streams.StreamHandle: ... + + def get_stream_handle( + self, + workflow_id: str | temporalio.streams.StreamRef | None = None, + *, + run_id: str | None = None, + activity_id: str | None = None, + stream_id: str | None = None, + ) -> temporalio.streams.StreamHandle: + """Get a handle on a stream from the provider registered on this client. + + Mirrors :py:meth:`get_workflow_handle`: without ``run_id`` the handle + follows the workflow's execution chain across continue-as-new, with + one it is pinned to that run. With ``activity_id`` the handle is on + the streams that activity owns: a standalone activity's when + ``workflow_id`` is left out, and ``run_id`` then pins the activity's + run, or an activity that ``workflow_id`` scheduled. With + ``stream_id`` it is on a standalone stream, one with an id of its own + and no owner, which :py:meth:`create_stream` made; it takes no other + argument. A :py:class:`temporalio.streams.StreamRef` in place of + ``workflow_id`` opens the stream the ref names, whatever owns it, and + takes no other argument either; a topic the ref names is what the + handle's calls address when they name none. The provider is the one + registered with ``plugins=[provider]`` at :py:meth:`connect`, or + passed as ``stream_provider``. A ``read`` starts at + :py:data:`temporalio.streams.BEGINNING`, at + :py:data:`temporalio.streams.END` or at the last ``N`` records with + ``last=N``, and resumes only after a cursor it was handed. See + :py:mod:`temporalio.streams`. + + Args: + workflow_id: Workflow ID whose stream to get a handle to, the + workflow that scheduled ``activity_id``, or a + :py:class:`temporalio.streams.StreamRef` naming the stream. + run_id: Run ID to pin the handle to. + activity_id: Activity ID whose own streams to get a handle to. + stream_id: ID of the standalone stream to get a handle to. + + Returns: + The stream handle. + + Raises: + ValueError: No owner was named, or a ref or ``stream_id`` was + given together with another argument. + temporalio.streams.StreamUnsupportedError: No stream provider is + registered on this client, or it cannot hold a stream an + activity owns or a stream without an owner. + """ + provider = self._config.get("stream_provider") + if provider is None: + raise temporalio.streams.StreamUnsupportedError( + "no stream provider is registered on this client; connect with " + "plugins=[provider]" + ) + if isinstance(workflow_id, temporalio.streams.StreamRef): + if run_id is not None or activity_id is not None or stream_id is not None: + raise ValueError( + "a StreamRef names the stream in full, so it takes no run_id, " + "activity_id or stream_id" + ) + return open_ref(provider, self, workflow_id) + if stream_id is not None: + if workflow_id is not None or run_id is not None or activity_id is not None: + raise ValueError( + "stream_id names a standalone stream, which has no workflow_id, " + "run_id or activity_id" + ) + return provider.get_standalone_stream_handle(self, stream_id) + if activity_id is not None: + return provider.get_activity_stream_handle( + self, activity_id, workflow_id=workflow_id, run_id=run_id + ) + if workflow_id is None: + raise ValueError( + "name the workflow_id, the activity_id or the stream_id to address" + ) + return provider.get_stream_handle(self, workflow_id, run_id=run_id) + + async def create_stream( + self, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> temporalio.streams.StreamHandle: + """Create a standalone stream and get a handle on it. + + A standalone stream has an id of its own and no owner, so it is + created here on purpose rather than by its first write, and it is + sealed on purpose with the handle's ``close()``, after which appends + are refused and the retained records stay readable. The three policy + arguments bound what it retains: records older than ``retention``, + beyond the newest ``max_records`` or past ``max_bytes`` of stored + records are dropped, and ``None`` leaves a 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. Another process + reaches the stream with ``get_stream_handle(stream_id=...)`` or with + the handle's ``ref()``. + + Args: + stream_id: ID of the stream to create. + retention: How long a record is kept. + max_records: How many of the newest records are kept. + max_bytes: How many bytes of records are kept. + + Returns: + A handle on the new or existing stream. + + Raises: + ValueError: ``stream_id`` is empty, a bound is not positive, or + the stream exists with a different policy. + temporalio.streams.StreamUnsupportedError: No stream provider is + registered on this client, or it cannot hold a stream without + an owner. + """ + provider = self._config.get("stream_provider") + if provider is None: + raise temporalio.streams.StreamUnsupportedError( + "no stream provider is registered on this client; connect with " + "plugins=[provider]" + ) + if not stream_id: + raise ValueError("stream_id must not be empty") + return await provider.create_standalone_stream( + self, + stream_id, + retention=retention, + max_records=max_records, + max_bytes=max_bytes, + ) + def get_workflow_handle_for( self, workflow: ( @@ -3348,6 +3510,7 @@ class ClientConnectConfig(TypedDict, total=False): grpc_compression: GrpcCompression payload_limits: PayloadLimitsConfig header_codec_behavior: HeaderCodecBehavior + stream_provider: temporalio.streams.StreamProvider | None class ClientConfig(TypedDict, total=False): @@ -3362,3 +3525,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..12a6576ed --- /dev/null +++ b/temporalio/streams/__init__.py @@ -0,0 +1,156 @@ +"""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.** :meth:`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 + :class:`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. +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 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 :func:`temporalio.workflow.stream_reader` and +:func:`temporalio.workflow.stream_writer`. An activity uses +:func:`temporalio.activity.stream_handle`: an activity a workflow scheduled +reaches that workflow's stream pinned to its run, a standalone activity +reaches its own, and ``scope="activity"`` gives the first kind its own +streams too. Any process holding a client uses +:meth:`temporalio.client.Client.get_stream_handle`, which mirrors +``get_workflow_handle`` and takes an ``activity_id`` for an activity's +streams. 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 RefHandle, StreamOwnerKind, StreamRef +from temporalio.streams._topic import StreamTopic, resolve_topic, topic + +__all__ = [ + "BEGINNING", + "CONTENT_HASH_KEY", + "END", + "Cursor", + "ReadSource", + "RecordKind", + "RefHandle", + "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..b27c839e5 --- /dev/null +++ b/temporalio/streams/_policy.py @@ -0,0 +1,79 @@ +"""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 collections.abc import Callable +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, warn: Callable[[str], None] | None = None) -> None: + """Start with no producer seen, saying anything odd through ``warn``.""" + self._attempts: dict[str, int] = {} + self._warn = warn + + 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. + + An attempt that goes backwards supersedes nothing either, and is said + rather than passed off as ordinary data: attempts only ever rise on one + producer, so a lower one means a store reordered two generations, and a + consumer reading it as the current answer would render a stale one. + """ + if not producer_id or attempt <= 0: + return None + seen = self._attempts.get(producer_id, 0) + if attempt < seen and self._warn is not None: + self._warn( + f"stream record on {topic!r} at {previous} is from attempt {attempt} of " + f"producer {producer_id!r}, behind attempt {seen}, which this reader has " + "already delivered: the store handed back two generations out of order" + ) + 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..18bdc9e67 --- /dev/null +++ b/temporalio/streams/_provider.py @@ -0,0 +1,455 @@ +"""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) is + written once and returns the position the original landed at. 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 + :meth:`temporalio.client.Client.get_stream_handle` or + :func:`temporalio.activity.stream_handle` and names the topic on each + call, as on any handle. + """ + + @overload + def read( + self, + *, + topic: StreamTopic[T], + after: Cursor = ..., + last: int | None = None, + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, + *, + topic: str, + after: Cursor = ..., + last: int | None = None, + result_type: type[T], + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, + *, + topic: str, + after: Cursor = ..., + last: int | None = None, + result_type: None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: ... + + def read( + self, + *, + topic: str | StreamTopic[Any], + 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. + + ``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 can ``aclose()`` it and + release whatever the provider parked against the store. + + 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]) -> Cursor: + """The cursor of the newest record on ``topic``, or ``BEGINNING`` when empty. + + 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, producer_id: str = ..., attempt: int = ... + ) -> StreamProducer[Any]: ... + + def producer( + self, + *, + topic: str | StreamTopic[Any], + producer_id: str = "", + attempt: int = 0, + ) -> StreamProducer[Any]: + """A producer on ``topic``. + + 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 the owner alone, 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. + """ + ... + + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> StreamHandle: + """A handle on the streams an activity owns. + + Without ``workflow_id`` the activity is a standalone one, an execution + of its own, and ``run_id`` pins one run of it. With ``workflow_id`` it + is an activity that workflow scheduled, and ``run_id`` pins the + workflow's run. Either way these streams are apart from any + workflow's: a topic here and the same topic on the workflow's handle + are two streams. A retry of the activity writes to the same streams, + and a read ends when the activity reaches a terminal status, not when + an attempt fails. + + Raises: + StreamUnsupportedError: The provider's store cannot hold a stream + an activity owns. + """ + ... + + 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. + """ + ... diff --git a/temporalio/streams/_record.py b/temporalio/streams/_record.py new file mode 100644 index 000000000..ed499a5d8 --- /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 = -1 + """The producer's position within its attempt, or -1 when unnumbered.""" + 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..2b31b6c32 --- /dev/null +++ b/temporalio/streams/_ref.py @@ -0,0 +1,249 @@ +"""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 collections.abc import AsyncGenerator +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Literal + +from temporalio.streams._record import BEGINNING, Cursor, StreamRecord +from temporalio.streams._topic import StreamTopic, resolve_topic + +if TYPE_CHECKING: + from temporalio.client import Client + from temporalio.streams._provider import ( + StreamHandle, + StreamProducer, + StreamProvider, + ) + +__all__ = ["RefHandle", "StreamOwnerKind", "StreamRef", "open_ref"] + +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 when the ref names one; a ref may name the owner alone, + and the reader names the topic on each call, as on any handle. + + 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 + :meth:`temporalio.client.Client.get_stream_handle` and + :func:`temporalio.activity.stream_handle` open one directly. + """ + + kind: StreamOwnerKind + topic: str | None = None + 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 self.topic is not None and not self.topic: + raise ValueError("a StreamRef's topic name must not be empty") + 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 + the owner alone. + """ + return cls("workflow", _name(topic), 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. + """ + return cls( + "activity", + _name(topic), + 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``.""" + return cls("standalone", _name(topic), stream_id=stream_id) + + def with_topic(self, topic: str | StreamTopic[Any] | None) -> StreamRef: + """The same owner, naming ``topic`` instead, or the owner alone for ``None``.""" + return dataclasses.replace(self, topic=_name(topic)) + + +def _name(topic: str | StreamTopic[Any] | None) -> str | None: + """The topic's plain name, or ``None`` for a ref that names the owner alone.""" + if topic is None: + return None + name, _ = resolve_topic(topic) + return name + + +def open_ref(provider: StreamProvider, client: Client, ref: StreamRef) -> RefHandle: + """The handle ``ref`` names, on ``provider``. + + The ref's kind picks the provider call that opens the owner. A topic the + ref names becomes the handle's default, so a call that names none + addresses the stream the ref names; a ref that names the owner alone + leaves the topic to each call. A provider that cannot host that owner + kind raises :class:`temporalio.streams.StreamUnsupportedError` from the + call that would have opened it. + """ + if ref.kind == "workflow": + assert ref.workflow_id is not None + handle = provider.get_stream_handle(client, ref.workflow_id, run_id=ref.run_id) + elif ref.kind == "activity": + assert ref.activity_id is not None + handle = provider.get_activity_stream_handle( + client, ref.activity_id, workflow_id=ref.workflow_id, run_id=ref.run_id + ) + else: + assert ref.stream_id is not None + handle = provider.get_standalone_stream_handle(client, ref.stream_id) + return RefHandle(handle, ref) + + +class RefHandle: + """A provider's handle whose default topic is the one a ref names. + + The handle :func:`open_ref`, :meth:`temporalio.client.Client.get_stream_handle` + and :func:`temporalio.activity.stream_handle` return for a + :class:`StreamRef`. It is a :class:`temporalio.streams.StreamHandle` whose + ``read``, ``latest`` and ``producer`` may leave the topic out. + + Every call passes through unchanged when it names a topic; one that + names none gets the ref's, and is refused when the ref names none + either. The wrapper exists so a ref can address a stream without every + provider learning about refs. + """ + + def __init__(self, inner: StreamHandle, ref: StreamRef) -> None: + """Wrap ``inner``, the provider's handle on the owner ``ref`` names.""" + self._inner = inner + self._ref = ref + + def _topic(self, topic: str | StreamTopic[Any] | None) -> str | StreamTopic[Any]: + if topic is not None: + return topic + if self._ref.topic is None: + raise ValueError( + "this handle was opened from a ref that names no topic, so the call " + "has to name one" + ) + return self._ref.topic + + 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]: + """Read ``topic``, or the ref's topic without one; see :meth:`StreamHandle.read`.""" + # The protocol's overloads each take one shape of topic and + # result_type; a passthrough hands over whatever it was given. + inner: Any = self._inner + return inner.read( + topic=self._topic(topic), after=after, last=last, result_type=result_type + ) + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The newest cursor on ``topic``, or on the ref's topic without one.""" + return await self._inner.latest(topic=self._topic(topic)) + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> StreamProducer[Any]: + """A producer on ``topic``, or on the ref's topic without one.""" + inner: Any = self._inner + return inner.producer( + topic=self._topic(topic), producer_id=producer_id, attempt=attempt + ) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """The ref this handle was opened from, or one to ``topic`` of the same owner.""" + if topic is None: + return self._ref + return self._inner.ref(topic=topic) + + async def close(self) -> None: + """Seal the stream, as :meth:`StreamHandle.close` does.""" + await self._inner.close() diff --git a/temporalio/streams/_topic.py b/temporalio/streams/_topic.py new file mode 100644 index 000000000..173291fad --- /dev/null +++ b/temporalio/streams/_topic.py @@ -0,0 +1,92 @@ +"""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. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any, Generic, TypeVar, overload + +__all__ = ["StreamTopic", "resolve_topic", "topic"] + +T = TypeVar("T") + + +@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 :func:`temporalio.workflow.stream_reader`, + :func:`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], 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. + + 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 + 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..c6eaa3f02 --- /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 = -1, +) -> 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(warn) + + 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..7aaa5a9b1 --- /dev/null +++ b/temporalio/streams/providers/__init__.py @@ -0,0 +1,87 @@ +"""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 configure_client(self, config: ClientConfig) -> ClientConfig: + """Set this provider as the client's ``stream_provider``.""" + config["stream_provider"] = self + 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``.""" + config["stream_provider"] = self + return config + + def configure_replayer(self, config: ReplayerConfig) -> ReplayerConfig: + """Set this provider as the replayer's ``stream_provider``.""" + config["stream_provider"] = self + 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..de6d2b43e --- /dev/null +++ b/temporalio/streams/providers/memory.py @@ -0,0 +1,809 @@ +"""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. +- An activity's own streams are keyed by the activity, not by its run. A + standalone activity's read ends when describing it shows a terminal + status. A workflow's activity is described through its workflow: the read + ends once the activity has been seen pending and is no longer, or the + workflow closed, so a read opened after the activity already finished + waits for the workflow. +- It keeps every record until :meth:`MemoryStreams.truncate` drops the + oldest ones, which stands in for a store's retention in tests. +- A standalone stream lives here with its policy and a sealed flag. + ``retention``, ``max_records`` and ``max_bytes`` are applied when a record + is appended, so a stream nobody writes to keeps records past their + retention. A read on it ends when it is sealed and the tail delivered, and + a read, ``latest`` or ``producer`` on a stream id that does not exist + raises :class:`temporalio.streams.StreamNotFoundError` at the call rather + than waiting for the stream to be created. +- 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 +import time +from collections.abc import AsyncGenerator +from dataclasses import dataclass +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 ActivityExecutionStatus, Client, WorkflowExecutionStatus +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import content_fingerprint, decode_body, encode_body +from temporalio.streams._errors import ( + StreamClosedError, + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) +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) + + +@dataclass(frozen=True) +class _Policy: + """What a standalone stream retains, applied as records are appended.""" + + retention: timedelta | None = None + max_records: int | None = None + max_bytes: int | None = None + + def __post_init__(self) -> None: + if self.retention is not None and self.retention <= timedelta(0): + raise ValueError("retention must be positive") + if self.max_records is not None and self.max_records <= 0: + raise ValueError("max_records must be positive") + if self.max_bytes is not None and self.max_bytes <= 0: + raise ValueError("max_bytes must be positive") + + +class _Standalone: + """One standalone stream: its policy, its seal and its topics.""" + + def __init__(self, policy: _Policy) -> None: + self.policy = policy + self.sealed = False + self.topics: dict[str, _Topic] = {} + + +class _Topic: + """One topic's records, and the waiters parked on its tail.""" + + def __init__(self, policy: _Policy | None = None, *, sealed: bool = False) -> None: + self.policy = policy + self.sealed = sealed + # 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] = [] + # When each retained record landed, for a retention policy. + self.stamps: list[float] = [] + # 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. + StreamClosedError: The stream was sealed. + """ + if self.sealed: + raise StreamClosedError("the stream is closed and takes no more records") + 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 + now = time.time() + self.records.extend(bodies) + self.stamps.extend([now] * len(bodies)) + if writer is not None: + self.seen[key] = (first, len(wires), content) + self._apply_policy(now) + self._wake_waiters() + return first, len(wires) + + def _wake_waiters(self) -> None: + waiters, self._waiters = self._waiters, [] + for loop, future in waiters: + loop.call_soon_threadsafe(_wake, future) + + def _apply_policy(self, now: float) -> None: + policy = self.policy + if policy is None: + return + drop = 0 + if policy.max_records is not None: + drop = max(drop, len(self.records) - policy.max_records) + if policy.max_bytes is not None: + held = sum(len(record) for record in self.records) + while drop < len(self.records) and held > policy.max_bytes: + held -= len(self.records[drop]) + drop += 1 + if policy.retention is not None: + floor = now - policy.retention.total_seconds() + while drop < len(self.records) and self.stamps[drop] < floor: + drop += 1 + if drop: + self._drop(drop) + + def seal(self) -> None: + """Take no more records, and let parked readers see the end.""" + self.sealed = True + self._wake_waiters() + + @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.""" + self._drop(max(0, len(self.records) - keep)) + + def _drop(self, count: int) -> None: + self.base += count + del self.records[:count] + del self.stamps[:count] + + 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 + self._sequence = 0 + 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 returns where the original landed; an empty call returns + the position of this producer's last record. + """ + if not values: + return self._last + return await self._write( + [ + to_wire( + self._converter.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 owner's stream from outside, with the shared reader rules. + + The owner is ``workflow_id``'s workflow, or with ``activity_id`` an + activity: a standalone one without ``workflow_id``, or one that workflow + scheduled. With ``stream_id`` the handle is on a standalone stream, which + has no owner. + """ + + def __init__( + self, + streams: MemoryStreams, + client: Client | None, + workflow_id: str | None, + run_id: str | None, + activity_id: str | None = None, + stream_id: str | None = None, + ) -> None: + """Address the owner's topics in ``streams``.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._activity_id = activity_id + self._stream_id = stream_id + self._seen_pending = False + self._converter = ( + client.data_converter + if client is not None + else temporalio.converter.DataConverter.default + ) + + def read( + self, + *, + topic: str | StreamTopic[Any], + after: Cursor = BEGINNING, + last: int | None = None, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` from where the read starts until the owner closes. + + ``END`` and ``last=`` are resolved by this call, against what the + topic holds when it is made. On a standalone stream the read ends + once the stream is sealed and the tail delivered. + """ + check_read_start(after, last) + name, result_type = resolve_topic(topic, result_type) + store = self._store(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(), + ) + + def _store(self, topic: str) -> _Topic: + if self._stream_id is not None: + return self._streams._standalone_topic(self._stream_id, topic) + if self._activity_id is not None: + return self._streams._activity_topic( + self._workflow_id, self._activity_id, topic + ) + assert self._workflow_id is not None + return self._streams._topic(self._workflow_id, topic) + + async def _closed(self) -> bool: + if self._stream_id is not None: + return self._streams._standalone_stream(self._stream_id).sealed + if self._client is None: + return False + try: + if self._activity_id is not None and self._workflow_id is None: + activity = await self._client.get_activity_handle( + self._activity_id, run_id=self._run_id + ).describe() + return activity.status != ActivityExecutionStatus.RUNNING + assert self._workflow_id is not None + description = await self._client.get_workflow_handle( + self._workflow_id, run_id=self._run_id + ).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + # A producer may write before the owner exists; there is + # nothing to follow yet, so keep waiting. + return False + raise + status = description.status + if self._activity_id is not None: + pending = any( + info.activity_id == self._activity_id + for info in description.raw_description.pending_activities + ) + if pending: + self._seen_pending = True + elif self._seen_pending: + return True + # The activity's streams are not the workflow's chain: a run that + # continued as new took its activities with it. + return status is not None and status != WorkflowExecutionStatus.RUNNING + 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]) -> Cursor: + """The cursor of the newest record on ``topic``, for following from now.""" + name, _ = resolve_topic(topic) + head = self._store(name).head + return mint_cursor(_PROVIDER, str(head - 1)) if head else BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any], + 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._store(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 owner's stream, pinned as this handle is.""" + if self._stream_id is not None: + return StreamRef.for_standalone(self._stream_id, topic=topic) + if self._activity_id is not None: + return StreamRef.for_activity( + self._activity_id, + workflow_id=self._workflow_id, + run_id=self._run_id, + topic=topic, + ) + assert self._workflow_id is not None + return StreamRef.for_workflow( + self._workflow_id, run_id=self._run_id, topic=topic + ) + + async def close(self) -> None: + """Seal a standalone stream. An owned stream ends with its owner, not by a caller.""" + if self._stream_id is None: + raise ValueError( + "only a standalone stream can be closed; this handle is on an owned " + "stream, which ends when its workflow or activity does" + ) + self._streams._seal(self._stream_id) + + +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] = {} + self._activity_topics: dict[tuple[str | None, str, str], _Topic] = {} + self._standalone: dict[str, _Standalone] = {} + + def reset(self) -> None: + """Drop every topic and every standalone stream. For tests.""" + self._topics.clear() + self._activity_topics.clear() + self._standalone.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) + + def get_activity_stream_handle( + self, + client: Client | None, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> MemoryStreamHandle: + """A handle on the topics ``activity_id`` owns, apart from any workflow's. + + As on :meth:`get_stream_handle`, ``client`` may be ``None``, and then + a read waits until the caller closes it. + """ + return MemoryStreamHandle(self, client, workflow_id, run_id, activity_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: + """Create the standalone stream ``stream_id``, or find it with the same policy. + + The policy is applied on every append to any of the stream's topics. + ``client`` may be ``None``, as on the other handles. + + Raises: + ValueError: ``stream_id`` is empty, a bound is not positive, or + the stream exists with a different policy. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + policy = _Policy(retention, max_records, max_bytes) + existing = self._standalone.get(stream_id) + if existing is None: + self._standalone[stream_id] = _Standalone(policy) + elif existing.policy != policy: + raise ValueError( + f"standalone stream {stream_id!r} exists with policy " + f"{existing.policy}, not {policy}" + ) + return MemoryStreamHandle(self, client, None, None, stream_id=stream_id) + + def get_standalone_stream_handle( + self, client: Client | None, stream_id: str + ) -> MemoryStreamHandle: + """A handle on the standalone stream ``stream_id``. + + Nothing is checked here: a ``read``, ``latest`` or ``producer`` on a + stream that was never created raises + :class:`temporalio.streams.StreamNotFoundError` at the call. + """ + return MemoryStreamHandle(self, client, None, None, stream_id=stream_id) + + 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 _activity_topic( + self, workflow_id: str | None, activity_id: str, topic: str + ) -> _Topic: + # Kept apart from the workflow topics, so no workflow id can name an + # activity's stream. + if not topic: + raise ValueError("topic must not be empty") + key = (workflow_id, activity_id, topic) + found = self._activity_topics.get(key) + if found is None: + found = self._activity_topics[key] = _Topic() + return found + + def _standalone_stream(self, stream_id: str) -> _Standalone: + stream = self._standalone.get(stream_id) + if stream is None: + raise StreamNotFoundError( + f"standalone stream {stream_id!r} does not exist; create it with " + "client.create_stream" + ) + return stream + + def _standalone_topic(self, stream_id: str, topic: str) -> _Topic: + stream = self._standalone_stream(stream_id) + if not topic: + raise ValueError("topic must not be empty") + found = stream.topics.get(topic) + if found is None: + found = stream.topics[topic] = _Topic(stream.policy, sealed=stream.sealed) + return found + + def _seal(self, stream_id: str) -> None: + stream = self._standalone_stream(stream_id) + stream.sealed = True + for store in stream.topics.values(): + store.seal() + + 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/_activity.py b/temporalio/worker/_activity.py index ded3047fc..10bf45bd2 100644 --- a/temporalio/worker/_activity.py +++ b/temporalio/worker/_activity.py @@ -33,6 +33,7 @@ import temporalio.common import temporalio.converter import temporalio.exceptions +import temporalio.streams from temporalio.converter import ( StorageDriverActivityInfo, StorageDriverStoreContext, @@ -63,9 +64,11 @@ def __init__( metric_meter: temporalio.common.MetricMeter, client: temporalio.client.Client, encode_headers: bool, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: self._bridge_worker = bridge_worker self._task_queue = task_queue + self._stream_provider = stream_provider self._activity_executor = activity_executor self._shared_state_manager = shared_state_manager self._running_activities: dict[bytes, _RunningActivity] = {} @@ -666,6 +669,9 @@ async def _execute_activity( runtime_metric_meter=None if sync_non_threaded else self._metric_meter, client=self._client if not running_activity.sync else None, cancellation_details=running_activity.cancellation_details, + stream_provider=( + self._stream_provider if not running_activity.sync else None + ), ) ) temporalio.activity.logger.debug("Starting activity") diff --git a/temporalio/worker/_replayer.py b/temporalio/worker/_replayer.py index 55cec4ca7..1387face7 100644 --- a/temporalio/worker/_replayer.py +++ b/temporalio/worker/_replayer.py @@ -18,6 +18,7 @@ import temporalio.client import temporalio.converter import temporalio.runtime +import temporalio.streams import temporalio.worker import temporalio.workflow @@ -58,13 +59,14 @@ def __init__( disable_safe_workflow_eviction: bool = False, header_codec_behavior: HeaderCodecBehavior = HeaderCodecBehavior.NO_CODEC, external_stream_backend: Any | None = None, + 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 @@ -96,6 +98,7 @@ def __init__( disable_safe_workflow_eviction=disable_safe_workflow_eviction, header_codec_behavior=header_codec_behavior, external_stream_backend=external_stream_backend, + stream_provider=stream_provider, ) self._initial_config = self._config.copy() self._default_workflow_logic_flags = set(_DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS) @@ -298,6 +301,7 @@ def on_eviction_hook( default_workflow_logic_flags=frozenset( self._default_workflow_logic_flags ), + stream_provider=self._config.get("stream_provider"), ) external_storage = data_converter.external_storage storage_driver_types = ( @@ -445,6 +449,7 @@ class ReplayerConfig(TypedDict, total=False): disable_safe_workflow_eviction: bool header_codec_behavior: HeaderCodecBehavior external_stream_backend: Any | None + stream_provider: temporalio.streams.StreamProvider | None @dataclass(frozen=True) diff --git a/temporalio/worker/_worker.py b/temporalio/worker/_worker.py index 4da7c13b3..04dfdd562 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, @@ -153,6 +154,7 @@ def __init__( disable_payload_error_limit: bool = False, max_workflow_task_external_storage_concurrency: int = _DEFAULT_WORKFLOW_TASK_EXTERNAL_STORAGE_CONCURRENCY, external_stream_backend: Any | None = None, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: """Create a worker to process workflows and/or activities. @@ -354,6 +356,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. """ @@ -404,6 +410,7 @@ def __init__( disable_payload_error_limit=disable_payload_error_limit, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, external_stream_backend=external_stream_backend, + stream_provider=stream_provider, ) plugins_from_client = cast( @@ -487,6 +494,11 @@ def _init_from_config(self, client: temporalio.client.Client, config: WorkerConf # Prepend applicable client interceptors to the given ones client_config = config["client"].config(active_config=True) # type: ignore[reportTypedDictNotRequiredAccess] + # A provider registered on the client serves its workers too, so one + # registration covers every context that asks for a stream. + stream_provider = config.get("stream_provider") or client_config.get( + "stream_provider" + ) interceptors_from_client = cast( list[Interceptor], [i for i in client_config["interceptors"] if isinstance(i, Interceptor)], @@ -524,6 +536,7 @@ def _init_from_config(self, client: temporalio.client.Client, config: WorkerConf interceptors=interceptors, metric_meter=self._runtime.metric_meter, client=client, + stream_provider=stream_provider, encode_headers=( client_config["header_codec_behavior"] == HeaderCodecBehavior.CODEC ), @@ -593,6 +606,7 @@ def check_activity(activity: str): != HeaderCodecBehavior.NO_CODEC, max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency, external_stream_backend=self._external_stream_backend, + stream_provider=stream_provider, ) tuner = config.get("tuner") @@ -1080,6 +1094,7 @@ class WorkerConfig(TypedDict, total=False): disable_payload_error_limit: bool max_workflow_task_external_storage_concurrency: int external_stream_backend: Any | None + stream_provider: temporalio.streams.StreamProvider | None def _warn_if_activity_executor_max_workers_is_inconsistent( diff --git a/temporalio/worker/_workflow.py b/temporalio/worker/_workflow.py index 9689dd4e7..2ea360f58 100644 --- a/temporalio/worker/_workflow.py +++ b/temporalio/worker/_workflow.py @@ -27,6 +27,7 @@ import temporalio.converter import temporalio.converter._extstore import temporalio.exceptions +import temporalio.streams import temporalio.workflow from temporalio.bridge.worker import PollShutdownError from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo @@ -38,6 +39,7 @@ _relax_sandbox_for_debugger, ) from ._interceptor import ( + ExecuteWorkflowInput, Interceptor, WorkflowInboundInterceptor, WorkflowInterceptorClassInput, @@ -59,6 +61,43 @@ LOG_PROTOS = False +class _StreamHooksInterceptor(WorkflowInboundInterceptor): + """Brackets the workflow function with the stream provider's lifecycle hooks. + + Installed by the worker when it has a stream provider, so no workflow + code has to call anything before it runs or before it returns. The finish + hook runs when the function returns, raises or continues as new, because + a provider that parked a reader against the run has to let go either way. + It does not run when the run is being evicted from the cache or when the + abandoned coroutine is collected: neither is the workflow ending, the + instance's state is not to be touched during eviction, and at collection + time the runtime on the thread belongs to whichever workflow happens to + be running, so the hook would act on that one. + """ + + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + runtime = temporalio.workflow._Runtime.current() + provider = runtime.workflow_streams().provider + provider.on_workflow_start() + try: + result = await self.next.execute_workflow(input) + except GeneratorExit: + raise + except BaseException: + if not _evicting(runtime): + await provider.on_workflow_finish() + raise + await provider.on_workflow_finish() + return result + + +def _evicting(runtime: temporalio.workflow._Runtime) -> bool: + # Eviction cancels the primary task the same way a workflow cancellation + # does; the flag the instance sets before cancelling is what tells them + # apart, and only the cancellation is a run ending. + return bool(getattr(runtime, "_deleting", False)) + + # Value was chosen abitrarily as a small number that allows some concurrency and prevents # large numbers of concurrent external storage operations causing resource contention. # This default limit is per workflow task activation and does not limit the total number @@ -165,6 +204,7 @@ def __init__( default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] | None = None, external_stream_backend: Any | None = None, client: Any = None, + stream_provider: temporalio.streams.StreamProvider | None = None, ) -> None: # Debug mode is enabled if specified or if the TEMPORAL_DEBUG env var is truthy debug_mode = debug_mode or bool(os.environ.get("TEMPORAL_DEBUG")) @@ -223,6 +263,11 @@ def __init__( __temporal_assert_local_activity_valid=assert_local_activity_valid, ) ) + self._stream_provider = stream_provider + if stream_provider is not None: + # Innermost, so the lifecycle hooks bracket the workflow function + # itself, after every user interceptor has done its own setup. + self._interceptor_classes.append(_StreamHooksInterceptor) # External Workflow Streams. The manager is per-Worker and owns the # backend connection and watcher tasks; it is created lazily on the @@ -981,6 +1026,7 @@ def _create_workflow_instance( default_workflow_logic_flags=frozenset(self._default_workflow_logic_flags), external_stream_runtime=runtime, external_streams_configured=self._external_streams_configured, + stream_provider=self._stream_provider, ) if defn.sandboxed: return self._workflow_runner.create_instance(det) diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index cd46317b4..0866b6c1a 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -59,7 +59,9 @@ import temporalio.converter import temporalio.exceptions import temporalio.nexus.system +import temporalio.streams import temporalio.workflow +import temporalio.workflow._streams from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo from temporalio.service import __version__ @@ -203,6 +205,7 @@ class WorkflowInstanceDetails: """ external_streams_configured: bool = False """Whether Workflow code may use the runtime for new subscriptions.""" + stream_provider: temporalio.streams.StreamProvider | None = None class WorkflowInstance(ABC): @@ -346,6 +349,8 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: det.worker_level_failure_exception_types ) self._patch_activation_callback = det.patch_activation_callback + self._stream_provider = det.stream_provider + self._streams: temporalio.workflow._streams._WorkflowStreams | None = None self._default_workflow_logic_flags = det.default_workflow_logic_flags self._subscribed_channels: set[str] = set() # Keyed by channel name; one subscription per channel per run @@ -2308,6 +2313,20 @@ async def workflow_start_nexus_operation( ) ) + def workflow_streams(self) -> temporalio.workflow._streams._WorkflowStreams: + if self._streams is None: + if self._stream_provider is None: + raise RuntimeError( + "no stream provider is configured on this worker; pass one with " + "Worker(plugins=[provider]) or stream_provider=" + ) + # The workflow half is made per instance, so whatever it keeps + # dies with the instance the way handlers do. + self._streams = temporalio.workflow._streams._WorkflowStreams( + self._stream_provider.workflow_provider() + ) + return self._streams + def workflow_time_ns(self) -> int: return self._time_ns diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index edb054bd6..d975ac912 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -153,6 +153,12 @@ logger, unsafe, ) +from ._streams import ( + StreamReader, + StreamWriter, + stream_reader, + stream_writer, +) from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowConfig, @@ -262,6 +268,10 @@ "Notification", "linked_channel", "subscribe_channel", + "StreamReader", + "StreamWriter", + "stream_reader", + "stream_writer", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index d53ba79b2..6349afcb4 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -25,6 +25,7 @@ from ._channels import ChannelSubscription from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle + from ._streams import _WorkflowStreams from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowHandle, @@ -476,6 +477,9 @@ async def workflow_start_nexus_operation( summary: str | None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_streams(self) -> _WorkflowStreams: ... + @abstractmethod def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ... diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py new file mode 100644 index 000000000..550398869 --- /dev/null +++ b/temporalio/workflow/_streams.py @@ -0,0 +1,331 @@ +"""Streams from inside workflow code. + +.. warning:: + This module is experimental and may change in future versions. + +The reader and writer here are the same on every provider. They convert +values, synthesize supersession and buffer nothing the provider did not hand +them; everything provider-specific sits behind the ``ReadSource`` and +``WriteSink`` that the provider's workflow half opens. The provider itself +comes from the worker, the way the payload converter does. +""" + +from __future__ import annotations + +import asyncio +from collections import deque +from collections.abc import AsyncIterator, Callable +from typing import Any, Generic, TypeVar, cast, overload + +from temporalio.streams._provider import ReadSource, WorkflowStreamProvider, WriteSink +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import RecordDecoder, to_wire +from temporalio.workflow._context import _Runtime, payload_converter +from temporalio.workflow._exceptions import ReadOnlyContextError +from temporalio.workflow._sandbox import logger + +__all__ = ["StreamReader", "StreamWriter", "stream_reader", "stream_writer"] + +T = TypeVar("T") + + +class _WorkflowStreams: + """The stream state a workflow instance carries: its provider half and open readers.""" + + def __init__(self, provider: WorkflowStreamProvider) -> None: + self.provider = provider + self.readers: dict[str, StreamReader[Any]] = {} + + +class StreamReader(Generic[T]): + """Reads one topic of this workflow's stream from inside workflow code. + + The reader is the async iterator: ``async for record in reader`` yields + every kind of record, including the supersession the reader synthesizes + when a producer's newer attempt appears. Check ``record.kind``, or + iterate :meth:`values` when the workflow only wants data. Iteration ends + when the reader is closed or the provider ends the subscription. One loop + per reader: two loops on one reader share its buffer and interleave. + """ + + def __init__( + self, + source: ReadSource, + *, + topic: str, + result_type: type | None, + after: Cursor, + on_close: Callable[[], None], + ) -> None: + """Prefer :func:`temporalio.workflow.stream_reader`.""" + self._source = source + self._topic = topic + self._result_type = result_type + self._decoder = RecordDecoder( + payload_converter(), result_type, after=after, warn=logger.warning + ) + self._pending: deque[StreamRecord[T]] = deque() + self._lock = asyncio.Lock() + self._closed = False + self._ended = False + self._on_close = on_close + + @property + def topic(self) -> str: + """The name of the topic this reader is subscribed to.""" + return self._topic + + def __aiter__(self) -> StreamReader[T]: + """The reader is its own iterator.""" + return self + + async def __anext__(self) -> StreamRecord[T]: + """The next record, waiting for one to arrive.""" + while True: + if self._pending: + return self._pending.popleft() + if self._closed or self._ended: + raise StopAsyncIteration + await self._fill() + + async def _fill(self) -> None: + # Two loops on one reader must not race the source, so one batch + # fetch is in flight at a time and the second loop takes what the + # first one buffered. + async with self._lock: + if self._pending or self._closed or self._ended: + return + try: + batch = await self._source.next_batch() + except StopAsyncIteration: + # The provider ended the subscription. Iteration stops rather + # than raising, so a workflow that reads to the end of a + # finished stream leaves the loop instead of failing its task. + self._ended = True + return + for cursor, wire in batch: + self._pending.extend(self._decoder.decode(cursor, wire)) + + async def values(self) -> AsyncIterator[T]: + """Iterate the data values, dropping control records.""" + async for record in self: + if record.kind is RecordKind.DATA: + yield cast("T", record.value) + + def close(self) -> None: + """End the subscription. Idempotent. + + A later :func:`temporalio.workflow.stream_reader` on the same topic + opens a new subscription, which is a new command. + """ + if self._closed: + return + self._closed = True + self._source.close() + self._on_close() + + +class StreamWriter(Generic[T]): + """Publishes to one topic of this workflow's stream. + + A workflow can only publish transactionally to its own stream, on every + provider. Writing to somebody else's stream is an activity's job, and it + gets the weaker guarantee that goes with doing I/O. The type parameter is + the topic definition's value type; a writer on a string-named topic takes + any value. + """ + + def __init__(self, sink: WriteSink, topic: str) -> None: + """Prefer :func:`temporalio.workflow.stream_writer`.""" + self._sink = sink + self._topic = topic + self._finished = False + + @property + def topic(self) -> str: + """The name of the topic this writer is bound to.""" + return self._topic + + def publish(self, value: T) -> None: + """Append ``value`` to this topic. + + Synchronous, because there is nothing to wait for inside a task: the + record becomes visible when this Workflow Task is accepted, and never + at all if the task fails, so a reader cannot see a decision the + workflow did not commit. A :class:`temporalio.common.RawValue` passes + through pre-encoded. + + Raises: + ValueError: :meth:`finish` was already called on this writer. + temporalio.workflow.ReadOnlyContextError: Called from a query or an + update validator, which commit nothing. + """ + if self._finished: + raise ValueError(f"topic {self._topic!r} was already finished") + # A query is answered on a task that carries the answer and nothing else, so a + # record published here could only ever be dropped. Said at the call rather than + # thrown away later without a word. + if _Runtime.current().workflow_is_read_only(): + raise ReadOnlyContextError( + "While in read-only function, action attempted: publish to a stream" + ) + self._sink.publish( + to_wire( + payload_converter(), + topic=self._topic, + kind=RecordKind.DATA, + value=value, + ) + ) + + def finish(self) -> None: + """Write ``FINISH`` for this workflow on this topic. Idempotent. + + Says this workflow has nothing more to send on the topic. It does not + say the workflow succeeded, and it does not end anyone's read. + """ + if self._finished: + return + self._finished = True + self._sink.publish( + to_wire(payload_converter(), topic=self._topic, kind=RecordKind.FINISH) + ) + + +@overload +def stream_reader( + topic: StreamTopic[T], *, after: Cursor = ..., last: int | None = None +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, + *, + result_type: type[T], + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, + *, + result_type: None = None, + after: Cursor = ..., + last: int | None = None, +) -> StreamReader[Any]: ... + + +def stream_reader( + topic: str | StreamTopic[Any], + *, + result_type: type | None = None, + after: Cursor = BEGINNING, + last: int | None = None, +) -> StreamReader[Any]: + """Subscribe this workflow to ``topic`` of its own stream. + + ``topic`` is a :func:`temporalio.streams.topic` definition, which carries + the record type, or a plain string with ``result_type=`` for a name + decided at runtime. One subscription per topic per run. A second call + for the same topic returns the reader already open on it, so records go + to whichever loop pulls first; such a call may pass no ``after``, no + ``last`` and no different type. Adding a reader on a new topic is a new command, so + gate it with :func:`temporalio.workflow.patched` as you would a timer. A + reader in a successor run starts a new subscription: nothing crosses + continue-as-new implicitly. + + Args: + topic: The topic, relative to this workflow's stream. + result_type: The value type for a string-named topic, used as the + decode hint. :class:`temporalio.common.RawValue` returns the + payload untouched. + after: Resume strictly after this record. ``BEGINNING`` starts at + the oldest record the topic still holds and + :data:`temporalio.streams.END` at whatever is appended after the + subscription is registered. Honoured on the first subscription + of a run, because after that the recorded observations decide. + last: Start at the newest ``last`` records instead, or at all of + them when there are fewer. Records of every kind count. Exclusive + with a cursor in ``after``. Where it lands is resolved once, + outside the workflow, and replay reproduces it. + + Raises: + ValueError: ``topic`` is empty, ``result_type`` was passed with a + definition, ``last`` is not positive or came with a cursor, or a + reader on the topic is already open and this call asked for a + different position or type. + temporalio.streams.StreamCursorError: ``after`` was minted by another + provider. + temporalio.streams.StreamUnsupportedError: The provider cannot start + where ``END`` or ``last`` asks. + """ + check_read_start(after, last) + name, result_type = resolve_topic(topic, result_type) + state: _WorkflowStreams = _Runtime.current().workflow_streams() + existing = state.readers.get(name) + if existing is not None: + if ( + after != BEGINNING + or last is not None + or result_type is not existing._result_type + ): + raise ValueError( + f"topic {name!r} already has a reader in this run; a second " + "stream_reader shares it and takes no after=, last= or other type" + ) + return existing + # Passed only when given, so a provider written before last= existed + # still serves every read it can. + source = ( + state.provider.open_reader(name, after=after) + if last is None + else state.provider.open_reader(name, after=after, last=last) + ) + + def forget() -> None: + state.readers.pop(name, None) + + # The decoder positions a synthesized record at the one before it. Where + # END or last= lands is not known here, so BEGINNING stands in: resuming + # from it may deliver a record twice, where END would skip one. + previous = BEGINNING if after == END or last is not None else after + reader: StreamReader[Any] = StreamReader( + source, topic=name, result_type=result_type, after=previous, on_close=forget + ) + state.readers[name] = reader + return reader + + +@overload +def stream_writer(topic: StreamTopic[T]) -> StreamWriter[T]: ... + + +@overload +def stream_writer(topic: str) -> StreamWriter[Any]: ... + + +def stream_writer(topic: str | StreamTopic[Any]) -> StreamWriter[Any]: + """Publish to ``topic`` of this workflow's stream. + + Args: + topic: A :func:`temporalio.streams.topic` definition, whose value type + the writer's ``publish`` takes, or a plain string for a name + decided at runtime. Encoding follows each published value. + + Raises: + ValueError: ``topic`` is empty. + """ + name, _ = resolve_topic(topic) + provider = _Runtime.current().workflow_streams().provider + return StreamWriter(provider.open_writer(name), name) 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..4dc059e13 --- /dev/null +++ b/tests/streams/conftest.py @@ -0,0 +1,74 @@ +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", + ) + config.addinivalue_line( + "markers", + "hosts_standalone_streams: the case needs a stream with an id of its own and " + "no owner", + ) + config.addinivalue_line( + "markers", + "wakes_by_notification: the case needs an outside append to wake a parked " + "workflow reader through the server", + ) + config.addinivalue_line( + "markers", + "wakes_by_linked_notification: the case needs an outside append to wake a " + "parked workflow reader through the channel linked to its workflow", + ) + config.addinivalue_line( + "markers", + "needs_channel_server: the case needs a server that serves notification " + "channels, named with -E host:port", + ) + config.addinivalue_line( + "markers", + "needs_linked_server: the case needs a server that serves channels linked " + "to a workflow, named with -E host:port; the case skips itself on one " + "with only independent channels", + ) + config.addinivalue_line( + "markers", + "needs_unsubscribe_server: the case needs a server that accepts the " + "unsubscribe-notification-channel command, named with -E host:port; an " + "older channel server fails the Workflow Task that carries it", + ) + + +#: The environments whose server the suite starts for itself. None of them +#: accepts the subscribe-notification-channel command. +_ENVIRONMENTS_WITHOUT_CHANNELS = ("local", "time-skipping", "envconfig") + +#: The markers naming a server capability the suite's own servers lack. +_CHANNEL_SERVER_MARKERS = ( + "needs_channel_server", + "needs_linked_server", + "needs_unsubscribe_server", +) + + +def pytest_collection_modifyitems( + config: pytest.Config, items: list[pytest.Item] +) -> None: + if config.getoption("--workflow-environment") not in _ENVIRONMENTS_WITHOUT_CHANNELS: + return + skip = pytest.mark.skip( + reason="needs a server that serves notification channels; name one with -E" + ) + for item in items: + if any(item.get_closest_marker(marker) for marker in _CHANNEL_SERVER_MARKERS): + item.add_marker(skip) diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py new file mode 100644 index 000000000..bd0a675ac --- /dev/null +++ b/tests/streams/test_activity_streams.py @@ -0,0 +1,263 @@ +"""Streams owned by activities, and which stream an activity's accessor reaches. + +``activity.stream_handle()`` resolves by a static rule, never by what exists: +an activity a workflow scheduled reaches its workflow's stream, a standalone +activity reaches its own, and ``scope="activity"`` gives an activity a +workflow scheduled its own streams. An activity's own streams are one per +activity execution, so a retry writes to the same stream and a reader sees +the attempt change as ``SUPERSEDED``. + +The memory provider always runs. A storage provider adds itself to +``SETUPS`` behind its own ``STREAMS_LIVE`` gate: its setup receives the +environment's client and hands back the provider and a client with it +registered, which the workers and the reads in these cases share. +""" + +from __future__ import annotations + +import asyncio +import uuid +from collections.abc import AsyncIterator, Callable +from dataclasses import dataclass +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.common import RetryPolicy +from temporalio.streams import ( + BEGINNING, + RecordKind, + StreamProvider, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +from tests.helpers import new_worker + +TOKENS = topic("tokens", dict) + + +@dataclass +class ActivitySetup: + """A provider under test, registered on the client the cases use.""" + + name: str + provider: StreamProvider + client: Client + + +async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: + provider = MemoryStreams(poll_interval=timedelta(milliseconds=50)) + config = client.config() + config["plugins"] = [provider] + yield ActivitySetup("memory", provider, Client(**config)) + provider.reset() + + +SETUPS: dict[str, Callable[[Client], AsyncIterator[ActivitySetup]]] = { + "memory": _memory_setup +} + + +@pytest.fixture(params=sorted(SETUPS)) +async def setup( + request: pytest.FixtureRequest, client: Client, env: WorkflowEnvironment +) -> AsyncIterator[ActivitySetup]: + if env.supports_time_skipping: + pytest.skip("the time-skipping test server has no standalone activities") + async for found in SETUPS[request.param](client): + yield found + + +async def read_all(records: Any, timeout: float = 30.0) -> list: + """Read until the stream ends, which is when its owner does.""" + + async def _collect() -> list: + return [record async for record in records] + + return await asyncio.wait_for(_collect(), timeout) + + +def summary(records: list) -> list[tuple[Any, ...]]: + return [ + (r.kind, r.attempt, r.value["token"] if r.kind is RecordKind.DATA else None) + for r in records + ] + + +@activity.defn +async def write_by_default(label: str) -> str: + # No arguments: where this lands depends only on where the activity runs. + producer = activity.stream_handle().producer(topic=TOKENS) + await producer.append({"token": label}) + await producer.finish() + return activity.info().activity_id + + +@activity.defn +async def write_to_own_streams(label: str) -> str: + producer = activity.stream_handle(scope="activity").producer(topic=TOKENS) + await producer.append({"token": label}) + await producer.finish() + return activity.info().activity_id + + +@workflow.defn +class RunsOneActivity: + """Runs one activity by name and returns what it returned.""" + + @workflow.run + async def run(self, name: str, label: str) -> str: + return await workflow.execute_activity( + name, + label, + activity_id="streamer", + start_to_close_timeout=timedelta(seconds=30), + ) + + +async def test_workflow_activity_defaults_to_its_workflow(setup: ActivitySetup): + client = setup.client + workflow_id = f"streams-wfa-{uuid.uuid4().hex}" + async with new_worker( + client, RunsOneActivity, activities=[write_by_default] + ) as worker: + await client.execute_workflow( + RunsOneActivity.run, + args=["write_by_default", "to the workflow"], + id=workflow_id, + task_queue=worker.task_queue, + ) + records = await read_all( + client.get_stream_handle(workflow_id).read(topic=TOKENS) + ) + assert summary(records) == [ + (RecordKind.DATA, 1, "to the workflow"), + (RecordKind.FINISH, 1, None), + ] + assert all(r.producer_id == "streamer" for r in records) + own = client.get_stream_handle(workflow_id, activity_id="streamer") + assert await own.latest(topic=TOKENS) == BEGINNING + + +async def test_scope_activity_gives_a_workflow_activity_its_own_streams( + setup: ActivitySetup, +): + client = setup.client + workflow_id = f"streams-wfa-own-{uuid.uuid4().hex}" + async with new_worker( + client, RunsOneActivity, activities=[write_to_own_streams] + ) as worker: + await client.execute_workflow( + RunsOneActivity.run, + args=["write_to_own_streams", "to the activity"], + id=workflow_id, + task_queue=worker.task_queue, + ) + own = client.get_stream_handle(workflow_id, activity_id="streamer") + assert summary(await read_all(own.read(topic=TOKENS))) == [ + (RecordKind.DATA, 1, "to the activity"), + (RecordKind.FINISH, 1, None), + ] + # The workflow's topic of the same name is a different stream. + workflow_stream = client.get_stream_handle(workflow_id) + assert await workflow_stream.latest(topic=TOKENS) == BEGINNING + + +async def test_standalone_activity_defaults_to_its_own_stream(setup: ActivitySetup): + client = setup.client + activity_id = f"streams-saa-{uuid.uuid4().hex}" + async with new_worker(client, activities=[write_by_default]) as worker: + handle = await client.start_activity( + write_by_default, + "standalone", + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=30), + ) + assert await handle.result() == activity_id + stream = client.get_stream_handle(activity_id=activity_id) + # The read ends by itself: the activity reached a terminal status. + records = await read_all(stream.read(topic=TOKENS)) + assert summary(records) == [ + (RecordKind.DATA, 1, "standalone"), + (RecordKind.FINISH, 1, None), + ] + assert all(r.producer_id == activity_id for r in records) + + +@activity.defn +async def fail_once_after_writing() -> None: + attempt = activity.info().attempt + producer = activity.stream_handle().producer(topic=TOKENS) + await producer.append({"token": f"attempt {attempt}"}) + if attempt == 1: + raise RuntimeError("the first attempt fails after writing") + await producer.finish() + + +async def test_a_retry_inherits_the_stream_and_supersedes(setup: ActivitySetup): + client = setup.client + activity_id = f"streams-saa-retry-{uuid.uuid4().hex}" + async with new_worker(client, activities=[fail_once_after_writing]) as worker: + handle = await client.start_activity( + fail_once_after_writing, + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=30), + retry_policy=RetryPolicy( + initial_interval=timedelta(milliseconds=200), maximum_attempts=2 + ), + ) + await handle.result() + records = await read_all( + client.get_stream_handle(activity_id=activity_id).read(topic=TOKENS) + ) + # One stream across both attempts, with the takeover reported between. + assert summary(records) == [ + (RecordKind.DATA, 1, "attempt 1"), + (RecordKind.SUPERSEDED, 2, None), + (RecordKind.DATA, 2, "attempt 2"), + (RecordKind.FINISH, 2, None), + ] + superseded = records[1].supersession + assert superseded is not None + assert (superseded.previous_attempt, superseded.attempt) == (1, 2) + + +@activity.defn +async def ask_for_misaddressed_handles() -> list[str]: + errors: list[str] = [] + try: + activity.stream_handle(scope="workflow") + except RuntimeError as error: + errors.append(f"RuntimeError: {error}") + try: + activity.stream_handle("some-workflow", scope="activity") + except ValueError as error: + errors.append(f"ValueError: {error}") + return errors + + +async def test_a_standalone_activity_has_no_workflow_to_address(setup: ActivitySetup): + client = setup.client + async with new_worker(client, activities=[ask_for_misaddressed_handles]) as worker: + errors = await client.execute_activity( + ask_for_misaddressed_handles, + id=f"streams-saa-errors-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=30), + ) + assert len(errors) == 2 + assert errors[0].startswith("RuntimeError: this activity belongs to no workflow") + assert errors[1].startswith("ValueError: scope='activity'") + + +async def test_get_stream_handle_needs_an_owner(setup: ActivitySetup): + with pytest.raises( + ValueError, match="workflow_id, the activity_id or the stream_id" + ): + setup.client.get_stream_handle() 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_stream_accessors.py b/tests/streams/test_stream_accessors.py new file mode 100644 index 000000000..5717a995c --- /dev/null +++ b/tests/streams/test_stream_accessors.py @@ -0,0 +1,236 @@ +"""The accessors each context asks for its stream, with one registration. + +The provider is registered on the client alone. A worker built from that +client inherits it, an activity on that worker reaches its own workflow's +stream through ``activity.stream_handle()``, and any code holding the client +reaches a stream through ``client.get_stream_handle()``. Without a +registration, both say so with the same error. +""" + +from __future__ import annotations + +import asyncio +import uuid +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.streams import ( + BEGINNING, + END, + RecordKind, + StreamClosedError, + StreamRef, + StreamUnsupportedError, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import ActivityEnvironment, WorkflowEnvironment +from tests.helpers import new_worker + +INPUTS = topic("inputs", dict) +DECISIONS = topic("decisions", dict) + + +@pytest.fixture +def provider(env: WorkflowEnvironment): # pyright: ignore[reportUnusedFunction] + if env.supports_time_skipping: + pytest.skip( + "the memory provider polls on a timer, which time skipping turns into a spin" + ) + streams = MemoryStreams() + yield streams + streams.reset() + + +def _with_provider(client: Client, provider: MemoryStreams) -> Client: + # The same connection, with the provider registered the way an + # application registers it: once, on the client. + config = client.config() + config["plugins"] = [provider] + return Client(**config) + + +@activity.defn +async def emit(count: int) -> str: + # No workflow id and no run id: the handle is this activity's own + # workflow, pinned to its run, and the producer's identity is the + # activity's. + model = activity.stream_handle().producer(topic=INPUTS) + for n in range(count): + await model.append({"n": n}) + await model.finish() + return activity.info().activity_id + + +@workflow.defn +class Echo: + """Runs the emitting activity and echoes what arrives on ``inputs``.""" + + @workflow.run + async def run(self, count: int) -> list[Any]: + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + emitting = workflow.start_activity( + emit, count, start_to_close_timeout=timedelta(seconds=30) + ) + seen: list[Any] = [] + async for record in inputs: + if record.kind is RecordKind.FINISH: + seen.append(("finish", record.producer_id)) + break + assert record.value is not None + seen.append(record.value["n"]) + decisions.publish({"echo": record.value["n"]}) + decisions.finish() + producer_id = await emitting + return [*seen, ("activity", producer_id)] + + +async def test_one_registration_on_the_client_serves_every_context( + client: Client, provider: MemoryStreams +): + registered = _with_provider(client, provider) + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + # No plugins on the worker: it inherits the client's provider. + async with new_worker(registered, Echo, activities=[emit]) as worker: + handle = await registered.start_workflow( + Echo.run, 2, id=workflow_id, task_queue=worker.task_queue + ) + result = await handle.result() + # The activity wrote as itself onto its own workflow's topic. + assert result[:2] == [0, 1] + assert result[2] == ["finish", result[3][1]] + + stream = registered.get_stream_handle(workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + assert await asyncio.wait_for(read_everything(), 30) == [ + (RecordKind.DATA, {"echo": 0}), + (RecordKind.DATA, {"echo": 1}), + (RecordKind.FINISH, None), + ] + + +async def test_get_stream_handle_needs_a_registered_provider(client: Client): + with pytest.raises(StreamUnsupportedError, match="plugins="): + client.get_stream_handle("wf") + + +async def test_stream_handle_needs_a_provider_on_the_worker(): + async def ask() -> None: + activity.stream_handle() + + with pytest.raises(StreamUnsupportedError, match="plugins="): + await ActivityEnvironment().run(ask) + + +async def test_an_activity_owned_stream_takes_every_read_start( + client: Client, provider: MemoryStreams +): + registered = _with_provider(client, provider) + stream = registered.get_stream_handle(activity_id=f"act-{uuid.uuid4().hex}") + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + + async def first(records: Any) -> Any: + async for record in records: + await records.aclose() + return record.value + return None + + assert await first(stream.read(topic=INPUTS, after=BEGINNING)) == {"n": 1} + assert await first(stream.read(topic=INPUTS, last=1)) == {"n": 3} + at_end = stream.read(topic=INPUTS, after=END) + await producer.append({"n": 4}) + assert await asyncio.wait_for(first(at_end), 10) == {"n": 4} + + +NOTES = topic("notes", dict) + + +@activity.defn +async def append_to_ref(ref: StreamRef) -> str: + # The ref arrived as an argument and names the stream in full, so it + # takes no scope; the handle it opens writes to the ref's topic. + try: + activity.stream_handle(ref, scope="activity") + except ValueError as error: + refused = str(error) + else: + refused = "accepted" + await activity.stream_handle(ref).producer().append({"via": "activity"}) + return refused + + +@workflow.defn +class PublishToRef: + """Hands a stream ref to an activity and returns it, both as plain data.""" + + @workflow.run + async def run(self, ref: StreamRef) -> tuple[StreamRef, str]: + refused = await workflow.execute_activity( + append_to_ref, ref, start_to_close_timeout=timedelta(seconds=30) + ) + return ref, refused + + +async def test_a_ref_travels_as_data_and_opens_from_every_context( + client: Client, provider: MemoryStreams +): + registered = _with_provider(client, provider) + shared = await registered.create_stream(f"shared-{uuid.uuid4().hex}") + ref = shared.ref(topic=NOTES) + async with new_worker( + registered, PublishToRef, activities=[append_to_ref] + ) as worker: + returned, refused = await registered.execute_workflow( + PublishToRef.run, + ref, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + # The workflow argument and result went through the data converter. + assert returned == ref + assert "takes no run_id or scope" in refused + # The activity wrote to the referenced topic, and the client reads it + # from the ref without naming the topic. Sealing the stream is what ends + # a full read of it. + await shared.close() + records = [ + (r.topic, r.value) async for r in registered.get_stream_handle(ref).read(last=1) + ] + assert records == [("notes", {"via": "activity"})] + with pytest.raises(ValueError, match="takes no run_id"): + registered.get_stream_handle(ref, run_id="r") + + +async def test_the_client_creates_reaches_and_closes_a_standalone_stream( + client: Client, provider: MemoryStreams +): + registered = _with_provider(client, provider) + stream_id = f"shared-{uuid.uuid4().hex}" + created = await registered.create_stream(stream_id, max_records=3) + producer = created.producer(topic=NOTES, producer_id="writer", attempt=1) + await producer.append({"n": 1}) + + reached = registered.get_stream_handle(stream_id=stream_id) + assert await reached.latest(topic=NOTES) != BEGINNING + await reached.close() + with pytest.raises(StreamClosedError): + await producer.append({"n": 2}) + assert [r.value async for r in reached.read(topic=NOTES)] == [{"n": 1}] + + with pytest.raises(ValueError, match="no workflow_id"): + registered.get_stream_handle(stream_id=stream_id, workflow_id="wf") + with pytest.raises(ValueError, match="empty"): + await registered.create_stream("") + with pytest.raises(StreamUnsupportedError, match="plugins="): + await client.create_stream(stream_id) + with pytest.raises(StreamUnsupportedError, match="plugins="): + client.get_stream_handle(stream_id=stream_id) diff --git a/tests/streams/test_stream_hooks.py b/tests/streams/test_stream_hooks.py new file mode 100644 index 000000000..9f98f75f6 --- /dev/null +++ b/tests/streams/test_stream_hooks.py @@ -0,0 +1,139 @@ +"""The lifecycle interceptor's rules about when the finish hook runs. + +Driven directly, with a stand-in runtime on the loop, because the two cases +that matter here are the ones a workflow test cannot stage on purpose: a +coroutine collected after its worker went away, and a run evicted from the +cache. Both used to reach the finish hook, and at collection time the hook +acted on whichever workflow was running on the thread. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.worker._interceptor import ( + ExecuteWorkflowInput, + WorkflowInboundInterceptor, +) +from temporalio.worker._workflow import _StreamHooksInterceptor + + +class _Provider: + def __init__(self) -> None: + self.calls: list[str] = [] + + def on_workflow_start(self) -> None: + self.calls.append("start") + + async def on_workflow_finish(self) -> None: + self.calls.append("finish") + + +class _Streams: + def __init__(self, provider: _Provider) -> None: + self.provider = provider + + +class _FakeRuntime: + """Only what the interceptor reads: the stream state and the eviction flag.""" + + def __init__(self, provider: _Provider, *, deleting: bool = False) -> None: + self._streams = _Streams(provider) + self._deleting = deleting + + def workflow_streams(self) -> _Streams: + return self._streams + + +class _Body(WorkflowInboundInterceptor): + """The workflow function's stand-in: returns, raises or parks forever.""" + + def __init__(self, outcome: Any) -> None: # type: ignore[reportMissingSuperCall] + self._outcome = outcome + + async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any: + del input + if self._outcome is _PARK: + await asyncio.Event().wait() + if isinstance(self._outcome, BaseException): + raise self._outcome + return self._outcome + + +_PARK = object() + + +async def _unused_run_fn() -> None: + pass + + +_INPUT = ExecuteWorkflowInput(type=object, run_fn=_unused_run_fn, args=(), headers={}) + + +@pytest.fixture +async def provider() -> Any: + fake = _Provider() + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake)) # type: ignore[arg-type] + yield fake + workflow._Runtime.set_on_loop(loop, None) + + +def _evicting(fake: _Provider) -> None: + loop = asyncio.get_running_loop() + workflow._Runtime.set_on_loop(loop, _FakeRuntime(fake, deleting=True)) # type: ignore[arg-type] + + +async def test_the_finish_hook_runs_on_return(provider: _Provider): + assert ( + await _StreamHooksInterceptor(_Body("done")).execute_workflow(_INPUT) == "done" + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_when_the_function_raises(provider: _Provider): + with pytest.raises(RuntimeError, match="boom"): + await _StreamHooksInterceptor(_Body(RuntimeError("boom"))).execute_workflow( + _INPUT + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_on_continue_as_new(provider: _Provider): + error = workflow.ContinueAsNewError.__new__(workflow.ContinueAsNewError) + with pytest.raises(workflow.ContinueAsNewError): + await _StreamHooksInterceptor(_Body(error)).execute_workflow(_INPUT) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_runs_on_a_workflow_cancellation(provider: _Provider): + # A cancelled primary task is the run ending, so the provider lets go. + with pytest.raises(asyncio.CancelledError): + await _StreamHooksInterceptor(_Body(asyncio.CancelledError())).execute_workflow( + _INPUT + ) + assert provider.calls == ["start", "finish"] + + +async def test_the_finish_hook_does_not_run_when_the_coroutine_is_collected( + provider: _Provider, +): + coroutine = _StreamHooksInterceptor(_Body(_PARK)).execute_workflow(_INPUT) + # Run up to the park, the way a worker that shut down without evicting + # leaves the primary task, then close it as garbage collection would. + coroutine.send(None) + coroutine.close() + assert provider.calls == ["start"] + + +async def test_the_finish_hook_does_not_run_during_eviction(provider: _Provider): + _evicting(provider) + with pytest.raises(asyncio.CancelledError): + await _StreamHooksInterceptor(_Body(asyncio.CancelledError())).execute_workflow( + _INPUT + ) + assert provider.calls == ["start"] diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py new file mode 100644 index 000000000..a5c32adcb --- /dev/null +++ b/tests/streams/test_streams_conformance.py @@ -0,0 +1,1184 @@ +"""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, the client the cases should use, a ``host`` that starts the +workflow owning a stream when the store lives inside a running workflow, 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 datetime import timedelta +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.client import Client +from temporalio.common import Execution, ExecutionType, RawValue +from temporalio.contrib.external_workflow_streams._wake import ChannelSupport +from temporalio.converter import ( + DataConverter, + ExternalStorage, + PayloadCodec, + StorageDriver, + StorageDriverClaim, + StorageDriverRetrieveContext, + StorageDriverStoreContext, +) +from temporalio.streams import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamClosedError, + StreamCursorError, + StreamHandle, + StreamNotFoundError, + StreamProducerError, + StreamProvider, + StreamRef, + StreamUnsupportedError, + Supersession, + topic, +) +from temporalio.streams._ref import RefHandle, open_ref +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +from tests.contrib.external_workflow_streams.conftest import server_channel_support +from tests.helpers import new_worker + +# 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 + client: Client | None = None + 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.""" + host: Callable[[str], Awaitable[None]] | None = None + """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" + 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.""" + hosts_standalone_streams: bool = True + """The store holds a stream with an id of its own and no owner.""" + waits_for_standalone_creation: bool = False + """A read on a standalone stream id that does not exist yet parks until + the first write instead of raising ``StreamNotFoundError``.""" + bounds_standalone_bytes: bool = True + """A standalone stream's policy can bound the bytes it keeps.""" + 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.""" + wakes_by_notification: bool = False + """An outside append wakes a parked workflow reader through the server, + rather than the reader finding the record on a timer of its own.""" + wakes_by_linked_notification: bool = False + """The wake above reaches a workflow-owned stream's reader through the + channel linked to its workflow, so the reader subscribes to nothing.""" + + async def open( + self, + workflow_id: str, + *, + run_id: str | None = None, + client: Client | None = None, + ) -> StreamHandle: + if self.host is not None: + await self.host(workflow_id) + 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) + if self.client is not None: + # A storage provider's setup registers the provider on the client, + # so the cases go through the accessor an application uses. + return self.client.get_stream_handle(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, + ) + + async def open_ref(self, ref: StreamRef) -> RefHandle: + if self.client is not None: + return self.client.get_stream_handle(ref) + return open_ref(self.provider, None, ref) # type: ignore[arg-type] + + async def create_stream( + self, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> StreamHandle: + if self.client is not None: + return await self.client.create_stream( + stream_id, + retention=retention, + max_records=max_records, + max_bytes=max_bytes, + ) + return await self.provider.create_standalone_stream( + None, # type: ignore[arg-type] + stream_id, + retention=retention, + max_records=max_records, + max_bytes=max_bytes, + ) + + async def open_standalone(self, stream_id: str) -> StreamHandle: + if self.client is not None: + return self.client.get_stream_handle(stream_id=stream_id) + return self.provider.get_standalone_stream_handle( + None, # type: ignore[arg-type] + stream_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, + 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, + "hosts_standalone_streams": lambda case: case.hosts_standalone_streams, + "wakes_by_notification": lambda case: case.wakes_by_notification, + "wakes_by_linked_notification": lambda case: case.wakes_by_linked_notification, +} + + +@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 drain(records: Any, timeout: float = 5.0) -> list: + """Every record until the read ends on its own.""" + + async def _collect() -> list: + return [record async for record in records] + + return await asyncio.wait_for(_collect(), timeout) + + +def new_stream_id() -> str: + return f"stream-{uuid.uuid4().hex}" + + +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] == [0, 1, 2] + 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"}] + + +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_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 owner alone; the reader names one. + assert stream.ref().topic is None + 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(case.client or 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(case.client or 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"}] + + +async def test_a_ref_opens_the_stream_it_names(case: ProviderCase): + workflow_id = new_workflow_id() + stream = await case.open(workflow_id) + await stream.producer(topic=A, producer_id="model", attempt=1).append({"n": 1}) + ref = stream.ref(topic=A) + + # The receiver names no topic: the ref carried it, so every call on the + # handle it opened addresses topic ``a`` of that workflow. + opened = await case.open_ref(ref) + records = await take(opened.read(result_type=dict), 1) + assert [(r.topic, r.value) for r in records] == [("a", {"n": 1})] + assert await opened.latest() == records[0].cursor + assert opened.ref() == ref + await opened.producer(producer_id="tool", attempt=1).append({"n": 2}) + assert [r.value for r in await take(stream.read(topic=A), 2)] == [ + {"n": 1}, + {"n": 2}, + ] + # Naming a topic on the opened handle addresses that topic instead. + assert await opened.latest(topic=B) == BEGINNING + assert opened.ref(topic=B) == stream.ref(topic=B) + + +@pytest.mark.hosts_standalone_streams +async def test_a_standalone_stream_is_read_from_another_handle(case: ProviderCase): + stream_id = new_stream_id() + created = await case.create_stream(stream_id) + producer = created.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + + # Any process reaches the stream by its id, or by a ref the creator + # handed out; nothing about the stream depends on who created it. + other = await case.open_standalone(stream_id) + records = await take(other.read(topic=OUT), 2) + assert [r.value for r in records] == [{"n": 1}, {"n": 2}] + assert await other.latest(topic=OUT) == records[1].cursor + assert other.ref(topic=OUT) == StreamRef.for_standalone(stream_id, topic="out") + via_ref = await case.open_ref(created.ref(topic=OUT)) + assert [r.value for r in await take(via_ref.read(), 2)] == [{"n": 1}, {"n": 2}] + + +@pytest.mark.hosts_standalone_streams +async def test_a_missing_standalone_stream_is_not_found(case: ProviderCase): + if case.waits_for_standalone_creation: + pytest.skip( + f"the {case.name} provider waits for a standalone stream to be created" + ) + # get_stream_handle(stream_id=) creates nothing: the stream has to have + # been created on purpose, and a use before that says so. + stream = await case.open_standalone(new_stream_id()) + with pytest.raises(StreamNotFoundError): + await stream.latest(topic=OUT) + with pytest.raises(StreamNotFoundError): + await take(stream.read(topic=OUT), 1) + with pytest.raises(StreamNotFoundError): + await stream.producer(topic=OUT, producer_id="writer", attempt=1).append( + {"n": 1} + ) + + +@pytest.mark.hosts_standalone_streams +async def test_creating_a_standalone_stream_is_idempotent_for_one_policy( + case: ProviderCase, +): + stream_id = new_stream_id() + first = await case.create_stream(stream_id, max_records=10) + # The same id and policy again is the same stream, not an error, so a + # retried create is harmless. + again = await case.create_stream(stream_id, max_records=10) + await first.producer(topic=OUT, producer_id="writer", attempt=1).append({"n": 1}) + assert [r.value for r in await take(again.read(topic=OUT), 1)] == [{"n": 1}] + # A different policy on an existing id is a mistake, not a change. + with pytest.raises(ValueError): + await case.create_stream(stream_id, max_records=5) + for bad in (dict(max_records=0), dict(max_bytes=-1), dict(retention=timedelta(0))): + with pytest.raises(ValueError): + await case.create_stream(new_stream_id(), **bad) # type: ignore[arg-type] + + +@pytest.mark.hosts_standalone_streams +async def test_closing_a_standalone_stream_ends_reads_and_refuses_appends( + case: ProviderCase, +): + stream_id = new_stream_id() + stream = await case.create_stream(stream_id) + producer = stream.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": 1}) + # A reader parked on the tail before the close has to learn of it. + other = await case.open_standalone(stream_id) + parked = asyncio.ensure_future(drain(other.read(topic=OUT), timeout=10)) + await asyncio.sleep(0.2) + await producer.append({"n": 2}) + + await stream.close() + assert [r.value for r in await parked] == [{"n": 1}, {"n": 2}] + # Sealed: the tail stays readable, and a read opened now ends by itself. + assert [r.value for r in await drain(stream.read(topic=OUT))] == [ + {"n": 1}, + {"n": 2}, + ] + with pytest.raises(StreamClosedError): + await producer.append({"n": 3}) + with pytest.raises(StreamClosedError): + await other.producer(topic=A, producer_id="late", attempt=1).append({"n": 3}) + # Closing again is not an error. + await stream.close() + + +@pytest.mark.hosts_standalone_streams +async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCase): + by_count = await case.create_stream(new_stream_id(), max_records=2) + producer = by_count.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + # BEGINNING is the oldest record still held, which the policy decided. + kept = await take(by_count.read(topic=OUT), 2) + assert [r.value for r in kept] == [{"n": 3}, {"n": 4}] + + if case.bounds_standalone_bytes: + by_bytes = await case.create_stream(new_stream_id(), max_bytes=700) + producer = by_bytes.producer(topic=OUT, producer_id="writer", attempt=1) + for n in range(3): + await producer.append({"n": n, "blob": "x" * 500}) + kept = await take(by_bytes.read(topic=OUT), 1) + assert kept[0].value is not None and kept[0].value["n"] == 2 + else: + with pytest.raises(StreamUnsupportedError): + await case.create_stream(new_stream_id(), max_bytes=700) + + if not case.trims_open_stream_by_age: + return + by_age = await case.create_stream( + new_stream_id(), retention=timedelta(milliseconds=200) + ) + producer = by_age.producer(topic=OUT, producer_id="writer", attempt=1) + await producer.append({"n": "old"}) + await asyncio.sleep(0.3) + await producer.append({"n": "new"}) + kept = await take(by_age.read(topic=OUT), 1) + assert [r.value for r in kept] == [{"n": "new"}] + + +@workflow.defn +class PublishFromConstructor: + """Publishes once from its ``@workflow.init`` constructor and once from ``run``.""" + + @workflow.init + def __init__(self) -> None: + workflow.stream_writer(OUT).publish({"from": "init"}) + + @workflow.run + async def run(self) -> None: + workflow.stream_writer(OUT).publish({"from": "run"}) + + +async def test_a_publish_from_the_constructor_is_delivered( + case: ProviderCase, client: Client, env: WorkflowEnvironment +): + if env.supports_time_skipping and case.client is None: + pytest.skip("the memory provider polls on a timer, which time skipping spins") + # A storage provider's setup registers it on its client, which a worker + # inherits; the memory provider is handed to the worker directly. + worker_client = case.client or client + plugins = [] if case.client is not None else [case.provider] + workflow_id = new_workflow_id() + async with new_worker( + worker_client, PublishFromConstructor, plugins=plugins + ) as worker: + handle = await worker_client.start_workflow( + PublishFromConstructor.run, id=workflow_id, task_queue=worker.task_queue + ) + await asyncio.wait_for(handle.result(), 30) + stream = case.provider.get_stream_handle(worker_client, workflow_id) + records = await take(stream.read(topic=OUT), 2, 30) + assert [r.value for r in records] == [{"from": "init"}, {"from": "run"}] + + +@workflow.defn +class ReadUntilFinished: + """Reads ``OUT`` until its producer finishes. + + Nothing but an outside append moves it, so what wakes it between appends + is the transport under test. + """ + + @workflow.run + async def run(self) -> list[Any]: + seen: list[Any] = [] + async for record in workflow.stream_reader(OUT): + if record.kind is RecordKind.FINISH: + break + seen.append(record.value) + return seen + + +@pytest.mark.wakes_by_notification +@pytest.mark.needs_channel_server +async def test_an_outside_producer_wakes_the_reader_through_the_channel( + case: ProviderCase, client: Client +): + """The channel path, on the public surface. + + The reader's run is subscribed to the stream's channel on the completion + that ends the task that opened the reader, after that task's marker, so + the task stays retained and parks as it would on a server without + channels. The producer's append notifies the channel, and the server + wakes the run with a Workflow Task whose scheduled event carries the + notification. History then holds the subscription and no Signal. + """ + worker_client = case.client or client + support = await server_channel_support(worker_client) + if support is ChannelSupport.NONE: + pytest.skip("the server does not implement notification channels") + if support is ChannelSupport.LINKED: + pytest.skip("a workflow-owned stream listens on its linked channel there") + handle = await _read_two_woken_from_outside(case, worker_client) + events = [e async for e in handle.fetch_history_events()] + assert _signalled(events) == [], ( + "a Signal woke the reader, so the channel was not the transport" + ) + subscribed = _subscribed(events) + assert len(subscribed) == 1, "the run subscribes once per channel" + assert _preceded_by_a_marker(events, _subscribed_event_index(events)), ( + "the subscription did not wait for the completion that ends the task" + ) + notified = _notified(events) + assert notified, "no Workflow Task was scheduled with a notification" + assert {n.channel for n in notified} == set(subscribed) + assert not any(n.HasField("linked_to") for n in notified) + + +@pytest.mark.wakes_by_linked_notification +@pytest.mark.needs_linked_server +async def test_an_outside_producer_wakes_the_reader_through_its_linked_channel( + case: ProviderCase, client: Client +): + """The linked kind, on the public surface. + + The stream's channel lives in the reading workflow's own state, so the run + subscribes to nothing; the producer's append notifies the channel by the + owner's id, and the server wakes the owner with a Workflow Task whose + scheduled event carries the notification naming it. History then holds + neither a Signal nor a subscription. + """ + worker_client = case.client or client + if await server_channel_support(worker_client) is not ChannelSupport.LINKED: + pytest.skip("the server does not serve channels linked to a workflow") + handle = await _read_two_woken_from_outside(case, worker_client) + events = [e async for e in handle.fetch_history_events()] + assert _signalled(events) == [], ( + "a Signal woke the reader, so the channel was not the transport" + ) + assert _subscribed(events) == [], "the owner is the listener by construction" + notified = _notified(events) + assert notified, "no Workflow Task was scheduled with a notification" + owners = {Execution.from_proto(n.linked_to) for n in notified} + assert {(owner.type, owner.business_id) for owner in owners} == { + (ExecutionType.WORKFLOW, handle.id) + } + assert len({n.channel for n in notified}) == 1 + + +@workflow.defn +class ReadTwoThenClose: + """Reads two values of ``OUT``, closes the reader short of ``FINISH``, waits. + + The timer after the close is what makes the close leave on a completion + the run survives: the channel has to leave on that completion, not with + the run. + """ + + @workflow.run + async def run(self) -> list[Any]: + reader = workflow.stream_reader(OUT) + seen: list[Any] = [] + async for record in reader: + seen.append(record.value) + if len(seen) == 2: + break + reader.close() + await workflow.sleep(1) + return seen + + +@workflow.defn +class ReadUntilFinishedThenClose: + """Reads ``OUT`` to its producer's ``FINISH``, closes the reader there, waits.""" + + @workflow.run + async def run(self) -> list[Any]: + reader = workflow.stream_reader(OUT) + seen: list[Any] = [] + async for record in reader: + if record.kind is RecordKind.FINISH: + break + seen.append(record.value) + reader.close() + await workflow.sleep(1) + return seen + + +@workflow.defn +class OpenAndCloseBesideTheRead: + """Opens and closes a reader on ``A`` in the task that opens the ``OUT`` reader.""" + + @workflow.run + async def run(self) -> list[Any]: + workflow.stream_reader(A).close() + seen: list[Any] = [] + async for record in workflow.stream_reader(OUT): + if record.kind is RecordKind.FINISH: + break + seen.append(record.value) + return seen + + +async def _independent_channels_or_skip(worker_client: Client) -> None: + support = await server_channel_support(worker_client) + if support is ChannelSupport.NONE: + pytest.skip("the server does not implement notification channels") + if support is ChannelSupport.LINKED: + pytest.skip("a workflow-owned stream listens on its linked channel there") + + +@pytest.mark.wakes_by_notification +@pytest.mark.needs_unsubscribe_server +async def test_a_reader_closed_short_of_finish_leaves_its_channel( + case: ProviderCase, client: Client +): + """Closing the reader ends the run's subscription on the completion that + leaves, after the progress marker and before the workflow's own command. + + The producer finishes only after the run has returned, so the reader + never saw ``FINISH``; it left because the workflow closed it. + """ + worker_client = case.client or client + await _independent_channels_or_skip(worker_client) + handle = await _read_two_woken_from_outside( + case, worker_client, ReadTwoThenClose, finish_before_result=False + ) + _assert_the_channel_left_after_the_marker( + [e async for e in handle.fetch_history_events()] + ) + + +@pytest.mark.wakes_by_notification +@pytest.mark.needs_unsubscribe_server +async def test_a_reader_closed_at_finish_leaves_its_channel( + case: ProviderCase, client: Client +): + """The same leaving, on the completion that consumed the producer's ``FINISH``.""" + worker_client = case.client or client + await _independent_channels_or_skip(worker_client) + handle = await _read_two_woken_from_outside( + case, worker_client, ReadUntilFinishedThenClose + ) + _assert_the_channel_left_after_the_marker( + [e async for e in handle.fetch_history_events()] + ) + + +@pytest.mark.wakes_by_notification +@pytest.mark.needs_unsubscribe_server +async def test_a_channel_opened_and_closed_in_one_task_is_never_subscribed( + case: ProviderCase, client: Client +): + """Only the latest report of a task counts, so a reader that came and went + inside it costs the server nothing: no subscription, no unsubscription.""" + worker_client = case.client or client + await _independent_channels_or_skip(worker_client) + handle = await _read_two_woken_from_outside( + case, worker_client, OpenAndCloseBesideTheRead + ) + events = [e async for e in handle.fetch_history_events()] + subscribed = _subscribed(events) + assert len(subscribed) == 1, "only the reader that stayed open subscribes" + assert {n.channel for n in _notified(events)} == set(subscribed) + assert _unsubscribed(events) == [], "the run ended with its reader open" + + +def _subscribed_event_index(events: Sequence[Any]) -> int: + [index] = [ + i + for i, e in enumerate(events) + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + return index + + +def _preceded_by_a_marker(events: Sequence[Any], index: int) -> bool: + """Whether the event at ``index`` follows the progress marker of its completion. + + Core issues the channel commands after the external stream marker, so an + event right after a marker landed on the completion that ended a task + rather than on a task of its own. + """ + return events[index - 1].HasField("marker_recorded_event_attributes") + + +def _assert_the_channel_left_after_the_marker(events: Sequence[Any]) -> None: + [channel] = _subscribed(events) + assert _unsubscribed(events) == [channel], "the channel leaves once" + [(index, leaving)] = [ + (i, e) + for i, e in enumerate(events) + if e.HasField("workflow_notification_channel_unsubscribed_event_attributes") + ] + joined = events[_subscribed_event_index(events)] + attributes = leaving.workflow_notification_channel_unsubscribed_event_attributes + assert attributes.subscribed_event_id == joined.event_id + assert _preceded_by_a_marker(events, index), ( + "the unsubscribe follows the progress marker of the leaving completion" + ) + assert events[index + 1].HasField("timer_started_event_attributes"), ( + "the leaving completion carried the workflow's own command after it" + ) + assert not any( + e.workflow_task_scheduled_event_attributes.notifications + for e in events[index + 1 :] + if e.HasField("workflow_task_scheduled_event_attributes") + ), "a notification reached the run after it left the channel" + + +async def _read_two_woken_from_outside( + case: ProviderCase, + worker_client: Client, + workflow_class: Any = ReadUntilFinished, + *, + finish_before_result: bool = True, +): + """Runs a reader with two appends spaced past its idle timeout. + + ``finish_before_result`` says whether the producer's ``FINISH`` is what + lets the workflow return, or is written only once it has. + """ + plugins = [] if case.client is not None else [case.provider] + workflow_id = new_workflow_id() + async with new_worker(worker_client, workflow_class, plugins=plugins) as worker: + handle = await worker_client.start_workflow( + workflow_class.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = case.provider.get_stream_handle(worker_client, workflow_id) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + # Spaced past the reader's idle timeout, so the reader parks between + # appends and only a wake from outside can move it. + for n in (1, 2): + await producer.append({"n": n}) + await asyncio.sleep(2) + if finish_before_result: + await producer.finish() + assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}, {"n": 2}] + if not finish_before_result: + await producer.finish() + return handle + + +def _signalled(events: Sequence[Any]) -> list[Any]: + return [ + e for e in events if e.HasField("workflow_execution_signaled_event_attributes") + ] + + +def _subscribed(events: Sequence[Any]) -> list[str]: + return [ + e.workflow_notification_channel_subscribed_event_attributes.channel + for e in events + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + + +def _unsubscribed(events: Sequence[Any]) -> list[str]: + return [ + e.workflow_notification_channel_unsubscribed_event_attributes.channel + for e in events + if e.HasField("workflow_notification_channel_unsubscribed_event_attributes") + ] + + +def _notified(events: Sequence[Any]) -> list[Any]: + return [ + notification + for e in events + if e.HasField("workflow_task_scheduled_event_attributes") + for notification in e.workflow_task_scheduled_event_attributes.notifications + ] diff --git a/tests/streams/test_streams_internals.py b/tests/streams/test_streams_internals.py new file mode 100644 index 000000000..8a6d60ecc --- /dev/null +++ b/tests/streams/test_streams_internals.py @@ -0,0 +1,269 @@ +"""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 + + +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_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 + registered = provider.configure_client(ClientConfig()) # type: ignore[typeddict-item] + assert registered.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" + # Without a topic a ref names the owner alone on this chain. + assert standalone == StreamRef("standalone", None, 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" diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py new file mode 100644 index 000000000..03b65b0da --- /dev/null +++ b/tests/streams/test_streams_workflow.py @@ -0,0 +1,554 @@ +"""Workflow-side conformance for the stream contract. + +Runs the reader and writer inside a real workflow on the memory provider +with a warm cache, and states the two rules about Workflow Tasks as tests: a +publish commits with its task (rule 1), and reads are recorded observations +that replay re-supplies (rule 2). The memory provider keeps neither and says +so in its docstring, so those two are strict expected failures here. A +storage provider that runs this module turns them into passes; that is the +measurement they exist for. + +The rest is what the portable surface promises on every provider: the +lifecycle hooks the worker calls, one subscription per topic per run, a read +that ends when the chain closes, a handle that follows continue-as-new, and +a topic shared by the workflow and an outside producer. +""" + +from __future__ import annotations + +import asyncio +import uuid +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.client import Client +from temporalio.streams import ( + BEGINNING, + END, + Cursor, + ReadSource, + RecordKind, + StreamCursorError, + WriteSink, + topic, +) +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +from temporalio.worker import Replayer +from tests.helpers import new_worker +from tests.streams.test_streams_conformance import take + +INPUTS = topic("inputs", dict) +DECISIONS = topic("decisions", dict) + + +@pytest.fixture +def provider(env: WorkflowEnvironment): # pyright: ignore[reportUnusedFunction] + if env.supports_time_skipping: + pytest.skip( + "the memory provider polls on a timer, which time skipping turns into a spin" + ) + streams = MemoryStreams() + yield streams + streams.reset() + + +@workflow.defn +class ContractLoop: + """Reads ``inputs``, publishes a decision per value, reports control records.""" + + @workflow.run + async def run(self) -> list[dict[str, Any]]: + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + try: + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "replaced": record.supersession.previous_attempt, + "attempt": record.supersession.attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decisions.publish({"decided": record.value["n"]}) + trace.append( + { + "kind": "decision", + "n": record.value["n"], + "attempt": record.attempt, + } + ) + finally: + inputs.close() + decisions.finish() + # Twice on purpose: a finished topic stays finished, with one marker. + decisions.finish() + return trace + + +async def _run_the_loop( + client: Client, provider: MemoryStreams +) -> tuple[Any, list[dict[str, Any]]]: + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, ContractLoop, plugins=[provider]) as worker: + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"n": 1}, {"n": 2}) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"n": 3}) + await second.finish() + trace = await handle.result() + return handle, trace + + +async def test_workflow_reads_decides_and_publishes( + client: Client, provider: MemoryStreams +): + handle, trace = await _run_the_loop(client, provider) + assert trace == [ + {"kind": "decision", "n": 1, "attempt": 1}, + {"kind": "decision", "n": 2, "attempt": 1}, + {"kind": "superseded", "replaced": 1, "attempt": 2}, + {"kind": "decision", "n": 3, "attempt": 2}, + {"kind": "finish", "producer": "model"}, + ] + + # The outside view of what the workflow published, on its own topic. + stream = provider.get_stream_handle(client, handle.id) + records = await take(stream.read(topic=DECISIONS), 5) + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"decided": 1}), + (RecordKind.DATA, {"decided": 2}), + (RecordKind.DATA, {"retracting_attempt": 1}), + (RecordKind.DATA, {"decided": 3}), + (RecordKind.FINISH, None), + ] + assert all(r.producer_id == "" and r.topic == DECISIONS.name for r in records) + # The second finish() wrote nothing: the marker is the newest record. + assert await stream.latest(topic=DECISIONS) == records[-1].cursor + + +async def test_read_ends_when_the_workflow_closes_and_the_tail_is_delivered( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + stream = provider.get_stream_handle(client, handle.id) + + async def read_everything() -> list[Any]: + return [r.value async for r in stream.read(topic=DECISIONS)] + + # No count and no early break: the read ends by itself once the workflow + # is closed and everything it retained has been handed over. + values = await asyncio.wait_for(read_everything(), 30) + assert values == [ + {"decided": 1}, + {"decided": 2}, + {"retracting_attempt": 1}, + {"decided": 3}, + None, + ] + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 2: reading is a recorded observation, so replaying the history " + "with the store gone must re-supply the same records; the memory " + "provider reads live process memory instead" + ), +) +async def test_replay_without_the_store_resupplies_the_records( + client: Client, provider: MemoryStreams +): + handle, _ = await _run_the_loop(client, provider) + history = await handle.fetch_history() + + provider.reset() + replayer = Replayer(workflows=[ContractLoop], plugins=[provider]) + await replayer.replay_workflow(history) + + +# Run ids whose first workflow task already failed, shared with the workflow +# thread so the retry can tell it is the retry. Outside the sandbox on +# purpose: the sandbox re-imports this module per run and would hide the set. +_failed_once: set[str] = set() + + +@workflow.defn(sandboxed=False) +class PublishThenFail: + """Publishes, then fails its first workflow task; the retry publishes again.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + run_id = workflow.info().run_id + committed = run_id in _failed_once + decisions.publish({"committed": committed}) + if not committed: + _failed_once.add(run_id) + raise RuntimeError("the first task fails after publishing") + decisions.finish() + + +@pytest.mark.xfail( + strict=True, + reason=( + "rule 1: a publish commits with its workflow task, so no reader sees " + "a record from a task that failed; the memory provider makes it " + "visible at publish time" + ), +) +async def test_a_failed_task_publishes_nothing(client: Client, provider: MemoryStreams): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, PublishThenFail, plugins=[provider]) as worker: + handle = await client.start_workflow( + PublishThenFail.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + records = await take(stream.read(topic=DECISIONS), 2, timeout=30) + await handle.result() + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"committed": True}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class Relay: + """Publishes one record per run and continues as new once.""" + + @workflow.run + async def run(self, run: int) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"run": run}) + if run == 0: + workflow.continue_as_new(run + 1) + decisions.finish() + + +class _RecordingHalf: + """A workflow half that logs the hooks the worker calls, then delegates.""" + + def __init__(self, inner: Any, calls: list[tuple[str, str]]) -> None: + self._inner = inner + self._calls = calls + + def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + return self._inner.open_reader(topic, after=after) + + def open_writer(self, topic: str) -> WriteSink: + return self._inner.open_writer(topic) + + def on_workflow_start(self) -> None: + self._calls.append(("start", workflow.info().run_id)) + + async def on_workflow_finish(self) -> None: + self._calls.append(("finish", workflow.info().run_id)) + + +class HookedMemory(MemoryStreams): + """The memory provider with its lifecycle hooks made visible.""" + + def __init__(self) -> None: + super().__init__() + self.calls: list[tuple[str, str]] = [] + + def workflow_provider(self) -> Any: + return _RecordingHalf(super().workflow_provider(), self.calls) + + +@pytest.mark.usefixtures("provider") +async def test_the_worker_calls_the_lifecycle_hooks_around_every_run(client: Client): + hooked = HookedMemory() + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[hooked]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + # Start before the function, finish after it, on both runs: the finish + # hook runs on the continue-as-new exit too, so a provider that parked + # something against the first run can let go before the successor starts. + kinds = [kind for kind, _ in hooked.calls] + assert kinds == ["start", "finish", "start", "finish"] + runs = [run_id for _, run_id in hooked.calls] + assert runs[0] == runs[1] and runs[2] == runs[3] and runs[0] != runs[2] + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, Relay, plugins=[provider]) as worker: + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + + async def read_everything() -> list[Any]: + return [(r.kind, r.value) async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + await handle.result() + # The chain is followed: the successor's records arrive on the same read, + # and the read ends only when the last run of the chain is closed. + assert records == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + + +@workflow.defn +class SharedReaders: + """Opens the same topic twice and pulls from both readers in turn.""" + + @workflow.run + async def run(self) -> list[Any]: + first = workflow.stream_reader(INPUTS) + second = workflow.stream_reader(INPUTS) + trace: list[Any] = ["shared" if first is second else "separate"] + try: + workflow.stream_reader(INPUTS, after=Cursor("memory:0")) + except ValueError: + trace.append("after-rejected") + try: + workflow.stream_reader(INPUTS.name, result_type=list) + except ValueError: + trace.append("type-rejected") + trace.append((await first.__anext__()).value) + trace.append((await second.__anext__()).value) + first.close() + return trace + + +async def test_a_second_reader_on_a_topic_shares_the_subscription( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, SharedReaders, plugins=[provider]) as worker: + handle = await client.start_workflow( + SharedReaders.run, id=workflow_id, task_queue=worker.task_queue + ) + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="model", attempt=1).append( + {"n": 1}, {"n": 2} + ) + assert await handle.result() == [ + "shared", + "after-rejected", + "type-rejected", + {"n": 1}, + {"n": 2}, + ] + + +@workflow.defn +class ForeignCursor: + """Resumes from a cursor another provider minted.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS, after=Cursor("elsewhere:1")) + except StreamCursorError: + return "refused" + return "accepted" + + +async def test_the_workflow_reader_refuses_a_foreign_cursor( + client: Client, provider: MemoryStreams +): + async with new_worker(client, ForeignCursor, plugins=[provider]) as worker: + result = await client.execute_workflow( + ForeignCursor.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert result == "refused" + + +@workflow.defn +class NoProvider: + """Opens a stream on a worker that has no provider.""" + + @workflow.run + async def run(self) -> str: + try: + workflow.stream_reader(INPUTS) + except RuntimeError as error: + return str(error) + return "opened" + + +async def test_a_worker_without_a_provider_says_so(client: Client): + async with new_worker(client, NoProvider) as worker: + result = await client.execute_workflow( + NoProvider.run, + id=f"streams-wf-{uuid.uuid4().hex}", + task_queue=worker.task_queue, + ) + assert "no stream provider is configured" in result + + +@workflow.defn +class OneLine: + """Publishes one record and finishes the topic.""" + + @workflow.run + async def run(self) -> None: + decisions = workflow.stream_writer(DECISIONS) + decisions.publish({"from": "workflow"}) + decisions.finish() + + +async def test_an_outside_producer_and_the_workflow_share_a_topic( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, OneLine, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=DECISIONS, producer_id="tool", attempt=1).append( + {"from": "producer"} + ) + handle = await client.start_workflow( + OneLine.run, id=workflow_id, task_queue=worker.task_queue + ) + await handle.result() + + async def read_everything() -> list[Any]: + return [r async for r in stream.read(topic=DECISIONS)] + + records = await asyncio.wait_for(read_everything(), 30) + # Both writers land on one topic in one order, each under its own + # identity: the producer's records carry its id, the workflow's carry none. + assert [(r.producer_id, r.kind, r.value) for r in records] == [ + ("tool", RecordKind.DATA, {"from": "producer"}), + ("", RecordKind.DATA, {"from": "workflow"}), + ("", RecordKind.FINISH, None), + ] + + +@workflow.defn +class PublishFromQuery: + """Publishes from a query handler, which commits nothing.""" + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: False) + + @workflow.query + def peek(self) -> str: + workflow.stream_writer(DECISIONS).publish({"from": "query"}) + return "unreachable" + + +async def test_a_publish_from_a_query_handler_is_refused_at_the_call( + client: Client, provider: MemoryStreams +): + # A query is answered on a task that carries the answer alone, so the record + # could only ever be dropped on the way out. The call says so instead. + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, PublishFromQuery, plugins=[provider]) as worker: + handle = await client.start_workflow( + PublishFromQuery.run, id=workflow_id, task_queue=worker.task_queue + ) + try: + with pytest.raises(Exception) as caught: + await handle.query(PublishFromQuery.peek) + assert "publish to a stream" in str(caught.value) + + stream = provider.get_stream_handle(client, workflow_id) + assert await stream.latest(topic=DECISIONS) == BEGINNING + finally: + await handle.terminate() + + +@workflow.defn +class NewestTwo: + """Starts at the newest two records on ``inputs`` and returns their values.""" + + @workflow.run + async def run(self) -> list[Any]: + reader = workflow.stream_reader(INPUTS, last=2) + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == 2: + reader.close() + return values + + +async def test_a_workflow_reader_starts_at_the_last_n_records( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, NewestTwo, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + await stream.producer(topic=INPUTS, producer_id="tool", attempt=1).append( + {"n": 1}, {"n": 2}, {"n": 3}, {"n": 4} + ) + handle = await client.start_workflow( + NewestTwo.run, id=workflow_id, task_queue=worker.task_queue + ) + assert await handle.result() == [3, 4] + + +@workflow.defn +class FromNow: + """Follows ``inputs`` from when it subscribes and returns the first value.""" + + @workflow.run + async def run(self) -> Any: + reader = workflow.stream_reader(INPUTS, after=END) + async for value in reader.values(): + reader.close() + return value["n"] + return None + + +async def test_a_workflow_reader_at_end_skips_what_was_there( + client: Client, provider: MemoryStreams +): + workflow_id = f"streams-wf-{uuid.uuid4().hex}" + async with new_worker(client, FromNow, plugins=[provider]) as worker: + stream = provider.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "old"}) + handle = await client.start_workflow( + FromNow.run, id=workflow_id, task_queue=worker.task_queue + ) + result = asyncio.ensure_future(handle.result()) + # The subscription starts when the workflow runs, which the test does + # not observe, so appends keep coming until the workflow takes one. + for _ in range(150): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({result}, timeout=0.2) + if done: + break + assert await asyncio.wait_for(result, 10) == "new" + + +def test_a_workflow_reader_start_names_one_place(): + # Checked before the reader needs a running workflow, so a mistake says + # what it is rather than that there is no workflow. + with pytest.raises(ValueError, match="either after= or last="): + workflow.stream_reader(INPUTS, after=END, last=1) + with pytest.raises(ValueError, match="positive"): + workflow.stream_reader(INPUTS, last=0)