diff --git a/temporalio/client_stream.py b/temporalio/client_stream.py index 1f02ab128..a7b5f8d32 100644 --- a/temporalio/client_stream.py +++ b/temporalio/client_stream.py @@ -18,37 +18,76 @@ - **This client is for use outside a Workflow.** Workflow code publishes and consumes with ``workflow.append_stream_records`` and ``workflow.read_stream_records`` instead. -- **No TLS or API-key support**, for the same reason: the channel is built - here rather than by the machinery that normally handles that. +- **The channel mirrors the client's connection rather than sharing it.** + :class:`Connection` reads a :class:`temporalio.service.ConnectConfig` and + opens a ``grpcio`` channel with the same target, TLS material, API key, + headers and keep-alive, so a client connected to Temporal Cloud reaches the + stream service the same way. What it cannot mirror is noted on that class. A failed call raises :class:`temporalio.streams.StreamNotFoundError` when the server answers ``NOT_FOUND``, :class:`temporalio.streams.StreamProducerError` when it refuses a producer -sequence it already holds, and :class:`temporalio.service.RPCError` -otherwise, never the transport's own exception type. +sequence it already holds, :class:`temporalio.streams.StreamCursorError` when +it refuses a read below the retention floor, +:class:`temporalio.streams.StreamClosedError` when it refuses an append to a +sealed stream, and :class:`temporalio.service.RPCError` otherwise, never the +transport's own exception type. :func:`translate_error` is the one place that +decides. + +A failure sdk-core would retry is retried here, on the same codes and with the +same default :class:`temporalio.service.RetryConfig`, because this channel is +not Core's and gets none of its retrying. ``RESOURCE_EXHAUSTED`` backs off +longer than the rest, as in Core, so a caller the server is throttling does not +add to the load. A call the server cannot tell from its own repeat, an append +without a producer id or a create, is retried only on ``RESOURCE_EXHAUSTED``, +which the server sends before it does anything. The budget is bounded; a caller +that wants a shorter one cancels, as with ``asyncio.timeout``, and the +cancellation lands whether an attempt or a wait between attempts is in progress. """ from __future__ import annotations import asyncio +import logging +import random +import time import weakref from collections.abc import AsyncIterator, Awaitable, Callable, Sequence from dataclasses import dataclass -from typing import Any, TypeVar +from datetime import timedelta +from typing import TYPE_CHECKING, Any, TypeVar -import google.protobuf.duration_pb2 import grpc import grpc.aio from google.protobuf.message import Message import temporalio.api.streamservice.v1 as stream +from temporalio.api.common.v1 import GrpcStatus +from temporalio.api.enums.v1 import ResourceExhaustedCause +from temporalio.api.errordetails.v1 import ResourceExhaustedFailure from temporalio.api.stream.v1 import StreamRecord, StreamStartPosition from temporalio.api.streamservice.v1 import service_pb2_grpc -from temporalio.service import RPCError, RPCStatusCode -from temporalio.streams import StreamNotFoundError, StreamProducerError +from temporalio.service import ( + ConnectConfig, + RetryConfig, + RPCError, + RPCStatusCode, + TLSConfig, + __version__, +) +from temporalio.streams import ( + StreamClosedError, + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) + +if TYPE_CHECKING: + from temporalio.client import Client __all__ = [ "Appended", + "Connection", "Page", "StreamClient", "StreamEntry", @@ -56,21 +95,293 @@ "WorkflowStreamHandle", "close_shared_clients", "shared_client", + "shared_key", + "translate_error", ] _T = TypeVar("_T") -# The server refuses a producer sequence it already holds with a message and -# no typed detail, so the phrase is the only thing to match on. Both refusals -# it sends carry it: a repeat with different content, and one behind the -# sequence it accepted last. -_PRODUCER_CONFLICT = "producer sequence" -# The status the server puts on a producer conflict; the older one is kept -# so a server built before the reason tokens still gets the typed error. -_PRODUCER_CONFLICT_CODES = ( - grpc.StatusCode.FAILED_PRECONDITION, - grpc.StatusCode.INVALID_ARGUMENT, +logger = logging.getLogger(__name__) + +# A refusal the caller has to act on is a FAILED_PRECONDITION whose message +# begins with a reason token and ": ", since the service carries no typed +# detail for these yet. A repeat with different content and one behind the +# sequence the server accepted last are both a producer error; a read below +# the retention floor is a cursor error. +_REASON_SEPARATOR = ": " +_REASONS: dict[str, type[Exception]] = { + "STREAM_PRODUCER_CONFLICT": StreamProducerError, + "STREAM_PRODUCER_STALE_SEQUENCE": StreamProducerError, + "STREAM_CURSOR_BELOW_FLOOR": StreamCursorError, + "STREAM_CLOSED": StreamClosedError, + # A create of an id that exists with another policy is the caller's + # mistake, which the interface contract spells as ValueError. + "STREAM_POLICY_MISMATCH": ValueError, +} +# The phrases a server built before the tokens existed sends for the same +# refusals, so a reader of either server gets the typed error. The sealed +# stream's message is matched whole. +_PRODUCER_PHRASE = "producer sequence" +_CURSOR_PHRASE = "below the stream's floor" +_CLOSED_PHRASE = "stream is closed" + +# The codes sdk-core retries. +_RETRYABLE = frozenset( + { + grpc.StatusCode.DATA_LOSS, + grpc.StatusCode.INTERNAL, + grpc.StatusCode.UNKNOWN, + grpc.StatusCode.RESOURCE_EXHAUSTED, + grpc.StatusCode.ABORTED, + grpc.StatusCode.OUT_OF_RANGE, + grpc.StatusCode.UNAVAILABLE, + } ) +# What the server sends before it does anything, so a call that cannot be told +# from its own repeat is still safe to make again on it. +_REFUSED = frozenset({grpc.StatusCode.RESOURCE_EXHAUSTED}) +# A message over the channel's limit comes back as RESOURCE_EXHAUSTED and is +# the same size every time. +_TOO_LARGE = ( + "grpc: received message larger than max", + "grpc: message after decompression larger than max", + "grpc: received message after decompression larger than max", +) +# The floor under a throttled call's wait, sdk-core's own. +_THROTTLE = RetryConfig( + initial_interval_millis=1000, + multiplier=2.0, + max_interval_millis=10000, + max_elapsed_time_millis=None, + max_retries=0, +) + + +class _Backoff: + """Exponential backoff over a :class:`RetryConfig`, with sdk-core's arithmetic.""" + + def __init__(self, config: RetryConfig) -> None: + self._config = config + self._started = time.monotonic() + self._interval = config.initial_interval_millis / 1000 + self._failures = 0 + + def next(self) -> float | None: + """Seconds to wait before the next attempt, or ``None`` once the budget is spent.""" + config = self._config + self._failures += 1 + if config.max_retries and self._failures >= config.max_retries: + return None + base = self._interval + self._interval = min( + base * config.multiplier, config.max_interval_millis / 1000 + ) + spread = base * config.randomization_factor + delay = max(base + random.uniform(-spread, spread), 0.0) + if config.max_elapsed_time_millis is not None and ( + time.monotonic() - self._started + delay + > config.max_elapsed_time_millis / 1000 + ): + return None + return delay + + +def _raw_status(error: grpc.aio.AioRpcError) -> bytes: + # The aio metadata iterates as (key, value) pairs at runtime, whatever + # shape the stubs give its items. + trailing: Any = error.trailing_metadata() + for item in trailing or (): + key, value = item[0], item[1] + if key == "grpc-status-details-bin" and isinstance(value, bytes): + return value + return b"" + + +def _exhausted_cause(error: grpc.aio.AioRpcError) -> int | None: + """The cause the server attached to a ``RESOURCE_EXHAUSTED``, when it attached one.""" + raw = _raw_status(error) + if not raw: + return None + status = GrpcStatus() + status.ParseFromString(raw) + for detail in status.details: + if detail.Is(ResourceExhaustedFailure.DESCRIPTOR): + failure = ResourceExhaustedFailure() + detail.Unpack(failure) + return failure.cause + return None + + +def _worth_waiting_out(error: grpc.aio.AioRpcError) -> bool: + """Whether a ``RESOURCE_EXHAUSTED`` is load the server will shed, rather than a limit.""" + if (error.details() or "").startswith(_TOO_LARGE): + return False + # A stream's budget refuses an append for as long as the stream is that + # full, which no wait changes. + return _exhausted_cause(error) != ( + ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_PERSISTENCE_STORAGE_LIMIT + ) + + +def _retry_after( + error: grpc.aio.AioRpcError, + backoff: _Backoff, + throttle: _Backoff, + *, + idempotent: bool, +) -> float | None: + """Seconds to wait before making the call again, or ``None`` to raise it.""" + code = error.code() + if code not in _RETRYABLE or (not idempotent and code not in _REFUSED): + return None + throttled = code is grpc.StatusCode.RESOURCE_EXHAUSTED + if throttled and not _worth_waiting_out(error): + return None + delay = backoff.next() + if delay is None: + return None + if throttled: + delay = max(delay, throttle.next() or 0.0) + return delay + + +class _Headers(grpc.aio.UnaryUnaryClientInterceptor): + """Attaches the connection's headers to every call, as Core's interceptor does.""" + + def __init__(self, headers: Sequence[tuple[str, str | bytes]]) -> None: + self._headers = headers + + async def intercept_unary_unary( # type: ignore[override] + self, + continuation: Callable[[grpc.aio.ClientCallDetails, Any], Awaitable[Any]], + client_call_details: grpc.aio.ClientCallDetails, + request: Any, + ) -> Any: + # The aio metadata iterates as (key, value) pairs at runtime, whatever + # shape the stubs give its items. + given: Any = client_call_details.metadata + metadata = grpc.aio.Metadata(*(given or ())) + for key, value in self._headers: + # A header the caller set on the call wins over the connection's. + if key not in metadata: + metadata.add(key, value) + details = client_call_details._replace(metadata=metadata) # type: ignore[attr-defined] + return await continuation(details, request) + + +@dataclass(frozen=True) +class Connection: + """How a stream channel reaches a frontend, taken from a client's connection. + + :meth:`from_config` reads what ``Client.connect`` was given and this opens + a ``grpc.aio`` channel that behaves the same way: the target, TLS with the + same root CA, client certificate and key, the API key as a bearer + ``authorization`` header, the client's default headers and keep-alive. + Two clients with the same settings yield equal connections, which is what + lets them share one channel per namespace. + + Two things ``grpcio`` cannot express the way sdk-core does. It has one + override for both the TLS server name it sends and the name it verifies, + so ``verification_server_name`` takes that override when set and + ``domain`` otherwise, while ``domain`` alone still sets the HTTP/2 + authority. And it reads the settings once, when the channel is opened, so + an API key or header updated on the client afterwards reaches the stream + channel only through a new connection. + """ + + target: str + secure: bool + server_root_ca_cert: bytes | None = None + client_cert: bytes | None = None + client_private_key: bytes | None = None + server_name: str | None = None + authority: str | None = None + headers: tuple[tuple[str, str | bytes], ...] = () + keep_alive: tuple[int, int] | None = None + http_proxy: str | None = None + + @staticmethod + def from_config(config: ConnectConfig) -> Connection: + """Read a :class:`temporalio.service.ConnectConfig` the way the bridge does.""" + target = config.target_host + tls: TLSConfig | None = None + if "://" in target: + # The bridge still accepts a URL with a scheme; the scheme decides. + scheme, _, target = target.partition("://") + secure = scheme == "https" + if isinstance(config.tls, TLSConfig): + tls = config.tls + elif isinstance(config.tls, TLSConfig): + secure, tls = True, config.tls + elif config.tls: + secure = True + else: + # TLS is on by default when an API key is given and tls was left unset. + secure = config.tls is None and config.api_key is not None + + headers: list[tuple[str, str | bytes]] = [ + ("client-name", "temporal-python"), + ("client-version", __version__), + ] + given = {key.lower() for key in config.rpc_metadata} + if config.api_key is not None and "authorization" not in given: + headers.append(("authorization", f"Bearer {config.api_key}")) + headers.extend(config.rpc_metadata.items()) + + proxy = config.http_connect_proxy_config + http_proxy: str | None = None + if proxy is not None: + auth = ( + f"{proxy.basic_auth[0]}:{proxy.basic_auth[1]}@" + if proxy.basic_auth + else "" + ) + http_proxy = f"http://{auth}{proxy.target_host}" + + keep_alive = config.keep_alive_config + return Connection( + target=target, + secure=secure, + server_root_ca_cert=tls.server_root_ca_cert if tls else None, + client_cert=tls.client_cert if tls else None, + client_private_key=tls.client_private_key if tls else None, + server_name=((tls.verification_server_name or tls.domain) if tls else None), + authority=tls.domain if tls else None, + headers=tuple(headers), + keep_alive=( + (keep_alive.interval_millis, keep_alive.timeout_millis) + if keep_alive + else None + ), + http_proxy=http_proxy, + ) + + def channel(self) -> grpc.aio.Channel: + """Open a channel with these settings. Nothing is sent until the first call.""" + options: list[tuple[str, Any]] = [] + if self.keep_alive is not None: + options.append(("grpc.keepalive_time_ms", self.keep_alive[0])) + options.append(("grpc.keepalive_timeout_ms", self.keep_alive[1])) + if self.server_name: + options.append(("grpc.ssl_target_name_override", self.server_name)) + if self.authority: + options.append(("grpc.default_authority", self.authority)) + if self.http_proxy: + options.append(("grpc.http_proxy", self.http_proxy)) + # The stubs do not know the aio interceptor base as a ClientInterceptor. + interceptors: Any = [_Headers(self.headers)] if self.headers else None + if not self.secure: + return grpc.aio.insecure_channel( + self.target, options=options, interceptors=interceptors + ) + credentials = grpc.ssl_channel_credentials( + root_certificates=self.server_root_ca_cert, + private_key=self.client_private_key, + certificate_chain=self.client_cert, + ) + return grpc.aio.secure_channel( + self.target, credentials, options=options, interceptors=interceptors + ) @dataclass(frozen=True) @@ -172,56 +483,140 @@ def _to_public(record: stream.StreamRecord) -> StreamEntry: return StreamEntry(record=out, offset=record.offset) -def _translate(error: grpc.aio.AioRpcError) -> Exception: - code = error.code() - details = error.details() or code.name +def translate_error( + code: grpc.StatusCode, details: str, raw_status: bytes = b"" +) -> Exception: + """The SDK error for one failed call, from its status code and message. + + ``NOT_FOUND`` is :class:`temporalio.streams.StreamNotFoundError`. A + ``FAILED_PRECONDITION`` whose message begins with a reason token is the + error the token names: a producer refusal, a repeat with different content + or one behind the sequence the server accepted last, is + :class:`temporalio.streams.StreamProducerError`, because the caller asked + to be deduplicated and could not be; a read below the retention floor is + :class:`temporalio.streams.StreamCursorError`. Everything else is + :class:`temporalio.service.RPCError` with the code and the raw status. + """ + details = details or code.name if code is grpc.StatusCode.NOT_FOUND: return StreamNotFoundError(details) - if code in _PRODUCER_CONFLICT_CODES and _PRODUCER_CONFLICT in details: - # A producer sequence the store already holds, either with different - # content or behind the one it accepted last. The caller asked to be - # deduplicated and could not be, which is a condition of its own. The - # server reports it as a failed precondition; one built before the - # reason tokens existed called it an invalid argument. + if code is grpc.StatusCode.FAILED_PRECONDITION: + token, separator, _ = details.partition(_REASON_SEPARATOR) + typed = _REASONS.get(token) if separator else None + if typed is not None: + return typed(details) + if _CURSOR_PHRASE in details: + return StreamCursorError(details) + if details == _CLOSED_PHRASE: + return StreamClosedError(details) + if code is grpc.StatusCode.INVALID_ARGUMENT and _PRODUCER_PHRASE in details: return StreamProducerError(details) - raw = b"" - # The aio metadata iterates as (key, value) pairs at runtime, whatever - # shape the stubs give its items. - trailing: Any = error.trailing_metadata() - for item in trailing or (): - key, value = item[0], item[1] - if key == "grpc-status-details-bin" and isinstance(value, bytes): - raw = value - return RPCError(details, RPCStatusCode(code.value[0]), raw) + return RPCError(details, RPCStatusCode(code.value[0]), raw_status) -async def _call(method: Callable[[Any], Awaitable[_T]], request: Any) -> _T: - """Make one stub call, translating the transport's failure to the SDK's.""" - try: - return await method(request) - except grpc.aio.AioRpcError as error: - raise _translate(error) from error +def _translate(error: grpc.aio.AioRpcError) -> Exception: + return translate_error(error.code(), error.details() or "", _raw_status(error)) + + +async def _call( + method: Callable[[Any], Awaitable[_T]], + request: Any, + *, + retry_config: RetryConfig | None = None, + idempotent: bool = True, +) -> _T: + """Make one stub call, retried as sdk-core would, translating the failure to the SDK's. + + ``idempotent`` is false for a call the server cannot tell from its own + repeat, which is then retried only on a refusal the server sent before it + did anything. The request is the same object on every attempt, so a + numbered append is deduplicated by the server whichever attempt landed. + """ + config = retry_config or RetryConfig() + backoff = _Backoff(config) + throttle = _Backoff(_THROTTLE) + attempts = 0 + while True: + attempts += 1 + try: + return await method(request) + except grpc.aio.AioRpcError as error: + delay = _retry_after(error, backoff, throttle, idempotent=idempotent) + if delay is None: + raise _translate(error) from error + _log_retry(error, attempts, config) + await asyncio.sleep(delay) + + +def _log_retry(error: grpc.aio.AioRpcError, attempts: int, config: RetryConfig) -> None: + # Quiet at first and louder once half the budget is gone, as sdk-core does, + # so a single throttled call is not a warning but a struggling one is. + level = logging.DEBUG + if config.max_retries and attempts * 2 >= config.max_retries: + level = logging.WARNING + logger.log( + level, + "stream call failed with %s on attempt %d, retrying: %s", + error.code().name, + attempts, + error.details(), + ) class StreamClient: """Creates and opens streams on a namespace.""" - def __init__(self, channel: Any, namespace: str) -> None: - """Wrap an existing ``grpc.aio`` channel. Prefer :meth:`connect`.""" + def __init__( + self, + channel: Any, + namespace: str, + *, + retry_config: RetryConfig | None = None, + ) -> None: + """Wrap an existing ``grpc.aio`` channel. Prefer :meth:`connect`. + + ``retry_config`` is the policy every call made through this client + retries under; ``None`` is the SDK's default. + """ self._channel = channel self._namespace = namespace + self._retry_config = retry_config # The generated stub is typed for a synchronous channel. This client # drives it over ``grpc.aio``, where every call is awaited. self._stub: Any = service_pb2_grpc.StreamServiceStub(channel) @staticmethod - def connect(target_host: str, namespace: str = "default") -> StreamClient: - """Open a channel to a frontend. + def connect( + target_host: str, + namespace: str = "default", + *, + retry_config: RetryConfig | None = None, + ) -> StreamClient: + """Open a plaintext channel to a frontend, for a local server. Separate from ``Client.connect`` because this does not share the - connection the rest of the SDK uses. + connection the rest of the SDK uses; :meth:`for_connection` opens + one with a client's settings. """ - return StreamClient(grpc.aio.insecure_channel(target_host), namespace) + return StreamClient( + grpc.aio.insecure_channel(target_host), + namespace, + retry_config=retry_config, + ) + + @staticmethod + def for_connection( + connection: Connection, + namespace: str = "default", + *, + retry_config: RetryConfig | None = None, + ) -> StreamClient: + """Open a channel the way ``connection`` describes. + + ``retry_config`` left ``None`` is the SDK's default; pass the client's + own to retry as its other calls do. + """ + return StreamClient(connection.channel(), namespace, retry_config=retry_config) async def close(self) -> None: """Close the underlying channel.""" @@ -231,23 +626,31 @@ async def create( self, stream_id: str, *, - retention: float | None = None, + retention: float | timedelta | None = None, max_items: int | None = None, + max_bytes: int | None = None, ) -> StreamHandle: """Create a stream and return a handle to it. - ``retention`` is how long a closed stream stays readable, in seconds. - ``max_items`` caps how many records remain readable, dropping the - oldest, which bounds storage for a stream nobody truncates. + ``retention`` is the age past which an open stream's records are + reclaimed and how long a closed stream stays readable, in seconds or + as a ``timedelta``. ``max_items`` caps how many records remain + readable, dropping the oldest, which bounds storage for a stream + nobody truncates. ``max_bytes`` caps the bytes held; an append that + would cross it is refused rather than reclaiming anything. """ lifecycle = stream.StreamLifecycle() if retention is not None: - lifecycle.retention.CopyFrom( - google.protobuf.duration_pb2.Duration(seconds=int(retention)) - ) + if not isinstance(retention, timedelta): + retention = timedelta(seconds=retention) + lifecycle.retention.FromTimedelta(retention) if max_items is not None: lifecycle.max_items = max_items + if max_bytes is not None: + lifecycle.max_bytes = max_bytes + # A create that landed but was not answered would be refused as a + # repeat, so it goes again only on a refusal. response = await _call( self._stub.CreateStream, stream.CreateStreamRequest( @@ -257,17 +660,22 @@ async def create( lifecycle=lifecycle, ) ), + retry_config=self._retry_config, + idempotent=False, ) return StreamHandle( self._stub, self._namespace, stream_id, run_id=response.frontend_response.run_id, + retry_config=self._retry_config, ) def get(self, stream_id: str) -> StreamHandle: """Open an existing stream without a round trip.""" - return StreamHandle(self._stub, self._namespace, stream_id) + return StreamHandle( + self._stub, self._namespace, stream_id, retry_config=self._retry_config + ) def workflow_stream( self, workflow_id: str, name: str = "", *, owner_run_id: str = "" @@ -289,7 +697,12 @@ def workflow_stream( :meth:`WorkflowStreamHandle.pin` for why a follower wants that. """ return WorkflowStreamHandle( - self._stub, self._namespace, workflow_id, name, owner_run_id + self._stub, + self._namespace, + workflow_id, + name, + owner_run_id, + retry_config=self._retry_config, ) def activity_stream( @@ -317,6 +730,7 @@ def activity_stream( name, run_id, activity_id=activity_id, + retry_config=self._retry_config, ) @@ -324,7 +738,13 @@ class StreamHandle: """A handle to one standalone stream.""" def __init__( - self, stub: Any, namespace: str, stream_id: str, run_id: str = "" + self, + stub: Any, + namespace: str, + stream_id: str, + run_id: str = "", + *, + retry_config: RetryConfig | None = None, ) -> None: """Prefer :meth:`StreamClient.get` or :meth:`StreamClient.create`.""" self._stub = stub @@ -333,6 +753,7 @@ def __init__( # Passing this back saves the server resolving the current run on every # call, which is otherwise a persistence lookup per request. self._run_id = run_id + self._retry_config = retry_config @property def id(self) -> str: @@ -349,9 +770,11 @@ async def append( Supplying ``producer_id`` and ``sequence`` makes the append idempotent: a retry with the same pair returns the original offsets rather than - appending twice, and says so. Without them the append is - at-least-once, which is only the right trade when duplicates are - harmless. + appending twice, and says so. That is also what lets a failed append + be made again here on every code sdk-core retries. Without them the + append is at-least-once, which is only the right trade when duplicates + are harmless, and it is made again only on a refusal the server sent + before it did anything. """ response = await _call( self._stub.AddMessages, @@ -365,6 +788,8 @@ async def append( sequence=sequence, ) ), + retry_config=self._retry_config, + idempotent=bool(producer_id), ) return _appended(response.frontend_response) @@ -448,6 +873,7 @@ async def poll( wait_new_messages=wait, ) ), + retry_config=self._retry_config, ) return _page(response.frontend_response) @@ -467,6 +893,7 @@ async def truncate(self, new_base_offset: int) -> None: new_base_offset=new_base_offset, ) ), + retry_config=self._retry_config, ) async def finish_writing(self, producer_id: str) -> None: @@ -480,6 +907,7 @@ async def finish_writing(self, producer_id: str) -> None: producer_id=producer_id, ) ), + retry_config=self._retry_config, ) async def close(self) -> None: @@ -491,6 +919,7 @@ async def close(self) -> None: namespace=self._namespace, stream_id=self._id ) ), + retry_config=self._retry_config, ) async def describe(self) -> stream.StreamState: @@ -502,6 +931,7 @@ async def describe(self) -> stream.StreamState: namespace=self._namespace, stream_id=self._id ) ), + retry_config=self._retry_config, ) return response.frontend_response.state @@ -525,6 +955,7 @@ def __init__( owner_run_id: str = "", *, activity_id: str = "", + retry_config: RetryConfig | None = None, ) -> None: """Prefer :meth:`StreamClient.workflow_stream` or :meth:`StreamClient.activity_stream`.""" self._stub = stub @@ -533,6 +964,7 @@ def __init__( self._name = name self._owner_run_id = owner_run_id self._activity_id = activity_id + self._retry_config = retry_config @property def workflow_id(self) -> str: @@ -616,6 +1048,8 @@ async def append( sequence=sequence, ) ), + retry_config=self._retry_config, + idempotent=bool(producer_id), ) return _appended(response.frontend_response) @@ -687,6 +1121,7 @@ async def poll( wait_new_messages=wait, ) ), + retry_config=self._retry_config, ) return _page(response.frontend_response) @@ -705,6 +1140,7 @@ async def describe(self) -> stream.StreamState: stream_name=self._name, ) ), + retry_config=self._retry_config, ) return response.frontend_response.state @@ -728,33 +1164,55 @@ def _page(out: stream.PollMessagesOutput) -> Page: ) -# One channel per loop, target and namespace, shared by every handle in the +# One channel per loop, connection and namespace, shared by every handle in the # process. A channel is multiplexed and long lived, and callers open a handle # per subscription, which would otherwise be a connection per subscription. The # loop is the key because a grpc.aio channel belongs to the loop that made it, # and it is held weakly so a loop that is gone cannot lend its channel to a # successor that happens to reuse its id. +SharedKey = tuple[Connection, str] + _shared: weakref.WeakKeyDictionary[ - asyncio.AbstractEventLoop, dict[tuple[str, str], StreamClient] + asyncio.AbstractEventLoop, dict[SharedKey, StreamClient] ] = weakref.WeakKeyDictionary() -def shared_client(target_host: str, namespace: str) -> StreamClient: - """The process-wide client for ``target_host`` and ``namespace`` on this loop.""" +def shared_key(client: Client) -> SharedKey: + """What names the shared channel ``client`` reaches the stream service through.""" + return Connection.from_config(client.service_client.config), client.namespace + + +def shared_client(client: Client) -> StreamClient: + """The process-wide stream client on this loop for ``client``'s connection and namespace. + + Opened with the client's connection settings and its ``retry_config``, so + a call on it authenticates and retries as the client's other calls do. + """ per_loop = _shared.setdefault(asyncio.get_running_loop(), {}) - key = (target_host, namespace) + key = shared_key(client) existing = per_loop.get(key) if existing is None: - existing = per_loop[key] = StreamClient.connect(target_host, namespace) + existing = per_loop[key] = StreamClient.for_connection( + key[0], key[1], retry_config=client.service_client.config.retry_config + ) return existing -async def close_shared_clients() -> None: - """Close every shared client this loop opened. +async def close_shared_clients(*keys: SharedKey) -> None: + """Close the shared clients this loop opened for ``keys``, or all of them. - For a process that is done with streams, and for tests, which open a - loop per case and would otherwise leave a channel behind on each. + A provider closes the ones it opened, named by :func:`shared_key`: + another provider on the same loop may still be reading through a channel + of its own, and taking that out from under it is not this one's to do. + With no keys it closes every one, which is what a process finished with + streams wants, and what a test that opened a loop of its own wants. """ - per_loop = _shared.pop(asyncio.get_running_loop(), {}) - for client in per_loop.values(): + loop = asyncio.get_running_loop() + if not keys: + per_loop = _shared.pop(loop, {}) + closing = list(per_loop.values()) + else: + per_loop = _shared.get(loop, {}) + closing = [per_loop.pop(key) for key in keys if key in per_loop] + for client in closing: await client.close() diff --git a/temporalio/contrib/langgraph/_plugin.py b/temporalio/contrib/langgraph/_plugin.py index 03881ca2a..1819c646d 100644 --- a/temporalio/contrib/langgraph/_plugin.py +++ b/temporalio/contrib/langgraph/_plugin.py @@ -104,11 +104,13 @@ class LangGraphPlugin(SimplePlugin): inherited default of either form. streaming_topic: When set, ``langgraph.config.get_stream_writer()`` inside a node publishes to this topic on the workflow's - :class:`WorkflowStream`. The workflow must construct - ``WorkflowStream()`` in its ``@workflow.init`` (the plugin's + :class:`temporalio.contrib.workflow_streams.WorkflowStream`. The + workflow must construct ``WorkflowStream()`` in its + ``@workflow.init`` (the plugin's interceptor verifies this on workflow start). Nodes with ``execute_in='activity'`` publish through - :class:`WorkflowStreamClient` (signal); nodes with + :class:`temporalio.contrib.workflow_streams.WorkflowStreamClient` + (signal); nodes with ``execute_in='workflow'`` publish synchronously to the in-workflow stream (no signal). streaming_batch_interval: How often the activity-side stream diff --git a/temporalio/contrib/server_streams/__init__.py b/temporalio/contrib/server_streams/__init__.py new file mode 100644 index 000000000..f6ac83556 --- /dev/null +++ b/temporalio/contrib/server_streams/__init__.py @@ -0,0 +1,460 @@ +"""Server-side streams behind the Workflow Streams API. + +This is the same surface as :mod:`temporalio.contrib.workflow_streams`, backed +by a Temporal-owned log instead of by Signals and Updates. An application swaps +the import and keeps its code: publishing from a Workflow is still a plain +call, publishing from an Activity is still a buffered handle, and a consumer +still subscribes by topic from an offset. An Activity's appends carry its own +id and attempt, so a retried Activity's repeat is deduplicated by the server +rather than written twice. + +What changes is underneath. A publish is a Workflow Command whose payload never +enters History, so History gets one fixed-size event per Workflow Task rather +than a Signal per batch. A consumer reads the log directly rather than +long-polling an Update, so there is no per-Workflow limit on how many can read +at once, and a closed Workflow stays readable until its stream's retention +expires. + +As in the shipped feature, a topic here is a label on a record in the +Workflow's one default stream, which a consumer filters on. The provider in +:mod:`temporalio.streams.providers.native` keeps one stream per topic instead; +the two do not share a log. + +Prototype support for AI-198. It needs a server built from that branch, and it +opens its own gRPC channel because sdk-core does not know the stream service +yet, which is also why it does not support TLS or API keys. +""" + +from __future__ import annotations + +import asyncio +import builtins +from collections.abc import AsyncIterator, Awaitable, Callable, Sequence +from dataclasses import dataclass +from datetime import timedelta +from typing import Any, Generic, TypeVar, overload + +from temporalio import activity, workflow +from temporalio.api.stream.v1 import StreamRecord, StreamRecordKind +from temporalio.client import Client, WorkflowExecutionDescription +from temporalio.client_stream import StreamClient, WorkflowStreamHandle, shared_client +from temporalio.common import RawValue +from temporalio.converter import PayloadCodec, PayloadConverter + +__all__ = [ + "RawPage", + "TopicHandle", + "WorkflowStream", + "WorkflowStreamClient", + "WorkflowStreamItem", + "WorkflowTopicHandle", +] + +T = TypeVar("T") + +DEFAULT_BATCH_INTERVAL = timedelta(milliseconds=50) + + +@dataclass +class RawPage: + """One read whose items are still the stored records. + + ``closed`` with ``next_offset >= head_offset`` is the end of the stream. + The bodies are as the server holds them, codec included, for a caller + that forwards records rather than using them. + """ + + items: list[WorkflowStreamItem[StreamRecord]] + next_offset: int + head_offset: int + closed: bool + + +@dataclass +class WorkflowStreamItem(Generic[T]): + """One item read from a workflow's stream. + + ``offset`` is where the item sits in the whole stream, so it is what a + consumer hands back to resume. A topic filter leaves gaps in it. + """ + + topic: str + data: T + offset: int = 0 + + +def _record(converter: PayloadConverter, topic: str, value: Any) -> StreamRecord: + """The record one published value becomes. + + The body is the value's payload, so the encoding metadata the consumer + needs to decode into a type travels with it and a + :class:`temporalio.common.RawValue` passes through pre-encoded. + """ + record = StreamRecord(topic=topic, kind=StreamRecordKind.STREAM_RECORD_KIND_DATA) + record.body.CopyFrom(converter.to_payloads([value])[0]) + return record + + +def _decode( + converter: PayloadConverter, record: StreamRecord, as_type: type | None +) -> Any: + if as_type is None: + return converter.from_payloads([record.body])[0] + return converter.from_payloads([record.body], [as_type])[0] + + +class WorkflowTopicHandle(Generic[T]): + """A topic on the stream the running Workflow owns.""" + + def __init__(self, topic: str, value_type: type[T]) -> None: + """Prefer :meth:`WorkflowStream.topic`.""" + self._name = topic + self._type = value_type + + @property + def name(self) -> str: + """The topic name this handle is bound to.""" + return self._name + + @property + def type(self) -> type[T]: + """The value type this handle is bound to.""" + return self._type + + def publish(self, value: T | RawValue) -> None: + """Append ``value`` to the Workflow's stream on this topic. + + Returns at once. There is nothing to await: the Workflow Task's + publishes become one command the server applies in the task's own + commit, so it costs this Workflow no round trip and no extra + transition. The Worker's payload codec applies to the body as it does + to any other payload the Workflow sends. + """ + workflow._append_stream_records( + [_record(workflow.payload_converter(), self._name, value)] + ) + + +class WorkflowStream: + """The stream the running Workflow owns, from inside it. + + Construct in ``@workflow.init``. Unlike the Signals-and-Updates + implementation this holds no state of its own: the log lives on the server, + so there is nothing here for replay to reconstruct. + """ + + def __init__(self, prior_state: Any = None) -> None: + """Take the same argument as the Workflow Streams version and ignore it. + + That version carried the log across a continue-as-new, because the log + was Workflow state. Here it is not, so there is nothing to carry. + """ + self._prior_state = prior_state + + @overload + def topic(self, name: str) -> WorkflowTopicHandle[Any]: ... + + @overload + def topic(self, name: str, *, type: type[T]) -> WorkflowTopicHandle[T]: ... + + def topic(self, name: str, *, type: type = object) -> WorkflowTopicHandle[Any]: + """Bind a topic on this Workflow's stream.""" + return WorkflowTopicHandle(name, type) + + +class TopicHandle(Generic[T]): + """A topic on a Workflow's stream, from outside that Workflow.""" + + def __init__( + self, client: WorkflowStreamClient, topic: str, value_type: type[T] + ) -> None: + """Prefer :meth:`WorkflowStreamClient.topic`.""" + self._client = client + self._name = topic + self._type = value_type + + @property + def name(self) -> str: + """The topic name this handle is bound to.""" + return self._name + + @property + def type(self) -> type[T]: + """The value type this handle is bound to.""" + return self._type + + def publish(self, value: T | RawValue, *, force_flush: bool = False) -> None: + """Buffer ``value`` for the next flush. + + Buffered rather than sent, because an append costs one transition on + the owning execution whatever its size. A token at a time would pay + that per token. + """ + self._client._buffer(self._name, value) + if force_flush: + self._client._flush_soon() + + def subscribe( + self, + *, + from_offset: int = 0, + # Spelled out because `type` in this class body is the property below. + result_type: builtins.type | None = None, + poll_cooldown: timedelta | None = None, + ) -> AsyncIterator[WorkflowStreamItem[T]]: + """Read this topic from ``from_offset`` onwards.""" + return self._client.subscribe( + topics=[self._name], + from_offset=from_offset, + result_type=result_type or self._type, + poll_cooldown=poll_cooldown, + ) + + +class WorkflowStreamClient: + """Publishes to and reads from a Workflow's stream, from outside it.""" + + def __init__( + self, + handle: WorkflowStreamHandle, + converter: PayloadConverter, + batch_interval: timedelta = DEFAULT_BATCH_INTERVAL, + *, + codec: PayloadCodec | None = None, + describe: Callable[[], Awaitable[WorkflowExecutionDescription]] | None = None, + producer_id: str = "", + ) -> None: + """Prefer :meth:`create` or :meth:`from_within_activity`. + + ``codec`` is applied to every body this client sends and receives, so + a namespace whose payloads are encoded agrees with the Worker, whose + payload visitor applies the same codec to the Workflow's publishes. + ``describe`` is how an unpinned handle learns which run it follows; + see :meth:`WorkflowStreamHandle.pin`. ``producer_id`` is who the + appends are written as; without one they are at-least-once, because + the server has nothing to deduplicate a retry against. + """ + self._handle = handle + self._converter = converter + self._codec = codec + self._batch_interval = batch_interval + self._describe = describe + self._producer_id = producer_id + self._sequence = 0 + self._pending: tuple[list[StreamRecord], int] | None = None + self._buffered: list[tuple[str, Any]] = [] + self._flusher: asyncio.Task[None] | None = None + self._wake = asyncio.Event() + self._closed = False + + @classmethod + def create( + cls, + client: Client, + workflow_id: str, + *, + owner_run_id: str = "", + batch_interval: timedelta = DEFAULT_BATCH_INTERVAL, + producer_id: str = "", + ) -> WorkflowStreamClient: + """Open the stream owned by ``workflow_id``. + + Without ``owner_run_id`` the current run is looked up on the first + call and the handle pinned to it, so a reader following across a + continue-as-new sees the run end rather than being moved to the + successor's stream at a stale offset. Without ``producer_id`` the + appends are at-least-once; inside an Activity, + :meth:`from_within_activity` supplies one. + """ + return cls( + _stream_client(client).workflow_stream( + workflow_id, owner_run_id=owner_run_id + ), + client.data_converter.payload_converter, + batch_interval, + codec=client.data_converter.payload_codec, + describe=client.get_workflow_handle(workflow_id).describe, + producer_id=producer_id, + ) + + @classmethod + def from_within_activity( + cls, *, batch_interval: timedelta = DEFAULT_BATCH_INTERVAL + ) -> WorkflowStreamClient: + """Open the stream owned by the Workflow that scheduled this Activity.""" + info = activity.info() + if info.workflow_id is None: + raise RuntimeError( + "no Workflow stream to open: this Activity was not started by a " + "Workflow" + ) + # The Activity's output belongs to the run that scheduled it, and the + # Activity already knows which run that is. Its id and attempt are + # also what lets the server drop a batch a retried Activity re-sends. + return cls.create( + activity.client(), + info.workflow_id, + owner_run_id=info.workflow_run_id or "", + batch_interval=batch_interval, + producer_id=f"{info.activity_id}#{info.attempt}", + ) + + async def __aenter__(self) -> WorkflowStreamClient: + """Start the background flusher.""" + self._flusher = asyncio.create_task(self._run_flusher()) + return self + + async def __aexit__(self, *_exc: object) -> None: + """Drain what is buffered before letting the caller go. + + An Activity that returned with a batch still buffered would have + reported work its readers never saw. The flusher is asked to stop + rather than cancelled: a cancel landing inside its append would + unwind with the batch it had already taken off the buffer. + """ + self._closed = True + self._wake.set() + if self._flusher is not None: + await self._flusher + self._flusher = None + await self.flush() + + async def _pin(self) -> None: + if self._handle.owner_run_id or self._describe is None: + return + self._handle.pin((await self._describe()).run_id) + + @overload + def topic(self, name: str) -> TopicHandle[Any]: ... + + @overload + def topic(self, name: str, *, type: type[T]) -> TopicHandle[T]: ... + + def topic(self, name: str, *, type: type = object) -> TopicHandle[Any]: + """Bind a topic on this Workflow's stream.""" + return TopicHandle(self, name, type) + + async def get_offset(self) -> int: + """Where the stream currently ends. + + A reader that wants only what comes next starts here. + """ + await self._pin() + return (await self._handle.describe()).head_offset + + async def subscribe( + self, + *, + topics: Sequence[str] = (), + from_offset: int = 0, + result_type: type | None = None, + poll_cooldown: timedelta | None = None, + ) -> AsyncIterator[WorkflowStreamItem[Any]]: + """Yield items from ``from_offset`` as they arrive. + + ``poll_cooldown`` is accepted and ignored. It paced a client that had + to re-ask; the server parks this read until something arrives. + """ + del poll_cooldown + await self._pin() + async for entry in self._handle.follow(from_offset=from_offset, topics=topics): + record = await self._decoded(entry.record) + yield WorkflowStreamItem( + topic=record.topic, + data=_decode(self._converter, record, result_type), + offset=entry.offset, + ) + + async def poll_raw( + self, + *, + topics: Sequence[str] = (), + from_offset: int = 0, + wait: bool = True, + ) -> RawPage: + """One read, with the records left as they were stored. + + For a caller that forwards items on rather than using them. A gateway + would only have to encode again what this decoded. + """ + await self._pin() + page = await self._handle.poll( + from_offset=from_offset, topics=topics, wait=wait + ) + return RawPage( + items=[ + WorkflowStreamItem( + topic=entry.record.topic, data=entry.record, offset=entry.offset + ) + for entry in page.entries + ], + next_offset=page.next_offset, + head_offset=page.head_offset, + closed=page.closed, + ) + + async def flush(self) -> None: + """Append everything buffered as one batch. + + A batch whose append failed stays pending and goes out again on the + next flush under the sequence it already had, so an append the server + did accept is deduplicated and one it never saw still lands. Nothing + comes off the buffer until there is a batch to replace it with. + + Raises: + temporalio.streams.StreamProducerError: The server holds this + producer's sequence with different content. + """ + if self._pending is not None: + records, sequence = self._pending + else: + if not self._buffered: + return + await self._pin() + # Encoded before the buffer is cleared, so a converter failure + # leaves the values where the caller can still see them. + records = [ + await self._encoded(_record(self._converter, topic, value)) + for topic, value in self._buffered + ] + sequence = self._sequence + self._buffered = [] + self._pending = (records, sequence) + await self._pin() + await self._handle.append( + *records, producer_id=self._producer_id, sequence=sequence + ) + self._sequence = sequence + len(records) + self._pending = None + + def _buffer(self, topic: str, value: Any) -> None: + self._buffered.append((topic, value)) + + def _flush_soon(self) -> None: + self._wake.set() + + async def _run_flusher(self) -> None: + """Append on a fixed cadence, so a slow producer still gets delivered.""" + while not self._closed: + try: + await asyncio.wait_for( + self._wake.wait(), self._batch_interval.total_seconds() + ) + except asyncio.TimeoutError: + pass + self._wake.clear() + await self.flush() + + async def _encoded(self, record: StreamRecord) -> StreamRecord: + if self._codec is not None and record.HasField("body"): + record.body.CopyFrom((await self._codec.encode([record.body]))[0]) + return record + + async def _decoded(self, record: StreamRecord) -> StreamRecord: + if self._codec is not None and record.HasField("body"): + record.body.CopyFrom((await self._codec.decode([record.body]))[0]) + return record + + +def _stream_client(client: Client) -> StreamClient: + return shared_client(client) diff --git a/temporalio/streams/providers/native.py b/temporalio/streams/providers/native.py new file mode 100644 index 000000000..95bb866a3 --- /dev/null +++ b/temporalio/streams/providers/native.py @@ -0,0 +1,1007 @@ +"""The server-side (native) provider. + +Streams live on the Temporal server, beside the workflow that owns them. A +topic is one owned stream named after the topic, created by whoever touches it +first: the workflow publishes to it with a command the server applies in the +transaction that accepts the Workflow Task, subscribes to it by name and reads +the ranges the server delivers on its Workflow Tasks; outside code appends and +reads through the stream service, and the workflow's records and an outside +producer's land in one log in the order the server accepted them. The default +topic, :data:`temporalio.streams.DEFAULT_TOPIC`, is the server's default +stream: the server resolves an unnamed stream to that same name, so the +provider sends the name explicitly and a record's topic and its stream's name +never differ. + +An activity owns topics of its own, apart from its workflow's: a standalone +activity is its own owner on the server, and an activity a workflow scheduled +is addressed through that workflow. They are one stream per activity +execution, so a retry writes to the same stream, and the server ends them when +the activity reaches a terminal status. + +A cursor names the run as well as the offset, because an owned stream belongs +to one run and a successor's starts over at zero. A handle without a run id +reads run after run, learning from the poll that a run's stream is closed and +from the run's close event who came next; with a run id it is pinned. + +Prototype support for AI-198. It needs a server built from that branch and +reaches the stream service on a channel of its own, opened with the client's +connection settings (target, TLS, API key, headers, retries), because sdk-core +does not know the service yet. +""" + +from __future__ import annotations + +import asyncio +import logging +import time +from collections.abc import AsyncGenerator, Callable +from datetime import timedelta +from typing import Any, Generic, TypeVar + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.stream.v1 import StreamStartPosition +from temporalio.client import Client, WorkflowHistoryEventFilterType +from temporalio.client_stream import ( + Appended, + Page, + SharedKey, + StreamClient, + WorkflowStreamHandle, + close_shared_clients, + shared_client, + shared_key, +) +from temporalio.client_stream import StreamHandle as ServiceStreamHandle +from temporalio.converter import ( + DataConverter, + StorageDriverStoreContext, + StorageDriverWorkflowInfo, +) +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams._body import ( + CONTENT_HASH_KEY, + content_hash, + decode_body, + encode_body, +) +from temporalio.streams._errors import StreamCursorError, StreamNotFoundError +from temporalio.streams._provider import ReadSource, StreamHandle, WriteSink +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._ref import StreamRef, open_ref +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__ = [ + "NativeActivityStreamHandle", + "NativeProducer", + "NativeStandaloneStreamHandle", + "NativeStreamHandle", + "NativeStreams", +] + +T = TypeVar("T") + +_PROVIDER = "native" + +# Seconds between polls on a standalone stream that does not exist yet, when +# the server answers without parking. +_CREATE_WAIT_PACE = 1.0 + +logger = logging.getLogger(__name__) + + +def _require_topic(topic: str) -> None: + if not topic: + raise ValueError("topic must not be empty") + + +def _cursor(run_id: str, offset: int) -> Cursor: + # A standalone stream has no run; its id takes the run's place. + return mint_cursor(_PROVIDER, f"{run_id}:{offset}") + + +def _position(after: Cursor) -> tuple[str, int] | None: + """The ``(run id, offset)`` a cursor of this provider names, or ``None`` for BEGINNING.""" + token = cursor_position(after, provider=_PROVIDER) + if token is None: + return None + run_id, _, offset = token.rpartition(":") + try: + if not run_id: + raise ValueError + return run_id, int(offset) + except ValueError: + raise StreamCursorError( + f"cursor {after.token!r} does not name a run and an offset on the " + "native provider" + ) from None + + +def _fingerprint(record: WireRecord) -> WireRecord: + """Stamp ``record`` with the hash of its body as converted, before the worker's pass. + + The outside half gets the same stamp from :func:`encode_body`; a workflow's + own publish is encoded later, by the worker's payload pass, so the stamp + is taken here while the body is still what the converter produced. A + record without a body has nothing to compare. + """ + if not record.HasField("body"): + return record + record.metadata[CONTENT_HASH_KEY].CopyFrom( + Payload( + metadata={"encoding": b"binary/plain"}, + data=content_hash(record.body).encode(), + ) + ) + return record + + +class _NativeReadSource: + """One subscription of the running workflow, fed by delivered ranges.""" + + def __init__(self, stream_id: str, run_id: str) -> None: + self._stream_id = stream_id + self._run_id = run_id + self._closed = False + + async def next_batch(self) -> list[tuple[Cursor, WireRecord]]: + if self._closed: + raise StopAsyncIteration + delivered = await workflow._read_stream_records(self._stream_id) + if self._closed: + # Closed while this was parked; the buffer woke it with nothing. + raise StopAsyncIteration + return [(_cursor(self._run_id, item.offset), item.record) for item in delivered] + + def close(self) -> None: + """Stop reading, and stop keeping what the server keeps delivering. + + The server has no unsubscribe command, so ranges keep arriving on + every Workflow Task for the life of the run. What this ends is the + reading and the keeping: nothing further is held for this stream, so + a run that closes a reader early does not grow for the rest of its + life. The subscription itself, and the delivery it costs each task, + stay until the run ends. + """ + if self._closed: + return + self._closed = True + workflow._close_stream_records(self._stream_id) + + +class _NativeWriteSink: + def __init__(self, topic: str) -> None: + self._topic = topic + + def publish(self, record: WireRecord) -> None: + # Held by the runtime until the task completes, when the task's + # records on this topic become one command the server applies with + # the task: rule 1 through the server's own commit. The body is still + # plaintext here; the worker's payload pass runs after the task. + workflow._append_stream_records([_fingerprint(record)], stream_name=self._topic) + + +class _NativeWorkflowProvider: + """The workflow half: the server's commands and delivered ranges.""" + + def open_reader( + self, topic: str, *, after: Cursor, last: int | None = None + ) -> ReadSource: + check_read_start(after, last) + _require_topic(topic) + run_id = workflow.info().run_id + # The server resolves the position when it registers the subscription + # and records the offset on the subscribed event, so replay never + # resolves it again. + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif after == BEGINNING: + start = StreamStartPosition(earliest=True) + else: + named = _position(after) + assert named is not None + if named[0] != run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run; a run's stream is its own" + ) + start = StreamStartPosition(offset=named[1] + 1) + workflow._subscribe_stream(topic, start=start) + return _NativeReadSource(topic, run_id) + + def open_writer(self, topic: str) -> WriteSink: + _require_topic(topic) + return _NativeWriteSink(topic) + + def on_workflow_start(self) -> None: + pass + + async def on_workflow_finish(self) -> None: + pass + + +class NativeProducer(Generic[T]): + """Appends to a topic from outside workflow code. + + Every append is visible as soon as the server accepts it. That is the + point for an activity streaming model output, and it is why an activity + carries its own identity: the retry of a failed attempt has no commit + boundary to sort it out afterwards. + """ + + def __init__( + self, + handle: WorkflowStreamHandle | _StandaloneTarget, + pin: Any, + converter: DataConverter, + store_target: Callable[[str], StorageDriverWorkflowInfo], + topic: str, + producer_id: str, + attempt: int, + ) -> None: + """Bind this producer to ``topic`` on the stream ``handle`` names. + + ``converter`` is the client's data converter, applied to every body + as the worker applies it to a workflow's own records, with the + plaintext hash stamped first. ``store_target`` names the execution an + offloaded body is stored under, given the run the producer pinned. + """ + self._handle = handle + self._pin = pin + self._converter = converter + self._store_target = store_target + self._bound: DataConverter | None = None + self._topic = topic + self._producer_id = producer_id + self._attempt = attempt + # One-based, because zero on the wire says the producer does not + # number its records and this one does. + self._sequence = 1 + self._last = BEGINNING + + @property + def producer_id(self) -> str: + """Who this producer writes as.""" + return self._producer_id + + @property + def attempt(self) -> int: + """The generation this producer is writing.""" + return self._attempt + + @property + def _writer(self) -> str: + # The server dedupes on this and the sequence. The attempt is part of + # it so a retried append is dropped while a new generation writing + # different words at the same sequence is not. + 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, because the server + answers a deduplicated batch with the original offsets; an empty call + returns the position of this producer's last record. + """ + if not values: + return self._last + return await self._write( + [ + to_wire( + self._converter.payload_converter, + topic=self._topic, + kind=RecordKind.DATA, + value=value, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence + index, + ) + for index, value in enumerate(values) + ] + ) + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic.""" + await self._write( + [ + to_wire( + self._converter.payload_converter, + topic=self._topic, + kind=RecordKind.FINISH, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + ) + ] + ) + + async def _write(self, records: list[WireRecord]) -> Cursor: + # Pinned before the first write, so every cursor this producer hands + # out names the run its records landed in. + if not self._handle.owner_run_id: + self._handle.pin(await self._pin()) + if self._bound is None: + # An offloaded body is stored under the execution that owns the + # stream, as the worker stores a workflow's own. + self._bound = self._converter._with_store_context( + StorageDriverStoreContext( + target=self._store_target(self._handle.owner_run_id) + ) + ) + for record in records: + await encode_body(self._bound, record) + appended = await self._handle.append( + *records, producer_id=self._writer, sequence=self._sequence + ) + self._sequence += len(records) + self._last = _cursor(self._handle.owner_run_id, appended.next_offset - 1) + return self._last + + +class NativeStreamHandle: + """One workflow's topics from outside, over the stream service.""" + + def __init__( + self, + client: Client, + workflow_id: str, + run_id: str | None, + *, + opened: set[SharedKey] | None = None, + ) -> None: + """Address ``workflow_id``'s topics, pinned to ``run_id`` when one is given. + + ``opened`` is where this handle records the shared channel it used, so + the provider that made it closes that one and no other. + """ + self._client = client + self._workflow_id = workflow_id + self._run_id = run_id + self._opened = set() if opened is None else opened + self._data_converter = client.data_converter + self._converter = client.data_converter.payload_converter + self._streams: StreamClient | None = None + + def _service(self) -> StreamClient: + # Resolved on first use, because the shared channel belongs to the + # running loop and a handle may be made before there is one. + if self._streams is None: + self._streams = shared_client(self._client) + self._opened.add(shared_key(self._client)) + return self._streams + + def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: + return self._service().workflow_stream( + self._workflow_id, topic, owner_run_id=run_id + ) + + def read( + self, + *, + topic: str | StreamTopic[Any] | None = None, + after: Cursor = BEGINNING, + last: int | None = None, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` from where the read starts until the chain, or the pinned run, closes. + + ``BEGINNING`` is the oldest record the chain's first retained run + still holds. ``END`` and ``last=`` start on the current run, or the + pinned one: an earlier run of a chain has ended and holds neither the + tail nor the newest records. The server resolves each on the first + poll, in the read that serves it. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + # Parsed here so a foreign cursor fails this call, not the first + # iteration of the generator. + named = None if after == END else _position(after) + if named is not None and self._run_id is not None and named[0] != self._run_id: + raise StreamCursorError( + f"cursor {after.token!r} names another run than this handle is pinned to" + ) + start: StreamStartPosition | None = None + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif named is None: + start = StreamStartPosition(earliest=True) + return self._read(topic, named, start, after, result_type) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + start: StreamStartPosition | None, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + decoder: RecordDecoder | None = None + offset = 0 + if named is not None: + run_id, offset = named[0], named[1] + 1 + elif start is not None and start.WhichOneof("position") != "earliest": + run_id = self._run_id or await self._current_run() + else: + run_id = self._run_id or await self._first_run() + while True: + stream = self._stream(topic, run_id) + while True: + page = await stream.poll(from_offset=offset, start=start) + if decoder is None: + decoder = RecordDecoder( + self._converter, + result_type, + after=self._previous(run_id, page, start, after), + warn=logger.warning, + ) + start = None + for entry in page.entries: + record = await decode_body(self._data_converter, entry.record) + for out in decoder.decode(_cursor(run_id, entry.offset), record): + yield out + offset = page.next_offset + # On a pinned stream the server reports the run's end as closed, + # and a closed stream is finished once its head is delivered. + if page.closed and offset >= page.head_offset: + break + if self._run_id is not None: + return + successor = await self._successor(run_id) + if successor is None: + return + run_id, offset = successor, 0 + + @staticmethod + def _previous( + run_id: str, page: Page, start: StreamStartPosition | None, after: Cursor + ) -> Cursor: + """The position before the first record a read yields. + + A synthesized record is positioned there. After ``END`` or ``last=`` + it is only known once the server resolved the start, from the first + page. It names a run so a chain-following resume stays on this one. + """ + if start is None or start.WhichOneof("position") == "earliest": + return after + first = page.entries[0].offset if page.entries else page.next_offset + return _cursor(run_id, first - 1) + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The cursor of the newest record on ``topic``, naming the run it was read from. + + An empty topic on the chain's first run is the beginning of the + stream; on a successor it is a position of its own, because + ``BEGINNING`` would send a chain-following read back to the first run. + """ + topic, _ = resolve_topic(topic) + run_id = self._run_id or await self._current_run() + try: + head = (await self._stream(topic, run_id).describe()).head_offset + except StreamNotFoundError: + # A topic nobody has written yet does not exist on the server, + # which is the same answer as an empty one. + head = 0 + if head > 0: + return _cursor(run_id, head - 1) + if self._run_id is None and await self._predecessor(run_id) is not None: + return _cursor(run_id, -1) + return BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> NativeProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return NativeProducer( + self._stream(topic, self._run_id or ""), + self._current_run, + self._data_converter, + self._store_target, + topic, + producer_id, + attempt, + ) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this workflow's stream, pinned as this handle is.""" + return StreamRef.for_workflow( + self._workflow_id, run_id=self._run_id, topic=topic + ) + + async def close(self) -> None: + """Refuse: a workflow's stream ends with the workflow. + + Raises: + ValueError: Always; only a standalone stream is closed by hand. + """ + raise ValueError( + "a workflow's stream ends with its workflow; only a standalone stream " + "can be closed" + ) + + def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: + """The execution an offloaded body of this owner's stream is stored under.""" + return StorageDriverWorkflowInfo( + namespace=self._client.namespace, id=self._workflow_id, run_id=run_id + ) + + async def _current_run(self) -> str: + try: + description = await self._client.get_workflow_handle( + self._workflow_id + ).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"workflow {self._workflow_id!r} was not found" + ) from error + raise + assert description.run_id is not None + return description.run_id + + async def _first_run(self) -> str: + """The oldest retained run of the chain, walking back from the latest.""" + run_id = await self._current_run() + while True: + previous = await self._predecessor(run_id) + if previous is None: + return run_id + run_id = previous + + async def _predecessor(self, run_id: str) -> str | None: + handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + try: + async for event in handle.fetch_history_events(page_size=1): + attributes = event.workflow_execution_started_event_attributes + return attributes.continued_execution_run_id or None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + # The run's History is gone: the chain's retained part starts here. + return None + + async def _successor(self, run_id: str) -> str | None: + handle = self._client.get_workflow_handle(self._workflow_id, run_id=run_id) + events = handle.fetch_history_events( + event_filter_type=WorkflowHistoryEventFilterType.CLOSE_EVENT + ) + try: + async for event in events: + if event.HasField( + "workflow_execution_continued_as_new_event_attributes" + ): + attributes = ( + event.workflow_execution_continued_as_new_event_attributes + ) + return attributes.new_execution_run_id or None + except RPCError as error: + if error.status != RPCStatusCode.NOT_FOUND: + raise + return None + + +class NativeActivityStreamHandle(NativeStreamHandle): + """The topics one activity owns, from outside, over the stream service. + + An activity's streams belong to one activity execution, not to a chain of + runs: a retry writes to the same stream and a read ends when the activity + reaches a terminal status. So the handle pins the execution on first use + and never follows a successor. A standalone activity is its own owner; an + activity a workflow scheduled is reached through that workflow's run. + """ + + def __init__( + self, + client: Client, + activity_id: str, + workflow_id: str | None, + run_id: str | None, + *, + opened: set[SharedKey] | None = None, + ) -> None: + """Address ``activity_id``'s topics, pinned to ``run_id`` when one is given.""" + super().__init__(client, workflow_id or "", run_id, opened=opened) + self._activity_id = activity_id + + def _stream(self, topic: str, run_id: str) -> WorkflowStreamHandle: + return self._service().activity_stream( + self._activity_id, topic, workflow_id=self._workflow_id, run_id=run_id + ) + + def _store_target(self, run_id: str) -> StorageDriverWorkflowInfo: + # A workflow's activity stores under that workflow, as the worker does + # for its activities; a standalone activity has no workflow to name. + if self._workflow_id: + return super()._store_target(run_id) + return StorageDriverWorkflowInfo(namespace=self._client.namespace) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this activity's streams, pinned as this handle is.""" + return StreamRef.for_activity( + self._activity_id, + workflow_id=self._workflow_id or None, + run_id=self._run_id, + topic=topic, + ) + + async def close(self) -> None: + """Refuse: an activity's streams end with the activity. + + Raises: + ValueError: Always; only a standalone stream is closed by hand. + """ + raise ValueError( + "an activity's streams end with the activity; only a standalone stream " + "can be closed" + ) + + async def _current_run(self) -> str: + if self._workflow_id: + return await super()._current_run() + try: + description = await self._client.get_activity_handle( + self._activity_id + ).describe() + except RPCError as error: + if error.status == RPCStatusCode.NOT_FOUND: + raise StreamNotFoundError( + f"activity {self._activity_id!r} was not found" + ) from error + raise + assert description.activity_run_id is not None + return description.activity_run_id + + async def _first_run(self) -> str: + return await self._current_run() + + async def _predecessor(self, run_id: str) -> str | None: + return None + + async def _successor(self, run_id: str) -> str | None: + return None + + +class _StandaloneTarget: + """A standalone stream as the producer addresses an owned one. + + The producer pins a run before its first write and names it in every + cursor it hands out. A standalone stream has no run to pin, so its id + stands in that place from the start and nothing is ever resolved. + """ + + def __init__(self, handle: ServiceStreamHandle, stream_id: str) -> None: + self._handle = handle + self.owner_run_id = stream_id + + def pin(self, run_id: str) -> None: + self.owner_run_id = run_id + + async def append( + self, *records: WireRecord, producer_id: str = "", sequence: int = 0 + ) -> Appended: + return await self._handle.append( + *records, producer_id=producer_id, sequence=sequence + ) + + +class NativeStandaloneStreamHandle: + """One standalone stream's topics from outside, over the stream service. + + A standalone stream has an id of its own and no owner, so there is no + chain to follow and no run to pin; a cursor names the stream id where an + owned stream's names a run, and a cursor from another stream is refused. + Its topics share one log, so a topic read is the server's filter over it. + + A read on an id nobody has created yet parks on the server until the + stream appears, so a reader can attach before the producer's first write. + ``latest`` and a producer's append on such an id raise + :class:`temporalio.streams.StreamNotFoundError`: they have nothing to wait + on. The policy is the server's lifecycle: ``retention`` is the age past + which an open stream's records are reclaimed and how long a closed one + stays readable, ``max_records`` drops the oldest records as new ones land, + and ``max_bytes`` refuses an append that would take the held bytes past it + rather than reclaiming anything. + """ + + def __init__( + self, client: Client, stream_id: str, *, opened: set[SharedKey] | None = None + ) -> None: + """Address the standalone stream ``stream_id``, which must exist or be created later.""" + self._client = client + self._stream_id = stream_id + self._opened = set() if opened is None else opened + self._data_converter = client.data_converter + self._converter = client.data_converter.payload_converter + self._streams: StreamClient | None = None + + @property + def stream_id(self) -> str: + """The id of the stream this handle is on.""" + return self._stream_id + + def _service(self) -> StreamClient: + if self._streams is None: + self._streams = shared_client(self._client) + self._opened.add(shared_key(self._client)) + return self._streams + + def _stream(self) -> ServiceStreamHandle: + return self._service().get(self._stream_id) + + def read( + self, + *, + topic: str | StreamTopic[Any] | None = None, + after: Cursor = BEGINNING, + last: int | None = None, + result_type: type | None = None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + """Yield records on ``topic`` from where the read starts until the stream is closed and drained. + + The server resolves ``BEGINNING``, ``END`` and ``last=`` on the first + poll. On an id that does not exist yet the poll parks until the stream + is created, so the read is open before the first write. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + named = None if after == END else _position(after) + if named is not None and named[0] != self._stream_id: + raise StreamCursorError( + f"cursor {after.token!r} names another stream than {self._stream_id!r}" + ) + start: StreamStartPosition | None = None + if last is not None: + start = StreamStartPosition(last_n=last) + elif after == END: + start = StreamStartPosition(tail=True) + elif named is None: + start = StreamStartPosition(earliest=True) + return self._read(topic, named, start, after, result_type) + + async def _read( + self, + topic: str, + named: tuple[str, int] | None, + start: StreamStartPosition | None, + after: Cursor, + result_type: type | None, + ) -> AsyncGenerator[StreamRecord[Any], None]: + stream = self._stream() + decoder: RecordDecoder | None = None + offset = named[1] + 1 if named is not None else 0 + while True: + asked = time.monotonic() + try: + page = await stream.poll( + from_offset=offset, start=start, topics=[topic] + ) + except StreamNotFoundError: + # The server parks a poll on an id nobody has created for its + # wait budget and answers NOT_FOUND when that runs out. The + # stream may still be created, so the read keeps waiting; one + # that has delivered before is gone for good. A server that + # answers at once does not park, so the wait is paced here. + if decoder is not None: + raise + if time.monotonic() - asked < _CREATE_WAIT_PACE: + await asyncio.sleep(_CREATE_WAIT_PACE) + continue + if decoder is None: + decoder = RecordDecoder( + self._converter, + result_type, + after=NativeStreamHandle._previous( + self._stream_id, page, start, after + ), + warn=logger.warning, + ) + start = None + for entry in page.entries: + record = await decode_body(self._data_converter, entry.record) + for out in decoder.decode( + _cursor(self._stream_id, entry.offset), record + ): + yield out + offset = page.next_offset + if page.closed and offset >= page.head_offset: + return + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The cursor of the newest record on the stream, or ``BEGINNING`` when empty. + + Topics share the stream's offsets, so the newest record may be on + another topic; a read after this cursor still yields only what lands + on ``topic`` later. + + Raises: + StreamNotFoundError: The stream does not exist. + """ + resolve_topic(topic) + head = (await self._stream().describe()).head_offset + return _cursor(self._stream_id, head - 1) if head > 0 else BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> NativeProducer[Any]: + """A producer on ``topic``; inside an activity its identity is the activity's.""" + topic, _ = resolve_topic(topic) + producer_id, attempt = producer_identity(producer_id, attempt) + return NativeProducer( + _StandaloneTarget(self._stream(), self._stream_id), + self._no_run, + self._data_converter, + self._store_target, + topic, + producer_id, + attempt, + ) + + async def _no_run(self) -> str: + return self._stream_id + + def _store_target(self, _run_id: str) -> StorageDriverWorkflowInfo: + # No workflow owns the stream, so an offloaded body has only the + # namespace to be stored under. + return StorageDriverWorkflowInfo(namespace=self._client.namespace) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A ref to ``topic`` of this stream.""" + return StreamRef.for_standalone(self._stream_id, topic=topic) + + async def close(self) -> None: + """Seal the stream: appends are refused from now on and the tail stays readable. Idempotent.""" + await self._stream().close() + + +class NativeStreams(ProviderPlugin): + """The server-side provider. + + Takes no options: the streams are on the server the client is already + connected to. Construct one, pass it to the worker as a plugin and open + handles from it anywhere else; :meth:`close` releases the channels this + provider opened to the stream service. + """ + + def __init__(self) -> None: + """Create the provider.""" + # What this provider's handles opened, so closing it leaves another + # provider's channels on the same loop alone. + self._opened: set[SharedKey] = set() + + def workflow_provider(self) -> _NativeWorkflowProvider: + """The workflow half, over the server's commands and delivered ranges.""" + return _NativeWorkflowProvider() + + def get_stream_handle( + self, + client: Client, + workflow_id: str | StreamRef, + *, + run_id: str | None = None, + ) -> StreamHandle: + """A handle on ``workflow_id``'s topics; without ``run_id`` it follows the chain. + + A :class:`temporalio.streams.StreamRef` in place of the id opens the + stream it names, whatever its owner kind, with the ref's topic as the + handle's default. + """ + if isinstance(workflow_id, StreamRef): + if run_id is not None: + raise ValueError("a StreamRef names the run itself; pass no run_id") + return open_ref(self, client, workflow_id) + return NativeStreamHandle(client, workflow_id, run_id, opened=self._opened) + + 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, + ) -> NativeStandaloneStreamHandle: + """Create the standalone stream ``stream_id`` on the server and return a handle on it. + + The three bounds are the server's lifecycle: ``retention`` reclaims + records older than it on an open stream and times a closed one's + deletion, ``max_records`` drops the oldest records as new ones land, + and ``max_bytes`` refuses an append that would take the held bytes past + it. A create of an id that exists with the same policy returns a handle + on it; with another policy the server refuses it as a ``ValueError``. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + if retention is not None and retention <= timedelta(0): + raise ValueError("retention must be positive") + if max_records is not None and max_records <= 0: + raise ValueError("max_records must be positive") + if max_bytes is not None and max_bytes <= 0: + raise ValueError("max_bytes must be positive") + streams = shared_client(client) + self._opened.add(shared_key(client)) + try: + await streams.create( + stream_id, + retention=retention, + max_items=max_records, + max_bytes=max_bytes, + ) + except RPCError as error: + # The id exists with this policy; a policy that differs arrives + # typed, as the ValueError the contract names. + if error.status != RPCStatusCode.ALREADY_EXISTS: + raise + return NativeStandaloneStreamHandle(client, stream_id, opened=self._opened) + + def get_standalone_stream_handle( + self, client: Client, stream_id: str + ) -> NativeStandaloneStreamHandle: + """A handle on the standalone stream ``stream_id``. + + Nothing here creates the stream. A read on an id that does not exist + yet parks on the server until it is created; ``latest`` and a + producer's append raise :class:`temporalio.streams.StreamNotFoundError`. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + return NativeStandaloneStreamHandle(client, stream_id, opened=self._opened) + + def get_activity_stream_handle( + self, + client: Client, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> NativeActivityStreamHandle: + """A handle on the topics ``activity_id`` owns, apart from any workflow's. + + Without ``workflow_id`` the activity is a standalone one and ``run_id`` + pins its run; with one it is that workflow's activity and ``run_id`` + pins the workflow's run. + """ + return NativeActivityStreamHandle( + client, activity_id, workflow_id, run_id, opened=self._opened + ) + + async def close(self) -> None: + """Close the channels this provider opened to the stream service. + + The application calls this; no worker or client owns the provider's + lifetime, because one provider serves the workers built from a client + and every handle opened outside them. + """ + await close_shared_clients(*self._opened) + self._opened.clear() diff --git a/temporalio/worker/_command_aware_visitor.py b/temporalio/worker/_command_aware_visitor.py index 7c03c2cd4..fa6d5ba18 100644 --- a/temporalio/worker/_command_aware_visitor.py +++ b/temporalio/worker/_command_aware_visitor.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from temporalio.api.enums.v1.command_type_pb2 import CommandType +from temporalio.api.stream.v1 import StreamRecord from temporalio.bridge._visitor import PayloadVisitor from temporalio.bridge._visitor_functions import VisitorFunctions from temporalio.bridge.proto.workflow_activation.workflow_activation_pb2 import ( @@ -26,6 +27,7 @@ StartChildWorkflowExecution, WorkflowCommand, ) +from temporalio.streams._body import CONTENT_HASH_KEY @dataclass(frozen=True) @@ -116,6 +118,18 @@ async def _visit_coresdk_workflow_commands_ScheduleNexusOperation( with current_command(CommandType.COMMAND_TYPE_SCHEDULE_NEXUS_OPERATION, o.seq): await super()._visit_coresdk_workflow_commands_ScheduleNexusOperation(fs, o) + async def _visit_temporal_api_stream_v1_StreamRecord( + self, fs: VisitorFunctions, o: StreamRecord + ) -> None: + if o.HasField("body"): + await self._visit_temporal_api_common_v1_Payload(fs, o.body) + for key, value in o.metadata.items(): + # The declared content hash is the server's dedupe identity and is + # read as sent, so neither the codec nor the store may rewrite it. + if key == CONTENT_HASH_KEY: + continue + await self._visit_temporal_api_common_v1_Payload(fs, value) + async def _visit_coresdk_workflow_commands_WorkflowCommand( self, fs: VisitorFunctions, o: WorkflowCommand ) -> None: diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index c53ea6f88..ce2aff001 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -47,6 +47,7 @@ import temporalio.api.common.v1 import temporalio.api.enums.v1 import temporalio.api.sdk.v1 +import temporalio.api.stream.v1 import temporalio.bridge.proto.activity_result import temporalio.bridge.proto.child_workflow import temporalio.bridge.proto.common @@ -284,6 +285,168 @@ def create_instance(self, det: WorkflowInstanceDetails) -> WorkflowInstance: _ExceptionHandler: TypeAlias = Callable[[asyncio.AbstractEventLoop, _Context], Any] +# Match the server's per-batch limits. A record over its limit is refused where +# it is published, because a rejected command would be reissued on every +# replay; a task's records are split into commands that fit the batch limits. +# +# Copied rather than learned: the activation does not carry them and the +# server does not report them, so this is a second copy of a number somebody +# else owns. If the server lowers one, or makes it per namespace, the split +# here stops fitting and the command is rejected on every replay, which is +# the failure the split exists to avoid. Carrying them on the activation is +# what would fix that, and it needs a Core and server change. +_STREAM_CONTINUITY_REMEDY = ( + "This fails the Workflow Task and will keep failing it, because the range " + "is recorded as consumed and will not be sent again. Reset the workflow to " + "before the subscription to start its stream reading over, or terminate it " + "if its output is no longer wanted." +) + +_MAX_STREAM_RECORDS_PER_BATCH = 1000 +_MAX_STREAM_RECORD_BYTES = 1 << 20 +_MAX_STREAM_BATCH_BYTES = 2 << 20 + + +def _stream_batches( + records: Sequence[temporalio.api.stream.v1.StreamRecord], +) -> Iterator[list[temporalio.api.stream.v1.StreamRecord]]: + """Split one task's records for a stream into batches the server accepts.""" + batch: list[temporalio.api.stream.v1.StreamRecord] = [] + size = 0 + for record in records: + record_size = record.ByteSize() + if batch and ( + len(batch) >= _MAX_STREAM_RECORDS_PER_BATCH + or size + record_size > _MAX_STREAM_BATCH_BYTES + ): + yield batch + batch, size = [], 0 + batch.append(record) + size += record_size + if batch: + yield batch + + +def _is_completion_command( + command: temporalio.bridge.proto.workflow_commands.WorkflowCommand, +) -> bool: + return ( + command.HasField("complete_workflow_execution") + or command.HasField("continue_as_new_workflow_execution") + or command.HasField("fail_workflow_execution") + or command.HasField("cancel_workflow_execution") + ) + + +class _StreamBuffer: + """Holds the stream ranges delivered to a workflow so far. + + Delivery is driven by the server, not by whether workflow code happens to be + reading. A range arrives once, is recorded in History as consumed, and is + never sent again, so anything not yet read has to be kept here rather than + dropped. + """ + + def __init__(self, stream_id: str = "") -> None: + self._stream_id = stream_id + self._records: list[temporalio.workflow._DeliveredStreamRecord] = [] + self._waiters: list[asyncio.Future] = [] + # Where the next range has to start. Unknown until the first one + # arrives, because a subscription may start wherever the stream is and + # the server is the one that resolves that. + self._next_offset: int | None = None + self._closed = False + + @property + def closed(self) -> bool: + """Whether workflow code has said it wants no more of this stream.""" + return self._closed + + def close(self) -> None: + """Stop keeping what arrives, and let go of what is held. Idempotent. + + There is no unsubscribe command, so the server keeps delivering for + the life of the run. Holding those records would grow the instance + without bound for a reader nobody will read again. Dropping them is + replay-safe because the close happens at the same point of the same + workflow code every time, so the same ranges are dropped; continuity + is still tracked, so a range that repeats or skips is still caught. + """ + self._closed = True + self._records = [] + waiters, self._waiters = self._waiters, [] + for waiter in waiters: + if not waiter.done(): + waiter.set_result(None) + + def extend( + self, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + from_offset: int = 0, + to_offset: int | None = None, + ) -> None: + if to_offset is None: + to_offset = from_offset + len(records) + # A range is recorded as consumed once and never resent, so one that + # repeats, skips or mis-sizes would hand the workflow duplicate or + # shifted records with nothing to say so. Failing the task is what + # makes the fault visible. + if to_offset - from_offset != len(records): + raise RuntimeError( + f"stream {self._stream_id!r} delivered {len(records)} records " + f"for offsets [{from_offset}, {to_offset}). {_STREAM_CONTINUITY_REMEDY}" + ) + if self._next_offset is not None and from_offset != self._next_offset: + raise RuntimeError( + f"stream {self._stream_id!r} delivered offsets [{from_offset}, " + f"{to_offset}) but the last range ended at {self._next_offset}. " + f"{_STREAM_CONTINUITY_REMEDY}" + ) + self._next_offset = to_offset + # An empty range still counts as a delivery, but there is nothing to + # hand a reader, so only a non-empty one wakes anyone. A closed buffer + # counts the range and keeps nothing: continuity is still checked + # above, and nobody is left to read what it held. + if not records or self._closed: + return + # Offsets are dense inside a delivered range and the range arrives in + # order, so counting from its start is the position rather than an + # estimate of it. The per-record field is not on the activation, and a + # reader that has to resume elsewhere needs a position it can name. + for index, record in enumerate(records): + # Copied so the buffer outlives the activation that carried it. + kept = temporalio.api.stream.v1.StreamRecord() + kept.CopyFrom(record) + self._records.append( + temporalio.workflow._DeliveredStreamRecord( + record=kept, offset=from_offset + index + ) + ) + waiters, self._waiters = self._waiters, [] + for waiter in waiters: + if not waiter.done(): + waiter.set_result(None) + + def take(self) -> list[temporalio.workflow._DeliveredStreamRecord]: + taken, self._records = self._records, [] + return taken + + def put_back( + self, records: Sequence[temporalio.workflow._DeliveredStreamRecord] + ) -> None: + """Return an unread tail to the front of the buffer.""" + self._records[:0] = records + + def wait_future(self) -> asyncio.Future: + loop = asyncio.get_event_loop() + fut = loop.create_future() + self._waiters.append(fut) + return fut + + def __len__(self) -> int: + return len(self._records) + + class _WorkflowInstanceImpl( # type: ignore[reportImplicitAbstractClass] WorkflowInstance, temporalio.workflow._Runtime, asyncio.AbstractEventLoop ): @@ -423,6 +586,16 @@ def __init__(self, det: WorkflowInstanceDetails) -> None: str, list[temporalio.bridge.proto.workflow_activation.SignalWorkflow] ] = {} + # Stream ranges delivered to this workflow, keyed by stream id. Ranges + # arrive whether or not anything is reading yet, because the server has + # already recorded them as consumed and will not send them again. + self._stream_buffers: dict[str, _StreamBuffer] = {} + # Records this task's publishes append, by stream, until the task + # completes and they become commands. + self._stream_appends: dict[ + str, list[temporalio.api.stream.v1.StreamRecord] + ] = {} + # When we evict, we have to mark the workflow as deleting so we don't # add any commands and we swallow exceptions on tear down self._deleting = False @@ -483,6 +656,9 @@ def activate( temporalio.bridge.proto.workflow_completion.WorkflowActivationCompletion() ) self._current_completion.successful.SetInParent() + # A failed task's publishes never became commands, so nothing carries + # over into this one. + self._stream_appends = {} self._current_activation_error: Exception | None = None self._deployment_version_for_current_task = ( @@ -574,6 +750,9 @@ def activate( ) activation_err = None + if activation_err is None and not self._deleting: + self._flush_stream_appends() + # Apply versioning behavior if one was established if self._versioning_behavior: self._current_completion.successful.versioning_behavior = ( @@ -638,6 +817,8 @@ def _apply( ) -> None: if job.HasField("cancel_workflow"): self._apply_cancel_workflow(job.cancel_workflow) + elif job.HasField("deliver_stream_records"): + self._apply_deliver_stream_records(job.deliver_stream_records) elif job.HasField("do_update"): self._apply_do_update(job.do_update) elif job.HasField("fire_timer"): @@ -1191,6 +1372,19 @@ def _apply_resolve_signal_external_workflow( else: fut.set_result(None) + def _apply_deliver_stream_records( + self, + job: temporalio.bridge.proto.workflow_activation.DeliverStreamRecords, + ) -> None: + buffer = self._stream_buffers.get(job.stream_id) + if buffer is None: + # Nothing subscribed. The range is already recorded as consumed and + # will not be sent again, so buffering it is the only way a + # subscription made later in the same task still sees it. + buffer = _StreamBuffer(job.stream_id) + self._stream_buffers[job.stream_id] = buffer + buffer.extend(job.records, job.from_offset, job.to_offset) + def _apply_signal_workflow( self, job: temporalio.bridge.proto.workflow_activation.SignalWorkflow ) -> None: @@ -1357,6 +1551,85 @@ def workflow_get_current_deployment_version( def get_info(self) -> temporalio.workflow.Info: return self._info + def workflow_subscribe_stream( + self, + stream_name_or_id: str, + start: temporalio.api.stream.v1.StreamStartPosition, + ) -> None: + # Reissued on every replay, so the buffer has to exist before the first + # range arrives and the command has to be harmless the second time. A + # repeat subscription leaves the server-side cursor where it is. + self._stream_buffers.setdefault( + stream_name_or_id, _StreamBuffer(stream_name_or_id) + ) + command = self._add_command() + command.subscribe_stream.stream_name_or_id = stream_name_or_id + command.subscribe_stream.start_position.CopyFrom(start) + + def workflow_append_stream_records( + self, + stream_name: str, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + ) -> None: + self._assert_not_read_only("append stream records") + if not records: + raise ValueError("append_stream_records needs at least one record") + kept: list[temporalio.api.stream.v1.StreamRecord] = [] + for record in records: + if record.ByteSize() > _MAX_STREAM_RECORD_BYTES: + raise ValueError( + f"a stream record is limited to {_MAX_STREAM_RECORD_BYTES} " + f"bytes, got {record.ByteSize()}" + ) + copy = temporalio.api.stream.v1.StreamRecord() + copy.CopyFrom(record) + # The workflow is the producer here, whatever the caller set. + copy.producer_id = "" + kept.append(copy) + # Held until the task completes, so a task's publishes on one stream + # become one command and one History event however many there were. + self._stream_appends.setdefault(stream_name, []).extend(kept) + + def _flush_stream_appends(self) -> None: + appends, self._stream_appends = self._stream_appends, {} + if not appends: + return + commands = self._current_completion.successful.commands + # Ahead of any command that ends the run, because the server accepts + # nothing after one of those. + insert_at = len(commands) + for index, command in enumerate(commands): + if _is_completion_command(command): + insert_at = index + break + for stream_name, records in appends.items(): + for batch in _stream_batches(records): + command = temporalio.bridge.proto.workflow_commands.WorkflowCommand() + command.append_stream_records.stream_name = stream_name + command.append_stream_records.records.extend(batch) + commands.insert(insert_at, command) + insert_at += 1 + + def workflow_close_stream_records(self, stream_id: str) -> None: + self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)).close() + + async def workflow_read_stream_records( + self, stream_id: str, max_records: int + ) -> list[temporalio.workflow._DeliveredStreamRecord]: + # Ranges arrive on Workflow Tasks, and a query activation carries none, + # so without this the read waits on a future nothing can resolve and the + # query times out with nothing to say why. + self._assert_not_read_only("read stream") + buffer = self._stream_buffers.setdefault(stream_id, _StreamBuffer(stream_id)) + while not len(buffer) and not buffer.closed: + await buffer.wait_future() + taken = buffer.take() + if max_records and len(taken) > max_records: + # Put the tail back rather than dropping it: nothing will resend it. + buffer.put_back(taken[max_records:]) + taken = taken[:max_records] + return taken + def workflow_get_current_history_length(self) -> int: return self._current_history_length diff --git a/temporalio/workflow/__init__.py b/temporalio/workflow/__init__.py index 212625511..98a125dcb 100644 --- a/temporalio/workflow/__init__.py +++ b/temporalio/workflow/__init__.py @@ -63,7 +63,7 @@ linked_channel, subscribe_channel, ) -from ._context import ( +from ._context import ( # noqa: F401 Info, ParentInfo, RootInfo, @@ -103,6 +103,21 @@ uuid7, wait_condition, ) +from ._context import ( + _append_stream_records as _append_stream_records, +) +from ._context import ( + _close_stream_records as _close_stream_records, +) +from ._context import ( + _DeliveredStreamRecord as _DeliveredStreamRecord, +) +from ._context import ( + _read_stream_records as _read_stream_records, +) +from ._context import ( + _subscribe_stream as _subscribe_stream, +) from ._definition import ( DynamicWorkflowConfig, _Definition, diff --git a/temporalio/workflow/_context.py b/temporalio/workflow/_context.py index a04791baf..8b688db19 100644 --- a/temporalio/workflow/_context.py +++ b/temporalio/workflow/_context.py @@ -14,6 +14,7 @@ from nexusrpc import InputT, OutputT import temporalio.api.common.v1 +import temporalio.api.stream.v1 import temporalio.common import temporalio.converter @@ -323,6 +324,28 @@ def workflow_get_current_deployment_version( self, ) -> temporalio.common.WorkerDeploymentVersion | None: ... + @abstractmethod + def workflow_subscribe_stream( + self, + stream_name_or_id: str, + start: temporalio.api.stream.v1.StreamStartPosition, + ) -> None: ... + + @abstractmethod + def workflow_append_stream_records( + self, + stream_name: str, + records: Sequence[temporalio.api.stream.v1.StreamRecord], + ) -> None: ... + + @abstractmethod + def workflow_close_stream_records(self, stream_id: str) -> None: ... + + @abstractmethod + async def workflow_read_stream_records( + self, stream_id: str, max_records: int + ) -> list[_DeliveredStreamRecord]: ... + @abstractmethod def workflow_get_current_history_length(self) -> int: ... @@ -1043,6 +1066,110 @@ async def sleep( ) +def _subscribe_stream( # type: ignore[reportUnusedFunction] + stream_name_or_id: str, + *, + start: temporalio.api.stream.v1.StreamStartPosition | None = None, +) -> None: + """Subscribe this workflow to a server-side stream. + + From here on its Workflow Tasks carry the ranges it has not consumed yet, + and :func:`_read_stream_records` returns them. Safe to call again: a second + subscription to a stream this run already consumes does not move its + cursor, though it does write one event. Calling it on every replay is + harmless because replay matches the command to the event already recorded. + + Only the name or id and the start position go to the server. The rest of the + stream's addressing is resolved there, because a workflow cannot look it up + without doing I/O and a value it carried would be a reading rather than a + fact. A name this workflow has not written yet names a stream it owns, and + subscribing creates it. + + Args: + stream_name_or_id: Stream to consume: the name of one this workflow + owns, or the id of a standalone stream. The server tries them in + that order. + start: Where to start: an absolute offset, the oldest record the + stream holds, the tail as of registration, or the last N records. + Omitted, it is the oldest record held. The server resolves it + once and records the offset, so replay does not resolve it again. + """ + if start is None: + start = temporalio.api.stream.v1.StreamStartPosition(earliest=True) + _Runtime.current().workflow_subscribe_stream(stream_name_or_id, start) + + +def _append_stream_records( # type: ignore[reportUnusedFunction] + records: Sequence[temporalio.api.stream.v1.StreamRecord], + *, + stream_name: str = "", +) -> None: + """Publish records to a server-side stream this workflow owns. + + Returns at once. The records become one command when this Workflow Task + completes, so they are visible when the task is accepted and never if it + fails. Their bodies go to the stream's own log rather than into History, + which gets one fixed-size event naming the offset range, so a task that + publishes a thousand records costs History the same as one that publishes + one. Readers do not have to exist yet, and adding one costs the writer + nothing. + + Args: + records: Records to append, in order. The server stores each with an + empty ``producer_id``, because the workflow is the producer. + stream_name: Name of a stream this workflow owns, created on first + use. Empty means the workflow's default output stream. A workflow + cannot append to a stream another execution owns. + + Raises: + ValueError: ``records`` is empty or one of them is over the server's + per-record size limit. + """ + _Runtime.current().workflow_append_stream_records(stream_name, records) + + +@dataclass(frozen=True) +class _DeliveredStreamRecord: + """One record a consuming workflow was given, with where it sat.""" + + record: temporalio.api.stream.v1.StreamRecord + offset: int + """Its position in the whole stream, which is what a reader resumes from.""" + + +def _close_stream_records(stream_id: str) -> None: # type: ignore[reportUnusedFunction] + """Say this workflow wants no more of ``stream_id``. + + There is no unsubscribe command, so the server keeps delivering for the + life of the run; this drops what arrives instead of holding it for a + reader that has gone. Deterministic on replay, because the same workflow + code closes at the same point and the same ranges are dropped. + + Args: + stream_id: Stream to stop keeping records for. + """ + _Runtime.current().workflow_close_stream_records(stream_id) + + +async def _read_stream_records( # type: ignore[reportUnusedFunction] + stream_id: str, *, max_records: int = 0 +) -> list[_DeliveredStreamRecord]: + """Read the next records of a server-side stream this workflow consumes. + + Waits until at least one record is available. Ranges arrive on Workflow + Tasks, and only the offsets they covered are written to History, so this is + deterministic on replay: the server re-supplies the same ranges by reading + the stream again. Subscribe first with :func:`_subscribe_stream`; this only + reads what has already been delivered to this workflow. + + Args: + stream_id: Stream to read from. + max_records: Most records to return at once, or 0 for everything + available. + """ + return await _Runtime.current().workflow_read_stream_records(stream_id, max_records) + + async def wait_condition( fn: Callable[[], bool], *, diff --git a/tests/contrib/test_server_streams.py b/tests/contrib/test_server_streams.py new file mode 100644 index 000000000..540f7f7d7 --- /dev/null +++ b/tests/contrib/test_server_streams.py @@ -0,0 +1,161 @@ +"""The Workflow Streams surface over a server-side stream. + +Both producers and the consumer, since the point of the surface is that an +application does not have to know which of them wrote a given item. Needs a +Temporal server built from the AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/contrib/test_server_streams.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from dataclasses import dataclass +from datetime import timedelta + +import pytest + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.contrib.server_streams import WorkflowStream, WorkflowStreamClient +from temporalio.worker import Worker + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream service", +) + +TOPIC = "turn_events" + + +@dataclass +class Event: + source: str + text: str + + +@activity.defn +async def emit_from_activity(count: int) -> None: + async with WorkflowStreamClient.from_within_activity() as client: + events = client.topic(TOPIC, type=Event) + for i in range(count): + events.publish(Event(source="activity", text=f"token {i}")) + + +@workflow.defn +class Emitting: + def __init__(self) -> None: + self._events = WorkflowStream().topic(TOPIC, type=Event) + + @workflow.run + async def run(self, count: int) -> None: + self._events.publish(Event(source="workflow", text="turn started")) + await workflow.execute_activity( + emit_from_activity, + count, + start_to_close_timeout=timedelta(seconds=30), + ) + self._events.publish(Event(source="workflow", text="turn ended")) + + +async def test_both_producers_reach_one_subscriber() -> None: + client = await Client.connect(TARGET or "") + task_queue = "ss-tq-" + uuid.uuid4().hex[:8] + wf_id = "ss-wf-" + uuid.uuid4().hex[:8] + tokens = 4 + + seen: list[Event] = [] + offsets: list[int] = [] + + async with Worker( + client, + task_queue=task_queue, + workflows=[Emitting], + activities=[emit_from_activity], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + Emitting.run, tokens, id=wf_id, task_queue=task_queue + ) + + # Subscribed before the Workflow has published anything, which is what + # a consumer attaching to a session does. + stream = WorkflowStreamClient.create(client, wf_id) + + async def read() -> None: + async for item in stream.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ): + seen.append(item.data) + offsets.append(item.offset) + if len(seen) == tokens + 2: + return + + reading = asyncio.ensure_future(read()) + await asyncio.wait_for(handle.result(), timeout=60) + await asyncio.wait_for(reading, timeout=60) + + # The Workflow's own publishes bracket the Activity's, and both are on one + # log in the order the server took them. + assert [e.source for e in seen] == ["workflow"] + ["activity"] * tokens + [ + "workflow" + ] + assert seen[0].text == "turn started" + assert seen[-1].text == "turn ended" + assert offsets == list(range(tokens + 2)) + + # A consumer that arrives after the fact reads the same thing, and is not + # left tailing: the Workflow has ended, so its stream is finished and the + # subscription ends on its own. + late = WorkflowStreamClient.create(client, wf_id) + assert await late.get_offset() == tokens + 2 + replayed = [ + item.data.text + async for item in late.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ) + ] + assert replayed == [e.text for e in seen] + + +async def test_a_reader_resumes_from_an_offset_it_was_given() -> None: + client = await Client.connect(TARGET or "") + task_queue = "ss-resume-tq-" + uuid.uuid4().hex[:8] + wf_id = "ss-resume-wf-" + uuid.uuid4().hex[:8] + + async with Worker( + client, + task_queue=task_queue, + workflows=[Emitting], + activities=[emit_from_activity], + max_cached_workflows=0, + ): + await client.execute_workflow(Emitting.run, 3, id=wf_id, task_queue=task_queue) + + stream = WorkflowStreamClient.create(client, wf_id) + first = [ + item + async for item in stream.subscribe( + topics=[TOPIC], from_offset=0, result_type=Event + ) + ] + assert len(first) == 5 + + # Resuming past the second item skips exactly the two before it, so the + # offset a reader was handed is the position it means. + resumed = [ + item.data.text + async for item in stream.subscribe( + topics=[TOPIC], from_offset=first[1].offset + 1, result_type=Event + ) + ] + assert resumed == [item.data.text for item in first[2:]] + + # The raw page carries the stored records themselves. + raw = await stream.poll_raw(topics=[TOPIC], from_offset=0, wait=False) + assert [item.data.topic for item in raw.items] == [TOPIC] * 5 + assert raw.closed and raw.next_offset == raw.head_offset == 5 diff --git a/tests/contrib/test_server_streams_flush.py b/tests/contrib/test_server_streams_flush.py new file mode 100644 index 000000000..23a72106e --- /dev/null +++ b/tests/contrib/test_server_streams_flush.py @@ -0,0 +1,86 @@ +"""What the buffered client does with a batch, without a server. + +The parity claim in the module docstring is that an application swaps the +import and keeps its code, so the two things the shipped client does on every +flush have to hold here too: the append carries a producer identity, and a +batch whose append failed is kept for the retry rather than lost with it. +""" + +from __future__ import annotations + +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio.api.stream.v1 import StreamRecord +from temporalio.contrib.server_streams import WorkflowStreamClient +from temporalio.converter import DataConverter + + +class _Handle: + """A stream handle that records its appends and can fail the first one.""" + + owner_run_id = "run" + + def __init__(self, *, fail_first: bool = False) -> None: + self.sent: list[tuple[list[StreamRecord], str, int]] = [] + self._fail_first = fail_first + + def pin(self, run_id: str) -> None: + del run_id + + async def append( + self, *records: StreamRecord, producer_id: str = "", sequence: int = 0 + ) -> Any: + self.sent.append((list(records), producer_id, sequence)) + if self._fail_first: + self._fail_first = False + raise ConnectionResetError("the server took it, the reply was lost") + return None + + +def _client(handle: _Handle) -> WorkflowStreamClient: + return WorkflowStreamClient( + handle, # type: ignore[arg-type] + DataConverter.default.payload_converter, + timedelta(seconds=60), + producer_id="act#1", + ) + + +def _bodies(sent: tuple[list[StreamRecord], str, int]) -> list[bytes]: + return [record.body.data for record in sent[0]] + + +async def test_an_append_carries_who_wrote_it_and_where_it_sits() -> None: + handle = _Handle() + client = _client(handle) + client.topic("t").publish({"n": 1}) + client.topic("t").publish({"n": 2}) + await client.flush() + client.topic("t").publish({"n": 3}) + await client.flush() + + # Without an identity the append is at-least-once, which is what the + # shipped client refuses to be. + assert [(who, seq) for _, who, seq in handle.sent] == [("act#1", 0), ("act#1", 2)] + + +async def test_a_batch_whose_append_failed_goes_out_again() -> None: + handle = _Handle(fail_first=True) + client = _client(handle) + client.topic("t").publish({"n": 1}) + with pytest.raises(ConnectionResetError): + await client.flush() + # Taken off the buffer before the await, it would have gone with the + # failure. It goes out again under the sequence it already had, so a copy + # the server did accept is deduplicated and one it never saw lands. + await client.flush() + assert len(handle.sent) == 2 + assert _bodies(handle.sent[0]) == _bodies(handle.sent[1]) + assert [seq for _, _, seq in handle.sent] == [0, 0] + + client.topic("t").publish({"n": 2}) + await client.flush() + assert handle.sent[2][2] == 1, "the next batch continues past the first" diff --git a/tests/streams/conftest.py b/tests/streams/conftest.py index 9a26d2e73..114bca6ea 100644 --- a/tests/streams/conftest.py +++ b/tests/streams/conftest.py @@ -58,6 +58,11 @@ def pytest_configure(config: pytest.Config) -> None: "needs_unsubscribe_server: the case needs a server that accepts the " "unsubscribe-notification-channel command, named with -E host:port", ) + config.addinivalue_line( + "markers", + "needs_stream_channel_server: the case needs a server on which a native " + "stream notifies the channel named by the stream, named with -E host:port", + ) config.addinivalue_line( "markers", "needs_execution_server: the case needs a server that addresses a linked " @@ -102,6 +107,7 @@ def pytest_collection_modifyitems( for marker, what in ( ("needs_describe_server", "lists channel subscriptions on describe"), ("needs_unsubscribe_server", "accepts the unsubscribe command"), + ("needs_stream_channel_server", "notifies a stream's channel"), ("needs_execution_server", "addresses a linked channel by execution"), ): skips.append( diff --git a/tests/streams/test_activity_streams.py b/tests/streams/test_activity_streams.py index 61eaf7fea..f23aca85c 100644 --- a/tests/streams/test_activity_streams.py +++ b/tests/streams/test_activity_streams.py @@ -19,6 +19,7 @@ from __future__ import annotations import asyncio +import os import uuid from collections.abc import AsyncIterator, Callable from dataclasses import dataclass @@ -37,6 +38,7 @@ topic, ) from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.native import NativeStreams from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider from temporalio.testing import WorkflowEnvironment from tests.helpers import new_worker @@ -63,6 +65,21 @@ async def _memory_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider.reset() +async def _native_setup(client: Client) -> AsyncIterator[ActivitySetup]: + # The store is a server built from the stream-carrying branch, which the + # test environment's own server is not; TEMPORAL_ADDRESS names it. + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = NativeStreams() + config = client.config() + config["plugins"] = [provider] + yield ActivitySetup("native", provider, Client(**config)) + await provider.close() + + async def _workflow_streams_setup(client: Client) -> AsyncIterator[ActivitySetup]: provider = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) config = client.config() @@ -78,6 +95,8 @@ async def _workflow_streams_setup(client: Client) -> AsyncIterator[ActivitySetup "memory": _memory_setup, "workflow_streams": _workflow_streams_setup, } +if os.environ.get("STREAMS_LIVE") == "native": + SETUPS["native"] = _native_setup @pytest.fixture(params=sorted(SETUPS)) diff --git a/tests/streams/test_native_fingerprint.py b/tests/streams/test_native_fingerprint.py new file mode 100644 index 000000000..fccc54aaa --- /dev/null +++ b/tests/streams/test_native_fingerprint.py @@ -0,0 +1,151 @@ +"""The native provider stamps every body with the hash of its plaintext. + +The server deduplicates a producer's repeat on that hash when it is present, +so a codec that encrypts with a fresh nonce per call cannot turn a retry into +a divergent write. These pin where the stamp is taken: over the body as the +payload converter produced it, before the codec on the outside path and before +the worker's payload pass on the workflow path. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Sequence +from typing import Any + +import pytest + +import temporalio.converter +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.stream.v1 import StreamRecord +from temporalio.bridge.proto.workflow_completion import WorkflowActivationCompletion +from temporalio.bridge.worker import encode_completion +from temporalio.client_stream import Appended +from temporalio.converter import ( + DataConverter, + PayloadCodec, + StorageDriverWorkflowInfo, +) +from temporalio.streams import CONTENT_HASH_KEY +from temporalio.streams.providers import native +from temporalio.streams.providers.native import NativeProducer + + +class _NonceCodec(PayloadCodec): + """Encodes to something different every time, like a nonce-based cipher.""" + + def __init__(self) -> None: + self.calls = 0 + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + self.calls += 1 + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=f"{self.calls}:".encode() + p.data, + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + return [Payload(data=p.data.split(b":", 1)[1]) for p in payloads] + + +class _Handle: + """Records what a producer appends and answers as the server would.""" + + def __init__(self) -> None: + self.owner_run_id = "run" + self.appended: list[StreamRecord] = [] + + async def append(self, *records: StreamRecord, **_: Any) -> Appended: + self.appended.extend(records) + return Appended(first_offset=0, next_offset=len(records), count=len(records)) + + +def _plaintext_hash(value: Any) -> bytes: + payload = temporalio.converter.default().payload_converter.to_payloads([value])[0] + return hashlib.sha256(payload.SerializeToString()).hexdigest().encode() + + +def _target(_run_id: str) -> StorageDriverWorkflowInfo: + return StorageDriverWorkflowInfo(namespace="ns", id="wf") + + +async def test_an_outside_append_is_stamped_before_the_codec() -> None: + handle: Any = _Handle() + converter = DataConverter(payload_codec=_NonceCodec()) + producer: NativeProducer[Any] = NativeProducer( + handle, None, converter, _target, "t", "p", 1 + ) + + await producer.append({"n": 1}) + await producer.append({"n": 1}) + + first, second = handle.appended + # The bodies went out differently encoded and the stamps agree anyway. + assert first.body.data != second.body.data + assert first.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 1}) + assert second.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 1}) + assert first.metadata[CONTENT_HASH_KEY].metadata["encoding"] == b"binary/plain" + + +async def test_a_finish_record_carries_no_stamp() -> None: + handle: Any = _Handle() + producer: NativeProducer[Any] = NativeProducer( + handle, None, DataConverter.default, _target, "t", "p", 1 + ) + await producer.finish() + (record,) = handle.appended + assert CONTENT_HASH_KEY not in record.metadata + + +def test_a_workflow_publish_is_stamped_on_the_workflow_thread( + monkeypatch: pytest.MonkeyPatch, +) -> None: + staged: list[StreamRecord] = [] + + def stage(records: Sequence[StreamRecord], *, stream_name: str) -> None: + assert stream_name == "t" + staged.extend(records) + + monkeypatch.setattr(workflow, "_append_stream_records", stage) + converter = temporalio.converter.default().payload_converter + body = converter.to_payloads([{"n": 2}])[0] + record = StreamRecord(topic="t", body=body) + + native._NativeWriteSink("t").publish(record) + + (stamped,) = staged + assert stamped.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 2}) + # Two publishes of the same value stamp the same, which is what replay + # reissues. + native._NativeWriteSink("t").publish(StreamRecord(topic="t", body=body)) + assert staged[1].metadata[CONTENT_HASH_KEY] == stamped.metadata[CONTENT_HASH_KEY] + + +async def test_the_workers_payload_pass_leaves_the_stamp_alone() -> None: + """The codec and the store rewrite the body and the other metadata, never the hash. + + The server reads the declared hash as sent and refuses a value that is not + hex, so a codec that touched it would fail the workflow's own append. + """ + converter = DataConverter(payload_codec=_NonceCodec()) + body = converter.payload_converter.to_payloads([{"n": 3}])[0] + record = native._fingerprint(StreamRecord(topic="t", body=body)) + record.metadata["note"].CopyFrom(Payload(data=b"plain")) + completion = WorkflowActivationCompletion() + command = completion.successful.commands.add() + command.append_stream_records.stream_name = "t" + command.append_stream_records.records.append(record) + + await encode_completion( + completion, converter, encode_headers=False, storage_concurrency_limit=1 + ) + + sent = completion.successful.commands[0].append_stream_records.records[0] + assert sent.body.metadata["encoding"] == b"binary/nonce" + assert sent.metadata["note"].metadata["encoding"] == b"binary/nonce" + assert sent.metadata[CONTENT_HASH_KEY].data == _plaintext_hash({"n": 3}) + assert sent.metadata[CONTENT_HASH_KEY].metadata["encoding"] == b"binary/plain" diff --git a/tests/streams/test_native_handles.py b/tests/streams/test_native_handles.py new file mode 100644 index 000000000..318c30494 --- /dev/null +++ b/tests/streams/test_native_handles.py @@ -0,0 +1,102 @@ +"""What the native handles answer without a server. + +A ref names the owner as the handle addresses it, a close on an owned stream +is refused, a standalone handle refuses a cursor from another stream, and a +create refuses a policy the server cannot hold before any call is made. +""" + +from __future__ import annotations + +from datetime import timedelta +from types import SimpleNamespace +from typing import Any + +import pytest + +import temporalio.converter +from temporalio.streams import DEFAULT_TOPIC, StreamCursorError, StreamRef, topic +from temporalio.streams.providers.native import ( + NativeActivityStreamHandle, + NativeStandaloneStreamHandle, + NativeStreamHandle, + NativeStreams, + _cursor, +) + +OUT = topic("out", dict) +_CLIENT: Any = SimpleNamespace( + data_converter=temporalio.converter.default(), namespace="ns" +) + + +def test_a_workflow_handle_refs_its_owner_as_it_is_pinned() -> None: + following = NativeStreamHandle(_CLIENT, "wf", None) + assert following.ref() == StreamRef.for_workflow("wf") + assert following.ref(topic=OUT) == StreamRef.for_workflow("wf", topic="out") + pinned = NativeStreamHandle(_CLIENT, "wf", "run-1") + assert pinned.ref(topic="a") == StreamRef.for_workflow( + "wf", run_id="run-1", topic="a" + ) + + +def test_an_activity_handle_refs_the_activity_and_its_workflow() -> None: + scheduled = NativeActivityStreamHandle(_CLIENT, "act", "wf", "run-1") + assert scheduled.ref() == StreamRef.for_activity( + "act", workflow_id="wf", run_id="run-1" + ) + standalone = NativeActivityStreamHandle(_CLIENT, "act", None, None) + assert standalone.ref(topic=OUT) == StreamRef.for_activity("act", topic="out") + assert standalone.ref().workflow_id is None + + +def test_a_standalone_handle_refs_its_id() -> None: + handle = NativeStandaloneStreamHandle(_CLIENT, "s1") + assert handle.stream_id == "s1" + assert handle.ref() == StreamRef.for_standalone("s1") + assert handle.ref().topic == DEFAULT_TOPIC + assert handle.ref(topic=OUT) == StreamRef.for_standalone("s1", topic="out") + + +async def test_only_a_standalone_stream_can_be_closed() -> None: + with pytest.raises(ValueError, match="standalone"): + await NativeStreamHandle(_CLIENT, "wf", None).close() + with pytest.raises(ValueError, match="standalone"): + await NativeActivityStreamHandle(_CLIENT, "act", "wf", None).close() + + +def test_a_standalone_handle_refuses_another_streams_cursor() -> None: + handle = NativeStandaloneStreamHandle(_CLIENT, "s1") + with pytest.raises(StreamCursorError, match="another stream"): + handle.read(topic=OUT, after=_cursor("s2", 3)) + with pytest.raises(StreamCursorError): + handle.read(topic=OUT, after=_cursor("", 3)) + + +def test_a_ref_opens_the_handle_it_names_on_the_provider() -> None: + provider = NativeStreams() + opened = provider.get_stream_handle( + _CLIENT, StreamRef.for_standalone("s1", topic="out") + ) + assert opened.ref() == StreamRef.for_standalone("s1", topic="out") + assert opened.ref(topic="b") == StreamRef.for_standalone("s1", topic="b") + workflow = provider.get_stream_handle( + _CLIENT, StreamRef.for_workflow("wf", run_id="r") + ) + assert workflow.ref() == StreamRef.for_workflow("wf", run_id="r") + with pytest.raises(ValueError, match="run_id"): + provider.get_stream_handle(_CLIENT, StreamRef.for_workflow("wf"), run_id="r") + + +async def test_a_create_refuses_a_policy_the_server_cannot_hold() -> None: + provider = NativeStreams() + for bad in ( + dict(max_records=0), + dict(max_bytes=-1), + dict(retention=timedelta(0)), + ): + with pytest.raises(ValueError): + await provider.create_standalone_stream(_CLIENT, "s1", **bad) # type: ignore[arg-type] + with pytest.raises(ValueError): + await provider.create_standalone_stream(_CLIENT, "") + with pytest.raises(ValueError): + provider.get_standalone_stream_handle(_CLIENT, "") diff --git a/tests/streams/test_native_standalone_e2e.py b/tests/streams/test_native_standalone_e2e.py new file mode 100644 index 000000000..2b9a43773 --- /dev/null +++ b/tests/streams/test_native_standalone_e2e.py @@ -0,0 +1,95 @@ +"""Standalone streams on the native provider against a live server. + +The conformance suite covers what every provider owes a standalone stream. +These pin what native adds on top: a read on an id nobody has created yet +parks on the server and delivers the first record once the stream exists, a +ref carries the stream to another client, and a create of an id that exists +is answered from the stream's own policy. Needs a server built from the +AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7433 uv run pytest tests/streams/test_native_standalone_e2e.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid + +import pytest + +from temporalio.client import Client +from temporalio.streams import StreamClosedError, StreamNotFoundError, StreamRef, topic +from temporalio.streams.providers.native import NativeStreams +from tests.streams.test_streams_conformance import take + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream service", +) + +OUT = topic("out", dict) + + +async def _connect(provider: NativeStreams) -> Client: + return await Client.connect(TARGET or "", plugins=[provider]) + + +async def test_a_read_on_a_stream_not_yet_created_waits_for_it() -> None: + provider = NativeStreams() + client = await _connect(provider) + stream_id = "late-" + uuid.uuid4().hex[:8] + try: + reader = client.get_stream_handle(stream_id=stream_id) + # Nothing to wait on for these: the stream does not exist. + with pytest.raises(StreamNotFoundError): + await reader.latest(topic=OUT) + with pytest.raises(StreamNotFoundError): + await reader.producer(topic=OUT, producer_id="early", attempt=1).append( + {"n": 0} + ) + + parked = asyncio.ensure_future(take(reader.read(topic=OUT), 1, timeout=30)) + await asyncio.sleep(0.5) + assert not parked.done(), "the read is parked, not failed" + + created = await client.create_stream(stream_id) + await created.producer(topic=OUT, producer_id="writer", attempt=1).append( + {"n": 1} + ) + assert [r.value for r in await parked] == [{"n": 1}] + finally: + await provider.close() + + +async def test_a_ref_carries_a_standalone_stream_to_another_client() -> None: + provider = NativeStreams() + client = await _connect(provider) + other_provider = NativeStreams() + other = await _connect(other_provider) + stream_id = "ref-" + uuid.uuid4().hex[:8] + try: + created = await client.create_stream(stream_id, max_records=10) + await created.producer(topic=OUT, producer_id="writer", attempt=1).append( + {"n": 1} + ) + ref = created.ref(topic=OUT) + assert ref == StreamRef.for_standalone(stream_id, topic="out") + + opened = other.get_stream_handle(ref) + records = await take(opened.read(), 1) + assert [r.value for r in records] == [{"n": 1}] + assert await opened.latest() == records[0].cursor + + await opened.close() + with pytest.raises(StreamClosedError): + await created.producer(topic=OUT, producer_id="writer", attempt=2).append( + {"n": 2} + ) + # The tail stays readable and the read ends on its own. + assert [r.value async for r in created.read(topic=OUT)] == [{"n": 1}] + finally: + await provider.close() + await other_provider.close() diff --git a/tests/streams/test_nexus_consumer.py b/tests/streams/test_nexus_consumer.py index 6726411f8..afe081993 100644 --- a/tests/streams/test_nexus_consumer.py +++ b/tests/streams/test_nexus_consumer.py @@ -997,6 +997,7 @@ async def test_a_worker_hosted_consumer_is_reached_through_the_frontend( # --------------------------------------------------------------------------- +@pytest.mark.needs_stream_channel_server @pytest.mark.needs_native_provider async def test_a_native_standalone_stream_is_consumed_through_its_channel( client: Client, diff --git a/tests/streams/test_stream_channel_e2e.py b/tests/streams/test_stream_channel_e2e.py new file mode 100644 index 000000000..6a632556a --- /dev/null +++ b/tests/streams/test_stream_channel_e2e.py @@ -0,0 +1,254 @@ +"""A native stream notifies the channel named by the stream, on a live server. + +Every append and the close of a native stream notify the channel +:func:`temporalio.client.stream_channel` derives from the stream's ref, so a +client follows a stream the way it follows an external one: by polling the +channel or registering a callback on it. The workflow's own consumption of a +stream is untouched by this and is covered elsewhere. Needs a server on which +streams drive channels, named with ``-E host:port``; skipped otherwise. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import uuid +from datetime import timedelta +from typing import Any + +import pytest + +from temporalio import activity, workflow +from temporalio.client import ( + Callback, + ChannelAddress, + ChannelKind, + Client, + stream_channel, +) +from temporalio.common import Execution, ExecutionType +from temporalio.service import RPCError +from temporalio.streams import StreamRef, topic +from temporalio.streams.providers.native import NativeStreams +from tests.helpers import assert_eventually, new_worker + +OUT = topic("out", dict) + +pytestmark = pytest.mark.needs_stream_channel_server + + +async def _native(client: Client, provider: NativeStreams) -> Client: + """A client on the same server as ``client`` with the native provider on it.""" + return await Client.connect( + client.service_client.config.target_host, + namespace=client.namespace, + plugins=[provider], + ) + + +def _closed(client: Client, notification: Any) -> bool: + if "closed" not in notification.metadata: + return False + return client.data_converter.payload_converter.from_payload( + notification.metadata["closed"], bool + ) + + +async def test_a_standalone_stream_notifies_the_channel_named_by_its_id( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + stream_id = f"chan-{uuid.uuid4().hex[:8]}" + try: + created = await native.create_stream(stream_id) + address = stream_channel(created.ref(topic=OUT)) + # A standalone stream's topics share one stream, so the topic is not in + # the name and the channel is an independent one. + assert address == ChannelAddress(f"stream/{stream_id}", None) + producer = created.producer(topic=OUT, producer_id="writer", attempt=1) + + async def changes(count: int) -> list[Any]: + polled = await native.poll_channel(address.channel, wait=False) + assert [n.counter for n in polled] == list(range(1, count + 1)) + return polled + + # A standalone stream reaches its channel through a task of its own + # that hands over the latest change, so the notifications trail the + # appends and a burst arrives as its newest change. Each append waits + # for its notification so that every change is seen. + await producer.append({"n": 1}) + await assert_eventually(lambda: changes(1)) + await producer.append({"n": 2}) + polled = await assert_eventually(lambda: changes(2)) + for notification in polled: + assert notification.channel == address.channel + assert notification.linked_to is None + # The position is the head after the change, as a native cursor + # names one: the stream id in place of a run, then the offset. + assert notification.position.decode().startswith(f"{stream_id}:") + assert not _closed(native, notification) + await created.close() + [third] = await native.poll_channel( + address.channel, after_counter=2, wait=timedelta(seconds=10) + ) + assert third.counter == 3 + assert _closed(native, third) + description = await native.describe_channel(address.channel) + assert description.kind == ChannelKind.INDEPENDENT + assert description.latest is not None and description.latest.counter == 3 + assert description.retained_count == 3 + # A callback registers on the derived name like on any channel. + callback = Callback(url="http://localhost:1/never-called", headers={}) + listener_id = await native.register_channel_listener(address.channel, callback) + description = await native.describe_channel(address.channel) + assert [listener.callback for listener in description.listeners] == [callback] + await native.unregister_channel_listener(address.channel, listener_id) + finally: + await provider.close() + + +@workflow.defn +class WriteOnNudge: + """Writes one record at the start and one more per nudge, until ``rounds``.""" + + def __init__(self) -> None: + self._nudges = 0 + self._written = 0 + + @workflow.run + async def run(self, rounds: int) -> int: + writer = workflow.stream_writer(OUT) + writer.publish({"n": 0}) + while self._written < rounds: + await workflow.wait_condition(lambda: self._nudges > self._written) + self._written += 1 + writer.publish({"n": self._written}) + return self._written + + @workflow.signal + def nudge(self) -> None: + self._nudges += 1 + + +async def test_a_workflow_stream_notifies_the_channel_linked_to_its_owner( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + worker = new_worker(native, WriteOnNudge) + running = asyncio.create_task(worker.run()) + # More rounds than nudges: the linked ring dies with the run, so the run + # stays open until the test has read it and is ended by hand. + handle = await native.start_workflow( + WriteOnNudge.run, 10, id=f"wf-{uuid.uuid4()}", task_queue=worker.task_queue + ) + try: + address = stream_channel(StreamRef.for_workflow(handle.id, topic=OUT)) + assert address == ChannelAddress("stream/out", Execution.workflow(handle.id)) + assert address.workflow_id == handle.id + assert address == stream_channel( + native.get_stream_handle(handle.id).ref(topic=OUT) + ) + + async def changes(count: int) -> list[Any]: + polled = await native.poll_channel( + address.channel, workflow_id=address.workflow_id, wait=False + ) + assert [n.counter for n in polled] == list(range(1, count + 1)) + return polled + + # The first task's publish is one append, so one notification. + await assert_eventually(lambda: changes(1), timeout=timedelta(seconds=30)) + await handle.signal(WriteOnNudge.nudge) + await assert_eventually(lambda: changes(2), timeout=timedelta(seconds=30)) + await handle.signal(WriteOnNudge.nudge) + polled = await assert_eventually( + lambda: changes(3), timeout=timedelta(seconds=30) + ) + run_id = handle.first_execution_run_id + for notification in polled: + assert notification.linked_to is not None + assert notification.linked_to.business_id == handle.id + assert notification.position.decode().startswith(f"{run_id}:") + description = await native.describe_channel( + address.channel, workflow_id=address.workflow_id + ) + assert description.kind == ChannelKind.LINKED + assert description.latest is not None and description.latest.counter == 3 + finally: + with contextlib.suppress(RPCError): + await handle.terminate() + try: + await asyncio.wait_for(worker.shutdown(), 15) + except asyncio.TimeoutError: + running.cancel() + await asyncio.gather(running, return_exceptions=True) + await provider.close() + + +_released: dict[str, bool] = {} +"""Activity ids the test has let go, read by the activity in the same process.""" + + +@activity.defn +async def write_then_wait() -> str: + """Writes one record to the activity's own stream and lingers until released.""" + activity_id = activity.info().activity_id + producer = activity.stream_handle().producer(topic=OUT) + await producer.append({"n": 1}) + while not _released.get(activity_id): + activity.heartbeat() + await asyncio.sleep(0.2) + return activity_id + + +@pytest.mark.needs_execution_server +async def test_a_standalone_activity_stream_notifies_the_channel_linked_to_it( + client: Client, +): + provider = NativeStreams() + native = await _native(client, provider) + activity_id = f"act-{uuid.uuid4().hex[:8]}" + async with new_worker(native, activities=[write_then_wait]) as worker: + handle = await native.start_activity( + write_then_wait, + id=activity_id, + task_queue=worker.task_queue, + start_to_close_timeout=timedelta(seconds=60), + ) + try: + # A standalone activity is an execution of its own, so its stream + # notifies a channel linked to it, under the topic's name alone. + address = stream_channel(StreamRef.for_activity(activity_id, topic=OUT)) + assert address == ChannelAddress( + "stream/out", Execution.activity(activity_id) + ) + assert address.workflow_id is None + + async def changes(count: int) -> list[Any]: + polled = await native.poll_channel( + address.channel, execution=address.execution, wait=False + ) + assert [n.counter for n in polled] == list(range(1, count + 1)) + return polled + + [first] = await assert_eventually( + lambda: changes(1), timeout=timedelta(seconds=30) + ) + assert first.linked_to is not None + assert (first.linked_to.type, first.linked_to.business_id) == ( + ExecutionType.ACTIVITY, + activity_id, + ) + description = await native.describe_channel( + address.channel, execution=address.execution + ) + assert description.kind == ChannelKind.LINKED + assert description.linked_to is not None + assert description.linked_to.business_id == activity_id + finally: + _released[activity_id] = True + with contextlib.suppress(Exception): + await asyncio.wait_for(handle.result(), 30) + await provider.close() diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index eceb5af2b..08cd19611 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -76,6 +76,7 @@ from temporalio.streams._ref import open_ref from temporalio.streams.providers import workflow_streams from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.native import NativeStreams from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider from tests.helpers import new_worker @@ -260,6 +261,22 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: provider.reset() +@workflow.defn +class StreamHost: + """Owns a stream and lingers, so outside code has a running workflow to address.""" + + def __init__(self) -> None: + self._released = False + + @workflow.signal + def release(self) -> None: + self._released = True + + @workflow.run + async def run(self) -> None: + await workflow.wait_condition(lambda: self._released) + + @workflow.defn class TruncatingStreamHost: """A stream host whose log an update can truncate, the way a workflow's retention would.""" @@ -293,6 +310,46 @@ async def run(self) -> None: reader.close() +async def _native_case(client: Client) -> AsyncIterator[ProviderCase]: + # The store is a server built from the stream-carrying branch, which the + # test environment's own server is not; TEMPORAL_ADDRESS names it. + address = os.environ.get("TEMPORAL_ADDRESS") + if address: + client = await Client.connect( + address, namespace=os.environ.get("TEMPORAL_NAMESPACE", "default") + ) + provider = NativeStreams() + # Registered once, on the client: the host's worker inherits it and the + # cases open handles through client.get_stream_handle. + config = client.config() + config["plugins"] = [provider] + client = Client(**config) + hosts: dict[str, WorkflowHandle[Any, Any]] = {} + async with new_worker(client, StreamHost, DefaultTopicAnswer) as worker: + + async def host(workflow_id: str) -> None: + if workflow_id not in hosts: + hosts[workflow_id] = await client.start_workflow( + StreamHost.run, id=workflow_id, task_queue=worker.task_queue + ) + + # No truncate hook: a stream a workflow owns has no truncation call. + # BEGINNING on a truncated stream is covered on the stream client, + # whose standalone streams can be truncated. + yield ProviderCase( + "native", + provider, + client, + host=host, + task_queue=worker.task_queue, + waits_for_standalone_creation=True, + refuses_appends_past_byte_cap=True, + ) + for handle in hosts.values(): + await handle.terminate() + await provider.close() + + async def _workflow_streams_case(client: Client) -> AsyncIterator[ProviderCase]: # No STREAMS_LIVE gate: the store is the workflow's own History, which the # test environment's server provides. @@ -347,6 +404,8 @@ async def truncate(workflow_id: str, topic: str, keep: int) -> None: "memory": _memory_case, "workflow_streams": _workflow_streams_case, } +if os.environ.get("STREAMS_LIVE") == "native": + SETUPS["native"] = _native_case _CAPABILITIES = { "reports_positions": lambda case: case.reports_positions, diff --git a/tests/test_client_stream.py b/tests/test_client_stream.py index 89abb0559..779d1b188 100644 --- a/tests/test_client_stream.py +++ b/tests/test_client_stream.py @@ -26,7 +26,11 @@ ) from temporalio.client_stream import StreamClient, StreamHandle from temporalio.service import RPCError -from temporalio.streams import StreamNotFoundError, StreamProducerError +from temporalio.streams import ( + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") @@ -230,7 +234,7 @@ async def test_earliest_reads_from_the_floor_of_a_truncated_stream( # Offset zero is what a reader with no position used to send, and a # truncated stream no longer holds it. - with pytest.raises(RPCError, match="truncated"): + with pytest.raises(StreamCursorError, match="truncated"): await stream.read(from_offset=0) entries, next_offset = await stream.read(start=StreamStartPosition(earliest=True)) assert data(entries) == [b"c", b"d"] diff --git a/tests/test_client_stream_connection.py b/tests/test_client_stream_connection.py new file mode 100644 index 000000000..73e4c5314 --- /dev/null +++ b/tests/test_client_stream_connection.py @@ -0,0 +1,216 @@ +"""The stream channel is the client's connection, opened again with ``grpcio``. + +sdk-core does not know the stream service, so its channel cannot be shared. +What can be shared is the configuration: these pin that a TLS-configured +client yields a secure channel with the same material, that an API key rides +as the bearer header along with the client's other headers, and that the +client's ``retry_config`` is what the shared stream client retries under. +""" + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import grpc +import grpc.aio +import pytest + +import temporalio.api.streamservice.v1 as stream +from temporalio import client_stream +from temporalio.client_stream import Connection, StreamClient +from temporalio.service import ( + ConnectConfig, + HttpConnectProxyConfig, + KeepAliveConfig, + RetryConfig, + TLSConfig, + __version__, +) + +# The name the vendored stubs call the service by, server-internal as it is. +_SERVICE = "temporal.server.chasm.lib.stream.proto.v1.StreamService" + + +def _fake_client(config: ConnectConfig, namespace: str = "ns") -> Any: + return SimpleNamespace( + service_client=SimpleNamespace(config=config), namespace=namespace + ) + + +def test_a_plaintext_client_yields_an_insecure_channel( + monkeypatch: pytest.MonkeyPatch, +) -> None: + opened: dict[str, Any] = {} + + def insecure_channel(target: str, **kwargs: Any) -> str: + opened.update(target=target, **kwargs) + return "channel" + + monkeypatch.setattr(grpc.aio, "insecure_channel", insecure_channel) + connection = Connection.from_config(ConnectConfig(target_host="localhost:7233")) + assert not connection.secure + assert connection.channel() == "channel" + assert opened["target"] == "localhost:7233" + keep_alive = KeepAliveConfig.default + assert ("grpc.keepalive_time_ms", keep_alive.interval_millis) in opened["options"] + assert ("grpc.keepalive_timeout_ms", keep_alive.timeout_millis) in opened["options"] + + +def test_a_tls_client_yields_a_secure_channel_with_its_credentials( + monkeypatch: pytest.MonkeyPatch, +) -> None: + made: dict[str, Any] = {} + opened: dict[str, Any] = {} + + def ssl_channel_credentials(**kwargs: Any) -> str: + made.update(kwargs) + return "credentials" + + def secure_channel(target: str, credentials: Any, **kwargs: Any) -> str: + opened.update(target=target, credentials=credentials, **kwargs) + return "channel" + + monkeypatch.setattr(grpc, "ssl_channel_credentials", ssl_channel_credentials) + monkeypatch.setattr(grpc.aio, "secure_channel", secure_channel) + + connection = Connection.from_config( + ConnectConfig( + target_host="cloud.example:7233", + tls=TLSConfig( + server_root_ca_cert=b"root", + client_cert=b"cert", + client_private_key=b"key", + domain="cloud.example", + verification_server_name="pinned.test", + ), + ) + ) + assert connection.secure + assert connection.channel() == "channel" + assert made == { + "root_certificates": b"root", + "private_key": b"key", + "certificate_chain": b"cert", + } + assert opened["target"] == "cloud.example:7233" + assert opened["credentials"] == "credentials" + assert ("grpc.ssl_target_name_override", "pinned.test") in opened["options"] + assert ("grpc.default_authority", "cloud.example") in opened["options"] + + +def test_tls_is_on_by_default_with_an_api_key_and_off_when_refused() -> None: + with_key = Connection.from_config( + ConnectConfig(target_host="cloud.example:7233", api_key="secret") + ) + assert with_key.secure + assert with_key.server_root_ca_cert is None, "system roots" + refused = Connection.from_config( + ConnectConfig(target_host="localhost:7233", api_key="secret", tls=False) + ) + assert not refused.secure + assert ("authorization", "Bearer secret") in refused.headers + + +def test_a_scheme_in_the_target_decides_and_is_dropped() -> None: + connection = Connection.from_config( + ConnectConfig(target_host="https://cloud.example:7233") + ) + assert connection.secure + assert connection.target == "cloud.example:7233" + + +def test_a_proxy_and_the_clients_own_authorization_carry_over() -> None: + connection = Connection.from_config( + ConnectConfig( + target_host="localhost:7233", + api_key="ignored", + tls=False, + rpc_metadata={"Authorization": "Custom token", "x-tenant": "t1"}, + http_connect_proxy_config=HttpConnectProxyConfig( + target_host="proxy:3128", basic_auth=("user", "pass") + ), + ) + ) + # Core leaves a caller's own authorization header alone. + assert ("Authorization", "Custom token") in connection.headers + assert ("x-tenant", "t1") in connection.headers + assert not any(key == "authorization" for key, _ in connection.headers) + assert connection.http_proxy == "http://user:pass@proxy:3128" + + +class _Recorder: + """Answers describe and keeps the headers it arrived with.""" + + def __init__(self) -> None: + self.metadata: dict[str, str | bytes] = {} + + async def describe( + self, _request: Any, context: Any + ) -> stream.DescribeStreamResponse: + self.metadata = dict(context.invocation_metadata() or ()) + return stream.DescribeStreamResponse( + frontend_response=stream.DescribeStreamOutput( + state=stream.StreamState(head_offset=3) + ) + ) + + def register(self, server: grpc.aio.Server) -> None: + # The generated registration wants the whole servicer; one method is + # enough to see the headers. + handler: Any = grpc.unary_unary_rpc_method_handler( + self.describe, + request_deserializer=stream.DescribeStreamRequest.FromString, + response_serializer=stream.DescribeStreamResponse.SerializeToString, + ) + server.add_generic_rpc_handlers( + ( + grpc.method_handlers_generic_handler( + _SERVICE, {"DescribeStream": handler} + ), + ) + ) + + +async def test_an_api_key_client_sends_the_header_on_every_call() -> None: + recorder = _Recorder() + server = grpc.aio.server() + recorder.register(server) + port = server.add_insecure_port("127.0.0.1:0") + await server.start() + config = ConnectConfig( + target_host=f"127.0.0.1:{port}", + api_key="secret", + tls=False, + rpc_metadata={"x-tenant": "t1"}, + ) + streams = StreamClient.for_connection(Connection.from_config(config), "ns") + try: + state = await streams.get("s").describe() + assert state.head_offset == 3 + assert recorder.metadata["authorization"] == "Bearer secret" + assert recorder.metadata["x-tenant"] == "t1" + assert recorder.metadata["client-name"] == "temporal-python" + assert recorder.metadata["client-version"] == __version__ + finally: + await streams.close() + await server.stop(None) + + +async def test_the_shared_client_retries_under_the_clients_config() -> None: + retry = RetryConfig(max_retries=3) + client = _fake_client( + ConnectConfig(target_host="localhost:7233", retry_config=retry), "ns" + ) + try: + shared = client_stream.shared_client(client) + assert shared._retry_config is retry + assert shared is client_stream.shared_client(client), "one channel per key" + # A client with other credentials to the same host is not the same channel. + other = _fake_client( + ConnectConfig(target_host="localhost:7233", api_key="k", tls=False), "ns" + ) + assert client_stream.shared_client(other) is not shared + assert client_stream.shared_key(other) != client_stream.shared_key(client) + finally: + await client_stream.close_shared_clients() diff --git a/tests/test_client_stream_errors.py b/tests/test_client_stream_errors.py new file mode 100644 index 000000000..7b379f4c7 --- /dev/null +++ b/tests/test_client_stream_errors.py @@ -0,0 +1,114 @@ +"""What a refused stream call raises. + +The service carries no typed detail for its refusals, so a typed one arrives +as a ``FAILED_PRECONDITION`` with a reason token at the front of the message. +These pin the mapping: each token to its error, the phrases an older server +sends to the same errors, and an unrelated message on the same code to a +plain ``RPCError`` that keeps the code. +""" + +from __future__ import annotations + +import grpc +import pytest + +from temporalio.client_stream import translate_error +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import ( + StreamClosedError, + StreamCursorError, + StreamNotFoundError, + StreamProducerError, +) + + +@pytest.mark.parametrize( + "details", + [ + "STREAM_PRODUCER_CONFLICT: producer sequence 3 already used with different content", + "STREAM_PRODUCER_STALE_SEQUENCE: stale producer sequence 2, last accepted for " + 'producer "p" is 3', + ], +) +def test_a_producer_refusal_is_a_producer_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamProducerError) + assert str(error) == details + + +def test_an_older_servers_producer_refusal_is_a_producer_error() -> None: + # Before the tokens the refusal was an INVALID_ARGUMENT with the phrase. + details = 'stale producer sequence 2, last accepted for producer "p" is 3' + error = translate_error(grpc.StatusCode.INVALID_ARGUMENT, details) + assert isinstance(error, StreamProducerError) + + +@pytest.mark.parametrize( + "details", + [ + "STREAM_CURSOR_BELOW_FLOOR: offset 2 is below the stream's floor of 5", + "offset 2 is below the stream's floor of 5", + ], +) +def test_a_read_below_the_floor_is_a_cursor_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamCursorError) + assert str(error) == details + + +@pytest.mark.parametrize( + "details", + ["STREAM_CLOSED: stream is closed", "stream is closed"], +) +def test_an_append_on_a_sealed_stream_is_a_closed_error(details: str) -> None: + error = translate_error(grpc.StatusCode.FAILED_PRECONDITION, details) + assert isinstance(error, StreamClosedError) + assert str(error) == details + + +def test_a_create_with_another_policy_is_a_value_error() -> None: + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, + "STREAM_POLICY_MISMATCH: stream exists keeping 10 records, asked for 5", + ) + assert type(error) is ValueError + + +def test_not_found_is_a_not_found_error() -> None: + error = translate_error(grpc.StatusCode.NOT_FOUND, "no stream with id 's'") + assert isinstance(error, StreamNotFoundError) + + +def test_an_unrelated_message_keeps_its_code() -> None: + # The same codes with another message are the server's ordinary refusals. + for code, expected in [ + (grpc.StatusCode.INVALID_ARGUMENT, RPCStatusCode.INVALID_ARGUMENT), + (grpc.StatusCode.FAILED_PRECONDITION, RPCStatusCode.FAILED_PRECONDITION), + (grpc.StatusCode.UNAVAILABLE, RPCStatusCode.UNAVAILABLE), + ]: + error = translate_error(code, "no records to append", b"raw") + assert type(error) is RPCError + assert error.status == expected + assert error.raw_grpc_status == b"raw" + assert str(error) == "no records to append" + + +def test_a_token_needs_its_code_and_its_separator() -> None: + # A token on another code is prose, and so is one with no ": " after it. + error = translate_error( + grpc.StatusCode.INVALID_ARGUMENT, "STREAM_PRODUCER_CONFLICT: elsewhere" + ) + assert type(error) is RPCError + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, "STREAM_CURSOR_BELOW_FLOOR" + ) + assert type(error) is RPCError + error = translate_error( + grpc.StatusCode.FAILED_PRECONDITION, "STREAM_SOMETHING_ELSE: unknown token" + ) + assert type(error) is RPCError + + +def test_an_empty_message_reads_as_the_code() -> None: + error = translate_error(grpc.StatusCode.UNAVAILABLE, "") + assert str(error) == "UNAVAILABLE" diff --git a/tests/test_client_stream_owner.py b/tests/test_client_stream_owner.py index a87f9d782..a19e8dbc8 100644 --- a/tests/test_client_stream_owner.py +++ b/tests/test_client_stream_owner.py @@ -39,6 +39,7 @@ def _client(recorder: _Recorder) -> StreamClient: client = StreamClient.__new__(StreamClient) client._stub = recorder client._namespace = "ns" + client._retry_config = None return client diff --git a/tests/test_client_stream_retry.py b/tests/test_client_stream_retry.py new file mode 100644 index 000000000..8baa2939b --- /dev/null +++ b/tests/test_client_stream_retry.py @@ -0,0 +1,257 @@ +"""What the stream client retries, and what it raises at once. + +The stream service is reached over a channel of the client's own, outside +sdk-core, so the retry Core gives every other call is reproduced here. These +pin its edges with a scripted stub: a throttled read or numbered append goes +again, an append the server could not tell from its repeat does not, a code +Core would not retry is raised at once, and cancelling the caller lands during +the wait between attempts. +""" + +from __future__ import annotations + +import asyncio +from collections import defaultdict +from types import SimpleNamespace +from typing import Any + +import grpc +import grpc.aio +import pytest + +import temporalio.api.streamservice.v1 as stream +import temporalio.converter +from temporalio import client_stream +from temporalio.api.common.v1 import GrpcStatus +from temporalio.api.enums.v1 import ResourceExhaustedCause +from temporalio.api.errordetails.v1 import ResourceExhaustedFailure +from temporalio.api.stream.v1 import StreamRecord +from temporalio.client_stream import WorkflowStreamHandle, _to_service +from temporalio.service import RetryConfig, RPCError, RPCStatusCode +from temporalio.streams import StreamNotFoundError +from temporalio.streams._record import RecordKind +from temporalio.streams._wire import to_wire +from temporalio.streams.providers.native import NativeStreamHandle + +FAST = RetryConfig( + initial_interval_millis=1, + max_interval_millis=2, + max_elapsed_time_millis=5000, + max_retries=5, +) + + +@pytest.fixture(autouse=True) +def fast_throttle(monkeypatch: pytest.MonkeyPatch) -> None: + # The floor under a throttled wait is a second, which is right for a + # caller and wrong for a test. + monkeypatch.setattr( + client_stream, + "_THROTTLE", + RetryConfig( + initial_interval_millis=1, + max_interval_millis=2, + max_elapsed_time_millis=None, + max_retries=0, + ), + ) + + +class _Stub: + """Answers each method from a script of exceptions and responses, in order.""" + + def __init__(self, **scripts: list[Any]) -> None: + self.calls: dict[str, int] = defaultdict(int) + self._scripts = scripts + + def __getattr__(self, method: str) -> Any: + script = self._scripts[method] + + async def call(_request: Any, **_: Any) -> Any: + self.calls[method] += 1 + outcome = script.pop(0) + if isinstance(outcome, BaseException): + raise outcome + return outcome + + return call + + +def _error( + code: grpc.StatusCode, + details: str = "", + *, + cause: ResourceExhaustedCause.ValueType | None = None, +) -> grpc.aio.AioRpcError: + trailing = grpc.aio.Metadata() + if cause is not None: + status = GrpcStatus(code=code.value[0], message=details) + status.details.add().Pack(ResourceExhaustedFailure(cause=cause)) + trailing.add("grpc-status-details-bin", status.SerializeToString()) + return grpc.aio.AioRpcError(code, grpc.aio.Metadata(), trailing, details=details) + + +def _throttled() -> grpc.aio.AioRpcError: + return _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "service rate limit exceeded", + cause=ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_RPS_LIMIT, + ) + + +def _page(*records: stream.StreamRecord) -> stream.PollWorkflowMessagesResponse: + return stream.PollWorkflowMessagesResponse( + frontend_response=stream.PollMessagesOutput( + records=list(records), + next_offset=len(records), + head_offset=len(records), + closed=True, + run_id="run", + ) + ) + + +def _appended() -> stream.AddWorkflowMessagesResponse: + return stream.AddWorkflowMessagesResponse( + frontend_response=stream.AddMessagesOutput( + first_offset=0, next_offset=1, count=1 + ) + ) + + +def _handle(stub: _Stub, retry: RetryConfig = FAST) -> WorkflowStreamHandle: + return WorkflowStreamHandle(stub, "ns", "wf", "topic", "run", retry_config=retry) + + +async def test_a_throttled_poll_is_read_again() -> None: + stub = _Stub(PollWorkflowMessages=[_throttled(), _page()]) + page = await _handle(stub).poll() + assert page.closed + assert stub.calls["PollWorkflowMessages"] == 2 + + +async def test_a_throttled_numbered_append_goes_again() -> None: + stub = _Stub(AddWorkflowMessages=[_throttled(), _appended()]) + appended = await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert appended.next_offset == 1 + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_a_numbered_append_survives_an_ambiguous_failure() -> None: + # The server holds the producer's sequence, so whichever attempt landed, + # the repeat comes back with the original offsets. + stub = _Stub(AddWorkflowMessages=[_error(grpc.StatusCode.UNAVAILABLE), _appended()]) + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_an_unnumbered_append_is_not_made_again_after_an_ambiguous_failure() -> ( + None +): + stub = _Stub(AddWorkflowMessages=[_error(grpc.StatusCode.UNAVAILABLE), _appended()]) + with pytest.raises(RPCError) as raised: + await _handle(stub).append(StreamRecord()) + assert raised.value.status == RPCStatusCode.UNAVAILABLE + assert stub.calls["AddWorkflowMessages"] == 1 + + +async def test_an_unnumbered_append_goes_again_after_a_refusal() -> None: + # Throttling is answered before the handler runs, so nothing landed. + stub = _Stub(AddWorkflowMessages=[_throttled(), _appended()]) + await _handle(stub).append(StreamRecord()) + assert stub.calls["AddWorkflowMessages"] == 2 + + +async def test_a_code_core_would_not_retry_is_raised_at_once() -> None: + stub = _Stub( + PollWorkflowMessages=[_error(grpc.StatusCode.INVALID_ARGUMENT, "bad"), _page()] + ) + with pytest.raises(RPCError) as raised: + await _handle(stub).poll() + assert raised.value.status == RPCStatusCode.INVALID_ARGUMENT + assert stub.calls["PollWorkflowMessages"] == 1 + + stub = _Stub(PollWorkflowMessages=[_error(grpc.StatusCode.NOT_FOUND), _page()]) + with pytest.raises(StreamNotFoundError): + await _handle(stub).poll() + assert stub.calls["PollWorkflowMessages"] == 1 + + +async def test_the_budget_is_bounded() -> None: + stub = _Stub(PollWorkflowMessages=[_throttled() for _ in range(10)]) + with pytest.raises(RPCError) as raised: + await _handle( + stub, RetryConfig(initial_interval_millis=1, max_retries=3) + ).poll() + assert raised.value.status == RPCStatusCode.RESOURCE_EXHAUSTED + assert stub.calls["PollWorkflowMessages"] == 3 + + +async def test_a_full_stream_is_not_waited_out() -> None: + full = _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "stream holds 10 of its budget of 10 records", + cause=ResourceExhaustedCause.RESOURCE_EXHAUSTED_CAUSE_PERSISTENCE_STORAGE_LIMIT, + ) + stub = _Stub(AddWorkflowMessages=[full, _appended()]) + with pytest.raises(RPCError) as raised: + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert raised.value.status == RPCStatusCode.RESOURCE_EXHAUSTED + assert stub.calls["AddWorkflowMessages"] == 1 + + too_large = _error( + grpc.StatusCode.RESOURCE_EXHAUSTED, + "grpc: received message larger than max (5 vs. 4)", + ) + stub = _Stub(AddWorkflowMessages=[too_large, _appended()]) + with pytest.raises(RPCError): + await _handle(stub).append(StreamRecord(), producer_id="p", sequence=1) + assert stub.calls["AddWorkflowMessages"] == 1 + + +async def test_cancelling_the_caller_lands_during_the_wait( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + client_stream, + "_THROTTLE", + RetryConfig(initial_interval_millis=60_000, max_elapsed_time_millis=None), + ) + stub = _Stub(PollWorkflowMessages=[_throttled(), _page()]) + task = asyncio.ensure_future(_handle(stub).poll()) + for _ in range(10): + await asyncio.sleep(0) + assert stub.calls["PollWorkflowMessages"] == 1, "parked in the wait" + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + assert stub.calls["PollWorkflowMessages"] == 1 + + +async def test_a_native_read_survives_a_throttled_poll( + monkeypatch: pytest.MonkeyPatch, +) -> None: + converter = temporalio.converter.default() + record = _to_service( + to_wire( + converter.payload_converter, + topic="topic", + kind=RecordKind.DATA, + value="hello", + producer_id="p", + sequence=1, + ) + ) + stub = _Stub(PollWorkflowMessages=[_throttled(), _page(record)]) + client: Any = SimpleNamespace(data_converter=converter) + handle = NativeStreamHandle(client, "wf", "run") + + def open_stream(_topic: str, _run_id: str) -> WorkflowStreamHandle: + return _handle(stub) + + monkeypatch.setattr(handle, "_stream", open_stream) + + values = [item.value async for item in handle.read(topic="topic")] + + assert values == ["hello"] + assert stub.calls["PollWorkflowMessages"] == 2 diff --git a/tests/test_client_stream_sharing.py b/tests/test_client_stream_sharing.py new file mode 100644 index 000000000..55fd3e7e3 --- /dev/null +++ b/tests/test_client_stream_sharing.py @@ -0,0 +1,57 @@ +"""Who owns a shared channel, and who may close it. + +One channel is shared per loop, connection and namespace, so two providers in +one process can be using the same one. Closing a provider has to leave the +other's channels alone. +""" + +from __future__ import annotations + +import asyncio + +from temporalio import client_stream +from temporalio.client_stream import Connection, SharedKey +from temporalio.streams.providers.native import NativeStreams + + +class _FakeClient: + """Stands in for a StreamClient, which would want a real channel.""" + + def __init__(self) -> None: + self.closed = False + + async def close(self) -> None: + self.closed = True + + +def _key(target: str, namespace: str) -> SharedKey: + return Connection(target=target, secure=False), namespace + + +def _put(key: SharedKey) -> _FakeClient: + fake = _FakeClient() + per_loop = client_stream._shared.setdefault(asyncio.get_running_loop(), {}) + per_loop[key] = fake # type: ignore[assignment] + return fake + + +async def test_closing_named_clients_leaves_the_others_open() -> None: + mine = _put(_key("host-a:7233", "ns")) + theirs = _put(_key("host-b:7233", "ns")) + await client_stream.close_shared_clients(_key("host-a:7233", "ns")) + assert mine.closed + assert not theirs.closed, "another provider is still reading through it" + await client_stream.close_shared_clients() + assert theirs.closed + + +async def test_a_provider_closes_only_what_its_own_handles_opened() -> None: + mine = _put(_key("host-a:7233", "ns")) + theirs = _put(_key("host-b:7233", "ns")) + + provider = NativeStreams() + provider._opened.add(_key("host-a:7233", "ns")) + await provider.close() + assert mine.closed + assert not theirs.closed + await client_stream.close_shared_clients() diff --git a/tests/worker/test_workflow_stream.py b/tests/worker/test_workflow_stream.py new file mode 100644 index 000000000..b4dd6107b --- /dev/null +++ b/tests/worker/test_workflow_stream.py @@ -0,0 +1,353 @@ +"""In-workflow consumption and publication of a server-side stream. + +Ranges arrive on Workflow Tasks and only the offsets they covered are written to +History, so replay is served by the server reading the stream again. These tests +drive the buffering, the workflow-facing read and the per-task publish directly, +which is the part this SDK owns; the delivery decision itself lives in the +server and sdk-core. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +import temporalio.api.common.v1 +import temporalio.api.stream.v1 as api_stream +from temporalio.bridge.proto.workflow_commands import WorkflowCommand +from temporalio.worker._workflow_instance import ( + _MAX_STREAM_BATCH_BYTES, + _MAX_STREAM_RECORD_BYTES, + _MAX_STREAM_RECORDS_PER_BATCH, + _StreamBuffer, + _WorkflowInstanceImpl, +) +from temporalio.workflow import ReadOnlyContextError + + +def record(body: bytes, topic: str = "") -> api_stream.StreamRecord: + return api_stream.StreamRecord( + body=temporalio.api.common.v1.Payload(data=body), topic=topic + ) + + +def bodies(delivered: list[Any]) -> list[bytes]: + return [item.record.body.data for item in delivered] + + +async def test_buffer_hands_over_in_order() -> None: + buffer = _StreamBuffer() + buffer.extend([record(b"one"), record(b"two")]) + + assert bodies(buffer.take()) == [b"one", b"two"] + assert len(buffer) == 0 + + +# A reader that arrives before the data must not miss it, and one that arrives +# after must not block: the range is delivered once and never resent. +async def test_buffer_wakes_a_waiting_reader() -> None: + buffer = _StreamBuffer() + waiter = buffer.wait_future() + assert not waiter.done() + + buffer.extend([record(b"late")]) + await asyncio.wait_for(waiter, timeout=1) + assert bodies(buffer.take()) == [b"late"] + + +# An empty range is still a delivery the server recorded, but there is nothing +# to hand a reader, so it must not wake one into returning nothing. +async def test_empty_range_does_not_wake_a_reader() -> None: + buffer = _StreamBuffer() + waiter = buffer.wait_future() + + buffer.extend([]) + + assert not waiter.done() + assert len(buffer) == 0 + + +async def test_buffer_keeps_data_delivered_before_anyone_reads() -> None: + buffer = _StreamBuffer() + buffer.extend([record(b"early")]) + + # No waiting: the data is already there. + assert len(buffer) == 1 + waiter = buffer.wait_future() + assert not waiter.done(), "a fresh waiter is only resolved by new data" + assert bodies(buffer.take()) == [b"early"] + + +class _ReadOnlyStub: + """Enough of the workflow instance to drive the real read. + + The cap lives inside ``workflow_read_stream_records``, so a test that + reimplements it proves nothing about the code that ships. Borrowing the + method off the real class is what keeps the test on the shipped path. + """ + + workflow_read_stream_records = _WorkflowInstanceImpl.workflow_read_stream_records + + def __init__(self, read_only: bool = False) -> None: + self._stream_buffers: dict[str, _StreamBuffer] = {} + self._read_only = read_only + + def _assert_not_read_only(self, action: str) -> None: + if self._read_only: + raise ReadOnlyContextError(f"cannot {action} in a read-only context") + + async def read(self, stream: str, max_records: int) -> list[bytes]: + instance: Any = self + return bodies( + await _WorkflowInstanceImpl.workflow_read_stream_records( + instance, stream, max_records + ) + ) + + +@pytest.mark.parametrize("max_records", [1, 2, 5]) +async def test_read_respects_a_cap_without_losing_the_tail(max_records: int) -> None: + stub = _ReadOnlyStub() + expected = [b"a", b"b", b"c"] + stub._stream_buffers["s"] = _StreamBuffer() + stub._stream_buffers["s"].extend([record(b) for b in expected]) + + got = await stub.read("s", max_records) + + assert got == expected[:max_records] + # Whatever the cap left behind has to still be there: nothing resends it. + assert len(stub._stream_buffers["s"]) == max(0, len(expected) - max_records) + + if len(stub._stream_buffers["s"]): + rest = await stub.read("s", 0) + assert got + rest == expected + else: + assert got == expected + + +# A query activation carries no ranges, so a read there would wait on a future +# nothing can resolve and the query would time out saying nothing. +async def test_read_is_refused_in_a_read_only_context() -> None: + stub = _ReadOnlyStub(read_only=True) + stub._stream_buffers["s"] = _StreamBuffer() + stub._stream_buffers["s"].extend([record(b"a")]) + + with pytest.raises(ReadOnlyContextError, match="read stream"): + await stub.read("s", 0) + # Refused before the buffer was touched, so the range is still there for + # the task that is allowed to read it. + assert len(stub._stream_buffers["s"]) == 1 + + +# A range is recorded as consumed once and never resent, so the buffer is the +# only place a repeated, skipped or mis-sized delivery can still be noticed. +async def test_ranges_have_to_abut_the_last_one() -> None: + buffer = _StreamBuffer("s") + buffer.extend([record(b"a"), record(b"b")], 0, 2) + # An empty range moves the expectation too: the server recorded it. + buffer.extend([], 2, 2) + buffer.extend([record(b"c")], 2, 3) + assert [item.offset for item in buffer.take()] == [0, 1, 2] + + with pytest.raises(RuntimeError, match=r"\[2, 3\).*ended at 3"): + buffer.extend([record(b"c")], 2, 3) + with pytest.raises(RuntimeError, match=r"\[5, 6\).*ended at 3"): + buffer.extend([record(b"f")], 5, 6) + + +async def test_a_range_has_to_carry_as_many_records_as_it_spans() -> None: + buffer = _StreamBuffer("s") + with pytest.raises(RuntimeError, match=r"2 records for offsets \[0, 1\)"): + buffer.extend([record(b"a"), record(b"b")], 0, 1) + assert len(buffer) == 0 + + +class _Completion: + def __init__(self) -> None: + self.commands: list[WorkflowCommand] = [] + + +class _Successful: + def __init__(self) -> None: + self.successful = _Completion() + + +class _CommandStub: + """Drives the real publish path and keeps the commands it issued. + + The completion's command list is a plain list here, which supports the + same ``insert`` the protobuf container does. + """ + + workflow_append_stream_records = ( + _WorkflowInstanceImpl.workflow_append_stream_records + ) + _flush_stream_appends = _WorkflowInstanceImpl._flush_stream_appends + + def __init__(self) -> None: + self._stream_appends: dict[str, list[api_stream.StreamRecord]] = {} + self._current_completion = _Successful() + + def _assert_not_read_only(self, _action: str) -> None: + pass + + @property + def commands(self) -> list[WorkflowCommand]: + return self._current_completion.successful.commands + + def publish(self, *records: api_stream.StreamRecord, stream_name: str = "") -> None: + instance: Any = self + _WorkflowInstanceImpl.workflow_append_stream_records( + instance, stream_name, list(records) + ) + + def flush(self) -> None: + instance: Any = self + _WorkflowInstanceImpl._flush_stream_appends(instance) + + +# A task's publishes on one stream become one command, whatever their number, +# because the event the command produces is what bounds a workflow's History. +def test_a_tasks_publishes_on_one_stream_become_one_command() -> None: + stub = _CommandStub() + stub.publish(record(b"a"), record(b"b")) + stub.publish(record(b"c")) + stub.publish(record(b"d"), stream_name="other") + assert stub.commands == [] + + stub.flush() + + by_stream = { + command.append_stream_records.stream_name: [ + r.body.data for r in command.append_stream_records.records + ] + for command in stub.commands + } + assert by_stream == {"": [b"a", b"b", b"c"], "other": [b"d"]} + # Flushed once: a second flush has nothing left to say. + stub.flush() + assert len(stub.commands) == 2 + + +# The workflow is the producer of what it publishes, whatever the caller set. +def test_the_workflows_records_carry_no_producer() -> None: + stub = _CommandStub() + stub.publish(api_stream.StreamRecord(producer_id="someone", attempt=3)) + stub.flush() + assert stub.commands[0].append_stream_records.records[0].producer_id == "" + + +# The server accepts nothing after a command that ends the run, so the +# publishes have to go ahead of it. +def test_publishes_are_flushed_ahead_of_the_completion_command() -> None: + stub = _CommandStub() + stub.publish(record(b"a")) + done = WorkflowCommand() + done.complete_workflow_execution.SetInParent() + stub.commands.append(done) + + stub.flush() + + assert [c.WhichOneof("variant") for c in stub.commands] == [ + "append_stream_records", + "complete_workflow_execution", + ] + + +# The server refuses an oversized record, and a refused command is reissued on +# every replay, so the limit has to be applied before the record is buffered. +def test_a_record_over_the_server_limit_is_refused_before_the_command() -> None: + stub = _CommandStub() + stub.publish(record(b"x" * (_MAX_STREAM_RECORD_BYTES - 16))) + with pytest.raises(ValueError, match=f"{_MAX_STREAM_RECORD_BYTES} bytes"): + stub.publish(record(b"x" * (_MAX_STREAM_RECORD_BYTES + 1))) + stub.flush() + assert len(stub.commands) == 1 + + +# A task that publishes more than one batch holds is split into commands the +# server accepts, by count and by bytes, rather than refused. +def test_a_task_over_the_batch_limits_is_split_into_commands() -> None: + stub = _CommandStub() + stub.publish(*(record(b"x") for _ in range(_MAX_STREAM_RECORDS_PER_BATCH + 1))) + stub.flush() + assert [len(c.append_stream_records.records) for c in stub.commands] == [ + _MAX_STREAM_RECORDS_PER_BATCH, + 1, + ] + + stub = _CommandStub() + # Two records that fit one batch together, then one whose framing alone + # overflows the room they leave. + big = record(b"x" * (_MAX_STREAM_RECORD_BYTES - 16)) + room = _MAX_STREAM_BATCH_BYTES - 2 * big.ByteSize() + stub.publish(big, big, record(b"y" * room)) + stub.flush() + assert [len(c.append_stream_records.records) for c in stub.commands] == [2, 1] + + +async def test_a_closed_buffer_keeps_nothing_and_still_checks_continuity() -> None: + # There is no unsubscribe command, so the server keeps delivering for the + # life of the run. A reader that closed would otherwise grow the instance + # for the rest of it. + buffer = _StreamBuffer("s") + buffer.extend([record(b"one")], 0, 1) + assert len(buffer) == 1 + + buffer.close() + assert buffer.closed + assert len(buffer) == 0, "what it held is let go of, not kept for nobody" + + buffer.extend([record(b"two"), record(b"three")], 1, 3) + assert len(buffer) == 0 + # Continuity is still tracked across what it dropped, so a range that + # repeats or skips is caught rather than passing unnoticed. + with pytest.raises(RuntimeError, match="last range ended at 3"): + buffer.extend([record(b"four")], 9, 10) + buffer.extend([record(b"four")], 3, 4) + + +async def test_closing_a_buffer_wakes_a_reader_parked_on_it() -> None: + buffer = _StreamBuffer("s") + waiter = buffer.wait_future() + buffer.close() + # Woken rather than left parked: the reader has to unwind, and nothing + # will ever arrive for it again. + await asyncio.wait_for(waiter, 5) + assert waiter.done() + + +async def test_the_continuity_failure_says_what_to_do_about_it() -> None: + buffer = _StreamBuffer("s") + buffer.extend([record(b"one")], 0, 1) + with pytest.raises(RuntimeError) as failed: + buffer.extend([record(b"two")], 5, 6) + # The task fails and keeps failing, so the message has to name the way out. + assert "Reset the workflow" in str(failed.value) + + +def test_the_raw_workflow_stream_api_is_not_public() -> None: + # A second workflow surface taking raw stream ids and raw protos, beside + # the typed one, is not what an application should reach for. It stays + # reachable under its private name, which is what the provider and the + # contrib surface use. + import temporalio.workflow as wf + + for name in ( + "subscribe_stream", + "append_stream_records", + "read_stream_records", + "DeliveredStreamRecord", + ): + assert name not in wf.__all__ + assert not hasattr(wf, name) + for name in ( + "_subscribe_stream", + "_append_stream_records", + "_read_stream_records", + "_close_stream_records", + "_DeliveredStreamRecord", + ): + assert hasattr(wf, name) diff --git a/tests/worker/test_workflow_stream_e2e.py b/tests/worker/test_workflow_stream_e2e.py new file mode 100644 index 000000000..e2ca2d9da --- /dev/null +++ b/tests/worker/test_workflow_stream_e2e.py @@ -0,0 +1,485 @@ +"""The native provider inside a real workflow, against a server that has streams. + +The two rules the memory provider cannot keep are the measurement here: a +publish commits with its Workflow Task and never lands if the task fails, and +a read is a recorded observation the server re-supplies on replay. Needs a +Temporal server built from the AI-198 branch, because neither the stream +service nor the commands exist on a released one: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/worker/test_workflow_stream_e2e.py + +Skipped otherwise, rather than passing against a server that has no idea what a +stream is. +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.enums.v1 import EventType +from temporalio.api.stream.v1 import StreamRecord +from temporalio.client import Client +from temporalio.client_stream import StreamClient +from temporalio.streams import END, RecordKind +from temporalio.streams.providers.native import NativeStreams +from temporalio.worker import Worker +from tests.streams.test_streams_conformance import take + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream commands", +) + +INPUTS = "inputs" +DECISIONS = "decisions" + +EVENT_STREAM_SUBSCRIBED = EventType.EVENT_TYPE_WORKFLOW_STREAM_SUBSCRIBED +EVENT_STREAM_RECORDS_APPENDED = EventType.EVENT_TYPE_WORKFLOW_STREAM_RECORDS_APPENDED + + +async def _connect(provider: NativeStreams | None = None) -> Client: + # Registered once, on the client: the worker inherits it and the tests + # open handles through client.get_stream_handle. The contrib tests below + # need no provider. + return await Client.connect(TARGET or "", plugins=[provider] if provider else []) + + +async def _event_counts(client: Client, workflow_id: str) -> dict[Any, int]: + counts = { + EVENT_STREAM_RECORDS_APPENDED: 0, + EVENT_STREAM_SUBSCRIBED: 0, + EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED: 0, + } + async for event in client.get_workflow_handle(workflow_id).fetch_history_events(): + if event.event_type in counts: + counts[event.event_type] += 1 + return counts + + +@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, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + 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, + } + ) + 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"]}) + decisions.finish() + return trace + + +async def test_the_interface_loop_runs_on_the_server_with_a_cold_cache() -> None: + """Rule 2 on the native provider: every task replays from the server. + + With the cache off, each Workflow Task rebuilds the workflow from History + and the server re-supplies the ranges earlier tasks consumed, so the loop + completing at all means the same records came back in the same order. + """ + provider = NativeStreams() + client = await _connect(provider) + task_queue = "loop-tq-" + uuid.uuid4().hex[:8] + workflow_id = "loop-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[ContractLoop], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + ContractLoop.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(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 asyncio.wait_for(handle.result(), 60) + + assert trace == [ + {"kind": "decision", "n": 1}, + {"kind": "decision", "n": 2}, + {"kind": "superseded", "replaced": 1}, + {"kind": "decision", "n": 3}, + {"kind": "finish", "producer": "model"}, + ] + + # The read ends by itself: the workflow is closed and the tail + # delivered, with the workflow's own records carrying no producer. + async def read_everything() -> list[Any]: + return [ + (r.producer_id, r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + assert await asyncio.wait_for(read_everything(), 60) == [ + ("", RecordKind.DATA, {"decided": 1}), + ("", RecordKind.DATA, {"decided": 2}), + ("", RecordKind.DATA, {"retracting_attempt": 1}), + ("", RecordKind.DATA, {"decided": 3}), + ("", RecordKind.FINISH, None), + ] + counts = await _event_counts(client, workflow_id) + assert counts[EVENT_STREAM_SUBSCRIBED] == 1 + # Several tasks published, one event each; the loop spanned more than + # one task or the cold cache proved nothing. + assert counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] >= 2 + assert ( + 1 + <= counts[EVENT_STREAM_RECORDS_APPENDED] + <= counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] + ) + finally: + await provider.close() + + +# 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() + + +async def test_a_failed_task_publishes_nothing() -> None: + """Rule 1 on the native provider: the server applies the command with the task.""" + provider = NativeStreams() + client = await _connect(provider) + task_queue = "fail-tq-" + uuid.uuid4().hex[:8] + workflow_id = "fail-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[PublishThenFail], + ): + handle = await client.start_workflow( + PublishThenFail.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + records = await take( + stream.read(topic=DECISIONS, result_type=dict), 2, timeout=60 + ) + await handle.result() + assert [(r.kind, r.value) for r in records] == [ + (RecordKind.DATA, {"committed": True}), + (RecordKind.FINISH, None), + ] + finally: + await provider.close() + + +@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() + + +async def test_a_handle_without_a_run_id_reads_across_continue_as_new() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "relay-tq-" + uuid.uuid4().hex[:8] + workflow_id = "relay-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Relay]): + handle = await client.start_workflow( + Relay.run, 0, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + + async def read_everything() -> list[Any]: + return [ + (r.kind, r.value) + async for r in stream.read(topic=DECISIONS, result_type=dict) + ] + + records = await asyncio.wait_for(read_everything(), 60) + await handle.result() + # The chain is followed: the successor's record arrives on the same + # read, each cursor names its run, and the read ends with the chain. + assert records == [ + (RecordKind.DATA, {"run": 0}), + (RecordKind.DATA, {"run": 1}), + (RecordKind.FINISH, None), + ] + # Pinned to the last run, a handle sees that run's stream alone. + last_run = (await client.get_workflow_handle(workflow_id).describe()).run_id + pinned = client.get_stream_handle(workflow_id, run_id=last_run) + only_last = [ + r.value async for r in pinned.read(topic=DECISIONS, result_type=dict) + ] + assert only_last == [{"run": 1}, None] + finally: + await provider.close() + + +@workflow.defn +class PublishAndRead: + """Publishes through the low-level surface and reads its own records back.""" + + @workflow.run + async def run(self) -> list[str]: + workflow._append_stream_records( + [_record(b"alpha", "progress"), _record(b"beta", "progress")], + stream_name="output", + ) + workflow._append_stream_records([_record(b"gamma")], stream_name="output") + # A name this workflow has not written yet still names a stream it + # owns, so subscribing creates the one the later publish lands in. + workflow._subscribe_stream("output") + + received: list[str] = [] + while len(received) < 3: + for item in await workflow._read_stream_records("output"): + received.append(item.record.body.data.decode()) + return received + + +def _record(body: bytes, topic: str = "") -> StreamRecord: + return StreamRecord( + body=Payload(data=body, metadata={"encoding": b"binary/plain"}), topic=topic + ) + + +async def test_a_tasks_publishes_become_one_event() -> None: + client = await _connect() + task_queue = "publish-tq-" + uuid.uuid4().hex[:8] + workflow_id = "publish-wf-" + uuid.uuid4().hex[:8] + + async with Worker( + client, + task_queue=task_queue, + workflows=[PublishAndRead], + # Every task after the first is a replay, so completing at all means the + # reissued publish matched the event the first run wrote. + max_cached_workflows=0, + ): + result = await asyncio.wait_for( + client.execute_workflow( + PublishAndRead.run, id=workflow_id, task_queue=task_queue + ), + timeout=60, + ) + assert result == ["alpha", "beta", "gamma"] + + counts = await _event_counts(client, workflow_id) + # Two calls in one task carrying three records, so one event: the event is + # per task and stream, which is what makes publishing often free. + assert counts[EVENT_STREAM_RECORDS_APPENDED] == 1 + assert counts[EVENT_STREAM_SUBSCRIBED] == 1 + + +@workflow.defn +class ConsumeAcrossTasks: + """Reads a standalone stream over several Workflow Tasks, then reports what it saw. + + The turn structure is the point: each read that finds nothing blocks, + which ends a Workflow Task, so the run spans several. With the cache on, + every task after the first is sticky. + """ + + @workflow.run + async def run(self, stream_id: str, expected: int) -> list[str]: + workflow._subscribe_stream(stream_id) + seen: list[str] = [] + while len(seen) < expected: + for item in await workflow._read_stream_records(stream_id): + seen.append(item.record.body.data.decode()) + return seen + + +async def test_a_cached_workflow_consumes_across_sticky_tasks() -> None: + """The workflow cache stays on, which is what a real worker does. + + The server sends no replay slice for a sticky task, while the sticky + history still carries the previous task's consumed range, and the two + together have to agree on every task after the first consumed range. + """ + client = await _connect() + streams = StreamClient.connect(TARGET or "") + task_queue = "sticky-tq-" + uuid.uuid4().hex[:8] + workflow_id = "sticky-wf-" + uuid.uuid4().hex[:8] + stream_id = "sticky-src-" + uuid.uuid4().hex[:8] + + batches = [["a1", "a2"], ["b1"], ["c1", "c2", "c3"]] + expected = [tok for batch in batches for tok in batch] + + try: + await streams.create(stream_id) + async with Worker( + client, task_queue=task_queue, workflows=[ConsumeAcrossTasks] + ): + handle = await client.start_workflow( + ConsumeAcrossTasks.run, + args=[stream_id, len(expected)], + id=workflow_id, + task_queue=task_queue, + ) + + # Spaced out so the workflow drains, blocks and ends a task between + # them. Without the gap the appends coalesce into one task and the + # sticky path is never taken. + for batch in batches: + await streams.get(stream_id).append( + *[_record(t.encode()) for t in batch] + ) + await asyncio.sleep(0.4) + + assert await asyncio.wait_for(handle.result(), timeout=60) == expected + + counts = await _event_counts(client, workflow_id) + assert counts[EventType.EVENT_TYPE_WORKFLOW_TASK_COMPLETED] >= 3, ( + "sticky path not exercised" + ) + finally: + await streams.close() + + +@workflow.defn +class StartsWhenTold: + """Subscribes to ``inputs`` at a start given by a signal and returns what it read.""" + + def __init__(self) -> None: + self._start: str | None = None + + @workflow.signal + def begin(self, start: str) -> None: + self._start = start + + @workflow.run + async def run(self) -> list[Any]: + await workflow.wait_condition(lambda: self._start is not None) + if self._start == "end": + reader = workflow.stream_reader(INPUTS, result_type=dict, after=END) + want = 1 + else: + reader = workflow.stream_reader(INPUTS, result_type=dict, last=2) + want = 2 + values: list[Any] = [] + async for value in reader.values(): + values.append(value["n"]) + if len(values) == want: + break + return values + + +async def _subscribed_offsets(client: Client, workflow_id: str) -> list[int]: + return [ + event.workflow_stream_subscribed_event_attributes.start_offset + async for event in client.get_workflow_handle( + workflow_id + ).fetch_history_events() + if event.event_type == EVENT_STREAM_SUBSCRIBED + ] + + +async def test_a_workflow_reader_starts_at_the_last_n_records_on_the_server() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "last-tq-" + uuid.uuid4().hex[:8] + workflow_id = "last-wf-" + uuid.uuid4().hex[:8] + try: + # Cold, so every task replays and the recorded start is what places + # the reader, not a second resolution. + async with Worker( + client, + task_queue=task_queue, + workflows=[StartsWhenTold], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + StartsWhenTold.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}, {"n": 4}) + await handle.signal(StartsWhenTold.begin, "last") + assert await asyncio.wait_for(handle.result(), 60) == [3, 4] + assert await _subscribed_offsets(client, workflow_id) == [2] + finally: + await provider.close() + + +async def test_a_workflow_reader_at_end_skips_what_was_there_on_the_server() -> None: + provider = NativeStreams() + client = await _connect(provider) + task_queue = "end-tq-" + uuid.uuid4().hex[:8] + workflow_id = "end-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker( + client, + task_queue=task_queue, + workflows=[StartsWhenTold], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + StartsWhenTold.run, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "old"}, {"n": "old"}) + await handle.signal(StartsWhenTold.begin, "end") + result = asyncio.ensure_future(handle.result()) + # The subscription registers when the signalled task completes, + # which the test does not observe, so appends keep coming. + 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, 60) == ["new"] + offsets = await _subscribed_offsets(client, workflow_id) + assert len(offsets) == 1 and offsets[0] >= 2 + finally: + await provider.close() diff --git a/tests/worker/test_workflow_stream_payloads_e2e.py b/tests/worker/test_workflow_stream_payloads_e2e.py new file mode 100644 index 000000000..aa739fb1f --- /dev/null +++ b/tests/worker/test_workflow_stream_payloads_e2e.py @@ -0,0 +1,252 @@ +"""How a body's encoding meets the server: retry identity and external storage. + +A payload codec and an ``ExternalStorage`` driver both change the bytes a body +is stored as. The server must still recognize a retry, and a driver must be +applied on both halves of the native provider. Needs a Temporal server built +from the AI-198 branch: + + TEMPORAL_STREAM_TARGET=127.0.0.1:7333 uv run pytest tests/worker/test_workflow_stream_payloads_e2e.py +""" + +from __future__ import annotations + +import asyncio +import os +import uuid +from collections.abc import Sequence +from typing import Any + +import pytest + +from temporalio import workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.sdk.v1.external_storage_pb2 import ExternalStorageReference +from temporalio.client import Client +from temporalio.client_stream import StreamClient +from temporalio.converter import ( + DataConverter, + ExternalStorage, + JSONProtoPayloadConverter, + PayloadCodec, +) +from temporalio.streams import CONTENT_HASH_KEY, RecordKind +from temporalio.streams.providers.native import NativeStreams +from temporalio.worker import Worker +from tests.streams.test_streams_conformance import take +from tests.test_extstore import InMemoryTestDriver + +TARGET = os.environ.get("TEMPORAL_STREAM_TARGET") + +pytestmark = pytest.mark.skipif( + not TARGET, + reason="set TEMPORAL_STREAM_TARGET to a server with the stream commands", +) + +INPUTS = "inputs" +DECISIONS = "decisions" + + +class _NonceCodec(PayloadCodec): + """Encodes to different bytes on every call, as a nonce-based cipher does. + + The plaintext is kept in the clear behind a counter so the test can read + the stored bytes back and see that two encodings of one value differ. + """ + + def __init__(self) -> None: + self.calls = 0 + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + self.calls += 1 + return [ + Payload( + metadata={"encoding": b"binary/nonce"}, + data=f"{self.calls}:".encode() + p.SerializeToString(), + ) + for p in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + out: list[Payload] = [] + for p in payloads: + if p.metadata.get("encoding", b"") != b"binary/nonce": + out.append(p) + continue + out.append(Payload.FromString(p.data.split(b":", 1)[1])) + return out + + +@workflow.defn +class Echo: + """Reads ``inputs`` and publishes each value on ``decisions``, until told to stop.""" + + def __init__(self) -> None: + self._stop = False + + @workflow.signal + def stop(self) -> None: + self._stop = True + + @workflow.run + async def run(self, count: int) -> list[Any]: + inputs = workflow.stream_reader(INPUTS, result_type=dict) + decisions = workflow.stream_writer(DECISIONS) + seen: list[Any] = [] + async for value in inputs.values(): + decisions.publish({"echo": value}) + seen.append(value) + if len(seen) >= count: + break + decisions.finish() + await workflow.wait_condition(lambda: self._stop) + return seen + + +async def _connect(converter: DataConverter, provider: NativeStreams) -> Client: + plain = await Client.connect(TARGET or "") + config = plain.config() + config["data_converter"] = converter + config["plugins"] = [provider] + return Client(**config) + + +async def test_a_retried_append_under_a_nonce_codec_is_deduplicated() -> None: + """Retry identity is the plaintext hash, not the encoded bytes. + + Two producers with the same identity append the same value at the same + sequence, as a retry after a lost response would. The codec encodes each + differently, so the server sees two different bodies and one hash, and + writes the record once. + """ + provider = NativeStreams() + client = await _connect(DataConverter(payload_codec=_NonceCodec()), provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "nonce-tq-" + uuid.uuid4().hex[:8] + workflow_id = "nonce-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + first = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + retry = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + landed = await first.append({"n": 1}) + again = await retry.append({"n": 1}) + assert again == landed, "the repeat answers with the original position" + + echoed = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert [r.value for r in echoed] == [{"echo": {"n": 1}}] + await handle.signal(Echo.stop) + assert await asyncio.wait_for(handle.result(), 60) == [{"n": 1}] + + run_id = (await handle.describe()).run_id + assert run_id is not None + page = await raw.workflow_stream(workflow_id, INPUTS, owner_run_id=run_id).poll( + from_offset=0, wait=False + ) + bodies = [e.record for e in page.entries if e.record.HasField("body")] + assert len(bodies) == 1, "written once" + assert bodies[0].body.metadata["encoding"] == b"binary/nonce" + assert bodies[0].metadata[CONTENT_HASH_KEY].data + finally: + await raw.close() + await provider.close() + + +async def test_a_workflow_publish_carries_the_plaintext_hash() -> None: + """The workflow's own records are stamped too, from the workflow thread.""" + provider = NativeStreams() + client = await _connect(DataConverter.default, provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "stamp-tq-" + uuid.uuid4().hex[:8] + workflow_id = "stamp-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + await stream.producer(topic=INPUTS, producer_id="tool", attempt=1).append( + {"n": 7} + ) + await take(stream.read(topic=DECISIONS, result_type=dict), 1) + await handle.signal(Echo.stop) + await asyncio.wait_for(handle.result(), 60) + + run_id = (await handle.describe()).run_id + assert run_id is not None + page = await raw.workflow_stream( + workflow_id, DECISIONS, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + kinds = {e.record.kind for e in page.entries} + assert len(page.entries) == 2, "one decision and the finish" + data = [e.record for e in page.entries if e.record.HasField("body")] + assert len(data) == 1 and data[0].metadata[CONTENT_HASH_KEY].data + finish = [e.record for e in page.entries if not e.record.HasField("body")] + assert CONTENT_HASH_KEY not in finish[0].metadata + assert RecordKind.FINISH in {RecordKind(k) for k in kinds} + finally: + await raw.close() + await provider.close() + + +async def test_external_storage_applies_on_both_halves() -> None: + """A ``StorageDriver`` on the client offloads stream bodies on both paths. + + Every body is over the threshold. The outside producer's append is stored + through the driver before it reaches the server; the workflow retrieves it + on its task, publishes, and the worker's payload pass offloads that too, off + the workflow thread, so the outside reader retrieves it. Read raw, the + server holds references on both topics and no plaintext. + """ + driver = InMemoryTestDriver() + converter = DataConverter( + external_storage=ExternalStorage(drivers=[driver], payload_size_threshold=0) + ) + provider = NativeStreams() + client = await _connect(converter, provider) + raw = StreamClient.connect(TARGET or "") + task_queue = "offload-tq-" + uuid.uuid4().hex[:8] + workflow_id = "offload-wf-" + uuid.uuid4().hex[:8] + try: + async with Worker(client, task_queue=task_queue, workflows=[Echo]): + handle = await client.start_workflow( + Echo.run, 1, id=workflow_id, task_queue=task_queue + ) + stream = client.get_stream_handle(workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="tool", attempt=1) + await producer.append({"n": "plaintext-marker"}) + stores_after_append = driver._store_calls + assert stores_after_append >= 1, "the outside append offloaded" + + echoed = await take(stream.read(topic=DECISIONS, result_type=dict), 1) + assert [r.value for r in echoed] == [{"echo": {"n": "plaintext-marker"}}] + await handle.signal(Echo.stop) + assert await asyncio.wait_for(handle.result(), 60) == [ + {"n": "plaintext-marker"} + ] + + run_id = (await handle.describe()).run_id + assert run_id is not None + for topic in (INPUTS, DECISIONS): + page = await raw.workflow_stream( + workflow_id, topic, owner_run_id=run_id + ).poll(from_offset=0, wait=False) + bodies = [e.record.body for e in page.entries if e.record.HasField("body")] + assert len(bodies) == 1, topic + assert b"plaintext-marker" not in bodies[0].data, topic + reference = JSONProtoPayloadConverter().from_payload( + bodies[0], ExternalStorageReference + ) + assert reference.driver_name == driver.name(), topic + # The workflow's publish was offloaded by the worker, after the append. + assert driver._store_calls > stores_after_append + # Both halves retrieved: the worker on its task and the outside reader. + assert driver._retrieve_calls >= 2 + # The outside append was stored under the workflow that owns the stream. + targets = [ctx.target for ctx in driver._store_contexts if ctx.target] + assert any(t.id == workflow_id and t.run_id == run_id for t in targets) + finally: + await raw.close() + await provider.close()