From c27bfcf00e781b1ce5f79bdadf2200fdf22d5a3a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:11:59 -0700 Subject: [PATCH 01/26] Added the temporalio.streams interface on the external base. One stream interface a workflow reads, decides on and writes, with the provider registered once as a plugin on the client and every context asking for its stream the same way. This is the same surface as the interface chain, placed on Max's external base so the Redis provider above it has something to implement. --- .gitignore | 3 + CHANGELOG.md | 13 + streams_demo/agent_loop.py | 110 ++++++ streams_demo/provider_setup.py | 36 ++ streams_demo/run_demo.py | 172 ++++++++ temporalio/activity.py | 48 +++ temporalio/client/_client.py | 42 ++ temporalio/streams/__init__.py | 108 ++++++ temporalio/streams/_errors.py | 40 ++ temporalio/streams/_ids.py | 26 ++ temporalio/streams/_policy.py | 66 ++++ temporalio/streams/_provider.py | 285 ++++++++++++++ temporalio/streams/_record.py | 115 ++++++ temporalio/streams/_topic.py | 92 +++++ temporalio/streams/_wire.py | 181 +++++++++ temporalio/streams/providers/__init__.py | 87 +++++ temporalio/streams/providers/memory.py | 453 ++++++++++++++++++++++ temporalio/worker/_activity.py | 6 + temporalio/worker/_replayer.py | 7 +- temporalio/worker/_worker.py | 15 + temporalio/worker/_workflow.py | 46 +++ temporalio/worker/_workflow_instance.py | 19 + temporalio/workflow/__init__.py | 10 + temporalio/workflow/_context.py | 4 + temporalio/workflow/_streams.py | 279 +++++++++++++ tests/streams/__init__.py | 0 tests/streams/conftest.py | 8 + tests/streams/test_stream_accessors.py | 122 ++++++ tests/streams/test_stream_hooks.py | 139 +++++++ tests/streams/test_streams_conformance.py | 416 ++++++++++++++++++++ tests/streams/test_streams_workflow.py | 443 +++++++++++++++++++++ 31 files changed, 3390 insertions(+), 1 deletion(-) create mode 100644 streams_demo/agent_loop.py create mode 100644 streams_demo/provider_setup.py create mode 100644 streams_demo/run_demo.py create mode 100644 temporalio/streams/__init__.py create mode 100644 temporalio/streams/_errors.py create mode 100644 temporalio/streams/_ids.py create mode 100644 temporalio/streams/_policy.py create mode 100644 temporalio/streams/_provider.py create mode 100644 temporalio/streams/_record.py create mode 100644 temporalio/streams/_topic.py create mode 100644 temporalio/streams/_wire.py create mode 100644 temporalio/streams/providers/__init__.py create mode 100644 temporalio/streams/providers/memory.py create mode 100644 temporalio/workflow/_streams.py create mode 100644 tests/streams/__init__.py create mode 100644 tests/streams/conftest.py create mode 100644 tests/streams/test_stream_accessors.py create mode 100644 tests/streams/test_stream_hooks.py create mode 100644 tests/streams/test_streams_conformance.py create mode 100644 tests/streams/test_streams_workflow.py 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..6bbd734e5 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -35,6 +35,19 @@ 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. + `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..271f3e9e7 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -29,6 +29,7 @@ import temporalio.bridge.proto.activity_task import temporalio.common import temporalio.converter +import temporalio.streams from temporalio.converter._payload_converter import ( _TemporalTransferTypePayloadConverter, ) @@ -209,6 +210,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 +300,52 @@ def client() -> Client: return client +def stream_handle( + workflow_id: str | None = None, *, run_id: str | None = None +) -> temporalio.streams.StreamHandle: + """Return a stream handle from the provider the worker was given. + + With no arguments the handle is on this activity's own workflow, pinned to + the run the activity belongs to, so a producer opened from it writes onto + that run's stream and a read follows that run. Name a ``workflow_id`` to + address another workflow; ``run_id`` then pins the handle to one run and + its absence follows the execution chain. See :py:mod:`temporalio.streams`. + + Like :py:func:`client`, this is only available in ``async def`` + activities. + + Returns: + :py:class:`temporalio.streams.StreamHandle` for use in the current + activity. + + Raises: + temporalio.streams.StreamUnsupportedError: The worker has no stream + provider. Register one with ``Client.connect(plugins=[provider])`` + or ``Worker(plugins=[provider])``. + RuntimeError: When the client is not available, or when the activity + has no workflow and no ``workflow_id`` was given. + ValueError: ``run_id`` was given without ``workflow_id``. + """ + 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 workflow_id is None: + if run_id is not None: + raise ValueError("run_id needs a workflow_id") + info = context.info() + if info.workflow_id is None: + raise RuntimeError( + "this activity belongs to no workflow, so name the workflow_id to " + "address" + ) + workflow_id, run_id = info.workflow_id, info.workflow_run_id + return provider.get_stream_handle(client(), workflow_id, run_id=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 5acdfe476..79617dff4 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, @@ -157,6 +158,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. @@ -222,6 +224,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, @@ -258,6 +265,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__( @@ -271,6 +279,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. @@ -285,6 +294,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() @@ -897,6 +907,36 @@ def get_workflow_handle( result_type=result_type, ) + def get_stream_handle( + self, workflow_id: str, *, run_id: str | None = None + ) -> temporalio.streams.StreamHandle: + """Get a handle on a workflow's 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. The provider is the one registered with + ``plugins=[provider]`` at :py:meth:`connect`, or passed as + ``stream_provider``. See :py:mod:`temporalio.streams`. + + Args: + workflow_id: Workflow ID whose stream to get a handle to. + run_id: Run ID to pin the handle to. + + Returns: + The stream handle. + + Raises: + temporalio.streams.StreamUnsupportedError: No stream provider is + registered on this client. + """ + 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]" + ) + return provider.get_stream_handle(self, workflow_id, run_id=run_id) + def get_workflow_handle_for( self, workflow: ( @@ -3056,6 +3096,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): @@ -3070,3 +3111,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..b23c95be4 --- /dev/null +++ b/temporalio/streams/__init__.py @@ -0,0 +1,108 @@ +"""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. +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`, which is its own workflow pinned +to its run unless told otherwise. Any process holding a client uses +:meth:`temporalio.client.Client.get_stream_handle`, which mirrors +``get_workflow_handle``. The explicit form, +``provider.get_stream_handle(client, workflow_id)``, stays for a process that +talks to two stores. This module keeps the shared types, the errors and the +protocols a provider implements; nothing here that workflow code imports does +I/O. + +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. +""" + +from __future__ import annotations + +from temporalio.streams._errors import ( + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, + StreamUnsupportedError, +) +from temporalio.streams._provider import ( + ReadSource, + StreamHandle, + StreamProducer, + StreamProvider, + WorkflowStreamProvider, + WriteSink, +) +from temporalio.streams._record import ( + BEGINNING, + Cursor, + RecordKind, + StreamRecord, + Supersession, +) +from temporalio.streams._topic import StreamTopic, resolve_topic, topic + +__all__ = [ + "BEGINNING", + "Cursor", + "ReadSource", + "RecordKind", + "StreamCursorError", + "StreamError", + "StreamHandle", + "StreamNotFoundError", + "StreamProducer", + "StreamProducerError", + "StreamProvider", + "StreamRecord", + "StreamTopic", + "StreamUnsupportedError", + "Supersession", + "WorkflowStreamProvider", + "WriteSink", + "resolve_topic", + "topic", +] diff --git a/temporalio/streams/_errors.py b/temporalio/streams/_errors.py new file mode 100644 index 000000000..f00bf1c3a --- /dev/null +++ b/temporalio/streams/_errors.py @@ -0,0 +1,40 @@ +"""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__ = [ + "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 StreamUnsupportedError(StreamError): + """This provider does not offer the requested capability.""" diff --git a/temporalio/streams/_ids.py b/temporalio/streams/_ids.py new file mode 100644 index 000000000..2a76dcf19 --- /dev/null +++ b/temporalio/streams/_ids.py @@ -0,0 +1,26 @@ +"""The store key a provider derives from a workflow id and a topic. + +A workflow id may contain any character, ``:`` included, so joining the pair +with a bare ``:`` is ambiguous: ``("a:b", "c")`` and ``("a", "b:c")`` would +land in one store. Every provider that keys a store by the pair goes through +:func:`topic_key`, so they all agree and none of them collides. +""" + +from __future__ import annotations + +__all__ = ["topic_key"] + + +def _escape(component: str) -> str: + # Percent first, so an escaped component cannot be mistaken for one that + # already contained the escape. + return component.replace("%", "%25").replace(":", "%3A") + + +def topic_key(workflow_id: str, topic: str) -> str: + """The store key for ``topic`` of ``workflow_id``'s stream. + + Both components are percent-encoded before joining, so the only bare ``:`` + in the result is the separator. + """ + return f"{_escape(workflow_id)}:{_escape(topic)}" diff --git a/temporalio/streams/_policy.py b/temporalio/streams/_policy.py new file mode 100644 index 000000000..81949e4fc --- /dev/null +++ b/temporalio/streams/_policy.py @@ -0,0 +1,66 @@ +"""Turning a producer's newer attempt into something a reader can act on. + +An activity that streams half an answer and then fails leaves those records in +the stream. Its retry calls the model again and writes different words. No +provider can undo the first half, and a workflow that already acted on it has +committed that decision, so the honest thing is to tell the reader that a new +attempt began and let the application decide. + +This runs in the reader over records it already observed, so it costs no round +trip and replays without the provider being involved. +""" + +from __future__ import annotations + +from typing import Any + +from temporalio.streams._record import ( + Cursor, + RecordKind, + StreamRecord, + Supersession, +) + +__all__ = ["AttemptTracker"] + + +class AttemptTracker: + """Watches producer attempts on one subscription.""" + + def __init__(self) -> None: + """Start with no producer seen.""" + self._attempts: dict[str, int] = {} + + def note( + self, producer_id: str, attempt: int, *, topic: str, previous: Cursor + ) -> StreamRecord[Any] | None: + """A supersession record when this record starts a newer attempt. + + ``previous`` is the cursor of the last record delivered before the one + being noted, or the cursor the read started from. The synthesized + record carries it, so a consumer that checkpoints the supersession and + resumes after it is handed the new attempt's first record next rather + than skipping it. + + A producer that declares no attempt supersedes nothing, because there + is no generation to compare. That is the same answer as an unnumbered + record: the interface reports what it was told and invents nothing. + """ + if not producer_id or attempt <= 0: + return None + seen = self._attempts.get(producer_id, 0) + if attempt <= seen: + return None + self._attempts[producer_id] = attempt + if seen == 0: + return None + return StreamRecord( + kind=RecordKind.SUPERSEDED, + cursor=previous, + topic=topic, + producer_id=producer_id, + attempt=attempt, + supersession=Supersession( + producer_id=producer_id, previous_attempt=seen, attempt=attempt + ), + ) diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py new file mode 100644 index 000000000..8c142080d --- /dev/null +++ b/temporalio/streams/_provider.py @@ -0,0 +1,285 @@ +"""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`. +""" + +from __future__ import annotations + +from collections.abc import AsyncGenerator +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 + +__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 workflow's stream, addressed by topic, from outside workflow code. + + A handle follows the workflow's 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. + """ + + @overload + def read( + self, *, topic: StreamTopic[T], after: Cursor = ... + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, *, topic: str, after: Cursor = ..., result_type: type[T] + ) -> AsyncGenerator[StreamRecord[T], None]: ... + + @overload + def read( + self, *, topic: str, after: Cursor = ..., result_type: None = None + ) -> AsyncGenerator[StreamRecord[Any], None]: ... + + def read( + self, + *, + topic: str | StreamTopic[Any], + after: Cursor = BEGINNING, + 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. 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. 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, + or the topic is empty. + StreamCursorError: ``after`` came from another provider or names + a record no longer retained. Raised by this call, not by the + first iteration. + 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``. + """ + ... + + +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) -> ReadSource: + """Subscribe the running workflow to ``topic`` of its own stream. + + Raises: + StreamCursorError: ``after`` was minted by another provider. + """ + ... + + 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. + """ + + def workflow_provider(self) -> WorkflowStreamProvider: + """The half that serves one workflow instance on its thread.""" + ... + + def get_stream_handle( + self, client: Client, workflow_id: str, *, run_id: str | None = None + ) -> StreamHandle: + """A handle on ``workflow_id``'s stream. + + Without ``run_id`` it follows the execution chain, so a consumer keeps + reading across continue-as-new; with one it is pinned to that run. + """ + ... + + async def 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..fef63e51c --- /dev/null +++ b/temporalio/streams/_record.py @@ -0,0 +1,115 @@ +"""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", + "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.""" + + +@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/_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..52299f44e --- /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() + + 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..f175570a8 --- /dev/null +++ b/temporalio/streams/providers/memory.py @@ -0,0 +1,453 @@ +"""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. + +The outside surface (producer identity, retry deduplication, positions, +supersession, cursors) is faithful, which is what the conformance tests lean +on. One list per topic; a topic is written by the workflow and by outside +producers alike and read from either side. +""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import AsyncGenerator +from datetime import timedelta +from typing import Any, Generic, TypeVar + +from google.protobuf.message import DecodeError + +import temporalio.converter +from temporalio import workflow +from temporalio.client import Client, WorkflowExecutionStatus +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._errors import StreamCursorError +from temporalio.streams._ids import topic_key +from temporalio.streams._provider import ReadSource, WriteSink +from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import ( + RecordDecoder, + WireRecord, + cursor_position, + mint_cursor, + producer_identity, + to_wire, +) +from temporalio.streams.providers import ProviderPlugin + +__all__ = ["MemoryProducer", "MemoryStreamHandle", "MemoryStreams"] + +_PROVIDER = "memory" + +T = TypeVar("T") + +logger = logging.getLogger(__name__) + + +def _wake(future: asyncio.Future[None]) -> None: + if not future.done(): + future.set_result(None) + + +class _Topic: + """One topic's records, and the waiters parked on its tail.""" + + def __init__(self) -> None: + self.records: list[bytes] = [] + # Dedupe identity is (producer#attempt, first sequence of the append), + # the same pair the storage providers use, mapped to where the batch + # landed so a repeat can answer with the original position. + self.seen: dict[tuple[str, int], tuple[int, int]] = {} + # 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, + ) -> tuple[int, int]: + """Store ``wires`` and return where they landed as ``(first offset, count)``. + + With a ``writer``, a repeat of ``(writer, sequence)`` stores nothing + and returns where the original landed. + """ + key = (writer or "", sequence) + if writer is not None and key in self.seen: + return self.seen[key] + first = len(self.records) + self.records.extend(wire.SerializeToString() for wire in wires) + if writer is not None: + self.seen[key] = (first, len(wires)) + waiters, self._waiters = self._waiters, [] + for loop, future in waiters: + loop.call_soon_threadsafe(_wake, future) + return first, len(wires) + + async def wait_past(self, offset: int, timeout: float | None) -> None: + """Wait until a record exists at ``offset``, or ``timeout`` passes.""" + if len(self.records) > 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: + 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: + records = self._store.records + if len(records) > self._offset: + batch: list[tuple[Cursor, WireRecord]] = [] + for offset in range(self._offset, len(records)): + cursor = mint_cursor(_PROVIDER, str(offset)) + wire = _parse(cursor, records[offset], workflow.logger.warning) + if wire is not None: + batch.append((cursor, wire)) + self._offset = len(records) + 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) -> ReadSource: + start = self._streams._offset_after(after) + store = self._streams._topic(workflow.info().workflow_id, topic) + 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.PayloadConverter, + 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 self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.DATA, + value=value, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence + index, + ) + for index, value in enumerate(values) + ] + ) + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic.""" + self._write( + [ + to_wire( + self._converter, + topic=self._topic, + kind=RecordKind.FINISH, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + ) + ] + ) + + def _write(self, wires: list[WireRecord]) -> Cursor: + first, count = self._store.append( + wires, writer=self._writer, sequence=self._sequence + ) + self._sequence += len(wires) + self._last = mint_cursor(_PROVIDER, str(first + count - 1)) + return self._last + + +class MemoryStreamHandle: + """One workflow's stream from outside, with the shared reader rules.""" + + def __init__( + self, + streams: MemoryStreams, + client: Client | None, + workflow_id: str, + run_id: str | None, + ) -> None: + """Address ``workflow_id``'s topics in ``streams``.""" + self._streams = streams + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._converter = ( + client.data_converter.payload_converter + if client is not None + else temporalio.converter.DataConverter.default.payload_converter + ) + + def read( + self, + *, + topic: str | StreamTopic[Any], + after: Cursor = BEGINNING, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` after ``after`` until the workflow closes.""" + name, result_type = resolve_topic(topic, result_type) + store = self._streams._topic(self._workflow_id, name) + # Parsed here so a foreign cursor fails this call, not the first + # iteration of the generator. + start = self._streams._offset_after(after) + return self._read(store, start, after, result_type) + + async def _read( + self, + store: _Topic, + offset: int, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + decoder = RecordDecoder( + self._converter, result_type, after=after, warn=logger.warning + ) + closed = False + while True: + records = store.records + while offset < len(records): + cursor = mint_cursor(_PROVIDER, str(offset)) + wire = _parse(cursor, records[offset], logger.warning) + offset += 1 + if wire is None: + continue + for record in decoder.decode(cursor, wire): + yield record + if closed: + return + # One more pass after learning the workflow closed, so a record + # that landed between the scan and the describe is not lost. + closed = await self._closed() + if not closed: + await store.wait_past( + offset, + None + if self._client is None + else self._streams._poll.total_seconds(), + ) + + async def _closed(self) -> bool: + if self._client is None: + return False + handle = self._client.get_workflow_handle( + self._workflow_id, run_id=self._run_id + ) + try: + description = await handle.describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + # A producer may write before the workflow exists; there is + # nothing to follow yet, so keep waiting. + return False + raise + status = description.status + if status is None or status == WorkflowExecutionStatus.RUNNING: + return False + # Following the chain, a run that continued as new is not the end: + # the next describe without a run id finds its successor. + return not ( + self._run_id is None and status == WorkflowExecutionStatus.CONTINUED_AS_NEW + ) + + async def latest(self, *, topic: str | StreamTopic[Any]) -> Cursor: + """The cursor of the newest record on ``topic``, for following from now.""" + name, _ = resolve_topic(topic) + count = len(self._streams._topic(self._workflow_id, name).records) + return mint_cursor(_PROVIDER, str(count - 1)) if count 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._streams._topic(self._workflow_id, name) + producer_id, attempt = producer_identity(producer_id, attempt) + return MemoryProducer(store, self._converter, name, producer_id, attempt) + + +class MemoryStreams(ProviderPlugin): + """The in-memory provider, one list per topic. + + Construct one and pass the same instance to the worker and to the code + that opens handles; two instances share nothing. + """ + + def __init__( + self, *, poll_interval: timedelta = timedelta(milliseconds=100) + ) -> None: + """Create an empty provider. + + Args: + poll_interval: How often a workflow-side reader with nothing to + read checks again, and how often an outside reader asks + whether the workflow closed. + """ + self._poll = poll_interval + self._topics: dict[str, _Topic] = {} + + def reset(self) -> None: + """Drop every topic. For tests.""" + self._topics.clear() + + def workflow_provider(self) -> _MemoryWorkflowProvider: + """The workflow half, over this provider's topics.""" + return _MemoryWorkflowProvider(self) + + def get_stream_handle( + self, client: Client | None, workflow_id: str, *, run_id: str | None = None + ) -> MemoryStreamHandle: + """A handle on ``workflow_id``'s topics. + + ``client`` may be ``None`` here, unlike on a storage provider; then + the handle cannot see the workflow close and a read waits until the + caller closes it. + """ + return MemoryStreamHandle(self, client, workflow_id, run_id) + + async def 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 _offset_after(self, after: Cursor) -> int: + position = cursor_position(after, provider=_PROVIDER) + if position is None: + return 0 + try: + return int(position) + 1 + except ValueError: + raise StreamCursorError( + f"cursor {after.token!r} does not name a position on the memory provider" + ) from None 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 42fb07ff8..d646ebdac 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 @@ -158,6 +197,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")) @@ -216,6 +256,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 @@ -968,6 +1013,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 88e7a891b..992bf7ed9 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -58,7 +58,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__ @@ -198,6 +200,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): @@ -328,6 +331,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._primary_task: asyncio.Task[None] | None = None self._cancel_primary_task_pending = False @@ -2079,6 +2084,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 fa2681139..f63135bf4 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -147,6 +147,12 @@ logger, unsafe, ) +from ._streams import ( + StreamReader, + StreamWriter, + stream_reader, + stream_writer, +) from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowConfig, @@ -252,6 +258,10 @@ "SandboxImportNotificationPolicy", "logger", "unsafe", + "StreamReader", + "StreamWriter", + "stream_reader", + "stream_writer", "ChildWorkflowCancellationType", "ChildWorkflowConfig", "ChildWorkflowHandle", diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index b33f83150..25523836a 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -24,6 +24,7 @@ from ._activities import ActivityCancellationType, ActivityHandle from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent from ._nexus import NexusOperationCancellationType, NexusOperationHandle + from ._streams import _WorkflowStreams from ._workflow_ops import ( ChildWorkflowCancellationType, ChildWorkflowHandle, @@ -475,6 +476,9 @@ async def workflow_start_nexus_operation( summary: str | None, ) -> NexusOperationHandle[OutputT]: ... + @abstractmethod + def workflow_streams(self) -> _WorkflowStreams: ... + @abstractmethod def workflow_time_ns(self) -> int: ... diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py new file mode 100644 index 000000000..5d319d50a --- /dev/null +++ b/temporalio/workflow/_streams.py @@ -0,0 +1,279 @@ +"""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, Cursor, RecordKind, StreamRecord +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._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. + """ + if self._finished: + raise ValueError(f"topic {self._topic!r} was already finished") + 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 = ...) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, *, result_type: type[T], after: Cursor = ... +) -> StreamReader[T]: ... + + +@overload +def stream_reader( + topic: str, *, result_type: None = None, after: Cursor = ... +) -> StreamReader[Any]: ... + + +def stream_reader( + topic: str | StreamTopic[Any], + *, + result_type: type | None = None, + after: Cursor = BEGINNING, +) -> 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 neither ``after`` nor + a 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. Honoured on the first + subscription of a run, because after that the recorded + observations decide. + + Raises: + ValueError: ``topic`` is empty, ``result_type`` was passed with a + definition, 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. + """ + 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 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= or other type" + ) + return existing + source = state.provider.open_reader(name, after=after) + + def forget() -> None: + state.readers.pop(name, None) + + reader: StreamReader[Any] = StreamReader( + source, topic=name, result_type=result_type, after=after, 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..12663fa23 --- /dev/null +++ b/tests/streams/conftest.py @@ -0,0 +1,8 @@ +import pytest + + +def pytest_configure(config: pytest.Config) -> None: + config.addinivalue_line( + "markers", + "reports_positions: the case needs append() to return where records landed", + ) diff --git a/tests/streams/test_stream_accessors.py b/tests/streams/test_stream_accessors.py new file mode 100644 index 000000000..eb733ba5b --- /dev/null +++ b/tests/streams/test_stream_accessors.py @@ -0,0 +1,122 @@ +"""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 RecordKind, 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) 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..9002ea16a --- /dev/null +++ b/tests/streams/test_streams_conformance.py @@ -0,0 +1,416 @@ +"""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 the contract: the record on the wire, producer +identity, retry deduplication, positions, supersession, topic addressing, +cursor resumption, cursor ownership, and store keys that cannot collide. The +workflow-side handles and the two rules about Workflow Tasks live in +``test_streams_workflow``. +""" + +from __future__ import annotations + +import asyncio +import uuid +from collections.abc import AsyncIterator, Awaitable, Callable +from dataclasses import dataclass +from typing import Any + +import pytest + +from temporalio.api.common.v1 import Payload +from temporalio.client import Client +from temporalio.common import RawValue +from temporalio.converter import DataConverter +from temporalio.streams import ( + BEGINNING, + Cursor, + RecordKind, + StreamCursorError, + StreamHandle, + StreamProvider, + Supersession, + _ids, + _wire, + topic, +) +from temporalio.streams._policy import AttemptTracker +from temporalio.streams.providers.memory import MemoryStreams + +# Defined once and shared by every case, the way an application shares them +# between its workflow, its activities and its backend. +OUT = topic("out", dict) +A = topic("a", dict) +B = topic("b", dict) +Y = topic("y", dict) +XY = topic("x:y", dict) + + +@dataclass +class ProviderCase: + """One provider under test, and what the cases may ask of it.""" + + name: str + provider: StreamProvider + client: Client | None = None + reports_positions: bool = True + """``append()`` returns where the records landed.""" + host: Callable[[str], Awaitable[None]] | None = None + """Starts the workflow that owns ``workflow_id``'s stream, when a store needs one.""" + + async def open( + self, workflow_id: str, *, run_id: str | None = None + ) -> StreamHandle: + if self.host is not None: + await self.host(workflow_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 _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: + provider = MemoryStreams() + yield ProviderCase("memory", provider) + provider.reset() + + +SETUPS: dict[str, Callable[[Client], AsyncIterator[ProviderCase]]] = { + "memory": _memory_case +} + +_CAPABILITIES = { + "reports_positions": lambda case: case.reports_positions, +} + + +@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 + + +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") + + +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 diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py new file mode 100644 index 000000000..71f2bb8bf --- /dev/null +++ b/tests/streams/test_streams_workflow.py @@ -0,0 +1,443 @@ +"""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 ( + 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), + ] From dad0743ff7203c39bd2c2e8260ce6b1b9aedde6a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 08:12:04 -0700 Subject: [PATCH 02/26] Answered a legacy query without the stream snapshot beside it. A parked Run has no open Workflow Task, so the server dispatches a query on a task of its own and Core allows nothing beside the answer there. The instance still reported the registered wait set, Core refused the completion, and the query timed out. --- temporalio/worker/_workflow_instance.py | 14 ++++ .../test_worker_integration.py | 66 +++++++++++++++++++ 2 files changed, 80 insertions(+) diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index 992bf7ed9..64cca2ea4 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -87,6 +87,9 @@ # Set to true to log all cases where we're ignoring things during delete LOG_IGNORE_DURING_DELETE = False +# Core answers a query carrying this id on the query's own task, alone. +_LEGACY_QUERY_ID = "legacy_query" + def _is_workflow_terminal_command( command: temporalio.bridge.proto.workflow_commands.workflow_commands_pb2.WorkflowCommand, @@ -2643,6 +2646,17 @@ def _emit_external_stream_commands(self) -> None: if runtime is None or self._deleting: return + # A legacy query is answered on a task of its own, and Core refuses any + # other command beside that answer. The activation ran no Workflow code, + # so the wait set it would report is the one the retained task already + # holds, and there is no observation delta to commit. + if any( + command.HasField("respond_to_query") + and command.respond_to_query.query_id == _LEGACY_QUERY_ID + for command in self._current_completion.successful.commands + ): + return + # First, because the boundary a Continue-As-New header has to carry is # the one this activation is about to commit, and it only stops moving # here. diff --git a/tests/contrib/external_workflow_streams/test_worker_integration.py b/tests/contrib/external_workflow_streams/test_worker_integration.py index cf797719f..f9ad6d5d7 100644 --- a/tests/contrib/external_workflow_streams/test_worker_integration.py +++ b/tests/contrib/external_workflow_streams/test_worker_integration.py @@ -301,6 +301,72 @@ async def test_a_blocked_workflow_retains_its_task_rather_than_completing_it( await handle.terminate() +@workflow.defn +class ParkedQueryWorkflow: + """Blocks on an empty stream with a short idle timeout, so the task parks.""" + + def __init__(self) -> None: + self._seen: list[str] = [] + + @workflow.run + async def run(self) -> list[str]: + tokens = external_stream.with_options(idle_timeout=timedelta(seconds=1)).topic( + "tokens", type=str + ) + async for token in tokens.subscribe(): + self._seen.append(token) + break + return self._seen + + @workflow.query + def seen(self) -> list[str]: + return self._seen + + +async def test_a_query_against_a_parked_workflow_is_answered( + client: Client, backend: MemoryStreamBackend +) -> None: + """A parked Run has no open Workflow Task, so its query travels alone. + + The server dispatches it on a task of its own, and Core allows nothing + beside the answer on that task. The wait set is still registered from the + parked task, so reporting it again with the answer made Core refuse the + completion, and the query never returned. + """ + task_queue = f"tq-{uuid.uuid4()}" + async with Worker( + client, + task_queue=task_queue, + workflows=[ParkedQueryWorkflow], + external_stream_backend=backend, + ): + handle = await client.start_workflow( + ParkedQueryWorkflow.run, + id=f"wf-{uuid.uuid4()}", + task_queue=task_queue, + ) + # Past the idle timeout the retained task is reported and the Run + # waits for a wake with no task open. + await asyncio.sleep(3) + + assert await asyncio.wait_for(handle.query(ParkedQueryWorkflow.seen), 10) == [] + + description = await handle.describe() + key = StreamKey( + client.namespace, + handle.id, + description.raw_description.workflow_execution_info.first_run_id, + "tokens", + ) + await publish(backend, key, ["first"]) + assert await asyncio.wait_for(handle.result(), 60) == ["first"] + + events = [e async for e in handle.fetch_history_events()] + assert not any( + e.HasField("workflow_task_failed_event_attributes") for e in events + ), "a Workflow Task failed while the query was outstanding" + + async def test_a_workflow_without_a_configured_backend_says_so( client: Client, ) -> None: From 588b41c5a0dfd77a2d49c8bbfeaa0260e70dc6f6 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:49:34 -0700 Subject: [PATCH 03/26] Refused a stream publish from a query handler. 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. --- temporalio/workflow/_streams.py | 10 +++++++ tests/streams/test_streams_workflow.py | 36 ++++++++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index 5d319d50a..b99125ca3 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -22,6 +22,7 @@ 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"] @@ -157,9 +158,18 @@ def publish(self, value: T) -> None: 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(), diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py index 71f2bb8bf..05adf23db 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -25,6 +25,7 @@ from temporalio import workflow from temporalio.client import Client from temporalio.streams import ( + BEGINNING, Cursor, ReadSource, RecordKind, @@ -441,3 +442,38 @@ async def read_everything() -> list[Any]: ("", 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() From d69847fd4dbf9ae7879046bcc65adc467d148e2e Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 25 Sep 2026 14:49:34 -0700 Subject: [PATCH 04/26] Said so when a producer's attempt goes backwards. Attempts only rise on one producer, so a lower one means the store reordered two generations and the older half would render as the current answer. --- temporalio/streams/_policy.py | 17 +++++++++++++++-- temporalio/streams/_wire.py | 2 +- tests/streams/test_streams_conformance.py | 21 +++++++++++++++++++++ 3 files changed, 37 insertions(+), 3 deletions(-) diff --git a/temporalio/streams/_policy.py b/temporalio/streams/_policy.py index 81949e4fc..b27c839e5 100644 --- a/temporalio/streams/_policy.py +++ b/temporalio/streams/_policy.py @@ -12,6 +12,7 @@ from __future__ import annotations +from collections.abc import Callable from typing import Any from temporalio.streams._record import ( @@ -27,9 +28,10 @@ class AttemptTracker: """Watches producer attempts on one subscription.""" - def __init__(self) -> None: - """Start with no producer seen.""" + 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 @@ -45,10 +47,21 @@ def note( 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 diff --git a/temporalio/streams/_wire.py b/temporalio/streams/_wire.py index 52299f44e..c6eaa3f02 100644 --- a/temporalio/streams/_wire.py +++ b/temporalio/streams/_wire.py @@ -118,7 +118,7 @@ def __init__( self._result_type = result_type self._previous = after self._warn = warn - self._attempts = AttemptTracker() + self._attempts = AttemptTracker(warn) def decode(self, cursor: Cursor, wire: WireRecord) -> list[StreamRecord[Any]]: """The records to yield for one stored record, in order.""" diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 9002ea16a..a13d48f43 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -187,6 +187,27 @@ def test_supersession_is_synthesized_from_observations(): assert attempts.note("model", 2, topic="t", previous=Cursor("memory:1")) is None +def test_an_attempt_that_goes_backwards_is_said_rather_than_passed_off(): + # Attempts only rise on one producer, so a lower one means the store handed + # two generations back out of order. Yielded as data with no signal, a + # consumer renders the stale generation as the current answer. + said: list[str] = [] + attempts = AttemptTracker(said.append) + assert attempts.note("model", 2, topic="t", previous=BEGINNING) is None + assert attempts.note("model", 1, topic="t", previous=Cursor("memory:3")) is None + assert len(said) == 1 + assert "attempt 1" in said[0] and "behind attempt 2" in said[0] + assert "model" in said[0] + + +def test_a_repeat_of_the_current_attempt_is_not_worth_saying(): + said: list[str] = [] + attempts = AttemptTracker(said.append) + attempts.note("model", 1, topic="t", previous=BEGINNING) + assert attempts.note("model", 1, topic="t", previous=Cursor("memory:1")) is None + assert said == [] + + 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") From 375a352793b6fc6393a26a756a0b37376e02abe4 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 29 Sep 2026 16:57:55 -0700 Subject: [PATCH 05/26] Let activities own streams and resolve stream_handle() by a static rule. Mirrors the accessor layer of the interface chain on the external base, so the Redis provider above it can hold a stream an activity owns. A standalone activity now reaches its own stream, which used to raise, and scope="activity" gives a workflow's activity its own streams without changing what a bare call reaches. The rule never probes, because a stream is created by its first write and probing would split attempts across two streams. --- temporalio/activity.py | 85 ++++++++++++++++++----- temporalio/client/_client.py | 31 +++++++-- temporalio/streams/__init__.py | 9 ++- temporalio/streams/_provider.py | 32 ++++++++- temporalio/streams/providers/memory.py | 93 ++++++++++++++++++++++---- 5 files changed, 207 insertions(+), 43 deletions(-) diff --git a/temporalio/activity.py b/temporalio/activity.py index 271f3e9e7..701a4f3c6 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -20,6 +20,7 @@ from typing import ( TYPE_CHECKING, Any, + Literal, NoReturn, overload, ) @@ -301,30 +302,58 @@ def client() -> Client: def stream_handle( - workflow_id: str | None = None, *, run_id: str | None = None + workflow_id: str | 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. - With no arguments the handle is on this activity's own workflow, pinned to - the run the activity belongs to, so a producer opened from it writes onto - that run's stream and a read follows that run. Name a ``workflow_id`` to - address another workflow; ``run_id`` then pins the handle to one run and - its absence follows the execution chain. See :py:mod:`temporalio.streams`. + 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. 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. + 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. Register one with ``Client.connect(plugins=[provider])`` - or ``Worker(plugins=[provider])``. - RuntimeError: When the client is not available, or when the activity - has no workflow and no ``workflow_id`` was given. - ValueError: ``run_id`` was given without ``workflow_id``. + 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``, or + ``scope="activity"`` with one. """ context = _Context.current() provider = context.stream_provider @@ -333,17 +362,37 @@ def stream_handle( "no stream provider is configured on this worker; register one with " "Client.connect(plugins=[provider]) or Worker(plugins=[provider])" ) - if workflow_id is None: - if run_id is not None: - raise ValueError("run_id needs a workflow_id") - info = context.info() + 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" + "address, or leave scope unset for the activity's own streams" ) - workflow_id, run_id = info.workflow_id, info.workflow_run_id - return provider.get_stream_handle(client(), workflow_id, run_id=run_id) + 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: diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 79617dff4..ed3a90500 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -908,26 +908,37 @@ def get_workflow_handle( ) def get_stream_handle( - self, workflow_id: str, *, run_id: str | None = None + self, + workflow_id: str | None = None, + *, + run_id: str | None = None, + activity_id: str | None = None, ) -> temporalio.streams.StreamHandle: - """Get a handle on a workflow's stream from the provider registered on this client. + """Get a handle on a workflow's or an activity's 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. The provider is the one registered with - ``plugins=[provider]`` at :py:meth:`connect`, or passed as - ``stream_provider``. See :py:mod:`temporalio.streams`. + 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. The provider is + the one registered with ``plugins=[provider]`` at :py:meth:`connect`, + or passed as ``stream_provider``. See :py:mod:`temporalio.streams`. Args: - workflow_id: Workflow ID whose stream to get a handle to. + workflow_id: Workflow ID whose stream to get a handle to, or the + workflow that scheduled ``activity_id``. run_id: Run ID to pin the handle to. + activity_id: Activity ID whose own streams to get a handle to. Returns: The stream handle. Raises: + ValueError: Neither ``workflow_id`` nor ``activity_id`` was given. temporalio.streams.StreamUnsupportedError: No stream provider is - registered on this client. + registered on this client, or it cannot hold a stream an + activity owns. """ provider = self._config.get("stream_provider") if provider is None: @@ -935,6 +946,12 @@ def get_stream_handle( "no stream provider is registered on this client; connect with " "plugins=[provider]" ) + 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 or the activity_id to address") return provider.get_stream_handle(self, workflow_id, run_id=run_id) def get_workflow_handle_for( diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index b23c95be4..e4f9d1691 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -39,10 +39,13 @@ 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`, which is its own workflow pinned -to its run unless told otherwise. Any process holding a client 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``. The explicit form, +``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 diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py index 8c142080d..621915f99 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -85,10 +85,11 @@ async def finish(self) -> None: class StreamHandle(Protocol): - """One workflow's stream, addressed by topic, from outside workflow code. + """One owner's stream, addressed by topic, from outside workflow code. - A handle follows the workflow's execution chain unless it was opened with - a ``run_id``, in which case it is pinned to that run. A topic is a + The owner is a workflow or an activity. 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 @@ -275,6 +276,31 @@ def get_stream_handle( """ ... + 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 close(self) -> None: """Release what this provider holds for the process. diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index f175570a8..e3392e268 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -15,6 +15,12 @@ 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. The outside surface (producer identity, retry deduplication, positions, supersession, cursors) is faithful, which is what the conformance tests lean @@ -34,7 +40,7 @@ import temporalio.converter from temporalio import workflow -from temporalio.client import Client, WorkflowExecutionStatus +from temporalio.client import ActivityExecutionStatus, Client, WorkflowExecutionStatus from temporalio.service import RPCError, RPCStatusCode from temporalio.streams._errors import StreamCursorError from temporalio.streams._ids import topic_key @@ -278,20 +284,28 @@ def _write(self, wires: list[WireRecord]) -> Cursor: class MemoryStreamHandle: - """One workflow's stream from outside, with the shared reader rules.""" + """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. + """ def __init__( self, streams: MemoryStreams, client: Client | None, - workflow_id: str, + workflow_id: str | None, run_id: str | None, + activity_id: str | None = None, ) -> None: - """Address ``workflow_id``'s topics in ``streams``.""" + """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._seen_pending = False self._converter = ( client.data_converter.payload_converter if client is not None @@ -307,7 +321,7 @@ def read( ) -> AsyncGenerator[StreamRecord[Any], None]: """Yield records on ``topic`` after ``after`` until the workflow closes.""" name, result_type = resolve_topic(topic, result_type) - store = self._streams._topic(self._workflow_id, name) + store = self._store(name) # Parsed here so a foreign cursor fails this call, not the first # iteration of the generator. start = self._streams._offset_after(after) @@ -347,21 +361,46 @@ async def _read( else self._streams._poll.total_seconds(), ) + def _store(self, topic: str) -> _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._client is None: return False - handle = self._client.get_workflow_handle( - self._workflow_id, run_id=self._run_id - ) try: - description = await handle.describe() + 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 workflow exists; there is + # 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: @@ -373,7 +412,7 @@ async def _closed(self) -> bool: 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) - count = len(self._streams._topic(self._workflow_id, name).records) + count = len(self._store(name).records) return mint_cursor(_PROVIDER, str(count - 1)) if count else BEGINNING def producer( @@ -385,7 +424,7 @@ def producer( ) -> MemoryProducer[Any]: """A producer on ``topic``; inside an activity its identity is the activity's.""" name, _ = resolve_topic(topic) - store = self._streams._topic(self._workflow_id, name) + store = self._store(name) producer_id, attempt = producer_identity(producer_id, attempt) return MemoryProducer(store, self._converter, name, producer_id, attempt) @@ -409,10 +448,12 @@ def __init__( """ self._poll = poll_interval self._topics: dict[str, _Topic] = {} + self._activity_topics: dict[tuple[str | None, str, str], _Topic] = {} def reset(self) -> None: """Drop every topic. For tests.""" self._topics.clear() + self._activity_topics.clear() def workflow_provider(self) -> _MemoryWorkflowProvider: """The workflow half, over this provider's topics.""" @@ -429,6 +470,21 @@ def get_stream_handle( """ 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 close(self) -> None: """Nothing to release: the provider holds no connection.""" @@ -441,6 +497,19 @@ def _topic(self, workflow_id: str, topic: str) -> _Topic: 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 _offset_after(self, after: Cursor) -> int: position = cursor_position(after, provider=_PROVIDER) if position is None: From 350457282a2a5bbf547b9e79dbedc4b905325bcf Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Tue, 29 Sep 2026 16:57:55 -0700 Subject: [PATCH 06/26] Tested the activity stream accessor rule and retry supersession. Covers all three arms of the rule, a retry writing to the same stream, and the misaddressed calls, on the memory provider. A storage provider adds itself to the setups behind its own STREAMS_LIVE gate. --- tests/streams/test_activity_streams.py | 261 +++++++++++++++++++++++++ 1 file changed, 261 insertions(+) create mode 100644 tests/streams/test_activity_streams.py diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py new file mode 100644 index 000000000..8aa4a653a --- /dev/null +++ b/tests/streams/test_activity_streams.py @@ -0,0 +1,261 @@ +"""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 or the activity_id"): + setup.client.get_stream_handle() From 60e8e036498d1af21e3866b99372a62ce34bbe72 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:52:11 -0700 Subject: [PATCH 07/26] Mapped BEGINNING to the oldest record held and added END and last=N. BEGINNING is documented as the oldest record a stream still retains, which on a truncated stream is not offset zero. END follows the tail from when a read starts and last=N starts at the newest records; after= stays the only way to resume. (cherry picked from commit ec46a74e412902ee9e2f7fc1d31dbdfe3d00f1f8) --- temporalio/streams/__init__.py | 6 +- temporalio/streams/_provider.py | 52 +++++++++--- temporalio/streams/_record.py | 31 +++++++ temporalio/streams/providers/memory.py | 108 ++++++++++++++++++++----- 4 files changed, 166 insertions(+), 31 deletions(-) diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index e4f9d1691..6b6b61cfc 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -24,7 +24,9 @@ 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. + 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 @@ -81,6 +83,7 @@ ) from temporalio.streams._record import ( BEGINNING, + END, Cursor, RecordKind, StreamRecord, @@ -90,6 +93,7 @@ __all__ = [ "BEGINNING", + "END", "Cursor", "ReadSource", "RecordKind", diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py index 621915f99..378791bc8 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -99,17 +99,31 @@ class StreamHandle(Protocol): @overload def read( - self, *, topic: StreamTopic[T], after: Cursor = ... + self, + *, + topic: StreamTopic[T], + after: Cursor = ..., + last: int | None = None, ) -> AsyncGenerator[StreamRecord[T], None]: ... @overload def read( - self, *, topic: str, after: Cursor = ..., result_type: type[T] + 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 = ..., result_type: None = None + self, + *, + topic: str, + after: Cursor = ..., + last: int | None = None, + result_type: None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: ... def read( @@ -117,14 +131,20 @@ def read( *, 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. 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. The read ends when the owning + ``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 @@ -132,10 +152,14 @@ def read( Raises: ValueError: ``result_type`` was passed with a topic definition, - or the topic is empty. + 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. """ @@ -223,11 +247,21 @@ class WorkflowStreamProvider(Protocol): the definitions are resolved before it is called. """ - def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: + 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. """ ... diff --git a/temporalio/streams/_record.py b/temporalio/streams/_record.py index fef63e51c..ed499a5d8 100644 --- a/temporalio/streams/_record.py +++ b/temporalio/streams/_record.py @@ -12,6 +12,7 @@ __all__ = [ "BEGINNING", + "END", "Cursor", "RecordKind", "StreamRecord", @@ -80,6 +81,36 @@ def __str__(self) -> str: 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: diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index e3392e268..7040ea646 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -21,6 +21,8 @@ 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. The outside surface (producer identity, retry deduplication, positions, supersession, cursors) is faithful, which is what the conformance tests lean @@ -45,7 +47,14 @@ from temporalio.streams._errors import StreamCursorError from temporalio.streams._ids import topic_key from temporalio.streams._provider import ReadSource, WriteSink -from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +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, @@ -75,6 +84,10 @@ class _Topic: """One topic's records, and the waiters parked on its tail.""" def __init__(self) -> None: + # The retained records, the first of which sits at offset ``base``. + # Offsets are never reused, so a cursor keeps naming the same record + # after truncation drops the ones before it. + self.base = 0 self.records: list[bytes] = [] # Dedupe identity is (producer#attempt, first sequence of the append), # the same pair the storage providers use, mapped to where the batch @@ -101,7 +114,7 @@ def append( key = (writer or "", sequence) if writer is not None and key in self.seen: return self.seen[key] - first = len(self.records) + first = self.head self.records.extend(wire.SerializeToString() for wire in wires) if writer is not None: self.seen[key] = (first, len(wires)) @@ -110,9 +123,24 @@ def append( loop.call_soon_threadsafe(_wake, future) return first, len(wires) + @property + def head(self) -> int: + """The offset the next record lands at.""" + return self.base + len(self.records) + + def at(self, offset: int) -> bytes: + """The retained record at ``offset``.""" + return self.records[offset - self.base] + + def truncate(self, keep: int) -> None: + """Drop all but the newest ``keep`` records.""" + drop = max(0, len(self.records) - keep) + self.base += drop + del self.records[:drop] + async def wait_past(self, offset: int, timeout: float | None) -> None: """Wait until a record exists at ``offset``, or ``timeout`` passes.""" - if len(self.records) > offset: + if self.head > offset: return loop = asyncio.get_running_loop() future: asyncio.Future[None] = loop.create_future() @@ -148,15 +176,17 @@ def __init__(self, store: _Topic, start: int, poll: timedelta) -> None: async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: while not self._closed: - records = self._store.records - if len(records) > self._offset: + head = self._store.head + if head > self._offset: batch: list[tuple[Cursor, WireRecord]] = [] - for offset in range(self._offset, len(records)): + for offset in range(max(self._offset, self._store.base), head): cursor = mint_cursor(_PROVIDER, str(offset)) - wire = _parse(cursor, records[offset], workflow.logger.warning) + wire = _parse( + cursor, self._store.at(offset), workflow.logger.warning + ) if wire is not None: batch.append((cursor, wire)) - self._offset = len(records) + self._offset = head if batch: return batch continue @@ -183,9 +213,12 @@ class _MemoryWorkflowProvider: def __init__(self, streams: MemoryStreams) -> None: self._streams = streams - def open_reader(self, topic: str, *, after: Cursor) -> ReadSource: - start = self._streams._offset_after(after) + 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: @@ -317,15 +350,24 @@ def read( *, topic: str | StreamTopic[Any], after: Cursor = BEGINNING, + last: int | None = None, result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: - """Yield records on ``topic`` after ``after`` until the workflow closes.""" + """Yield records on ``topic`` from where the read starts until the workflow closes. + + ``END`` and ``last=`` are resolved by this call, against what the + topic holds when it is made. + """ + check_read_start(after, last) name, result_type = resolve_topic(topic, result_type) store = self._store(name) # Parsed here so a foreign cursor fails this call, not the first # iteration of the generator. - start = self._streams._offset_after(after) - return self._read(store, start, after, result_type) + 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, @@ -339,10 +381,14 @@ async def _read( ) closed = False while True: - records = store.records - while offset < len(records): + 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, records[offset], logger.warning) + wire = _parse(cursor, store.at(offset), logger.warning) offset += 1 if wire is None: continue @@ -412,8 +458,8 @@ async def _closed(self) -> bool: 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) - count = len(self._store(name).records) - return mint_cursor(_PROVIDER, str(count - 1)) if count else BEGINNING + head = self._store(name).head + return mint_cursor(_PROVIDER, str(head - 1)) if head else BEGINNING def producer( self, @@ -455,6 +501,15 @@ def reset(self) -> None: self._topics.clear() self._activity_topics.clear() + def truncate(self, workflow_id: str, topic: str, *, keep: int) -> None: + """Drop all but the newest ``keep`` records of a topic. For tests. + + Stands in for a store's retention: offsets are kept, so a cursor from + before still names its record, and a read from ``BEGINNING`` starts + at the oldest one left. + """ + self._topic(workflow_id, topic).truncate(keep) + def workflow_provider(self) -> _MemoryWorkflowProvider: """The workflow half, over this provider's topics.""" return _MemoryWorkflowProvider(self) @@ -510,13 +565,24 @@ def _activity_topic( found = self._activity_topics[key] = _Topic() return found - def _offset_after(self, after: Cursor) -> int: + 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 0 + return store.base try: - return int(position) + 1 + start = int(position) + 1 except ValueError: raise StreamCursorError( f"cursor {after.token!r} does not name a position on the memory provider" ) from None + if start < store.base: + raise StreamCursorError( + f"cursor {after.token!r} names a record no longer retained; the " + f"topic starts at offset {store.base}" + ) + return start From 9cc17dc7bdf75dde1bec2f39c492c74768b65467 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:52:11 -0700 Subject: [PATCH 08/26] Added conformance cases for each read start, on a truncated topic too. (cherry picked from commit 2e4298439c6bf3e5e036d889061537771a592875) --- tests/streams/conftest.py | 4 + tests/streams/test_streams_conformance.py | 97 ++++++++++++++++++++++- 2 files changed, 100 insertions(+), 1 deletion(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 12663fa23..7be000f3b 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -6,3 +6,7 @@ def pytest_configure(config: pytest.Config) -> None: "markers", "reports_positions: the case needs append() to return where records landed", ) + config.addinivalue_line( + "markers", + "truncates: the case needs a way to drop a topic's oldest records", + ) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index a13d48f43..9aa456676 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -34,6 +34,7 @@ from temporalio.converter import DataConverter from temporalio.streams import ( BEGINNING, + END, Cursor, RecordKind, StreamCursorError, @@ -67,6 +68,9 @@ class ProviderCase: """``append()`` returns where the records landed.""" 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.""" async def open( self, workflow_id: str, *, run_id: str | None = None @@ -87,7 +91,11 @@ async def open( async def _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: provider = MemoryStreams() - yield ProviderCase("memory", provider) + + async def truncate(workflow_id: str, topic: str, keep: int) -> None: + provider.truncate(workflow_id, topic, keep=keep) + + yield ProviderCase("memory", provider, truncate=truncate) provider.reset() @@ -97,6 +105,7 @@ async def _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, + "truncates": lambda case: case.truncate is not None, } @@ -435,3 +444,89 @@ async def test_a_definition_carries_its_type_once(case: ProviderCase): 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) From b81e932cc885169106de65bb0ed1d211d94f2612 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:55:12 -0700 Subject: [PATCH 09/26] Let stream_reader() start at END or at the last N records. A workflow could only start from BEGINNING or a cursor, so it had no way to follow from now. The provider resolves the start outside the workflow and records it, so replay reproduces it; last= is passed only when given, which keeps older providers working. (cherry picked from commit a3edfbbfbc80261e7b2b7bfb414c41902a02417b) --- temporalio/workflow/_streams.py | 72 +++++++++++++++++++------ tests/streams/test_streams_workflow.py | 75 ++++++++++++++++++++++++++ 2 files changed, 132 insertions(+), 15 deletions(-) diff --git a/temporalio/workflow/_streams.py b/temporalio/workflow/_streams.py index b99125ca3..550398869 100644 --- a/temporalio/workflow/_streams.py +++ b/temporalio/workflow/_streams.py @@ -18,7 +18,14 @@ from typing import Any, Generic, TypeVar, cast, overload from temporalio.streams._provider import ReadSource, WorkflowStreamProvider, WriteSink -from temporalio.streams._record import BEGINNING, Cursor, RecordKind, StreamRecord +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 @@ -194,18 +201,28 @@ def finish(self) -> None: @overload -def stream_reader(topic: StreamTopic[T], *, after: Cursor = ...) -> StreamReader[T]: ... +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 = ... + 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 = ... + topic: str, + *, + result_type: None = None, + after: Cursor = ..., + last: int | None = None, ) -> StreamReader[Any]: ... @@ -214,6 +231,7 @@ def stream_reader( *, result_type: type | None = None, after: Cursor = BEGINNING, + last: int | None = None, ) -> StreamReader[Any]: """Subscribe this workflow to ``topic`` of its own stream. @@ -221,8 +239,8 @@ def stream_reader( 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 neither ``after`` nor - a different type. Adding a reader on a new topic is a new command, so + 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. @@ -232,34 +250,58 @@ def stream_reader( 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. Honoured on the first - subscription of a run, because after that the recorded - observations decide. + 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, or a reader on the topic is already open and this - call asked for a different position or type. + 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 result_type is not existing._result_type: + 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= or other type" + "stream_reader shares it and takes no after=, last= or other type" ) return existing - source = state.provider.open_reader(name, after=after) + # 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=after, on_close=forget + source, topic=name, result_type=result_type, after=previous, on_close=forget ) state.readers[name] = reader return reader diff --git a/tests/streams/test_streams_workflow.py b/tests/streams/test_streams_workflow.py index 05adf23db..03b65b0da 100644 --- a/tests/streams/test_streams_workflow.py +++ b/tests/streams/test_streams_workflow.py @@ -26,6 +26,7 @@ from temporalio.client import Client from temporalio.streams import ( BEGINNING, + END, Cursor, ReadSource, RecordKind, @@ -477,3 +478,77 @@ async def test_a_publish_from_a_query_handler_is_refused_at_the_call( 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) From 24452e3aa7b2b7856c3b5814eb77ae949d82224d Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Mon, 28 Sep 2026 16:56:25 -0700 Subject: [PATCH 10/26] Documented the read starts on both accessors and tested an activity stream. (cherry picked from commit 06cfa0182b86bf9ebe4c8f63014ede3cce1de48e) --- temporalio/activity.py | 6 ++++-- temporalio/client/_client.py | 6 +++++- tests/streams/test_stream_accessors.py | 29 +++++++++++++++++++++++++- 3 files changed, 37 insertions(+), 4 deletions(-) diff --git a/temporalio/activity.py b/temporalio/activity.py index 701a4f3c6..1e4928a37 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -327,8 +327,10 @@ def stream_handle( 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. See - :py:mod:`temporalio.streams`. + 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. diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index ed3a90500..ad0bea575 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -923,7 +923,11 @@ def get_stream_handle( ``workflow_id`` is left out, and ``run_id`` then pins the activity's run, or an activity that ``workflow_id`` scheduled. The provider is the one registered with ``plugins=[provider]`` at :py:meth:`connect`, - or passed as ``stream_provider``. See :py:mod:`temporalio.streams`. + 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, or the diff --git a/tests/streams/test_stream_accessors.py b/tests/streams/test_stream_accessors.py index eb733ba5b..fed963aa6 100644 --- a/tests/streams/test_stream_accessors.py +++ b/tests/streams/test_stream_accessors.py @@ -18,7 +18,13 @@ from temporalio import activity, workflow from temporalio.client import Client -from temporalio.streams import RecordKind, StreamUnsupportedError, topic +from temporalio.streams import ( + BEGINNING, + END, + RecordKind, + StreamUnsupportedError, + topic, +) from temporalio.streams.providers.memory import MemoryStreams from temporalio.testing import ActivityEnvironment, WorkflowEnvironment from tests.helpers import new_worker @@ -120,3 +126,24 @@ async def ask() -> None: 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} From 4191d4d4f65880b58f3c0c99869f9f87aa6ebb3c Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:22:54 -0700 Subject: [PATCH 11/26] Ran stream bodies through the data converter with a plaintext hash. A provider owes a record body the codec and external storage the SDK gives every payload. The shared helper takes the retry fingerprint before either runs and stamps the plaintext hash on the record, so a nondeterministic codec cannot turn a retry into a divergent write and the store can compare retries without the plaintext. (cherry picked from commit 50b278bbb2739e975f1e48c872155363dfb17f13) --- temporalio/streams/__init__.py | 22 +- temporalio/streams/_body.py | 116 ++++++++++ temporalio/streams/_provider.py | 24 ++- temporalio/streams/providers/memory.py | 68 ++++-- tests/streams/test_memory_provider.py | 69 ++++++ tests/streams/test_streams_conformance.py | 136 +++++++++++- tests/streams/test_streams_internals.py | 248 ++++++++++++++++++++++ 7 files changed, 663 insertions(+), 20 deletions(-) create mode 100644 temporalio/streams/_body.py create mode 100644 tests/streams/test_memory_provider.py create mode 100644 tests/streams/test_streams_internals.py diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index 6b6b61cfc..10ca304ee 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -61,11 +61,26 @@ 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. +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 ( StreamCursorError, StreamError, @@ -93,6 +108,7 @@ __all__ = [ "BEGINNING", + "CONTENT_HASH_KEY", "END", "Cursor", "ReadSource", @@ -110,6 +126,10 @@ "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/_provider.py b/temporalio/streams/_provider.py index 378791bc8..6cee7e5e6 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -11,7 +11,10 @@ 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`. +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 @@ -294,6 +297,25 @@ class StreamProvider(Protocol): ``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. """ def workflow_provider(self) -> WorkflowStreamProvider: diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index 7040ea646..ab9d84202 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -23,6 +23,11 @@ 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. +- 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 @@ -44,7 +49,12 @@ from temporalio import workflow from temporalio.client import ActivityExecutionStatus, Client, WorkflowExecutionStatus from temporalio.service import RPCError, RPCStatusCode -from temporalio.streams._errors import StreamCursorError +from temporalio.streams._body import content_fingerprint, decode_body, encode_body +from temporalio.streams._errors import ( + StreamCursorError, + StreamProducerError, + StreamUnsupportedError, +) from temporalio.streams._ids import topic_key from temporalio.streams._provider import ReadSource, WriteSink from temporalio.streams._record import ( @@ -105,15 +115,34 @@ def append( *, 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)`` stores nothing - and returns where the original landed. + With a ``writer``, a repeat of ``(writer, sequence)`` carrying the same + content stores nothing and returns where the original landed. + ``content`` is the fingerprint the repeat is matched by; a producer + takes it over the records before their bodies are encoded, and + without one it is taken over ``wires`` as they are. + + Raises: + StreamProducerError: ``(writer, sequence)`` is held with different + content. """ key = (writer or "", sequence) - if writer is not None and key in self.seen: - return self.seen[key] + bodies = [wire.SerializeToString(deterministic=True) for wire in wires] + if content is None: + content = content_fingerprint(wires) + if writer is not None: + held = self.seen.get(key) + if held is not None: + first, count, seen_content = held + if seen_content != content: + raise StreamProducerError( + f"producer sequence {sequence} already used with different " + f"content by {writer!r}" + ) + return first, count first = self.head self.records.extend(wire.SerializeToString() for wire in wires) if writer is not None: @@ -237,7 +266,7 @@ class MemoryProducer(Generic[T]): def __init__( self, store: _Topic, - converter: temporalio.converter.PayloadConverter, + converter: temporalio.converter.DataConverter, topic: str, producer_id: str, attempt: int, @@ -277,10 +306,10 @@ async def append(self, *values: T) -> Cursor: """ if not values: return self._last - return self._write( + return await self._write( [ to_wire( - self._converter, + self._converter.payload_converter, topic=self._topic, kind=RecordKind.DATA, value=value, @@ -294,10 +323,10 @@ async def append(self, *values: T) -> Cursor: async def finish(self) -> None: """Write ``FINISH`` for this producer on this topic.""" - self._write( + await self._write( [ to_wire( - self._converter, + self._converter.payload_converter, topic=self._topic, kind=RecordKind.FINISH, producer_id=self._producer_id, @@ -307,9 +336,14 @@ async def finish(self) -> None: ] ) - def _write(self, wires: list[WireRecord]) -> Cursor: + 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 + wires, writer=self._writer, sequence=self._sequence, content=content ) self._sequence += len(wires) self._last = mint_cursor(_PROVIDER, str(first + count - 1)) @@ -340,9 +374,9 @@ def __init__( self._activity_id = activity_id self._seen_pending = False self._converter = ( - client.data_converter.payload_converter + client.data_converter if client is not None - else temporalio.converter.DataConverter.default.payload_converter + else temporalio.converter.DataConverter.default ) def read( @@ -377,7 +411,10 @@ async def _read( result_type: type | None, ) -> AsyncGenerator[StreamRecord[Any], None]: decoder = RecordDecoder( - self._converter, result_type, after=after, warn=logger.warning + self._converter.payload_converter, + result_type, + after=after, + warn=logger.warning, ) closed = False while True: @@ -392,6 +429,7 @@ async def _read( offset += 1 if wire is None: continue + await decode_body(self._converter, wire) for record in decoder.decode(cursor, wire): yield record if closed: diff --git a/tests/streams/test_memory_provider.py b/tests/streams/test_memory_provider.py new file mode 100644 index 000000000..3fd83a9fb --- /dev/null +++ b/tests/streams/test_memory_provider.py @@ -0,0 +1,69 @@ +"""What the reference provider does that the conformance suite cannot see. + +The conformance suite is the public surface, so it can say that closing a read +returns and that the topic still works afterwards, but not that the provider +let go of what the read parked on. That is this file: a few assertions against +``MemoryStreams`` internals, where holding on would leak quietly. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from temporalio.api.stream.v1 import StreamRecord +from temporalio.streams import CONTENT_HASH_KEY, StreamProducerError, topic +from temporalio.streams.providers.memory import MemoryStreams + +OUT = topic("out", dict) + + +async def test_closing_a_parked_read_drops_its_waiter(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-parked") # type: ignore[arg-type] + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}) + store = provider._topic("wf-parked", OUT.name) + + records = stream.read(topic=OUT) + await asyncio.wait_for(records.__anext__(), 5.0) + + pending = asyncio.ensure_future(records.__anext__()) + await asyncio.sleep(0.2) + assert store._waiters, "the read should be parked on the topic by now" + + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + # Nothing left behind: a reader that comes and goes must not grow this + # list for the life of the topic. + assert store._waiters == [] + await asyncio.wait_for(records.aclose(), 5.0) + + +async def test_a_divergent_retry_leaves_the_store_alone(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-divergent") # type: ignore[arg-type] + store = provider._topic("wf-divergent", OUT.name) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + await producer.append({"n": 1}) + assert len(store.records) == 1 + + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + with pytest.raises( + StreamProducerError, match="already used with different content" + ): + await retry.append({"n": 2}) + assert len(store.records) == 1 + + +async def test_a_stored_record_carries_the_plaintext_hash(): + provider = MemoryStreams() + stream = provider.get_stream_handle(None, "wf-hash") # type: ignore[arg-type] + await stream.producer(topic=OUT, producer_id="model", attempt=1).append({"n": 1}) + stored = StreamRecord.FromString(provider._topic("wf-hash", OUT.name).records[0]) + # What the store holds is the record after encode_body: the hash the + # server-side dedupe reads is on it, under the shared key. + assert stored.metadata[CONTENT_HASH_KEY].data.decode().isalnum() + assert len(stored.metadata[CONTENT_HASH_KEY].data) == 64 diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 9aa456676..f817208b1 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -21,8 +21,10 @@ from __future__ import annotations import asyncio +import dataclasses +import os import uuid -from collections.abc import AsyncIterator, Awaitable, Callable +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass from typing import Any @@ -31,7 +33,15 @@ from temporalio.api.common.v1 import Payload from temporalio.client import Client from temporalio.common import RawValue -from temporalio.converter import DataConverter +from temporalio.converter import ( + DataConverter, + ExternalStorage, + PayloadCodec, + StorageDriver, + StorageDriverClaim, + StorageDriverRetrieveContext, + StorageDriverStoreContext, +) from temporalio.streams import ( BEGINNING, END, @@ -73,10 +83,18 @@ class ProviderCase: for retention, or ``None`` when the provider offers no way to.""" async def open( - self, workflow_id: str, *, run_id: str | None = None + 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. @@ -89,6 +107,65 @@ async def open( ) +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() @@ -530,3 +607,56 @@ async def test_a_read_start_names_one_place(case: ProviderCase): 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_body_above_the_threshold_is_offloaded_and_read_back( + case: ProviderCase, client: Client +): + driver = RecordingDriver() + converter = dataclasses.replace( + DataConverter.default, + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=256), + ) + workflow_id = new_workflow_id() + stream = await case.open(workflow_id, client=_client_with(client, converter)) + producer = stream.producer(topic=OUT, producer_id="model", attempt=1) + small = {"n": 1} + large = {"blob": "x" * 1024} + await producer.append(small) + await producer.append(large) + # Only the body over the threshold left the record; the small one stayed + # inline, as it would on any other payload the SDK sends. + assert driver.stored == 1 + + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [small, large] + assert driver.retrieved == 1 + + +@pytest.mark.detects_divergent_retries +async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( + case: ProviderCase, client: Client +): + codec = NonceCodec() + converter = dataclasses.replace(DataConverter.default, payload_codec=codec) + workflow_id = new_workflow_id() + stream = await case.open(workflow_id, client=_client_with(client, converter)) + first = stream.producer(topic=OUT, producer_id="model", attempt=1) + landed = await first.append({"id": "r1"}) + assert codec.encoded == 1 + + # The codec produced different bytes for the retry. The provider matched + # it by the plaintext it converted, so it is the same append: stored + # once, answered with the original position. + retry = stream.producer(topic=OUT, producer_id="model", attempt=1) + again = await retry.append({"id": "r1"}) + if landed is not None: + assert again == landed + # And a retry that really does differ is still told apart. + divergent = stream.producer(topic=OUT, producer_id="model", attempt=1) + with pytest.raises(StreamProducerError): + await divergent.append({"id": "other"}) + + await first.append({"id": "r2"}) + records = await take(stream.read(topic=OUT), 2) + assert [r.value for r in records] == [{"id": "r1"}, {"id": "r2"}] diff --git a/tests/streams/test_streams_internals.py b/tests/streams/test_streams_internals.py new file mode 100644 index 000000000..d986e2084 --- /dev/null +++ b/tests/streams/test_streams_internals.py @@ -0,0 +1,248 @@ +"""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, + Supersession, + _ids, + _wire, + content_fingerprint, + content_hash, + decode_body, + encode_body, +) +from temporalio.streams._policy import AttemptTracker +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.worker import ReplayerConfig, WorkerConfig + + +def test_record_roundtrips_through_the_wire(): + converter = DataConverter.default.payload_converter + wire = _wire.to_wire( + converter, + topic="decisions", + kind=RecordKind.DATA, + value={"n": 1}, + producer_id="model", + attempt=3, + sequence=7, + ) + parsed = _wire.WireRecord.FromString(wire.SerializeToString()) + record = _wire.from_wire(converter, Cursor("memory:0"), parsed, dict) + assert ( + record.kind, + record.topic, + record.producer_id, + record.attempt, + record.sequence, + record.value, + ) == (RecordKind.DATA, "decisions", "model", 3, 7, {"n": 1}) + assert record.supersession is None + finish = _wire.to_wire(converter, topic="decisions", kind=RecordKind.FINISH) + assert not finish.HasField("body") + assert _wire.from_wire(converter, Cursor("memory:1"), finish, dict).value is None + + +def test_a_stored_supersession_is_not_a_record(): + converter = DataConverter.default.payload_converter + wire = _wire.WireRecord(topic="t", kind=int(RecordKind.SUPERSEDED)) # type: ignore[arg-type] + with pytest.raises(ValueError, match="synthesized"): + _wire.from_wire(converter, Cursor("memory:0"), wire, None) + + +def test_an_unset_kind_is_read_as_data(): + converter = DataConverter.default.payload_converter + wire = _wire.WireRecord(topic="t", body=converter.to_payloads([{"n": 1}])[0]) + record = _wire.from_wire(converter, Cursor("memory:0"), wire, dict) + assert record.kind is RecordKind.DATA + assert record.value == {"n": 1} + + +def test_supersession_is_synthesized_from_observations(): + attempts = AttemptTracker() + assert attempts.note("model", 1, topic="t", previous=BEGINNING) is None + superseded = attempts.note("model", 2, topic="t", previous=Cursor("memory:0")) + assert superseded is not None + assert superseded.kind is RecordKind.SUPERSEDED + assert superseded.supersession == Supersession("model", 1, 2) + assert superseded.value is None + # Positioned before the triggering record, so a resume after it delivers + # that record next. + assert superseded.cursor == Cursor("memory:0") + # The same attempt again is not a new generation. + assert attempts.note("model", 2, topic="t", previous=Cursor("memory:1")) is None + + +def test_topic_keys_cannot_collide(): + # A colon in a workflow id must not make two addresses one key. + assert _ids.topic_key("a:b", "c") != _ids.topic_key("a", "b:c") + assert _ids.topic_key("a%3Ab", "c") != _ids.topic_key("a:b", "c") + assert _ids.topic_key("wf", "inputs") == "wf:inputs" + + +def test_cursors_name_their_provider(): + assert _wire.cursor_position(BEGINNING, provider="memory") is None + assert _wire.cursor_position(Cursor("memory:42"), provider="memory") == "42" + with pytest.raises(StreamCursorError): + _wire.cursor_position(Cursor("redis:1700000000000-0"), provider="memory") + + +def test_registering_a_provider_twice_is_refused(): + # There is one slot on each of the three, and a user who passes a provider + # by hand and a provider plugin, or two provider plugins, meant both. + first, second = MemoryStreams(), MemoryStreams() + with pytest.raises(ValueError, match="already registered"): + second.configure_client(ClientConfig(stream_provider=first)) # type: ignore[typeddict-item] + with pytest.raises(ValueError, match="already registered"): + second.configure_worker(WorkerConfig(stream_provider=first)) # type: ignore[typeddict-item] + with pytest.raises(ValueError, match="already registered"): + second.configure_replayer(ReplayerConfig(stream_provider=first)) # type: ignore[typeddict-item] + + +def test_registering_the_same_provider_twice_is_fine(): + # A worker built from a client that already carries the plugin configures + # it again with the same object, which is not a conflict. + provider = MemoryStreams() + config = provider.configure_client(ClientConfig(stream_provider=provider)) # type: ignore[typeddict-item] + assert config.get("stream_provider") is provider + assert provider.configure_client(ClientConfig()).get("stream_provider") is provider # type: ignore[typeddict-item] + + +class _HoldEverything(StorageDriver): + """A driver that keeps every payload it is handed, in memory.""" + + def __init__(self) -> None: + self.held: list[bytes] = [] + + def name(self) -> str: + return "hold" + + async def store( + self, context: StorageDriverStoreContext, payloads: Sequence[Payload] + ) -> list[StorageDriverClaim]: + claims = [] + for payload in payloads: + claims.append(StorageDriverClaim(claim_data={"i": str(len(self.held))})) + self.held.append(payload.SerializeToString()) + return claims + + async def retrieve( + self, + context: StorageDriverRetrieveContext, + claims: Sequence[StorageDriverClaim], + ) -> list[Payload]: + return [Payload.FromString(self.held[int(c.claim_data["i"])]) for c in claims] + + +class _NonceCodec(PayloadCodec): + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=os.urandom(8) + p.SerializeToString(), + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [Payload.FromString(p.data[8:]) for p in payloads] + + +async def test_encode_body_stamps_the_plaintext_hash_and_offloads_the_body(): + driver = _HoldEverything() + converter = dataclasses.replace( + DataConverter.default, + payload_codec=_NonceCodec(), + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=0), + ) + wire = _wire.to_wire( + converter.payload_converter, topic="t", kind=RecordKind.DATA, value={"n": 1} + ) + plaintext = Payload() + plaintext.CopyFrom(wire.body) + + await encode_body(converter, wire) + # The hash is over what the converter produced, not over what the codec + # or the driver made of it, and it rides the record where the store can + # read it without the plaintext. + stamped = wire.metadata[CONTENT_HASH_KEY] + assert stamped.metadata["encoding"] == b"binary/plain" + assert stamped.data.decode() == content_hash(plaintext) + assert len(stamped.data) == 64 + # With a threshold of zero the body was offloaded: the record holds the + # claim and the driver holds the coded payload. + assert wire.body != plaintext + assert len(wire.body.external_payloads) == 1 + assert len(driver.held) == 1 + + await decode_body(converter, wire) + assert wire.body == plaintext + assert wire.metadata[CONTENT_HASH_KEY] == stamped + + # A record without a body has nothing to hash or offload. + finish = _wire.to_wire( + converter.payload_converter, topic="t", kind=RecordKind.FINISH + ) + await encode_body(converter, finish) + assert CONTENT_HASH_KEY not in finish.metadata + assert len(driver.held) == 1 + + +async def test_content_fingerprint_is_taken_before_the_codec(): + converter = dataclasses.replace(DataConverter.default, payload_codec=_NonceCodec()) + plain = converter.payload_converter + + def batch(*values: dict) -> list[_wire.WireRecord]: + return [ + _wire.to_wire(plain, topic="t", kind=RecordKind.DATA, value=v, sequence=i) + for i, v in enumerate(values, 1) + ] + + first, retry = batch({"n": 1}, {"n": 2}), batch({"n": 1}, {"n": 2}) + before = content_fingerprint(first) + assert before == content_fingerprint(retry) + # Different content, and the same content split differently, both differ. + assert before != content_fingerprint(batch({"n": 1}, {"n": 3})) + assert before != content_fingerprint(batch({"n": 1}) + batch({"n": 2})) + + for record in first + retry: + await encode_body(converter, record) + # The codec made the two batches' bytes differ; the identity taken first + # is what lets a store still recognise the retry. + assert first[0].body != retry[0].body + assert content_fingerprint(first) != content_fingerprint(retry) + # Decoding gives the converted bodies back; the hash stays stamped on the + # record, which is why the identity is taken before encoding, not after. + for record in first: + await decode_body(converter, record) + assert [r.body for r in first] == [r.body for r in batch({"n": 1}, {"n": 2})] + assert all(CONTENT_HASH_KEY in r.metadata for r in first) From f52bb8210c6a17cc2bd79927c97f6a05d503603d Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:22:54 -0700 Subject: [PATCH 12/26] Declared standalone streams and a handle close on the provider surface. A stream with an id of its own and no owner is created on purpose with a retention policy and sealed on purpose. The memory provider refuses both calls for now, and a workflow's handle refuses close, since its stream ends with the workflow. (cherry picked from commit ca0515cb33af15f6723905e2270d488c5b0ffc7a) --- temporalio/streams/__init__.py | 11 ++++ temporalio/streams/_errors.py | 8 +++ temporalio/streams/_provider.py | 68 +++++++++++++++++++++++ temporalio/streams/providers/memory.py | 48 ++++++++++++++-- tests/streams/test_streams_conformance.py | 10 ++++ 5 files changed, 141 insertions(+), 4 deletions(-) diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index 10ca304ee..9335e10bd 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -53,6 +53,15 @@ protocols a provider implements; nothing here that workflow code imports does I/O. + +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 @@ -82,6 +91,7 @@ encode_body, ) from temporalio.streams._errors import ( + StreamClosedError, StreamCursorError, StreamError, StreamNotFoundError, @@ -113,6 +123,7 @@ "Cursor", "ReadSource", "RecordKind", + "StreamClosedError", "StreamCursorError", "StreamError", "StreamHandle", diff --git a/temporalio/streams/_errors.py b/temporalio/streams/_errors.py index f00bf1c3a..a2f4fecf5 100644 --- a/temporalio/streams/_errors.py +++ b/temporalio/streams/_errors.py @@ -12,6 +12,7 @@ import temporalio.exceptions __all__ = [ + "StreamClosedError", "StreamCursorError", "StreamError", "StreamNotFoundError", @@ -36,5 +37,12 @@ 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/_provider.py b/temporalio/streams/_provider.py index 6cee7e5e6..088cec252 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -20,6 +20,7 @@ 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 @@ -204,6 +205,21 @@ def producer( """ ... + 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.""" @@ -316,6 +332,15 @@ class StreamProvider(Protocol): :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: @@ -357,6 +382,49 @@ def get_activity_stream_handle( """ ... + 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. diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index ab9d84202..40fffd253 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -23,6 +23,8 @@ 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. +- It does not host standalone streams; both standalone calls raise + :class:`temporalio.streams.StreamUnsupportedError`. - The outside path encodes and decodes bodies through the client's data converter, codec and external storage included, and fingerprints a retry over the converted bytes first. The workflow half has no client, so a @@ -101,8 +103,9 @@ def __init__(self) -> None: self.records: list[bytes] = [] # Dedupe identity is (producer#attempt, first sequence of the append), # the same pair the storage providers use, mapped to where the batch - # landed so a repeat can answer with the original position. - self.seen: dict[tuple[str, int], tuple[int, int]] = {} + # 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 @@ -144,9 +147,9 @@ def append( ) return first, count first = self.head - self.records.extend(wire.SerializeToString() for wire in wires) + self.records.extend(bodies) if writer is not None: - self.seen[key] = (first, len(wires)) + self.seen[key] = (first, len(wires), content) waiters, self._waiters = self._waiters, [] for loop, future in waiters: loop.call_soon_threadsafe(_wake, future) @@ -512,6 +515,13 @@ def producer( producer_id, attempt = producer_identity(producer_id, attempt) return MemoryProducer(store, self._converter, name, producer_id, attempt) + async def close(self) -> None: + """Refuse: a workflow's stream ends with the workflow, not by a caller.""" + raise ValueError( + "only a standalone stream can be closed; this handle is on a workflow's " + "stream, which ends when the workflow does" + ) + class MemoryStreams(ProviderPlugin): """The in-memory provider, one list per topic. @@ -578,6 +588,36 @@ def get_activity_stream_handle( """ 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: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the memory provider does not host standalone streams" + ) + + def get_standalone_stream_handle( + self, client: Client | None, stream_id: str + ) -> MemoryStreamHandle: + """Refuse: this provider keeps no stream without an owner. + + Raises: + StreamUnsupportedError: Always. + """ + raise StreamUnsupportedError( + "the memory provider does not host standalone streams" + ) + async def close(self) -> None: """Nothing to release: the provider holds no connection.""" diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index f817208b1..69773c2d5 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -609,6 +609,16 @@ async def test_a_read_start_names_one_place(case: ProviderCase): stream.read(topic=OUT, after=END, last=1) + + +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 ): From 1897c9184709e5a82351e2d627afdd8fccbdc583 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:22:54 -0700 Subject: [PATCH 13/26] Added StreamRef, a serializable name for one stream. A handle is bound to its client and provider, so a stream crosses a process boundary as its owner and topic in plain data, with no cursor and no provider name. The default converter carries it as JSON, so it can be a workflow argument, an activity result or a Nexus operation input or result. (cherry picked from commit 5d600c45a791330a99e3a86236fe546ce0e6ee34) --- temporalio/streams/__init__.py | 9 ++ temporalio/streams/_provider.py | 36 ++++-- temporalio/streams/_ref.py | 142 ++++++++++++++++++++++ temporalio/streams/providers/memory.py | 15 +++ tests/streams/test_streams_conformance.py | 38 +++++- tests/streams/test_streams_internals.py | 30 +++++ 6 files changed, 258 insertions(+), 12 deletions(-) create mode 100644 temporalio/streams/_ref.py diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index 9335e10bd..1eda4e030 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -53,6 +53,12 @@ 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 @@ -114,6 +120,7 @@ StreamRecord, Supersession, ) +from temporalio.streams._ref import StreamOwnerKind, StreamRef from temporalio.streams._topic import StreamTopic, resolve_topic, topic __all__ = [ @@ -128,10 +135,12 @@ "StreamError", "StreamHandle", "StreamNotFoundError", + "StreamOwnerKind", "StreamProducer", "StreamProducerError", "StreamProvider", "StreamRecord", + "StreamRef", "StreamTopic", "StreamUnsupportedError", "Supersession", diff --git a/temporalio/streams/_provider.py b/temporalio/streams/_provider.py index 088cec252..18bdc9e67 100644 --- a/temporalio/streams/_provider.py +++ b/temporalio/streams/_provider.py @@ -29,6 +29,7 @@ if TYPE_CHECKING: from temporalio.client import Client + from temporalio.streams._ref import StreamRef __all__ = [ "ReadSource", @@ -91,14 +92,21 @@ async def finish(self) -> None: class StreamHandle(Protocol): """One owner's stream, addressed by topic, from outside workflow code. - The owner is a workflow or an activity. 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. + 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 @@ -205,6 +213,18 @@ def producer( """ ... + 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. diff --git a/temporalio/streams/_ref.py b/temporalio/streams/_ref.py new file mode 100644 index 000000000..35b52ec49 --- /dev/null +++ b/temporalio/streams/_ref.py @@ -0,0 +1,142 @@ +"""A reference to one stream that crosses a process boundary as data. + +A handle is bound to a client and a provider, so it cannot be a workflow +argument, an activity result or a Nexus operation result. A +:class:`StreamRef` can: it names the owner and the topic, nothing more, and +whoever receives it opens the stream on the provider its own client carries. +It carries no cursor, because a position belongs to a reader, and no provider +name, because the same owner and topic name the same stream on every provider +a deployment runs. +""" + +from __future__ import annotations + +import dataclasses +from dataclasses import dataclass +from typing import Any, Literal + +from temporalio.streams._topic import StreamTopic, resolve_topic + +__all__ = ["StreamOwnerKind", "StreamRef"] + +StreamOwnerKind = Literal["workflow", "activity", "standalone"] +"""What owns a stream: a workflow, an activity, or the stream itself.""" + + +@dataclass(frozen=True) +class StreamRef: + """One stream, named by its owner and its topic. + + ``kind`` says what owns the stream. A ``"workflow"`` ref carries + ``workflow_id`` and, when pinned to one run, ``run_id``. An ``"activity"`` + ref carries ``activity_id``, plus ``workflow_id`` (and its ``run_id``) + when a workflow scheduled the activity; without ``workflow_id`` it is a + standalone activity, and ``run_id`` then pins one run of it. A + ``"standalone"`` ref carries ``stream_id`` and nothing else. ``topic`` is + the topic's name 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 diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index 40fffd253..c83893cbd 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -67,6 +67,7 @@ StreamRecord, check_read_start, ) +from temporalio.streams._ref import StreamRef from temporalio.streams._topic import StreamTopic, resolve_topic from temporalio.streams._wire import ( RecordDecoder, @@ -515,6 +516,20 @@ def producer( 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._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: """Refuse: a workflow's stream ends with the workflow, not by a caller.""" raise ValueError( diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 69773c2d5..cd3664a31 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -11,10 +11,15 @@ skipped with a reason on a provider whose ``append()`` learns positions at read time. -What this file pins down is the contract: the record on the wire, producer -identity, retry deduplication, positions, supersession, topic addressing, -cursor resumption, cursor ownership, and store keys that cannot collide. The -workflow-side handles and the two rules about Workflow Tasks live in +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``. """ @@ -49,7 +54,9 @@ RecordKind, StreamCursorError, StreamHandle, + StreamProducerError, StreamProvider, + StreamRef, Supersession, _ids, _wire, @@ -609,6 +616,29 @@ async def test_a_read_start_names_one_place(case: ProviderCase): 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): diff --git a/tests/streams/test_streams_internals.py b/tests/streams/test_streams_internals.py index d986e2084..8e3d09feb 100644 --- a/tests/streams/test_streams_internals.py +++ b/tests/streams/test_streams_internals.py @@ -32,6 +32,7 @@ Cursor, RecordKind, StreamCursorError, + StreamRef, Supersession, _ids, _wire, @@ -246,3 +247,32 @@ def batch(*values: dict) -> list[_wire.WireRecord]: await decode_body(converter, record) assert [r.body for r in first] == [r.body for r in batch({"n": 1}, {"n": 2})] assert all(CONTENT_HASH_KEY in r.metadata for r in first) + + +async def test_a_stream_ref_names_one_owner_and_travels_as_json(): + workflow = StreamRef.for_workflow("wf", run_id="r", topic="out") + activity = StreamRef.for_activity("act", workflow_id="wf", topic="progress") + standalone = StreamRef.for_standalone("shared") + assert workflow == StreamRef("workflow", "out", workflow_id="wf", run_id="r") + assert activity.kind == "activity" and activity.activity_id == "act" + assert standalone == StreamRef("standalone", "output", stream_id="shared") + assert standalone.with_topic("x").topic == "x" + + for bad in ( + dict(kind="workflow"), + dict(kind="workflow", workflow_id="wf", stream_id="s"), + dict(kind="activity", workflow_id="wf"), + dict(kind="standalone", stream_id="s", workflow_id="wf"), + dict(kind="standalone"), + dict(kind="nexus", stream_id="s"), + dict(kind="workflow", workflow_id="wf", topic=""), + ): + with pytest.raises(ValueError): + StreamRef(**bad) # type: ignore[arg-type] + + converter = DataConverter.default + for ref in (workflow, activity, standalone): + [carried] = await converter.decode(await converter.encode([ref]), [StreamRef]) + assert carried == ref + payload = (await converter.encode([standalone]))[0] + assert payload.metadata["encoding"] == b"json/plain" From 3019a5884f4c7eefcbd39915761b5782caab6eb5 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:30:38 -0700 Subject: [PATCH 14/26] Let the client and activity accessors open refs and standalone streams. A StreamRef in place of the workflow id opens the stream it names on whatever provider the client carries, and its topic becomes the handle's default. A standalone stream is created with client.create_stream and reached by its stream_id. (cherry picked from commit 61a58f78b953b5c5fdeac575c1c3c64e9d099d5b) --- CHANGELOG.md | 8 +- temporalio/activity.py | 21 ++++- temporalio/client/_client.py | 101 ++++++++++++++++++++++--- temporalio/streams/_ref.py | 100 +++++++++++++++++++++++- tests/streams/test_activity_streams.py | 4 +- 5 files changed, 216 insertions(+), 18 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 6bbd734e5..a1e7b5f40 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -44,7 +44,13 @@ to include examples, links to docs, or any other relevant information. 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. + 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. diff --git a/temporalio/activity.py b/temporalio/activity.py index 1e4928a37..b06e3c78a 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -34,6 +34,7 @@ from temporalio.converter._payload_converter import ( _TemporalTransferTypePayloadConverter, ) +from temporalio.streams._ref import open_ref from .types import CallableType @@ -302,13 +303,18 @@ def client() -> Client: def stream_handle( - workflow_id: str | None = None, + 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: @@ -336,7 +342,8 @@ def stream_handle( activities. Args: - workflow_id: Another workflow whose stream to address. + 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 @@ -354,8 +361,8 @@ def stream_handle( 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``, or - ``scope="activity"`` with one. + ValueError: ``run_id`` was given without ``workflow_id``, + ``scope="activity"`` with one, or a ref with either. """ context = _Context.current() provider = context.stream_provider @@ -364,6 +371,12 @@ def stream_handle( "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( diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index ad0bea575..01927bad5 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -41,6 +41,7 @@ ServiceClient, TLSConfig, ) +from temporalio.streams._ref import open_ref from ..common import HeaderCodecBehavior from ..types import ( @@ -909,40 +910,50 @@ def get_workflow_handle( def get_stream_handle( self, - workflow_id: str | None = None, + 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 workflow's or an activity's stream from the provider registered on this client. + """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. The provider is - the one registered with ``plugins=[provider]`` at :py:meth:`connect`, - or passed as ``stream_provider``. A ``read`` starts at + 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, or the - workflow that scheduled ``activity_id``. + 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: Neither ``workflow_id`` nor ``activity_id`` was given. + 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. + activity owns or a stream without an owner. """ provider = self._config.get("stream_provider") if provider is None: @@ -950,14 +961,84 @@ def get_stream_handle( "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 or the activity_id to address") + 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: ( diff --git a/temporalio/streams/_ref.py b/temporalio/streams/_ref.py index 35b52ec49..33349e431 100644 --- a/temporalio/streams/_ref.py +++ b/temporalio/streams/_ref.py @@ -12,12 +12,22 @@ from __future__ import annotations import dataclasses +from collections.abc import AsyncGenerator from dataclasses import dataclass -from typing import Any, Literal +from typing import TYPE_CHECKING, Any, Literal +from temporalio.streams._record import BEGINNING, Cursor, StreamRecord from temporalio.streams._topic import StreamTopic, resolve_topic -__all__ = ["StreamOwnerKind", "StreamRef"] +if TYPE_CHECKING: + from temporalio.client import Client + from temporalio.streams._provider import ( + StreamHandle, + StreamProducer, + StreamProvider, + ) + +__all__ = ["StreamOwnerKind", "StreamRef", "open_ref"] StreamOwnerKind = Literal["workflow", "activity", "standalone"] """What owns a stream: a workflow, an activity, or the stream itself.""" @@ -140,3 +150,89 @@ def _name(topic: str | StreamTopic[Any] | None) -> str | None: return None name, _ = resolve_topic(topic) return name + + +def open_ref(provider: StreamProvider, client: Client, ref: StreamRef) -> StreamHandle: + """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. + + 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: + 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]: + # 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: + 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]: + 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: + if topic is None: + return self._ref + return self._inner.ref(topic=topic) + + async def close(self) -> None: + await self._inner.close() diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index 8aa4a653a..bd0a675ac 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -257,5 +257,7 @@ async def test_a_standalone_activity_has_no_workflow_to_address(setup: ActivityS async def test_get_stream_handle_needs_an_owner(setup: ActivitySetup): - with pytest.raises(ValueError, match="workflow_id or the activity_id"): + with pytest.raises( + ValueError, match="workflow_id, the activity_id or the stream_id" + ): setup.client.get_stream_handle() From 4e20489e15af3701c2b0d055f18fbf97d22b179a Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:42:04 -0700 Subject: [PATCH 15/26] Hosted standalone streams in the memory provider. A standalone stream lives in the provider with its policy and a sealed flag. The policy trims on append, a seal wakes parked readers so their reads end, and a stream id that was never created is refused at the call, so the conformance suite can run the standalone cases without a server. (cherry picked from commit a95c913161d5211baab19cd5d7fba594566061f8) --- temporalio/streams/providers/memory.py | 180 +++++++++++++++++++++---- 1 file changed, 152 insertions(+), 28 deletions(-) diff --git a/temporalio/streams/providers/memory.py b/temporalio/streams/providers/memory.py index c83893cbd..b8483ab43 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -23,8 +23,13 @@ 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. -- It does not host standalone streams; both standalone calls raise - :class:`temporalio.streams.StreamUnsupportedError`. +- 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 @@ -41,7 +46,9 @@ 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 @@ -53,9 +60,10 @@ 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, - StreamUnsupportedError, ) from temporalio.streams._ids import topic_key from temporalio.streams._provider import ReadSource, WriteSink @@ -93,15 +101,45 @@ def _wake(future: asyncio.Future[None]) -> None: 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) -> None: + 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 @@ -132,7 +170,10 @@ def append( 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: @@ -148,13 +189,43 @@ def append( ) 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) - return first, len(wires) + + 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: @@ -167,9 +238,12 @@ def at(self, offset: int) -> bytes: def truncate(self, keep: int) -> None: """Drop all but the newest ``keep`` records.""" - drop = max(0, len(self.records) - keep) - self.base += drop - del self.records[:drop] + 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.""" @@ -359,7 +433,8 @@ class MemoryStreamHandle: The owner is ``workflow_id``'s workflow, or with ``activity_id`` an activity: a standalone one without ``workflow_id``, or one that workflow - scheduled. + scheduled. With ``stream_id`` the handle is on a standalone stream, which + has no owner. """ def __init__( @@ -369,6 +444,7 @@ def __init__( 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 @@ -376,6 +452,7 @@ def __init__( 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 @@ -391,10 +468,11 @@ def read( last: int | None = None, result_type: type | None = None, ) -> AsyncGenerator[StreamRecord[Any], None]: - """Yield records on ``topic`` from where the read starts until the workflow closes. + """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. + 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) @@ -450,6 +528,8 @@ async def _read( ) 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 @@ -458,6 +538,8 @@ def _store(self, topic: str) -> _Topic: 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: @@ -518,6 +600,8 @@ def producer( 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, @@ -531,11 +615,13 @@ def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: ) async def close(self) -> None: - """Refuse: a workflow's stream ends with the workflow, not by a caller.""" - raise ValueError( - "only a standalone stream can be closed; this handle is on a workflow's " - "stream, which ends when the workflow does" - ) + """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): @@ -558,11 +644,13 @@ def __init__( 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. For tests.""" + """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. @@ -612,26 +700,38 @@ async def create_standalone_stream( max_records: int | None = None, max_bytes: int | None = None, ) -> MemoryStreamHandle: - """Refuse: this provider keeps no stream without an owner. + """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: - StreamUnsupportedError: Always. + ValueError: ``stream_id`` is empty, a bound is not positive, or + the stream exists with a different policy. """ - raise StreamUnsupportedError( - "the memory provider does not host standalone streams" - ) + 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: - """Refuse: this provider keeps no stream without an owner. + """A handle on the standalone stream ``stream_id``. - Raises: - StreamUnsupportedError: Always. + Nothing is checked here: a ``read``, ``latest`` or ``producer`` on a + stream that was never created raises + :class:`temporalio.streams.StreamNotFoundError` at the call. """ - raise StreamUnsupportedError( - "the memory provider does not host standalone streams" - ) + return MemoryStreamHandle(self, client, None, None, stream_id=stream_id) async def close(self) -> None: """Nothing to release: the provider holds no connection.""" @@ -658,6 +758,30 @@ def _activity_topic( 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: From 333e3f7cb75a3b8e0b2cfbf19d8fed507e91e4c8 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 02:42:04 -0700 Subject: [PATCH 16/26] Tested refs, standalone streams and body encoding through the accessors. The conformance suite opens a ref, creates and seals a standalone stream and checks its retention policy on every provider that hosts one. The accessor tests carry a ref through a workflow argument and result and open it from the activity and the client. (cherry picked from commit 5287f6e53a6b5a512d3297ce39a6bcd7b70deba6) --- temporalio/activity.py | 18 ++ temporalio/client/_client.py | 20 ++ temporalio/streams/__init__.py | 3 +- temporalio/streams/_ref.py | 13 +- tests/streams/conftest.py | 5 + tests/streams/test_stream_accessors.py | 87 +++++++ tests/streams/test_streams_conformance.py | 285 +++++++++++++++------- 7 files changed, 333 insertions(+), 98 deletions(-) diff --git a/temporalio/activity.py b/temporalio/activity.py index b06e3c78a..c906c5f05 100644 --- a/temporalio/activity.py +++ b/temporalio/activity.py @@ -302,6 +302,24 @@ 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, *, diff --git a/temporalio/client/_client.py b/temporalio/client/_client.py index 01927bad5..568914cd3 100644 --- a/temporalio/client/_client.py +++ b/temporalio/client/_client.py @@ -908,6 +908,26 @@ 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, diff --git a/temporalio/streams/__init__.py b/temporalio/streams/__init__.py index 1eda4e030..12a6576ed 100644 --- a/temporalio/streams/__init__.py +++ b/temporalio/streams/__init__.py @@ -120,7 +120,7 @@ StreamRecord, Supersession, ) -from temporalio.streams._ref import StreamOwnerKind, StreamRef +from temporalio.streams._ref import RefHandle, StreamOwnerKind, StreamRef from temporalio.streams._topic import StreamTopic, resolve_topic, topic __all__ = [ @@ -130,6 +130,7 @@ "Cursor", "ReadSource", "RecordKind", + "RefHandle", "StreamClosedError", "StreamCursorError", "StreamError", diff --git a/temporalio/streams/_ref.py b/temporalio/streams/_ref.py index 33349e431..586eeb4ae 100644 --- a/temporalio/streams/_ref.py +++ b/temporalio/streams/_ref.py @@ -27,7 +27,7 @@ StreamProvider, ) -__all__ = ["StreamOwnerKind", "StreamRef", "open_ref"] +__all__ = ["RefHandle", "StreamOwnerKind", "StreamRef", "open_ref"] StreamOwnerKind = Literal["workflow", "activity", "standalone"] """What owns a stream: a workflow, an activity, or the stream itself.""" @@ -152,7 +152,7 @@ def _name(topic: str | StreamTopic[Any] | None) -> str | None: return name -def open_ref(provider: StreamProvider, client: Client, ref: StreamRef) -> StreamHandle: +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 @@ -173,12 +173,17 @@ def open_ref(provider: StreamProvider, client: Client, ref: StreamRef) -> Stream else: assert ref.stream_id is not None handle = provider.get_standalone_stream_handle(client, ref.stream_id) - return _RefHandle(handle, ref) + return RefHandle(handle, ref) -class _RefHandle: +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 diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 7be000f3b..d4dc1c010 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -10,3 +10,8 @@ def pytest_configure(config: pytest.Config) -> None: "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", + ) diff --git a/tests/streams/test_stream_accessors.py b/tests/streams/test_stream_accessors.py index fed963aa6..5717a995c 100644 --- a/tests/streams/test_stream_accessors.py +++ b/tests/streams/test_stream_accessors.py @@ -22,6 +22,8 @@ BEGINNING, END, RecordKind, + StreamClosedError, + StreamRef, StreamUnsupportedError, topic, ) @@ -147,3 +149,88 @@ async def first(records: Any) -> Any: 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_streams_conformance.py b/tests/streams/test_streams_conformance.py index cd3664a31..efc7f6861 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -31,6 +31,7 @@ import uuid from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass +from datetime import timedelta from typing import Any import pytest @@ -52,17 +53,17 @@ END, Cursor, RecordKind, + StreamClosedError, StreamCursorError, StreamHandle, + StreamNotFoundError, StreamProducerError, StreamProvider, StreamRef, Supersession, - _ids, - _wire, topic, ) -from temporalio.streams._policy import AttemptTracker +from temporalio.streams._ref import RefHandle, open_ref from temporalio.streams.providers.memory import MemoryStreams # Defined once and shared by every case, the way an application shares them @@ -88,6 +89,11 @@ class ProviderCase: 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``.""" async def open( self, @@ -113,6 +119,42 @@ async def open( 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.""" @@ -190,6 +232,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, "truncates": lambda case: case.truncate is not None, + "hosts_standalone_streams": lambda case: case.hosts_standalone_streams, } @@ -223,96 +266,17 @@ async def _collect() -> None: return out -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_an_attempt_that_goes_backwards_is_said_rather_than_passed_off(): - # Attempts only rise on one producer, so a lower one means the store handed - # two generations back out of order. Yielded as data with no signal, a - # consumer renders the stale generation as the current answer. - said: list[str] = [] - attempts = AttemptTracker(said.append) - assert attempts.note("model", 2, topic="t", previous=BEGINNING) is None - assert attempts.note("model", 1, topic="t", previous=Cursor("memory:3")) is None - assert len(said) == 1 - assert "attempt 1" in said[0] and "behind attempt 2" in said[0] - assert "model" in said[0] - - -def test_a_repeat_of_the_current_attempt_is_not_worth_saying(): - said: list[str] = [] - attempts = AttemptTracker(said.append) - attempts.note("model", 1, topic="t", previous=BEGINNING) - assert attempts.note("model", 1, topic="t", previous=Cursor("memory:1")) is None - assert said == [] - - -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") +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): @@ -700,3 +664,138 @@ async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( 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}] + + 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 + + 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"}] From f4e314dec4f4c4b16b22120e8f0ff97958317b2d Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:17:49 -0700 Subject: [PATCH 17/26] Fitted the ported interface commits to this chain. This chain has no default topic, so a ref may name the owner alone and the handle opened from a ref is a public `RefHandle` whose calls may leave the topic out; the two accessors return it for a ref. The memory topic drops a parked waiter on cancellation, which the ported provider test checks, and the plugin-registry refusal test targets a registry this chain does not have. --- temporalio/streams/_ref.py | 6 ++++++ temporalio/streams/providers/memory.py | 4 ++++ tests/streams/conftest.py | 5 +++++ tests/streams/test_streams_conformance.py | 3 +++ tests/streams/test_streams_internals.py | 23 +++++++---------------- 5 files changed, 25 insertions(+), 16 deletions(-) diff --git a/temporalio/streams/_ref.py b/temporalio/streams/_ref.py index 586eeb4ae..2b31b6c32 100644 --- a/temporalio/streams/_ref.py +++ b/temporalio/streams/_ref.py @@ -191,6 +191,7 @@ class RefHandle: """ 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 @@ -212,6 +213,7 @@ def read( 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 @@ -220,6 +222,7 @@ def read( ) 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( @@ -229,15 +232,18 @@ def producer( 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/providers/memory.py b/temporalio/streams/providers/memory.py index b8483ab43..de6d2b43e 100644 --- a/temporalio/streams/providers/memory.py +++ b/temporalio/streams/providers/memory.py @@ -255,6 +255,10 @@ async def wait_past(self, offset: int, timeout: float | None) -> None: 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] diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index d4dc1c010..f03dbc919 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -6,6 +6,11 @@ def pytest_configure(config: pytest.Config) -> None: "markers", "reports_positions: the case needs append() to return where records landed", ) + config.addinivalue_line( + "markers", + "detects_divergent_retries: the case needs append() to compare a repeat's " + "content with what the store holds", + ) config.addinivalue_line( "markers", "truncates: the case needs a way to drop a topic's oldest records", diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index efc7f6861..4dc09cd98 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -84,6 +84,8 @@ class ProviderCase: 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 @@ -231,6 +233,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: _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, } diff --git a/tests/streams/test_streams_internals.py b/tests/streams/test_streams_internals.py index 8e3d09feb..8a6d60ecc 100644 --- a/tests/streams/test_streams_internals.py +++ b/tests/streams/test_streams_internals.py @@ -43,7 +43,6 @@ ) from temporalio.streams._policy import AttemptTracker from temporalio.streams.providers.memory import MemoryStreams -from temporalio.worker import ReplayerConfig, WorkerConfig def test_record_roundtrips_through_the_wire(): @@ -117,25 +116,16 @@ def test_cursors_name_their_provider(): _wire.cursor_position(Cursor("redis:1700000000000-0"), provider="memory") -def test_registering_a_provider_twice_is_refused(): - # There is one slot on each of the three, and a user who passes a provider - # by hand and a provider plugin, or two provider plugins, meant both. - first, second = MemoryStreams(), MemoryStreams() - with pytest.raises(ValueError, match="already registered"): - second.configure_client(ClientConfig(stream_provider=first)) # type: ignore[typeddict-item] - with pytest.raises(ValueError, match="already registered"): - second.configure_worker(WorkerConfig(stream_provider=first)) # type: ignore[typeddict-item] - with pytest.raises(ValueError, match="already registered"): - second.configure_replayer(ReplayerConfig(stream_provider=first)) # type: ignore[typeddict-item] - - def test_registering_the_same_provider_twice_is_fine(): # A worker built from a client that already carries the plugin configures # it again with the same object, which is not a conflict. provider = MemoryStreams() - config = provider.configure_client(ClientConfig(stream_provider=provider)) # type: ignore[typeddict-item] + config = provider.configure_client( + ClientConfig(stream_provider=provider) # type: ignore[typeddict-item] + ) assert config.get("stream_provider") is provider - assert provider.configure_client(ClientConfig()).get("stream_provider") is provider # type: ignore[typeddict-item] + registered = provider.configure_client(ClientConfig()) # type: ignore[typeddict-item] + assert registered.get("stream_provider") is provider # type: ignore[typeddict-item] class _HoldEverything(StorageDriver): @@ -255,7 +245,8 @@ async def test_a_stream_ref_names_one_owner_and_travels_as_json(): standalone = StreamRef.for_standalone("shared") assert workflow == StreamRef("workflow", "out", workflow_id="wf", run_id="r") assert activity.kind == "activity" and activity.activity_id == "act" - assert standalone == StreamRef("standalone", "output", stream_id="shared") + # 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 ( From e1f115f722ff56ec301d493de721e3121a8f22f0 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:04:49 -0700 Subject: [PATCH 18/26] Declared the byte bound and age trim as standalone capabilities. A store may bound a standalone stream by record count and age but not by bytes, or apply retention only once the stream is closed. The two flags let such a provider say so and have the suite hold it to what it declares. (cherry picked from commit 440a18dbdeab1f3c8eb217c910a5d437b41621d4) --- tests/streams/test_streams_conformance.py | 13 ++++++++++++- 1 file changed, 12 insertions(+), 1 deletion(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 4dc09cd98..9b4575913 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -96,6 +96,11 @@ class ProviderCase: 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.""" async def open( self, @@ -223,7 +228,13 @@ async def _memory_case(_client: Client) -> AsyncIterator[ProviderCase]: async def truncate(workflow_id: str, topic: str, keep: int) -> None: provider.truncate(workflow_id, topic, keep=keep) - yield ProviderCase("memory", provider, truncate=truncate) + yield ProviderCase( + "memory", + provider, + truncate=truncate, + bounds_standalone_bytes=True, + trims_open_stream_by_age=True, + ) provider.reset() From 1aceee767b1704f465e9874d462f6a9c881415b7 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Wed, 30 Sep 2026 03:06:35 -0700 Subject: [PATCH 19/26] Held the retention case to the declared capabilities and used the case's client. A provider that cannot bound a standalone stream by bytes is now expected to refuse the policy, and one that trims by age only after close skips that check. The converter cases derive their client from the provider's own, so a live provider is not asked to reach the fixture's dev server. (cherry picked from commit f5a7e1bf7d0eff8485ce0f0f99162bacc8f8e73c) --- tests/streams/test_streams_conformance.py | 29 ++++++++++++++++------- 1 file changed, 20 insertions(+), 9 deletions(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 9b4575913..c557b9b69 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -60,6 +60,7 @@ StreamProducerError, StreamProvider, StreamRef, + StreamUnsupportedError, Supersession, topic, ) @@ -636,7 +637,9 @@ async def test_a_body_above_the_threshold_is_offloaded_and_read_back( external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=256), ) workflow_id = new_workflow_id() - stream = await case.open(workflow_id, client=_client_with(client, converter)) + 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} @@ -658,7 +661,9 @@ async def test_a_retry_through_a_nondeterministic_codec_still_deduplicates( codec = NonceCodec() converter = dataclasses.replace(DataConverter.default, payload_codec=codec) workflow_id = new_workflow_id() - stream = await case.open(workflow_id, client=_client_with(client, converter)) + 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 @@ -797,13 +802,19 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas kept = await take(by_count.read(topic=OUT), 2) assert [r.value for r in kept] == [{"n": 3}, {"n": 4}] - 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 - + 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) ) From 41ddb61cce3c3f6c8628afb91e9c60280817d8ea Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 13:25:51 -0700 Subject: [PATCH 20/26] Added a conformance case for a publish from the constructor. --- tests/streams/test_streams_conformance.py | 38 +++++++++++++++++++++++ 1 file changed, 38 insertions(+) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index c557b9b69..195950f82 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -36,6 +36,7 @@ import pytest +from temporalio import workflow from temporalio.api.common.v1 import Payload from temporalio.client import Client from temporalio.common import RawValue @@ -66,6 +67,8 @@ ) from temporalio.streams._ref import RefHandle, open_ref from temporalio.streams.providers.memory import MemoryStreams +from temporalio.testing import WorkflowEnvironment +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. @@ -824,3 +827,38 @@ async def test_a_standalone_stream_honors_its_retention_policy(case: ProviderCas 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"}] From c0a0d188ebfd94cf62732ff598ab7fa960097202 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 19:01:23 -0700 Subject: [PATCH 21/26] Added a conformance case for a reader woken through the channel. Gated by a wakes_by_notification capability, which the memory provider lacks since its reader polls on a timer. --- tests/streams/conftest.py | 5 ++ tests/streams/test_streams_conformance.py | 79 +++++++++++++++++++++++ 2 files changed, 84 insertions(+) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index f03dbc919..107134d94 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -20,3 +20,8 @@ def pytest_configure(config: pytest.Config) -> None: "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", + ) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 195950f82..f8a128fae 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -40,6 +40,7 @@ from temporalio.api.common.v1 import Payload from temporalio.client import Client from temporalio.common import RawValue +from temporalio.contrib.external_workflow_streams._wake import server_has_channels from temporalio.converter import ( DataConverter, ExternalStorage, @@ -105,6 +106,9 @@ class ProviderCase: 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.""" async def open( self, @@ -251,6 +255,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: "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, } @@ -862,3 +867,77 @@ async def test_a_publish_from_the_constructor_is_delivered( 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 +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 subscribes to the stream's channel on the task that opens + the reader, the producer's append notifies that 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 + if not await server_has_channels(worker_client): + pytest.skip("the server does not implement notification channels") + plugins = [] if case.client is not None else [case.provider] + workflow_id = new_workflow_id() + async with new_worker(worker_client, ReadUntilFinished, plugins=plugins) as worker: + handle = await worker_client.start_workflow( + ReadUntilFinished.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) + await producer.finish() + assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}, {"n": 2}] + + events = [e async for e in handle.fetch_history_events()] + signalled = [ + e for e in events if e.HasField("workflow_execution_signaled_event_attributes") + ] + assert signalled == [], ( + "a Signal woke the reader, so the channel was not the transport" + ) + subscribed = [ + e + for e in events + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + assert len(subscribed) == 1, "the run subscribes once per channel" + notified = [ + notification + for e in events + if e.HasField("workflow_task_scheduled_event_attributes") + for notification in e.workflow_task_scheduled_event_attributes.notifications + ] + assert notified, "no Workflow Task was scheduled with a notification" + channel = subscribed[ + 0 + ].workflow_notification_channel_subscribed_event_attributes.channel + assert {n.channel for n in notified} == {channel} From 58c87ba50a13e4523ba2b1fc741cf4975d7f943e Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 20:02:55 -0700 Subject: [PATCH 22/26] Marked the channel conformance case for a channel server. The marker skips it unless -E names a server, the way the main chain gates its live channel cases; the memory provider skips it by capability as well. --- tests/streams/conftest.py | 23 +++++++++++++++++++++++ tests/streams/test_streams_conformance.py | 1 + 2 files changed, 24 insertions(+) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 107134d94..b891540fc 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -25,3 +25,26 @@ def pytest_configure(config: pytest.Config) -> None: "wakes_by_notification: the case needs an outside append to wake a parked " "workflow reader through the server", ) + config.addinivalue_line( + "markers", + "needs_channel_server: the case needs a server that serves notification " + "channels, named with -E host:port", + ) + + +#: 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") + + +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 item.get_closest_marker("needs_channel_server"): + item.add_marker(skip) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index f8a128fae..5e4408b62 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -888,6 +888,7 @@ async def run(self) -> list[Any]: @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 ): From fb67a8f949e2e98ad44b5c08063fea50f2889821 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Thu, 1 Oct 2026 23:29:48 -0700 Subject: [PATCH 23/26] Added a conformance case for a reader woken through its linked channel. --- tests/streams/conftest.py | 15 ++++- tests/streams/test_streams_conformance.py | 82 ++++++++++++++++++----- 2 files changed, 80 insertions(+), 17 deletions(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index b891540fc..53aa935bd 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -25,11 +25,22 @@ def pytest_configure(config: pytest.Config) -> None: "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", + ) #: The environments whose server the suite starts for itself. None of them @@ -46,5 +57,7 @@ def pytest_collection_modifyitems( reason="needs a server that serves notification channels; name one with -E" ) for item in items: - if item.get_closest_marker("needs_channel_server"): + if item.get_closest_marker("needs_channel_server") or item.get_closest_marker( + "needs_linked_server" + ): item.add_marker(skip) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 5e4408b62..de62a66fe 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -40,7 +40,7 @@ from temporalio.api.common.v1 import Payload from temporalio.client import Client from temporalio.common import RawValue -from temporalio.contrib.external_workflow_streams._wake import server_has_channels +from temporalio.contrib.external_workflow_streams._wake import ChannelSupport from temporalio.converter import ( DataConverter, ExternalStorage, @@ -69,6 +69,7 @@ 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 @@ -109,6 +110,9 @@ class ProviderCase: 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, @@ -256,6 +260,7 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: "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, } @@ -900,8 +905,54 @@ async def test_an_outside_producer_wakes_the_reader_through_the_channel( notification. History then holds the subscription and no Signal. """ worker_client = case.client or client - if not await server_has_channels(worker_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" + 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" + assert {n.linked_to.workflow_id for n in notified} == {handle.id} + assert len({n.channel for n in notified}) == 1 + + +async def _read_two_woken_from_outside(case: ProviderCase, worker_client: Client): + """Runs the reader with two appends spaced past its idle timeout.""" plugins = [] if case.client is not None else [case.provider] workflow_id = new_workflow_id() async with new_worker(worker_client, ReadUntilFinished, plugins=plugins) as worker: @@ -917,28 +968,27 @@ async def test_an_outside_producer_wakes_the_reader_through_the_channel( await asyncio.sleep(2) await producer.finish() assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}, {"n": 2}] + return handle - events = [e async for e in handle.fetch_history_events()] - signalled = [ + +def _signalled(events: Sequence[Any]) -> list[Any]: + return [ e for e in events if e.HasField("workflow_execution_signaled_event_attributes") ] - assert signalled == [], ( - "a Signal woke the reader, so the channel was not the transport" - ) - subscribed = [ - e + + +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") ] - assert len(subscribed) == 1, "the run subscribes once per channel" - notified = [ + + +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 ] - assert notified, "no Workflow Task was scheduled with a notification" - channel = subscribed[ - 0 - ].workflow_notification_channel_subscribed_event_attributes.channel - assert {n.channel for n in notified} == {channel} From 790f8337834f8e58c6c65c804badaabfb0ac78b8 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:48:34 -0700 Subject: [PATCH 24/26] Added conformance cases for a reader leaving its channel. A reader closed short of FINISH and one closed at FINISH both end the run's subscription on the completion that leaves, after the progress marker, and a reader opened and closed inside one task never subscribes. The cases need a server that accepts the unsubscribe command, so they carry a marker of their own. --- tests/streams/conftest.py | 17 ++- tests/streams/test_streams_conformance.py | 177 +++++++++++++++++++++- 2 files changed, 186 insertions(+), 8 deletions(-) diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 53aa935bd..4dc059e13 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -41,12 +41,25 @@ def pytest_configure(config: pytest.Config) -> None: "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] @@ -57,7 +70,5 @@ def pytest_collection_modifyitems( reason="needs a server that serves notification channels; name one with -E" ) for item in items: - if item.get_closest_marker("needs_channel_server") or item.get_closest_marker( - "needs_linked_server" - ): + if any(item.get_closest_marker(marker) for marker in _CHANNEL_SERVER_MARKERS): item.add_marker(skip) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index de62a66fe..749b7c3c5 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -951,13 +951,169 @@ async def test_an_outside_producer_wakes_the_reader_through_its_linked_channel( assert len({n.channel for n in notified}) == 1 -async def _read_two_woken_from_outside(case: ProviderCase, worker_client: Client): - """Runs the reader with two appends spaced past its idle timeout.""" +@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 _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] = [ + e + for e in events + if e.HasField("workflow_notification_channel_subscribed_event_attributes") + ] + attributes = leaving.workflow_notification_channel_unsubscribed_event_attributes + assert attributes.subscribed_event_id == joined.event_id + assert events[index - 1].HasField("marker_recorded_event_attributes"), ( + "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, ReadUntilFinished, plugins=plugins) as worker: + async with new_worker(worker_client, workflow_class, plugins=plugins) as worker: handle = await worker_client.start_workflow( - ReadUntilFinished.run, id=workflow_id, task_queue=worker.task_queue + 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) @@ -966,8 +1122,11 @@ async def _read_two_woken_from_outside(case: ProviderCase, worker_client: Client for n in (1, 2): await producer.append({"n": n}) await asyncio.sleep(2) - await producer.finish() + 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 @@ -985,6 +1144,14 @@ def _subscribed(events: Sequence[Any]) -> list[str]: ] +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 From 662eba9c78892b6b39dfb2946fddc5d2c8996d03 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 03:49:59 -0700 Subject: [PATCH 25/26] Checked that the subscription waits for the completion that ends the task. The channel case now finds the subscribed event after a marker, which is where Core puts it, so the task that opened the reader stayed retained. --- tests/streams/test_streams_conformance.py | 36 ++++++++++++++++++----- 1 file changed, 28 insertions(+), 8 deletions(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 749b7c3c5..d1c54fb17 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -899,8 +899,10 @@ async def test_an_outside_producer_wakes_the_reader_through_the_channel( ): """The channel path, on the public surface. - The reader's run subscribes to the stream's channel on the task that opens - the reader, the producer's append notifies that channel, and the server + 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. """ @@ -917,6 +919,9 @@ async def test_an_outside_producer_wakes_the_reader_through_the_channel( ) 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) @@ -1069,6 +1074,25 @@ async def test_a_channel_opened_and_closed_in_one_task_is_never_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" @@ -1077,14 +1101,10 @@ def _assert_the_channel_left_after_the_marker(events: Sequence[Any]) -> None: for i, e in enumerate(events) if e.HasField("workflow_notification_channel_unsubscribed_event_attributes") ] - [joined] = [ - e - for e in events - if e.HasField("workflow_notification_channel_subscribed_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 events[index - 1].HasField("marker_recorded_event_attributes"), ( + 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"), ( From 58baa2ad899f648a41e3acbd0ebad82c58fc49ef Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Fri, 2 Oct 2026 17:50:58 -0700 Subject: [PATCH 26/26] Read the linked owner as an execution in the conformance case. The notified scheduled event names the owner as a `temporal.api.common.v1.Execution`, so the case checks the type and the business id through the public helper. --- tests/streams/test_streams_conformance.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index d1c54fb17..a5c32adcb 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -39,7 +39,7 @@ from temporalio import workflow from temporalio.api.common.v1 import Payload from temporalio.client import Client -from temporalio.common import RawValue +from temporalio.common import Execution, ExecutionType, RawValue from temporalio.contrib.external_workflow_streams._wake import ChannelSupport from temporalio.converter import ( DataConverter, @@ -952,7 +952,10 @@ async def test_an_outside_producer_wakes_the_reader_through_its_linked_channel( assert _subscribed(events) == [], "the owner is the listener by construction" notified = _notified(events) assert notified, "no Workflow Task was scheduled with a notification" - assert {n.linked_to.workflow_id for n in notified} == {handle.id} + 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