From 19cbdd720b23a81a3153dea9028599251a25f5bd Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:06:52 -0700 Subject: [PATCH 1/3] Added a Nexus operation that consumes a stream through its channel. It registers a callback listener on the stream's notification channel and reads on each delivery, so a caller waits on a stream without polling. --- CHANGELOG.md | 5 + temporalio/streams/providers/nexus.py | 625 +++++++++++++++++++++++++- 2 files changed, 628 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 38673bb6d..8f8a2a95a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -87,6 +87,11 @@ to include examples, links to docs, or any other relevant information. which the handler maps onto the store's accessor for that owner. Configure the front with `data_converter=` to run a payload codec on the caller side, so records are encoded before they leave the process. +- **Experimental**: `temporalio.streams.providers.nexus.stream_consumer_operation` builds an + asynchronous Nexus operation that consumes a stream. It registers a callback listener on the + stream's notification channel, reads on each delivery and completes when the stream closes. + `temporalio.streams.providers.nexus_consumer_service` hosts it on its own and needs the + `streams-nexus` extra. ### Changed diff --git a/temporalio/streams/providers/nexus.py b/temporalio/streams/providers/nexus.py index 1cc60d236..f4ad1f9ce 100644 --- a/temporalio/streams/providers/nexus.py +++ b/temporalio/streams/providers/nexus.py @@ -57,33 +57,57 @@ :class:`temporalio.streams.StreamError` the store raised, when the handler named one, and as :class:`temporalio.service.RPCError` otherwise, never as an HTTP or urllib exception. + +The third piece is the consumer side of a Nexus operation. +:class:`StreamConsumerOperation`, built with :func:`stream_consumer_operation`, +is an asynchronous operation handler whose input is a ``StreamRef`` and whose +result is what a consume function folded the stream's records into. It does +not poll. On start it registers a callback listener on the stream's +notification channel, with a URL the hosting process serves and a header that +names the operation, and reads the stream from its start through the client's +provider, the front when the client carries one. The server posts every +notification on the channel to that URL, the handler reads from its cursor to +the head, and the close completes the operation through the caller's +completion callback. The delivery is what the server's channel library posts: +``POST`` to the listener's URL, the ``Notification`` as protobuf JSON with +``Content-Type: application/json``, the listener's own headers and the +channel's name in ``Temporal-Notification-Channel``. A process that hosts the +handler feeds each such request to :meth:`StreamConsumerOperation.deliver`; +``nexus_consumer_service`` in this package is the standalone one. """ from __future__ import annotations import asyncio +import email.utils import http.client +import inspect import json import logging import time import urllib.error import urllib.request +import uuid import weakref from collections import OrderedDict from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import datetime, timedelta, timezone from typing import Any, Generic, NoReturn, TypeVar, cast import nexusrpc import nexusrpc.handler +from google.protobuf import json_format from google.protobuf.message import DecodeError +import temporalio.api.notification.v1 import temporalio.client +import temporalio.common import temporalio.converter +import temporalio.nexus from temporalio.api.common.v1 import Payload from temporalio.api.operatorservice.v1 import ListNexusEndpointsRequest -from temporalio.client import Client, ClientConfig +from temporalio.client import Callback, Client, ClientConfig from temporalio.common import RawValue from temporalio.service import ConnectConfig, RPCError, RPCStatusCode, ServiceClient from temporalio.streams._errors import ( @@ -120,16 +144,24 @@ TemporalStreams, ) from temporalio.streams.providers._nexus_generated import StreamRef as WireStreamRef +from temporalio.workflow import Notification __all__ = [ + "NOTIFICATION_CHANNEL_HEADER", + "STREAM_CONSUMER_TOKEN_HEADER", + "ConsumerState", + "Delivery", "NexusProducer", "NexusStreamHandle", "NexusStreams", + "StreamConsumerOperation", "TemporalStreamsHandler", "WireStreamRef", + "stream_consumer_operation", ] T = TypeVar("T") +S = TypeVar("S") # The reference's identity on the handler, topic included, so one parked read # and one producer state key on exactly what the wire names. @@ -1240,3 +1272,592 @@ async def _endpoint_id(self, client: Client | None) -> str: ) self._resolved = found.endpoints[0].id return self._resolved + + +# --------------------------------------------------------------------------- +# A Nexus operation as a stream consumer. +# --------------------------------------------------------------------------- + +NOTIFICATION_CHANNEL_HEADER = "Temporal-Notification-Channel" +"""The header the server adds to a channel's callback delivery, naming the channel.""" + +STREAM_CONSUMER_TOKEN_HEADER = "Temporal-Stream-Consumer-Token" +"""The header a consumer's listener registration carries, naming the operation. + +The server sends a listener's own headers back on every delivery, so this is +how one delivery URL serves any number of operations. +""" + +_COMPLETION_TIMEOUT = timedelta(seconds=30) +_CLOSED_KEY = "closed" +_FAILURE_CONTENT_TYPE = "application/json" +_CONTENT_TYPE_HEADER = "Content-Type" +_STATE_HEADER = "Nexus-Operation-State" +_TOKEN_HEADER = "Nexus-Operation-Token" +_START_TIME_HEADER = "Nexus-Operation-Start-Time" + + +ConsumeFunction = Callable[[StreamRecord[Any], S], "S | Awaitable[S]"] +ChannelRule = Callable[ + [StreamRef], + "temporalio.client.ChannelAddress | tuple[str, temporalio.common.Execution | None]", +] + + +def _address(where: Any) -> tuple[str, temporalio.common.Execution | None]: + """The channel name and the owner a rule answered with, as an address or a pair.""" + channel = getattr(where, "channel", None) + if isinstance(channel, str): + return channel, getattr(where, "execution", None) + channel, owner = where + return channel, owner + + +@dataclass(frozen=True) +class ConsumerState(Generic[S]): + """What a consumer operation holds for one token, read-only. + + .. warning:: + This API is experimental and unstable. + """ + + ref: StreamRef + """The stream being consumed.""" + + channel: str + """The channel the listener is registered on.""" + + owner: temporalio.common.Execution | None + """The execution a linked channel belongs to, ``None`` for an independent one.""" + + listener_id: str + """The listener id the server assigned.""" + + cursor: Cursor + """Where the next read starts. :data:`BEGINNING` until a record was consumed.""" + + value: S + """What the consume function has folded the records into so far.""" + + reads: int + """How many read passes reached the store, the start's included.""" + + deliveries: int + """How many notifications were delivered for this operation.""" + + records: int + """How many records the consume function has been handed.""" + + +@dataclass(frozen=True) +class Delivery: + """What one delivered notification led to. + + .. warning:: + This API is experimental and unstable. + """ + + token: str | None + """The operation the delivery named, when the headers carried one.""" + + known: bool + """Whether this process holds that operation. A delivery for one it does not is ignored.""" + + read: bool + """Whether the delivery led to a read. One that brought nothing new does not.""" + + records: int + """How many records that read handed to the consume function.""" + + completed: bool + """Whether the delivery closed the operation.""" + + +@dataclass +class _Consumption(Generic[S]): + ref: StreamRef + channel: str + owner: temporalio.common.Execution | None + listener_id: str + callback_url: str | None + callback_headers: dict[str, str] + started: datetime + value: S + cursor: Cursor = BEGINNING + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + reads: int = 0 + deliveries: int = 0 + records: int = 0 + finished: bool = False + done: bool = False + opening: asyncio.Task[None] | None = None + + +def _header(headers: Mapping[str, str], name: str) -> str | None: + wanted = name.lower() + for key, value in headers.items(): + if key.lower() == wanted: + return value + return None + + +def _parse_notification(body: bytes | str | Mapping[str, Any]) -> Notification: + """The notification a delivery carries, from the JSON the server posts.""" + text = json.dumps(body) if isinstance(body, Mapping) else body + if isinstance(text, bytes): + text = text.decode() + proto = json_format.Parse( + text, temporalio.api.notification.v1.Notification(), ignore_unknown_fields=True + ) + return Notification._from_proto(proto) + + +def _payload_content(payload: Payload) -> tuple[dict[str, str], bytes]: + """The HTTP content headers and body that carry ``payload`` to the server. + + The mapping is the server's: a plain JSON payload travels as JSON, a null + one as an empty body with no type, and anything else, a codec's output + included, as the serialized payload itself. + """ + metadata = {key: value.decode() for key, value in payload.metadata.items()} + encoding = metadata.get("encoding") + if set(metadata) == {"encoding"}: + if encoding == "json/plain": + return {_CONTENT_TYPE_HEADER: "application/json"}, payload.data + if encoding == "binary/plain": + return {_CONTENT_TYPE_HEADER: "application/octet-stream"}, payload.data + if encoding == "binary/null": + return {}, b"" + return ( + {_CONTENT_TYPE_HEADER: "application/x-temporal-payload"}, + payload.SerializeToString(), + ) + + +def _post_completion( + url: str, body: bytes, headers: Mapping[str, str], timeout: timedelta +) -> None: + """Post an operation completion to the caller's callback URL. + + Unlike :func:`_post` this sets no content type of its own, because a null + result travels without one. + + Raises: + _EndpointFailure: The callback answered with an error or not at all. + """ + try: + request = urllib.request.Request( + url, data=body or None, headers=dict(headers), method="POST" + ) + with urllib.request.urlopen(request, timeout=timeout.total_seconds()): + return + except urllib.error.HTTPError as error: + raise _EndpointFailure( + error.read().decode(errors="replace"), error.code + ) from error + except urllib.error.URLError as error: + raise _EndpointFailure( + f"completion callback unreachable at {url}: {error.reason}", None + ) from error + except (TimeoutError, http.client.HTTPException, OSError, ValueError) as error: + raise _EndpointFailure( + f"completion callback at {url} did not answer: {error!r}", None + ) from error + + +class StreamConsumerOperation(nexusrpc.handler.OperationHandler[StreamRef, S]): + """An asynchronous operation that consumes the stream its input names. + + The caller starts it with a :class:`temporalio.streams.StreamRef` and + awaits its result. On start the handler registers a callback listener on + the stream's notification channel, with the URL the hosting process + serves and this operation's token in a header, and reads the stream from + its start through the client's provider. Every notification the server + delivers leads to one read from the cursor to the head, however many + writes the channel folded into it, and a delivery that brings nothing new + reads nothing. The operation completes through the caller's completion + callback when the notification says the stream closed, when the read + reaches ``FINISH`` or when the store ends the read, with what the consume + function folded the records into. Cancel unregisters the listener and + reports the operation canceled. + + The consume function is a reducer: it takes a record and the value so far + and returns the new value, awaitable or not. It is handed every record + kind and narrows on ``kind`` itself. An exception from it fails the + operation. The cursor and the value live in this process keyed by the + operation token, so a process that loses them loses the operation; a + durable cursor is a follow-up. + + .. warning:: + This API is experimental and unstable. + """ + + def __init__( + self, + consume: ConsumeFunction[S], + *, + initial: Callable[[], S], + listener_url: str, + client: Client | None = None, + channel_for: ChannelRule = temporalio.client.stream_channel, + result_type: type | None = None, + token_header: str = STREAM_CONSUMER_TOKEN_HEADER, + ) -> None: + """Consume with ``consume``, served at ``listener_url``. + + Args: + consume: The reducer each record is handed with the value so far. + initial: Makes the value an operation starts with. + listener_url: Where the server posts the channel's notifications. + The hosting process serves it and feeds each request to + :meth:`deliver`. + client: Registers the listener and opens the stream. Leave it + unset in a handler hosted by a Temporal worker, where the + operation context's client is used. + channel_for: Names the channel for a ref and the execution a + linked channel belongs to, a workflow or a standalone + activity, as a :class:`temporalio.client.ChannelAddress` or a + ``(channel, execution)`` pair with ``None`` for an + independent channel. The default is + :func:`temporalio.client.stream_channel`, the server's rule + for native streams; a store that names channels its own + way passes its rule. + result_type: What record values are decoded as. + token_header: The header carrying the operation token on the + registration and so on every delivery. + """ + if not listener_url: + raise ValueError("listener_url must not be empty") + self._consume = consume + self._initial = initial + self._listener_url = listener_url + self._given_client = client + self._channel_for = channel_for + self._result_type = result_type + self._token_header = token_header + self._states: dict[str, _Consumption[S]] = {} + + @property + def tokens(self) -> list[str]: + """The operations this process holds, oldest first.""" + return list(self._states) + + def state(self, token: str) -> ConsumerState[S] | None: + """What is held for ``token``, or ``None`` once it completed or never was.""" + held = self._states.get(token) + if held is None: + return None + return ConsumerState( + ref=held.ref, + channel=held.channel, + owner=held.owner, + listener_id=held.listener_id, + cursor=held.cursor, + value=held.value, + reads=held.reads, + deliveries=held.deliveries, + records=held.records, + ) + + def _client(self) -> Client: + if self._given_client is not None: + return self._given_client + return temporalio.nexus.client() + + async def start( + self, ctx: nexusrpc.handler.StartOperationContext, input: StreamRef + ) -> nexusrpc.handler.StartOperationResultAsync: + """Register as the stream's listener and start reading it. + + The read runs after the start answers, so a long stream does not hold + the start request, and the first delivery waits on it. + """ + client = self._client() + token = uuid.uuid4().hex + channel, owner = _address(self._channel_for(input)) + # Registered before the first read: a record appended between the + # read and the registration would otherwise be missed, where one + # appended between the registration and the read is read twice at + # worst, once by the read and once by the delivery it provokes, and + # the cursor makes the second a no-op. + listener_id = await client.register_channel_listener( + channel, + Callback(url=self._listener_url, headers={self._token_header: token}), + execution=owner, + ) + state: _Consumption[S] = _Consumption( + ref=input, + channel=channel, + owner=owner, + listener_id=listener_id, + callback_url=ctx.callback_url, + callback_headers=dict(ctx.callback_headers), + started=datetime.now(timezone.utc), + value=self._initial(), + ) + self._states[token] = state + state.opening = asyncio.create_task(self._open(token, state, client)) + return nexusrpc.handler.StartOperationResultAsync(token) + + async def _open(self, token: str, state: _Consumption[S], client: Client) -> None: + try: + async with state.lock: + if state.done: + return + await self._catch_up(token, state, client) + if state.finished and not state.done: + await self._complete(token, state, client) + except Exception: + # The next delivery reads again from the cursor, so a failure here + # costs nothing but the records it would have consumed early. + logger.exception("the first read of %s failed", state.ref) + + async def deliver( + self, headers: Mapping[str, str], body: bytes | str | Mapping[str, Any] + ) -> Delivery: + """Act on one notification the server posted to the listener URL. + + ``headers`` are the request's, in any case; ``body`` is the + ``Notification`` as protobuf JSON, raw or already parsed. The delivery + is matched to its operation by the token header. Reading from the + cursor to the head happens under the operation's lock, so a delivery + that arrives during the first read waits for it and then finds the + cursor at the head. + + A failure to read raises, so the hosting process answers the server + with an error and the server retries the delivery; a failure in the + consume function fails the operation instead. + """ + token = _header(headers, self._token_header) + state = self._states.get(token) if token is not None else None + if token is None or state is None: + logger.warning( + "ignoring a channel delivery for an operation this process does not hold" + ) + return Delivery(token, known=False, read=False, records=0, completed=False) + client = self._client() + notification = _parse_notification(body) + closed = await self._closed(notification, client) + async with state.lock: + if state.done: + return Delivery(token, True, read=False, records=0, completed=True) + state.deliveries += 1 + reads_before = state.reads + records = await self._catch_up(token, state, client) + completed = False + if (closed or state.finished) and not state.done: + await self._complete(token, state, client) + completed = True + return Delivery( + token, + True, + read=state.reads > reads_before, + records=records, + completed=completed, + ) + + async def _closed(self, notification: Notification, client: Client) -> bool: + payload = notification.metadata.get(_CLOSED_KEY) + if payload is None: + return False + try: + converter = client.data_converter + if converter.payload_codec is not None: + [payload] = await converter.payload_codec.decode([payload]) + return bool(converter.payload_converter.from_payload(payload, bool)) + except Exception: + logger.warning( + "could not decode the notification's closed flag", exc_info=True + ) + return False + + async def _catch_up( + self, token: str, state: _Consumption[S], client: Client + ) -> int: + """Read from the cursor to the head, consuming, and return how many records.""" + handle = client.get_stream_handle(state.ref) + head = await handle.latest() + if head == state.cursor: + return 0 + state.reads += 1 + source = handle.read(after=state.cursor, result_type=self._result_type) + count = 0 + exhausted = True + try: + async for record in source: + try: + value = self._consume(record, state.value) + if inspect.isawaitable(value): + value = await value + except Exception as error: + await self._complete(token, state, client, error=error) + return count + state.value = cast(S, value) + state.cursor = record.cursor + state.records += 1 + count += 1 + if record.kind is RecordKind.FINISH: + state.finished = True + exhausted = False + break + if record.cursor == head: + exhausted = False + break + finally: + await source.aclose() + if exhausted: + # The store ended the read: nothing more will arrive on it. + state.finished = True + return count + + async def _complete( + self, + token: str, + state: _Consumption[S], + client: Client, + *, + error: BaseException | None = None, + ) -> None: + state.done = True + if error is None: + await self._post_result(token, state, client) + else: + await self._post_failure( + token, state, "failed", f"{type(error).__name__}: {error}" + ) + await self._unregister(client, state) + self._states.pop(token, None) + + async def cancel( + self, ctx: nexusrpc.handler.CancelOperationContext, token: str + ) -> None: + """Unregister the listener and report the operation canceled. + + Raises: + nexusrpc.HandlerError: ``NOT_FOUND`` when this process holds no + operation under ``token``. + """ + del ctx + state = self._states.get(token) + if state is None: + raise nexusrpc.HandlerError( + "no stream consumer operation is held under this token", + type=nexusrpc.HandlerErrorType.NOT_FOUND, + retryable_override=False, + ) + client = self._client() + if state.opening is not None and not state.opening.done(): + state.opening.cancel() + async with state.lock: + if state.done: + return + state.done = True + await self._unregister(client, state) + await self._post_failure( + token, state, "canceled", "the caller canceled the stream consumer" + ) + self._states.pop(token, None) + + async def close(self) -> None: + """Unregister every listener this process still holds, completing nothing. + + Call it when the hosting process stops, so the server does not keep + posting to a URL nobody serves. + """ + client = self._given_client + for token, state in list(self._states.items()): + if state.opening is not None and not state.opening.done(): + state.opening.cancel() + state.done = True + if client is not None: + await self._unregister(client, state) + self._states.pop(token, None) + + async def _unregister(self, client: Client, state: _Consumption[S]) -> None: + try: + await client.unregister_channel_listener( + state.channel, state.listener_id, execution=state.owner + ) + except RPCError: + logger.warning( + "could not unregister listener %s from channel %s", + state.listener_id, + state.channel, + exc_info=True, + ) + + async def _post_result( + self, token: str, state: _Consumption[S], client: Client + ) -> None: + converter = client.data_converter + [payload] = converter.payload_converter.to_payloads([state.value]) + if converter.payload_codec is not None: + [payload] = await converter.payload_codec.encode([payload]) + headers, body = _payload_content(payload) + await self._post(token, state, "succeeded", headers, body) + + async def _post_failure( + self, token: str, state: _Consumption[S], outcome: str, message: str + ) -> None: + await self._post( + token, + state, + outcome, + {_CONTENT_TYPE_HEADER: _FAILURE_CONTENT_TYPE}, + json.dumps({"message": message}).encode(), + ) + + async def _post( + self, + token: str, + state: _Consumption[S], + outcome: str, + content: Mapping[str, str], + body: bytes, + ) -> None: + if not state.callback_url: + logger.warning( + "operation %s %s with no completion callback to tell", token, outcome + ) + return + headers = { + **state.callback_headers, + _TOKEN_HEADER: token, + _STATE_HEADER: outcome, + _START_TIME_HEADER: email.utils.format_datetime(state.started, usegmt=True), + **content, + } + try: + await asyncio.to_thread( + _post_completion, state.callback_url, body, headers, _COMPLETION_TIMEOUT + ) + except _EndpointFailure: + logger.exception("could not complete operation %s as %s", token, outcome) + + +def stream_consumer_operation( + consume: ConsumeFunction[S], + *, + initial: Callable[[], S], + listener_url: str, + client: Client | None = None, + channel_for: ChannelRule = temporalio.client.stream_channel, + result_type: type | None = None, + token_header: str = STREAM_CONSUMER_TOKEN_HEADER, +) -> StreamConsumerOperation[S]: + """Turn ``consume`` into an asynchronous operation handler that consumes a stream. + + A service handler returns the result from an operation handler factory, + and the hosting process feeds the deliveries it receives at + ``listener_url`` to the same object's :meth:`StreamConsumerOperation.deliver`. + See :class:`StreamConsumerOperation` for the arguments and the contract. + """ + return StreamConsumerOperation( + consume, + initial=initial, + listener_url=listener_url, + client=client, + channel_for=channel_for, + result_type=result_type, + token_header=token_header, + ) From 164613216f5d841afdf1d91280e63dc1ed187a40 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:06:52 -0700 Subject: [PATCH 2/3] Added a standalone aiohttp host for the stream consumer operation. A process that is not a worker can serve the operation and receive the channel's callbacks. aiohttp comes in through the streams-nexus extra. --- pyproject.toml | 1 + .../providers/nexus_consumer_service.py | 509 ++++++++++++++++++ uv.lock | 6 +- 3 files changed, 515 insertions(+), 1 deletion(-) create mode 100644 temporalio/streams/providers/nexus_consumer_service.py diff --git a/pyproject.toml b/pyproject.toml index fcf82b526..f9baf3359 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,6 +57,7 @@ cloud-run-worker-otel = [ aioboto3 = ["aioboto3>=10.4.0", "types-aioboto3[s3]>=10.4.0"] google-genai = ["google-genai>=2.21.0,<3.0.0"] strands-agents = ["strands-agents>=1.51.0"] +streams-nexus = ["aiohttp>=3.9,<4"] [project.urls] Homepage = "https://github.com/temporalio/sdk-python" diff --git a/temporalio/streams/providers/nexus_consumer_service.py b/temporalio/streams/providers/nexus_consumer_service.py new file mode 100644 index 000000000..0aea10d46 --- /dev/null +++ b/temporalio/streams/providers/nexus_consumer_service.py @@ -0,0 +1,509 @@ +r"""A standalone Nexus handler process that consumes a stream through its channel. + +The Temporal frontend forwards a caller's ``StartOperation`` to an endpoint +whose target is an external URL, over the Nexus HTTP protocol. This module +serves that protocol with aiohttp for one :class:`nexusrpc.handler.Handler`, +start and cancel, and one more route the server's notification channels post +to, which it hands to the :class:`StreamConsumerOperation` instances it was +given. It is the shape the PoC demonstrates: a process with a URL of its own, +no worker, consuming a stream the server tells it about. + +Run it with one command, against a server that serves channels and a stream +front endpoint named ``streams``:: + + uv run python -m temporalio.streams.providers.nexus_consumer_service \\ + --address 127.0.0.1:7813 --http http://127.0.0.1:7823 \\ + --endpoint streams --port 8813 --register-endpoint consumers + +``--register-endpoint`` creates a Nexus endpoint whose target is this process, +so a workflow calls ``CollectStream.collect`` on it with a ``StreamRef`` and +gets the stream's values back when the stream closes. The server has to allow +this host in ``callback.allowedAddresses``. + +The wire details are the Nexus SDK's: a start is ``POST /{service}/{operation}`` +with the input as the body, the caller's completion callback in the ``callback`` +query parameter and its headers under the ``Nexus-Callback-`` prefix; an +asynchronous start answers ``201`` with the operation token; a cancel is +``POST /{service}/{operation}/cancel`` with the token in +``Nexus-Operation-Token`` and answers ``202``; a handler error answers the +status its type maps to with a JSON failure body. The body of a start is +turned into a payload the way the server does it, by content type, and a +synchronous result goes back the same way. +""" + +from __future__ import annotations + +import argparse +import asyncio +import logging +import re +import uuid +from collections.abc import Mapping, Sequence +from datetime import datetime, timedelta, timezone +from typing import Any + +import nexusrpc +import nexusrpc.handler +from aiohttp import web + +import temporalio.converter +from temporalio.api.common.v1 import Payload +from temporalio.api.nexus.v1 import EndpointSpec, EndpointTarget +from temporalio.api.operatorservice.v1 import CreateNexusEndpointRequest +from temporalio.client import Client +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._record import RecordKind, StreamRecord +from temporalio.streams._ref import StreamRef +from temporalio.streams.providers.nexus import ( + NexusStreams, + StreamConsumerOperation, + _payload_content, + stream_consumer_operation, +) + +__all__ = [ + "CollectStream", + "CollectStreamHandler", + "NexusHttpService", + "collect_values", + "main", +] + +logger = logging.getLogger(__name__) + +_CALLBACK_PREFIX = "nexus-callback-" +_CONTENT_PREFIX = "content-" +_REQUEST_ID_HEADER = "nexus-request-id" +_REQUEST_TIMEOUT_HEADER = "request-timeout" +_OPERATION_TOKEN_HEADER = "nexus-operation-token" +_OPERATION_STATE_HEADER = "nexus-operation-state" +_RETRYABLE_HEADER = "nexus-request-retryable" +_CALLBACK_QUERY = "callback" +_TOKEN_QUERY = "token" +_DEFAULT_DELIVERIES_PATH = "/deliveries" + +# The status each handler error type answers with, per the Nexus HTTP spec. +_ERROR_STATUS = { + nexusrpc.HandlerErrorType.BAD_REQUEST: 400, + nexusrpc.HandlerErrorType.UNAUTHENTICATED: 401, + nexusrpc.HandlerErrorType.UNAUTHORIZED: 403, + nexusrpc.HandlerErrorType.NOT_FOUND: 404, + nexusrpc.HandlerErrorType.REQUEST_TIMEOUT: 408, + nexusrpc.HandlerErrorType.CONFLICT: 409, + nexusrpc.HandlerErrorType.RESOURCE_EXHAUSTED: 429, + nexusrpc.HandlerErrorType.INTERNAL: 500, + nexusrpc.HandlerErrorType.NOT_IMPLEMENTED: 501, + nexusrpc.HandlerErrorType.UNAVAILABLE: 503, + nexusrpc.HandlerErrorType.UPSTREAM_TIMEOUT: 520, +} +_OPERATION_FAILED_STATUS = 424 + +_DURATION = re.compile(r"(\d+(?:\.\d+)?)(ns|us|µs|ms|s|m|h)") +_UNIT_SECONDS = { + "ns": 1e-9, + "us": 1e-6, + "µs": 1e-6, + "ms": 1e-3, + "s": 1.0, + "m": 60.0, + "h": 3600.0, +} + + +class _NeverCancelled(nexusrpc.handler.OperationTaskCancellation): + """The task cancellation of a request no worker is going to cancel.""" + + def is_cancelled(self) -> bool: + return False + + def cancellation_reason(self) -> str | None: + return None + + def wait_until_cancelled_sync(self, timeout: float | None = None) -> bool: + del timeout + return False + + async def wait_until_cancelled(self) -> None: + await asyncio.Event().wait() + + +def _deadline(timeout: str | None) -> datetime | None: + """The request deadline a Go duration string in ``Request-Timeout`` names.""" + if not timeout: + return None + seconds = 0.0 + matched = False + for amount, unit in _DURATION.findall(timeout): + seconds += float(amount) * _UNIT_SECONDS[unit] + matched = True + if not matched: + return None + return datetime.now(timezone.utc) + timedelta(seconds=seconds) + + +def _media_type(content_type: str) -> tuple[str, dict[str, str]]: + media, _, rest = content_type.partition(";") + params: dict[str, str] = {} + for part in rest.split(";"): + key, _, value = part.strip().partition("=") + if key: + params[key.strip().lower()] = value.strip().strip('"') + return media.strip().lower(), params + + +def _payload_from_content(headers: Mapping[str, str], body: bytes) -> Payload: + """The payload a start body is, mapped the way the server maps Nexus content.""" + content = { + key[len(_CONTENT_PREFIX) :]: value + for key, value in headers.items() + if key.startswith(_CONTENT_PREFIX) + } + content_type = content.pop("type", "") + content.pop("length", None) + if not content_type: + if not content and not body: + return Payload(metadata={"encoding": b"binary/null"}) + return _unknown(content, body) + media, params = _media_type(content_type) + if media == "application/x-temporal-payload": + return Payload.FromString(body) + if media == "application/json": + if params.get("format") == "protobuf" and params.get("message-type"): + return Payload( + metadata={ + "encoding": b"json/protobuf", + "messageType": params["message-type"].encode(), + }, + data=body, + ) + return Payload(metadata={"encoding": b"json/plain"}, data=body) + if media == "application/x-protobuf" and params.get("message-type"): + return Payload( + metadata={ + "encoding": b"binary/protobuf", + "messageType": params["message-type"].encode(), + }, + data=body, + ) + if media == "application/octet-stream": + return Payload(metadata={"encoding": b"binary/plain"}, data=body) + content["type"] = content_type + return _unknown(content, body) + + +def _unknown(content: Mapping[str, str], body: bytes) -> Payload: + metadata = {key: value.encode() for key, value in content.items()} + metadata["encoding"] = b"unknown/nexus-content" + return Payload(metadata=metadata, data=body) + + +class _PayloadSerializer: + """Hands a start's payload to the handler through the data converter.""" + + def __init__( + self, converter: temporalio.converter.DataConverter, payload: Payload + ) -> None: + self._converter = converter + self._payload = payload + + async def serialize(self, value: Any) -> nexusrpc.Content: + del value + raise NotImplementedError("the service serializes results itself") + + async def deserialize( + self, content: nexusrpc.Content, as_type: type[Any] | None = None + ) -> Any: + del content + payload = self._payload + if self._converter.payload_codec is not None: + [payload] = await self._converter.payload_codec.decode([payload]) + try: + [value] = self._converter.payload_converter.from_payloads( + [payload], [as_type] if as_type is not None else None + ) + except Exception as error: + raise nexusrpc.HandlerError( + f"invalid operation input: {error}", + type=nexusrpc.HandlerErrorType.BAD_REQUEST, + retryable_override=False, + ) from error + return value + + +def _failure(error: nexusrpc.HandlerError) -> web.Response: + headers = {} + if error.retryable_override is not None: + headers[_RETRYABLE_HEADER] = "true" if error.retryable_override else "false" + error_type = ( + error.type + if isinstance(error.type, nexusrpc.HandlerErrorType) + else nexusrpc.HandlerErrorType.UNKNOWN + ) + return web.json_response( + {"message": error.message}, + status=_ERROR_STATUS.get(error_type, 500), + headers=headers, + ) + + +class NexusHttpService: + """Serves one Nexus handler over HTTP, plus the route channel deliveries reach. + + ``handler`` is the Nexus SDK's dispatcher over the service handlers the + process hosts. ``consumers`` are the stream consumer operations among + them; a delivery is offered to each until one holds the operation it + names. ``data_converter`` turns start bodies into inputs and results + into bodies; the caller's converter when a client is at hand. + + .. warning:: + This API is experimental and unstable. + """ + + def __init__( + self, + handler: nexusrpc.handler.Handler, + consumers: Sequence[StreamConsumerOperation[Any]], + *, + data_converter: temporalio.converter.DataConverter | None = None, + deliveries_path: str = _DEFAULT_DELIVERIES_PATH, + ) -> None: + """Serve ``handler`` and route deliveries at ``deliveries_path`` to ``consumers``.""" + if deliveries_path.count("/") != 1 or not deliveries_path.startswith("/"): + raise ValueError( + "deliveries_path must be one path segment, so it cannot be mistaken " + "for a start request" + ) + self._handler = handler + self._consumers = list(consumers) + self._converter = data_converter or temporalio.converter.DataConverter.default + self._deliveries_path = deliveries_path + + @property + def deliveries_path(self) -> str: + """The route the channel's deliveries are served at.""" + return self._deliveries_path + + def application(self) -> web.Application: + """The aiohttp application serving the routes; run it with the usual runner.""" + app = web.Application() + app.add_routes( + [ + web.post(self._deliveries_path, self._deliver), + web.post("/{service}/{operation}", self._start), + web.post("/{service}/{operation}/cancel", self._cancel), + ] + ) + return app + + async def _start(self, request: web.Request) -> web.Response: + headers = {key.lower(): value for key, value in request.headers.items()} + body = await request.read() + ctx = nexusrpc.handler.StartOperationContext( + service=request.match_info["service"], + operation=request.match_info["operation"], + headers={ + key: value + for key, value in headers.items() + if not key.startswith((_CONTENT_PREFIX, _CALLBACK_PREFIX)) + }, + request_id=headers.get(_REQUEST_ID_HEADER) or uuid.uuid4().hex, + callback_url=request.query.get(_CALLBACK_QUERY), + callback_headers={ + key[len(_CALLBACK_PREFIX) :]: value + for key, value in headers.items() + if key.startswith(_CALLBACK_PREFIX) + }, + task_cancellation=_NeverCancelled(), + request_deadline=_deadline(headers.get(_REQUEST_TIMEOUT_HEADER)), + ) + input = nexusrpc.LazyValue( + serializer=_PayloadSerializer( + self._converter, _payload_from_content(headers, body) + ), + headers={}, + stream=None, + ) + try: + result = await self._handler.start_operation(ctx, input) + except nexusrpc.HandlerError as error: + return _failure(error) + except nexusrpc.OperationError as error: + return web.json_response( + {"message": str(error)}, + status=_OPERATION_FAILED_STATUS, + headers={_OPERATION_STATE_HEADER: error.state.value}, + ) + if isinstance(result, nexusrpc.handler.StartOperationResultAsync): + return web.json_response( + {"token": result.token, "state": "running"}, status=201 + ) + [payload] = self._converter.payload_converter.to_payloads([result.value]) + if self._converter.payload_codec is not None: + [payload] = await self._converter.payload_codec.encode([payload]) + content, data = _payload_content(payload) + return web.Response(status=200, body=data, headers=content) + + async def _cancel(self, request: web.Request) -> web.Response: + headers = {key.lower(): value for key, value in request.headers.items()} + token = headers.get(_OPERATION_TOKEN_HEADER) or request.query.get(_TOKEN_QUERY) + if not token: + return _failure( + nexusrpc.HandlerError( + "missing operation token", + type=nexusrpc.HandlerErrorType.BAD_REQUEST, + ) + ) + ctx = nexusrpc.handler.CancelOperationContext( + service=request.match_info["service"], + operation=request.match_info["operation"], + headers={ + key: value + for key, value in headers.items() + if key != _OPERATION_TOKEN_HEADER + }, + task_cancellation=_NeverCancelled(), + ) + try: + await self._handler.cancel_operation(ctx, token) + except nexusrpc.HandlerError as error: + return _failure(error) + return web.Response(status=202) + + async def _deliver(self, request: web.Request) -> web.Response: + body = await request.read() + headers = dict(request.headers) + try: + for consumer in self._consumers: + if (await consumer.deliver(headers, body)).known: + break + except Exception as error: + # A 5xx makes the server retry the delivery, which is what a + # failed read needs. + logger.exception("a channel delivery failed") + return web.json_response({"message": str(error)}, status=500) + return web.Response(status=200) + + +# --------------------------------------------------------------------------- +# The demo service: collect a stream's values into the operation's result. +# --------------------------------------------------------------------------- + + +def collect_values(record: StreamRecord[Any], values: list[Any]) -> list[Any]: + """The demo's consume function: keep every published value, in order.""" + if record.kind is RecordKind.DATA: + values.append(record.value) + return values + + +@nexusrpc.service +class CollectStream: + """Consumes the stream a ref names and answers with its values when it closes.""" + + collect: nexusrpc.Operation[StreamRef, list[Any]] + + +@nexusrpc.handler.service_handler(service=CollectStream) +class CollectStreamHandler: + """Serves :class:`CollectStream` out of one consumer operation.""" + + def __init__(self, consumer: StreamConsumerOperation[list[Any]]) -> None: + """Serve ``collect`` with ``consumer``, which the deliveries also reach.""" + self._consumer = consumer + + @nexusrpc.handler.operation_handler + def collect(self) -> nexusrpc.handler.OperationHandler[StreamRef, list[Any]]: + """The consumer operation, as the operation handler factory hands it out.""" + return self._consumer + + +async def _register_endpoint(client: Client, name: str, url: str) -> None: + try: + await client.operator_service.create_nexus_endpoint( + CreateNexusEndpointRequest( + spec=EndpointSpec( + name=name, + target=EndpointTarget(external=EndpointTarget.External(url=url)), + ) + ) + ) + logger.info("registered nexus endpoint %s -> %s", name, url) + except RPCError as error: + if error.status is not RPCStatusCode.ALREADY_EXISTS: + raise + logger.info("nexus endpoint %s exists; leaving it as it is", name) + + +async def main(argv: Sequence[str] | None = None) -> None: + """Serve the demo consumer until interrupted; see the module docstring.""" + parser = argparse.ArgumentParser(description=(__doc__ or "").split("\n\n")[0]) + parser.add_argument( + "--address", default="127.0.0.1:7233", help="the server's gRPC address" + ) + parser.add_argument("--namespace", default="default") + parser.add_argument( + "--http", + default="http://127.0.0.1:7243", + help="the server's Nexus HTTP ingress, where the stream front is read", + ) + parser.add_argument( + "--endpoint", required=True, help="the stream front's Nexus endpoint name" + ) + parser.add_argument("--host", default="127.0.0.1", help="the interface to serve on") + parser.add_argument("--port", type=int, default=8813, help="the port to serve on") + parser.add_argument( + "--listener-url", + default=None, + help="the URL the server posts deliveries to; defaults to this process's route", + ) + parser.add_argument( + "--register-endpoint", + default=None, + metavar="NAME", + help="create a Nexus endpoint of this name whose target is this process", + ) + args = parser.parse_args(argv) + logging.basicConfig( + level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s" + ) + + base_url = f"http://{args.host}:{args.port}" + client = await Client.connect( + args.address, + namespace=args.namespace, + plugins=[NexusStreams(endpoint=args.endpoint, http_address=args.http)], + ) + consumer = stream_consumer_operation( + collect_values, + initial=list, + listener_url=args.listener_url or base_url + _DEFAULT_DELIVERIES_PATH, + client=client, + ) + service = NexusHttpService( + nexusrpc.handler.Handler([CollectStreamHandler(consumer)]), + [consumer], + data_converter=client.data_converter, + ) + runner = web.AppRunner(service.application()) + await runner.setup() + try: + await web.TCPSite(runner, args.host, args.port).start() + if args.register_endpoint: + await _register_endpoint(client, args.register_endpoint, base_url) + logger.info( + "serving %s at %s, deliveries at %s", + CollectStream.__name__, + base_url, + consumer._listener_url, # pyright: ignore[reportPrivateUsage] + ) + await asyncio.Event().wait() + finally: + await consumer.close() + await runner.cleanup() + + +if __name__ == "__main__": + try: + asyncio.run(main()) + except KeyboardInterrupt: + pass diff --git a/uv.lock b/uv.lock index cee51de68..acd3814ef 100644 --- a/uv.lock +++ b/uv.lock @@ -4758,6 +4758,9 @@ opentelemetry = [ pydantic = [ { name = "pydantic" }, ] +streams-nexus = [ + { name = "aiohttp" }, +] strands-agents = [ { name = "strands-agents" }, ] @@ -4813,6 +4816,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "aioboto3", marker = "extra == 'aioboto3'", specifier = ">=10.4.0" }, + { name = "aiohttp", marker = "extra == 'streams-nexus'", specifier = ">=3.9,<4" }, { name = "deepagents", marker = "python_full_version >= '3.11' and extra == 'deepagents'", specifier = ">=0.6.12,<0.7" }, { name = "google-adk", marker = "extra == 'google-adk'", specifier = ">=2.8.0,<3" }, { name = "google-genai", marker = "extra == 'google-genai'", specifier = ">=2.21.0,<3.0.0" }, @@ -4844,7 +4848,7 @@ requires-dist = [ { name = "types-protobuf", specifier = ">=3.20,<8.0.0" }, { name = "typing-extensions", specifier = ">=4.2.0,<5" }, ] -provides-extras = ["grpc", "opentelemetry", "pydantic", "openai-agents", "google-adk", "langgraph", "langsmith", "deepagents", "lambda-worker-otel", "cloud-run-worker-otel", "aioboto3", "google-genai", "strands-agents"] +provides-extras = ["grpc", "opentelemetry", "pydantic", "openai-agents", "google-adk", "langgraph", "langsmith", "deepagents", "lambda-worker-otel", "cloud-run-worker-otel", "aioboto3", "google-genai", "strands-agents", "streams-nexus"] [package.metadata.requires-dev] dev = [ From 97afa38cc4868ddb6918174373d44c12066cbf37 Mon Sep 17 00:00:00 2001 From: Mohammad Dashti Date: Sat, 3 Oct 2026 02:06:52 -0700 Subject: [PATCH 3/3] Covered the Nexus stream consumer operation. Unit cases with a client double, and live cases that need a channel server named with -E. The native case stays gated until the union. --- tests/streams/conftest.py | 21 + tests/streams/test_nexus_consumer.py | 1033 ++++++++++++++++++++++++++ 2 files changed, 1054 insertions(+) create mode 100644 tests/streams/test_nexus_consumer.py diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 793bf2acd..9a26d2e73 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -2,6 +2,12 @@ _ENVIRONMENTS_WITHOUT_CHANNELS = ("local", "time-skipping", "envconfig") +# A native stream lives on the server, and only the native layers carry a +# provider that opens one; the providers here hold streams in memory or in a +# workflow's History. A case that produces to a native stream waits for the +# layer with that provider. +PINNED_LAYER_HOSTS_NATIVE_STREAMS = False + def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( @@ -58,6 +64,11 @@ def pytest_configure(config: pytest.Config) -> None: "channel by execution, a standalone activity's included, named with " "-E host:port", ) + config.addinivalue_line( + "markers", + "needs_native_provider: the case needs a provider that opens native streams " + "on the server, which only the native layers carry", + ) def pytest_collection_modifyitems( @@ -101,6 +112,16 @@ def pytest_collection_modifyitems( ), ) ) + if not PINNED_LAYER_HOSTS_NATIVE_STREAMS: + skips.append( + ( + "needs_native_provider", + pytest.mark.skip( + reason="this layer carries no provider for native streams; the " + "native layers and the union do" + ), + ) + ) for item in items: for marker, skip in skips: if item.get_closest_marker(marker): diff --git a/tests/streams/test_nexus_consumer.py b/tests/streams/test_nexus_consumer.py new file mode 100644 index 000000000..6726411f8 --- /dev/null +++ b/tests/streams/test_nexus_consumer.py @@ -0,0 +1,1033 @@ +"""The Nexus operation that consumes a stream through its notification channel. + +The unit cases drive :class:`StreamConsumerOperation` with a client stand-in +over the memory store and the deliveries a server would post, and catch what +it posts back as completions. The live cases need a server with notification +channels and a Nexus HTTP ingress, named with ``-E host:port`` and +``TEMPORAL_HTTP``; they put the memory store behind the Nexus front, because +this layer carries no store whose producers notify a channel themselves, so +the producer notifies the stream's channel by hand the way such a store +would. The consumer's HTTP port is ``STREAM_CONSUMER_PORT``, 8813 by default. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import email.utils +import json +import os +import uuid +from collections.abc import AsyncIterator, Callable, Mapping +from dataclasses import dataclass +from datetime import timedelta +from typing import Any, cast + +import nexusrpc +import nexusrpc.handler +import pytest +from aiohttp import web +from aiohttp.test_utils import TestClient, TestServer +from google.protobuf import json_format + +import temporalio.api.notification.v1 +import temporalio.converter +from temporalio import workflow +from temporalio.api.nexus.v1 import EndpointSpec, EndpointTarget +from temporalio.api.operatorservice.v1 import ( + CreateNexusEndpointRequest, + DeleteNexusEndpointRequest, +) +from temporalio.client import Callback, ChannelAddress, Client, stream_channel +from temporalio.common import Execution +from temporalio.exceptions import NexusOperationError +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import BEGINNING, StreamProvider, StreamRecord, StreamRef +from temporalio.streams._ref import open_ref +from temporalio.streams.providers import nexus +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.nexus import ( + STREAM_CONSUMER_TOKEN_HEADER, + NexusStreams, + StreamConsumerOperation, + TemporalStreamsHandler, + stream_consumer_operation, +) +from temporalio.streams.providers.nexus_consumer_service import ( + CollectStream, + CollectStreamHandler, + NexusHttpService, + _NeverCancelled, + collect_values, +) +from temporalio.worker import Worker +from tests.helpers import assert_eventually +from tests.streams.test_nexus_provider import _own_endpoint + +VALUES = "values" +LISTENER_URL = "http://127.0.0.1:8813/deliveries" +CALLBACK_URL = "http://caller.invalid/namespaces/default/nexus/callback" +CALLBACK_HEADERS = {"temporal-callback-token": "caller-token"} + +_HTTP = os.environ.get("TEMPORAL_HTTP", "http://127.0.0.1:7243") +_PORT = int(os.environ.get("STREAM_CONSUMER_PORT", "8813")) + + +# --------------------------------------------------------------------------- +# Stand-ins. +# --------------------------------------------------------------------------- + + +def _as_client(stand_in: object) -> Client: + """Pass a stand-in where a client is typed; the consumer uses only what both have.""" + return cast(Client, stand_in) + + +class _ClientStandIn: + """What the consumer asks of a client: a provider, a converter, the channel calls.""" + + def __init__(self, store: MemoryStreams) -> None: + self._store = store + self.data_converter = temporalio.converter.DataConverter.default + self.registered: list[tuple[str, Callback, Execution | None]] = [] + self.unregistered: list[tuple[str, str, Execution | None]] = [] + + def get_stream_handle(self, ref: StreamRef) -> Any: + # The memory store opens a handle without a client. + return open_ref(self._store, _as_client(None), ref) + + async def register_channel_listener( + self, channel: str, callback: Callback, *, execution: Execution | None = None + ) -> str: + self.registered.append((channel, callback, execution)) + return f"listener-{len(self.registered)}" + + async def unregister_channel_listener( + self, channel: str, listener_id: str, *, execution: Execution | None = None + ) -> None: + self.unregistered.append((channel, listener_id, execution)) + + +@dataclass +class _Posted: + url: str + headers: dict[str, str] + body: bytes + + @property + def state(self) -> str: + return self.headers["Nexus-Operation-State"] + + +@pytest.fixture +def posted(monkeypatch: pytest.MonkeyPatch) -> list[_Posted]: + """Catch the completions the consumer posts to the caller's callback.""" + caught: list[_Posted] = [] + + def record(url: str, body: bytes, headers: Mapping[str, str], timeout: Any) -> None: + del timeout + caught.append(_Posted(url, dict(headers), body)) + + monkeypatch.setattr(nexus, "_post_completion", record) + return caught + + +def _start_context( + callback_url: str | None = CALLBACK_URL, +) -> nexusrpc.handler.StartOperationContext: + return nexusrpc.handler.StartOperationContext( + service="CollectStream", + operation="collect", + headers={}, + request_id=uuid.uuid4().hex, + callback_url=callback_url, + callback_headers=dict(CALLBACK_HEADERS), + task_cancellation=_NeverCancelled(), + ) + + +def _cancel_context() -> nexusrpc.handler.CancelOperationContext: + return nexusrpc.handler.CancelOperationContext( + service="CollectStream", + operation="collect", + headers={}, + task_cancellation=_NeverCancelled(), + ) + + +def _notification( + channel: str, counter: int, *, position: bytes = b"", closed: bool = False +) -> bytes: + """The body the server posts: the notification as protobuf JSON.""" + proto = temporalio.api.notification.v1.Notification( + channel=channel, counter=counter, position=position + ) + if closed: + proto.metadata["closed"].CopyFrom( + temporalio.converter.DataConverter.default.payload_converter.to_payload( + True + ) + ) + return json_format.MessageToJson(proto).encode() + + +@dataclass +class _Rig: + store: MemoryStreams + client: _ClientStandIn + operation: StreamConsumerOperation[list[Any]] + ref: StreamRef + channel: str + producer: Any + + async def start(self) -> str: + result = await self.operation.start(_start_context(), self.ref) + assert isinstance(result, nexusrpc.handler.StartOperationResultAsync) + await self.settled(result.token) + return result.token + + async def settled(self, token: str) -> None: + """Wait for the read the start kicked off.""" + opening = self.operation._states[token].opening + assert opening is not None + await asyncio.wait_for(opening, 10) + + def headers(self, token: str) -> dict[str, str]: + return {STREAM_CONSUMER_TOKEN_HEADER: token} + + async def deliver(self, token: str, counter: int, *, closed: bool = False) -> Any: + return await self.operation.deliver( + self.headers(token), _notification(self.channel, counter, closed=closed) + ) + + +async def _rig( + consume: Callable[[StreamRecord[Any], list[Any]], list[Any]] = collect_values, +) -> _Rig: + store = MemoryStreams() + stream_id = f"consumed-{uuid.uuid4().hex}" + handle = await store.create_standalone_stream(None, stream_id) + ref = handle.ref(topic=VALUES) + client = _ClientStandIn(store) + operation = stream_consumer_operation( + consume, + initial=list, + listener_url=LISTENER_URL, + client=_as_client(client), + ) + return _Rig( + store, + client, + operation, + ref, + stream_channel(ref).channel, + handle.producer(topic=VALUES, producer_id="writer", attempt=1), + ) + + +# --------------------------------------------------------------------------- +# The helper's state machine. +# --------------------------------------------------------------------------- + + +def test_the_default_rule_is_the_servers_naming_of_a_streams_channel(): + # The consumer registers where the server's stream component notifies: + # a standalone stream's channel is independent, an owned stream's is + # linked to the owning workflow. + assert stream_channel(StreamRef.for_standalone("s-1", topic="values")) == ( + ChannelAddress("stream/s-1", None) + ) + assert stream_channel(StreamRef.for_workflow("wf", topic="t")) == ChannelAddress( + "stream/t", Execution.workflow("wf") + ) + assert stream_channel( + StreamRef.for_activity("act", workflow_id="wf", topic="t") + ) == ChannelAddress("stream/act/t", Execution.workflow("wf")) + # A standalone activity is an execution of its own, so its stream's + # channel is linked to it. + assert stream_channel(StreamRef.for_activity("act", topic="t")) == ( + ChannelAddress("stream/t", Execution.activity("act")) + ) + + +async def test_a_rule_may_answer_with_a_pair(posted: list[_Posted]): + # A store's rule that returns (channel, workflow_id) fits as well as an + # address does. + store = MemoryStreams() + handle = await store.create_standalone_stream(None, "paired") + client = _ClientStandIn(store) + operation = stream_consumer_operation( + collect_values, + initial=list, + listener_url=LISTENER_URL, + client=_as_client(client), + channel_for=lambda ref: (f"custom/{ref.stream_id}", None), + ) + result = await operation.start(_start_context(), handle.ref(topic=VALUES)) + assert isinstance(result, nexusrpc.handler.StartOperationResultAsync) + state = operation.state(result.token) + assert state is not None and (state.channel, state.owner) == ("custom/paired", None) + await operation.close() + assert client.unregistered == [("custom/paired", "listener-1", None)] + assert posted == [] + + +async def test_an_owned_stream_registers_on_the_owners_linked_channel( + posted: list[_Posted], +): + store = MemoryStreams() + ref = store.get_stream_handle(None, "wf-1").ref(topic=VALUES) + client = _ClientStandIn(store) + operation = stream_consumer_operation( + collect_values, + initial=list, + listener_url=LISTENER_URL, + client=_as_client(client), + ) + result = await operation.start(_start_context(), ref) + assert isinstance(result, nexusrpc.handler.StartOperationResultAsync) + state = operation.state(result.token) + owner = Execution.workflow("wf-1") + assert state is not None and (state.channel, state.owner) == ( + "stream/values", + owner, + ) + assert [(channel, owner) for channel, _, owner in client.registered] == [ + ("stream/values", owner) + ] + await operation.close() + assert client.unregistered == [("stream/values", "listener-1", owner)] + assert posted == [] + + +async def test_start_registers_the_listener_and_reads_what_was_there( + posted: list[_Posted], +): + rig = await _rig() + await rig.producer.append("a", "b") + token = await rig.start() + + assert rig.operation.tokens == [token] + assert rig.client.registered == [ + ( + rig.channel, + Callback(url=LISTENER_URL, headers={STREAM_CONSUMER_TOKEN_HEADER: token}), + None, + ) + ] + state = rig.operation.state(token) + assert state is not None + assert state.value == ["a", "b"] + assert (state.reads, state.deliveries, state.records) == (1, 0, 2) + assert state.cursor != BEGINNING + assert state.listener_id == "listener-1" + assert posted == [] + + +async def test_a_delivery_reads_from_the_cursor_once_and_in_order( + posted: list[_Posted], +): + rig = await _rig() + await rig.producer.append("a") + token = await rig.start() + + # Records after the registration arrive on the next delivery, each once. + await rig.producer.append("b", "c") + delivery = await rig.deliver(token, 1) + assert (delivery.known, delivery.read, delivery.records, delivery.completed) == ( + True, + True, + 2, + False, + ) + state = rig.operation.state(token) + assert state is not None and state.value == ["a", "b", "c"] + + # A delivery that brings nothing new reads nothing. + delivery = await rig.deliver(token, 1) + assert (delivery.read, delivery.records) == (False, 0) + state = rig.operation.state(token) + assert state is not None and (state.reads, state.deliveries) == (2, 2) + + # A burst the channel folded into one delivery costs one read. + await rig.producer.append("d", "e", "f") + delivery = await rig.operation.deliver( + {STREAM_CONSUMER_TOKEN_HEADER.lower(): token}, _notification(rig.channel, 4) + ) + assert (delivery.read, delivery.records) == (True, 3) + state = rig.operation.state(token) + assert state is not None + assert state.value == ["a", "b", "c", "d", "e", "f"] + assert (state.reads, state.deliveries, state.records) == (3, 3, 6) + assert posted == [] + + +async def test_the_close_completes_through_the_callers_callback( + posted: list[_Posted], +): + rig = await _rig() + await rig.producer.append("a") + token = await rig.start() + await rig.producer.append("b") + + delivery = await rig.deliver(token, 2, closed=True) + assert (delivery.read, delivery.records, delivery.completed) == (True, 1, True) + [completion] = posted + assert completion.url == CALLBACK_URL + assert completion.state == "succeeded" + assert completion.headers["temporal-callback-token"] == "caller-token" + assert completion.headers["Nexus-Operation-Token"] == token + assert completion.headers["Content-Type"] == "application/json" + assert email.utils.parsedate_to_datetime( + completion.headers["Nexus-Operation-Start-Time"] + ) + assert json.loads(completion.body) == ["a", "b"] + assert rig.client.unregistered == [(rig.channel, "listener-1", None)] + assert rig.operation.state(token) is None + assert rig.operation.tokens == [] + + # Once complete, the token is unknown and a late delivery is ignored. + delivery = await rig.deliver(token, 3) + assert (delivery.known, delivery.read) == (False, False) + assert len(posted) == 1 + + +async def test_a_finish_record_closes_the_operation(posted: list[_Posted]): + rig = await _rig() + token = await rig.start() + await rig.producer.append("only") + await rig.producer.finish() + + delivery = await rig.deliver(token, 1) + assert delivery.completed + [completion] = posted + assert completion.state == "succeeded" + assert json.loads(completion.body) == ["only"] + assert rig.client.unregistered == [(rig.channel, "listener-1", None)] + + +async def test_a_stream_finished_before_the_start_completes_at_once( + posted: list[_Posted], +): + rig = await _rig() + await rig.producer.append("a") + await rig.producer.finish() + token = await rig.start() + assert rig.operation.state(token) is None + [completion] = posted + assert completion.state == "succeeded" + assert json.loads(completion.body) == ["a"] + + +async def test_cancel_unregisters_and_reports_canceled(posted: list[_Posted]): + rig = await _rig() + token = await rig.start() + + await rig.operation.cancel(_cancel_context(), token) + assert rig.client.unregistered == [(rig.channel, "listener-1", None)] + [completion] = posted + assert completion.state == "canceled" + assert completion.headers["Nexus-Operation-Token"] == token + assert completion.headers["Content-Type"] == "application/json" + assert "canceled" in json.loads(completion.body)["message"] + assert rig.operation.state(token) is None + + with pytest.raises(nexusrpc.HandlerError) as unknown: + await rig.operation.cancel(_cancel_context(), token) + assert unknown.value.type is nexusrpc.HandlerErrorType.NOT_FOUND + + +async def test_a_delivery_for_an_operation_not_held_here_is_ignored( + posted: list[_Posted], +): + rig = await _rig() + await rig.producer.append("a") + delivery = await rig.operation.deliver({}, _notification(rig.channel, 1)) + assert (delivery.token, delivery.known) == (None, False) + delivery = await rig.deliver("nobody", 1) + assert (delivery.token, delivery.known, delivery.read) == ("nobody", False, False) + assert posted == [] + + +async def test_a_consume_failure_fails_the_operation(posted: list[_Posted]): + def fussy(record: StreamRecord[Any], values: list[Any]) -> list[Any]: + if record.value == "bad": + raise ValueError("cannot take bad") + return collect_values(record, values) + + rig = await _rig(fussy) + token = await rig.start() + await rig.producer.append("good", "bad", "later") + delivery = await rig.deliver(token, 1) + assert (delivery.records, delivery.completed) == (1, False) + [completion] = posted + assert completion.state == "failed" + assert json.loads(completion.body)["message"] == "ValueError: cannot take bad" + assert rig.client.unregistered == [(rig.channel, "listener-1", None)] + assert rig.operation.state(token) is None + + +async def test_a_read_failure_leaves_the_operation_for_the_retry( + posted: list[_Posted], monkeypatch: pytest.MonkeyPatch +): + rig = await _rig() + token = await rig.start() + await rig.producer.append("a") + handle = rig.client.get_stream_handle + + def broken(ref: StreamRef) -> Any: + del ref + raise RuntimeError("the store is away") + + monkeypatch.setattr(rig.client, "get_stream_handle", broken) + with pytest.raises(RuntimeError, match="away"): + await rig.deliver(token, 1) + monkeypatch.setattr(rig.client, "get_stream_handle", handle) + delivery = await rig.deliver(token, 1) + assert (delivery.read, delivery.records) == (True, 1) + assert posted == [] + + +async def test_close_lets_go_of_every_listener_without_completing( + posted: list[_Posted], +): + rig = await _rig() + first = await rig.start() + second = await rig.start() + assert rig.operation.tokens == [first, second] + await rig.operation.close() + assert rig.operation.tokens == [] + assert rig.client.unregistered == [ + (rig.channel, "listener-1", None), + (rig.channel, "listener-2", None), + ] + assert posted == [] + + +async def test_a_null_result_travels_without_a_content_type(posted: list[_Posted]): + def nothing(record: StreamRecord[Any], value: None) -> None: + del record + return value + + store = MemoryStreams() + handle = await store.create_standalone_stream(None, "silent") + ref = handle.ref(topic=VALUES) + client = _ClientStandIn(store) + operation = stream_consumer_operation( + nothing, + initial=lambda: None, + listener_url=LISTENER_URL, + client=_as_client(client), + ) + result = await operation.start(_start_context(), ref) + assert isinstance(result, nexusrpc.handler.StartOperationResultAsync) + await handle.producer(topic=VALUES, producer_id="w", attempt=1).finish() + await operation.deliver( + {STREAM_CONSUMER_TOKEN_HEADER: result.token}, + _notification(stream_channel(ref).channel, 1), + ) + [completion] = posted + assert completion.state == "succeeded" + assert "Content-Type" not in completion.headers + assert completion.body == b"" + + +async def test_the_listener_url_must_be_given(): + with pytest.raises(ValueError, match="listener_url"): + stream_consumer_operation(collect_values, initial=list, listener_url="") + + +# --------------------------------------------------------------------------- +# The standalone service speaks the Nexus HTTP protocol. +# --------------------------------------------------------------------------- + + +async def test_the_service_speaks_the_nexus_http_protocol(posted: list[_Posted]): + rig = await _rig() + service = NexusHttpService( + nexusrpc.handler.Handler([CollectStreamHandler(rig.operation)]), [rig.operation] + ) + converter = temporalio.converter.DataConverter.default.payload_converter + ref_body = converter.to_payload(rig.ref).data + async with TestClient(TestServer(service.application())) as http: + # A start carries the completion callback in the query and its + # headers under the callback prefix; an asynchronous start is a 201. + answer = await http.post( + "/CollectStream/collect", + params={"callback": CALLBACK_URL}, + data=ref_body, + headers={ + "Content-Type": "application/json", + "Nexus-Request-Id": "req-1", + "Nexus-Callback-Temporal-Callback-Token": "caller-token", + "Request-Timeout": "9.5s", + }, + ) + assert answer.status == 201 + token = (await answer.json())["token"] + assert rig.operation.tokens == [token] + held = rig.operation._states[token] + assert held.callback_url == CALLBACK_URL + assert held.callback_headers == {"temporal-callback-token": "caller-token"} + await rig.settled(token) + + # The delivery route reaches the operation the token names. + await rig.producer.append("a") + answer = await http.post( + service.deliveries_path, + data=_notification(rig.channel, 1), + headers={ + "Content-Type": "application/json", + "Temporal-Notification-Channel": rig.channel, + STREAM_CONSUMER_TOKEN_HEADER: token, + }, + ) + assert answer.status == 200 + state = rig.operation.state(token) + assert state is not None and state.value == ["a"] + + # A cancel names the token in its header and is accepted; a second one + # finds nothing and says so. + answer = await http.post( + "/CollectStream/collect/cancel", headers={"Nexus-Operation-Token": token} + ) + assert answer.status == 202 + assert [p.state for p in posted] == ["canceled"] + answer = await http.post( + "/CollectStream/collect/cancel", headers={"Nexus-Operation-Token": token} + ) + assert answer.status == 404 + assert "token" in (await answer.json())["message"] + + # What is not a stream reference is the caller's fault. + answer = await http.post( + "/CollectStream/collect", + data=b'{"kind": "nowhere"}', + headers={"Content-Type": "application/json"}, + ) + assert answer.status == 400 + answer = await http.post("/Other/op", data=b"{}") + assert answer.status == 404 + + +# --------------------------------------------------------------------------- +# Live: an external stream whose producer notifies the stream's channel. +# --------------------------------------------------------------------------- + + +@dataclass +class ConsumeInput: + """Which endpoint and service to call, which stream to hand it, and whether to cancel.""" + + endpoint: str + service: str + ref: StreamRef + cancel: bool = False + + +@workflow.defn +class ConsumeCaller: + """Starts the consumer operation with a ref and awaits its result.""" + + def __init__(self) -> None: + self._go = False + + @workflow.signal + def go(self) -> None: + """Let a cancelling run cancel now.""" + self._go = True + + @workflow.run + async def run(self, input: ConsumeInput) -> Any: + """Return the operation's result, or ``"canceled"`` after a cancel.""" + client = workflow.create_nexus_client( + service=input.service, endpoint=input.endpoint + ) + handle = await client.start_operation("collect", input.ref, output_type=list) + if not input.cancel: + return await handle + await workflow.wait_condition(lambda: self._go) + handle.cancel() + try: + await handle + except NexusOperationError: + return "canceled" + return "completed" + + +@contextlib.asynccontextmanager +async def _own_external_endpoint( + client: Client, name: str, url: str +) -> AsyncIterator[str]: + """Register a Nexus endpoint whose target is ``url`` for one test's life.""" + created = await client.operator_service.create_nexus_endpoint( + CreateNexusEndpointRequest( + spec=EndpointSpec( + name=name, + target=EndpointTarget(external=EndpointTarget.External(url=url)), + ) + ) + ) + try: + yield name + finally: + await client.operator_service.delete_nexus_endpoint( + DeleteNexusEndpointRequest( + id=created.endpoint.id, version=created.endpoint.version + ) + ) + + +@dataclass +class _Live: + store: StreamProvider + fronted: Client + consumer: StreamConsumerOperation[list[Any]] + endpoint: str + service: str + + +@contextlib.asynccontextmanager +async def _consumer_service( + client: Client, store: StreamProvider | None = None +) -> AsyncIterator[_Live]: + """Stand up the front over ``store`` and the standalone consumer behind an endpoint.""" + store = store or MemoryStreams() + uid = uuid.uuid4().hex + stream_handler = TemporalStreamsHandler(store, client) + handler_tq = f"streams-{uid}" + async with ( + _own_endpoint(client, f"streams-{uid}", handler_tq) as front_endpoint, + Worker(client, task_queue=handler_tq, nexus_service_handlers=[stream_handler]), + ): + front = NexusStreams( + endpoint=front_endpoint, http_address=_HTTP, read_wait=timedelta(seconds=2) + ) + config = client.config() + config["plugins"] = [front] + fronted = Client(**config) + consumer = stream_consumer_operation( + collect_values, + initial=list, + listener_url=f"http://127.0.0.1:{_PORT}/deliveries", + client=fronted, + ) + service = NexusHttpService( + nexusrpc.handler.Handler([CollectStreamHandler(consumer)]), + [consumer], + data_converter=fronted.data_converter, + ) + runner = web.AppRunner(service.application()) + await runner.setup() + await web.TCPSite(runner, "127.0.0.1", _PORT).start() + try: + async with _own_external_endpoint( + client, f"consumers-{uid}", f"http://127.0.0.1:{_PORT}" + ) as endpoint: + yield _Live(store, fronted, consumer, endpoint, CollectStream.__name__) + finally: + await consumer.close() + await runner.cleanup() + await stream_handler.close() + + +async def _listener_token(client: Client, channel: str) -> str: + """The operation token the one callback listener on ``channel`` carries.""" + + async def registered() -> str: + # The channel comes into being with the registration, so a describe + # that races it is answered with not found. + try: + description = await client.describe_channel(channel) + except RPCError as err: + assert err.status != RPCStatusCode.NOT_FOUND, "channel not created yet" + raise + assert len(description.listeners) == 1, description.listeners + [listener] = description.listeners + assert listener.callback is not None + # The server reports the registration's headers in lower case. + headers = { + key.lower(): value for key, value in listener.callback.headers.items() + } + return headers[STREAM_CONSUMER_TOKEN_HEADER.lower()] + + return await assert_eventually(registered, timeout=timedelta(seconds=30)) + + +async def _consumed( + consumer: StreamConsumerOperation[list[Any]], token: str, count: int +) -> None: + async def reached() -> None: + state = consumer.state(token) + assert state is not None and state.records == count, state + + await assert_eventually(reached, timeout=timedelta(seconds=30)) + + +@pytest.mark.needs_channel_server +async def test_a_standalone_service_consumes_an_external_stream_through_its_channel( + client: Client, +): + store = MemoryStreams() + stream_id = f"consumed-{uuid.uuid4().hex}" + handle = await store.create_standalone_stream(None, stream_id) + ref = handle.ref(topic=VALUES) + channel = stream_channel(ref).channel + producer = handle.producer(topic=VALUES, producer_id="writer", attempt=1) + + async with _consumer_service(client, store) as live: + # Written before anyone listens: the channel keeps the notification, + # and the registration is handed it once. + await producer.append("a", "b") + assert await client.notify_channel(channel, position=b"2", counter=2) == 0 + + async with Worker( + client, task_queue=f"callers-{stream_id}", workflows=[ConsumeCaller] + ) as worker: + run = await client.start_workflow( + ConsumeCaller.run, + ConsumeInput(live.endpoint, live.service, ref), + id=f"caller-{stream_id}", + task_queue=worker.task_queue, + ) + token = await _listener_token(client, channel) + await _consumed(live.consumer, token, 2) + + # Written after the registration: each notification is delivered + # and read from the cursor; the retained one brought nothing new. + await producer.append("c") + assert await client.notify_channel(channel, position=b"3", counter=3) == 1 + await _consumed(live.consumer, token, 3) + await producer.append("d", "e", "f") + for counter in (4, 5, 6): + await client.notify_channel( + channel, position=b"%d" % counter, counter=counter + ) + await _consumed(live.consumer, token, 6) + state = live.consumer.state(token) + assert state is not None + assert state.value == ["a", "b", "c", "d", "e", "f"] + # Every read moved the cursor, none repeated a record, and a burst + # of three cost at most three reads. + assert state.reads <= 1 + 1 + 3 + assert state.deliveries >= 2 + + # The close completes the operation with what was consumed. + await producer.append("g") + await producer.finish() + await client.notify_channel( + channel, position=b"7", counter=7, metadata={"closed": True} + ) + assert await asyncio.wait_for(run.result(), 60) == [ + "a", + "b", + "c", + "d", + "e", + "f", + "g", + ] + assert live.consumer.state(token) is None + assert (await client.describe_channel(channel)).listeners == [] + + +@pytest.mark.needs_channel_server +async def test_cancel_unregisters_the_listener(client: Client): + store = MemoryStreams() + stream_id = f"consumed-{uuid.uuid4().hex}" + handle = await store.create_standalone_stream(None, stream_id) + ref = handle.ref(topic=VALUES) + channel = stream_channel(ref).channel + + async with ( + _consumer_service(client, store) as live, + Worker( + client, task_queue=f"callers-{stream_id}", workflows=[ConsumeCaller] + ) as worker, + ): + run = await client.start_workflow( + ConsumeCaller.run, + ConsumeInput(live.endpoint, live.service, ref, cancel=True), + id=f"caller-{stream_id}", + task_queue=worker.task_queue, + ) + token = await _listener_token(client, channel) + await run.signal(ConsumeCaller.go) + assert await asyncio.wait_for(run.result(), 60) == "canceled" + assert live.consumer.state(token) is None + assert (await client.describe_channel(channel)).listeners == [] + + +# --------------------------------------------------------------------------- +# Live: the handler hosted in a Temporal worker, deliveries through the +# frontend's Nexus route of a companion operation. +# --------------------------------------------------------------------------- + + +@nexusrpc.service +class CollectInWorker: + """The consumer next to the synchronous operation the channel posts to.""" + + collect: nexusrpc.Operation[StreamRef, list[Any]] + deliver: nexusrpc.Operation[dict[str, Any], None] + + +@nexusrpc.handler.service_handler(service=CollectInWorker) +class CollectInWorkerHandler: + """Hosts the consumer in a worker; ``deliver`` is the listener's URL target.""" + + def __init__(self, consumer: StreamConsumerOperation[list[Any]]) -> None: + self._consumer = consumer + + @nexusrpc.handler.operation_handler + def collect(self) -> nexusrpc.handler.OperationHandler[StreamRef, list[Any]]: + return self._consumer + + @nexusrpc.handler.sync_operation + async def deliver( + self, ctx: nexusrpc.handler.StartOperationContext, input: dict[str, Any] + ) -> None: + # The server's delivery arrives as a start of this operation: the + # notification is the input and the listener's headers are the + # request's. + await self._consumer.deliver(ctx.headers, input) + + +@pytest.mark.needs_channel_server +async def test_a_worker_hosted_consumer_is_reached_through_the_frontend( + client: Client, +): + """The frontend takes the channel's post as a start of the companion operation. + + What the worker-hosted shape still lacks is the completion: the start + context a worker hands the handler names ``temporal://system`` as the + callback, which only a workflow or activity run the server completes + can reach, not an HTTP post. So this case checks the deliveries and + stops before the close; the completion through a run is a follow-up. + """ + store = MemoryStreams() + stream_id = f"consumed-{uuid.uuid4().hex}" + handle = await store.create_standalone_stream(None, stream_id) + ref = handle.ref(topic=VALUES) + channel = stream_channel(ref).channel + producer = handle.producer(topic=VALUES, producer_id="writer", attempt=1) + uid = uuid.uuid4().hex + stream_handler = TemporalStreamsHandler(store, client) + + async with ( + _own_endpoint(client, f"streams-{uid}", f"streams-{uid}") as front_endpoint, + Worker( + client, task_queue=f"streams-{uid}", nexus_service_handlers=[stream_handler] + ), + ): + front = NexusStreams( + endpoint=front_endpoint, http_address=_HTTP, read_wait=timedelta(seconds=2) + ) + config = client.config() + config["plugins"] = [front] + fronted = Client(**config) + # The listener's URL is the frontend's route to the companion + # operation, so the endpoint has to exist before the consumer does. + created = await client.operator_service.create_nexus_endpoint( + CreateNexusEndpointRequest( + spec=EndpointSpec( + name=f"hosted-{uid}", + target=EndpointTarget( + worker=EndpointTarget.Worker( + namespace=client.namespace, task_queue=f"hosted-{uid}" + ) + ), + ) + ) + ) + consumer = stream_consumer_operation( + collect_values, + initial=list, + listener_url=( + f"{_HTTP}/nexus/endpoints/{created.endpoint.id}/services/" + f"{CollectInWorker.__name__}/deliver" + ), + ) + try: + async with ( + Worker( + fronted, + task_queue=f"hosted-{uid}", + nexus_service_handlers=[CollectInWorkerHandler(consumer)], + ), + Worker( + client, task_queue=f"callers-{uid}", workflows=[ConsumeCaller] + ) as callers, + ): + await producer.append("a") + run = await client.start_workflow( + ConsumeCaller.run, + ConsumeInput(f"hosted-{uid}", CollectInWorker.__name__, ref), + id=f"caller-{uid}", + task_queue=callers.task_queue, + ) + token = await _listener_token(client, channel) + await _consumed(consumer, token, 1) + await producer.append("b", "c") + assert ( + await client.notify_channel(channel, position=b"3", counter=3) == 1 + ) + await _consumed(consumer, token, 3) + state = consumer.state(token) + assert state is not None + assert state.value == ["a", "b", "c"] + assert state.deliveries >= 1 + held = consumer._states[token] + assert held.callback_url is not None + assert held.callback_url.startswith("temporal://"), held.callback_url + await run.terminate() + finally: + await consumer.close() + await stream_handler.close() + await client.operator_service.delete_nexus_endpoint( + DeleteNexusEndpointRequest( + id=created.endpoint.id, version=created.endpoint.version + ) + ) + + +# --------------------------------------------------------------------------- +# Live, waiting for a server: a native standalone stream notifies its own +# channel, so nobody notifies by hand. +# --------------------------------------------------------------------------- + + +@pytest.mark.needs_native_provider +async def test_a_native_standalone_stream_is_consumed_through_its_channel( + client: Client, +): + provider: StreamProvider | None = client.config().get("stream_provider") + assert provider is not None, "the client needs the native provider registered" + stream_id = f"consumed-{uuid.uuid4().hex}" + handle = await client.create_stream(stream_id) + ref = handle.ref(topic=VALUES) + producer = handle.producer(topic=VALUES, producer_id="writer", attempt=1) + + async with _consumer_service(client, provider) as live: + await producer.append("a", "b") + async with Worker( + client, task_queue=f"callers-{stream_id}", workflows=[ConsumeCaller] + ) as worker: + run = await client.start_workflow( + ConsumeCaller.run, + ConsumeInput(live.endpoint, live.service, ref), + id=f"caller-{stream_id}", + task_queue=worker.task_queue, + ) + token = await _listener_token(client, stream_channel(ref).channel) + await _consumed(live.consumer, token, 2) + # The server notifies the stream's channel on every append and on + # the close, so the records reach the consumer without a notify. + await producer.append("c") + await _consumed(live.consumer, token, 3) + await producer.finish() + await handle.close() + assert await asyncio.wait_for(run.result(), 60) == ["a", "b", "c"] + assert ( + await client.describe_channel(stream_channel(ref).channel) + ).listeners == []