diff --git a/CHANGELOG.md b/CHANGELOG.md index 4be5a64a9..38673bb6d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -76,6 +76,17 @@ to include examples, links to docs, or any other relevant information. refuses a conflicting repeat, falling back to the shipped Signal on a workflow whose worker predates it. A workflow's activity keeps its own streams in the workflow's log under `activity//`. +- **Experimental**: `temporalio.streams.providers.nexus.NexusStreams` puts one + Nexus endpoint in front of a storage provider, so a caller reaches a stream + through the endpoint and never names the store, and + `TemporalStreamsHandler` serves that endpoint by fronting the provider's own + handles. Its contract is defined in `temporal_streams.nexusrpc.yaml` and the + bindings are generated from it; a record crosses as the serialized + `StreamRecord` proto, and both operations address a stream by a `StreamRef` + naming its owner (a workflow, an activity or a standalone stream) and topic, + which the handler maps onto the store's accessor for that owner. Configure + the front with `data_converter=` to run a payload codec on the caller side, + so records are encoded before they leave the process. ### Changed diff --git a/pyproject.toml b/pyproject.toml index cacec1ad0..fcf82b526 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -125,9 +125,11 @@ format = [ ] gen-docs = "uv run scripts/gen_docs.py" gen-nexus-system-api = "uv run scripts/gen_nexus_system_api.py" +gen-streams-nexus-api = "uv run scripts/gen_streams_nexus_api.py" gen-protos = [ { cmd = "uv run scripts/gen_protos.py" }, { ref = "gen-nexus-system-api" }, + { ref = "gen-streams-nexus-api" }, { cmd = "uv run scripts/gen_payload_visitor.py" }, { cmd = "uv run scripts/gen_bridge_client.py" }, { ref = "format" }, @@ -135,6 +137,7 @@ gen-protos = [ gen-protos-docker = [ { cmd = "uv run scripts/gen_protos_docker.py" }, { ref = "gen-nexus-system-api" }, + { ref = "gen-streams-nexus-api" }, { cmd = "uv run scripts/gen_payload_visitor.py" }, { cmd = "uv run scripts/gen_bridge_client.py" }, { ref = "format" }, @@ -198,16 +201,21 @@ exclude = [ 'temporalio/api', 'temporalio/bridge/proto', 'temporalio/nexus/system/workflow_service', + 'temporalio/streams/providers/_nexus_generated', ] [[tool.mypy.overrides]] module = "temporalio.nexus.system.workflow_service.*" ignore_errors = true +[[tool.mypy.overrides]] +module = "temporalio.streams.providers._nexus_generated.*" +ignore_errors = true + [tool.pydocstyle] convention = "google" # https://github.com/PyCQA/pydocstyle/issues/363#issuecomment-625563088 -match_dir = "^(?!(docs|scripts|tests|api|proto|system|\\.)).*" +match_dir = "^(?!(docs|scripts|tests|api|proto|system|_nexus_generated|\\.)).*" add_ignore = [ # We like to wrap at a certain number of chars, even long summary sentences. # https://github.com/PyCQA/pydocstyle/issues/184 @@ -241,6 +249,7 @@ privacy = [ "HIDDEN:temporalio.worker.workflow_sandbox.importer", "HIDDEN:temporalio.worker.workflow_sandbox.in_sandbox", "HIDDEN:**.*_pb2*", + "HIDDEN:temporalio.streams.providers._nexus_generated._definitions", ] project-name = "Temporal Python" sidebar-expand-depth = 2 diff --git a/scripts/gen_streams_nexus_api.py b/scripts/gen_streams_nexus_api.py new file mode 100644 index 000000000..2d4f16c6d --- /dev/null +++ b/scripts/gen_streams_nexus_api.py @@ -0,0 +1,95 @@ +import os +import shutil +import subprocess +import sys +from pathlib import Path + +base_dir = Path(__file__).parent.parent +providers_dir = base_dir / "temporalio" / "streams" / "providers" +contract_path = providers_dir / "temporal_streams.nexusrpc.yaml" +output_dir = providers_dir / "_nexus_generated" +# Pinned to the version CI installs, because nexgen's output changes between +# releases and check-protos compares against what is committed here. +NEX_GEN_VERSION = "0.2.4" + + +def nex_gen_command() -> list[str]: + if bin_path := os.environ.get("NEX_GEN_BIN"): + return [bin_path] + + if shutil.which("nexgen") is None: + subprocess.check_call( + [ + "cargo", + "install", + "--locked", + "nexgen", + "--version", + NEX_GEN_VERSION, + # Same build as the system API script installs, so one binary + # serves both and neither can overwrite the other's output. + "--features", + "advanced", + "--force", + ] + ) + return ["nexgen"] + + +def check_version(command: list[str]) -> None: + # A different release on PATH would regenerate different code locally + # and the drift would only show in CI, so refuse before writing anything. + reported = subprocess.check_output([*command, "--version"], text=True).strip() + found = reported.split()[-1] if reported else "" + if found != NEX_GEN_VERSION: + raise SystemExit( + f"found nexgen {found or '?'} at {command[0]}, but the stream contract is " + f"generated with {NEX_GEN_VERSION}. Install it with `cargo install --locked " + f"nexgen --version {NEX_GEN_VERSION} --features advanced` or point " + "NEX_GEN_BIN at that binary." + ) + + +def generate_streams_nexus_api() -> None: + if not contract_path.exists(): + raise RuntimeError(f"missing stream contract: {contract_path}") + + command = nex_gen_command() + check_version(command) + shutil.rmtree(output_dir, ignore_errors=True) + subprocess.check_call( + [ + *command, + "python", + str(contract_path), + "--output", + str(output_dir), + ] + ) + subprocess.check_call( + [ + sys.executable, + "-m", + "ruff", + "check", + "--select", + "I", + "--fix", + str(output_dir), + ] + ) + subprocess.check_call( + [ + sys.executable, + "-m", + "ruff", + "format", + str(output_dir), + ] + ) + + +if __name__ == "__main__": + print("Generating stream endpoint Nexus API...", file=sys.stderr) + generate_streams_nexus_api() + print("Done", file=sys.stderr) diff --git a/temporalio/streams/providers/_nexus_generated/__init__.py b/temporalio/streams/providers/_nexus_generated/__init__.py new file mode 100644 index 000000000..d855a6381 --- /dev/null +++ b/temporalio/streams/providers/_nexus_generated/__init__.py @@ -0,0 +1,25 @@ +# Generated by nexgen v0.2.4. DO NOT EDIT! + +from __future__ import annotations + +from ._definitions import Violation +from .models import ( + AppendInput, + AppendOutput, + ReadInput, + ReadOutput, + RecordWire, + StreamRef, +) +from .services import TemporalStreams + +__all__ = [ + "Violation", + "AppendInput", + "AppendOutput", + "ReadInput", + "ReadOutput", + "RecordWire", + "StreamRef", + "TemporalStreams", +] diff --git a/temporalio/streams/providers/_nexus_generated/_definitions.py b/temporalio/streams/providers/_nexus_generated/_definitions.py new file mode 100644 index 000000000..dc8de5bb5 --- /dev/null +++ b/temporalio/streams/providers/_nexus_generated/_definitions.py @@ -0,0 +1,542 @@ +# Generated by nexgen v0.2.4. DO NOT EDIT! + +from __future__ import annotations + +import base64 +import collections.abc +import dataclasses +import datetime +import json +import re +import typing + +import temporalio.converter +import temporalio.exceptions + +__all__ = [ + "Violation", + "_binary64", + "_check_contains", + "_check_date_time", + "_check_duration", + "_check_time", + "_check_unique_items", + "_collect", + "_json_values_equal", + "_format_base64", + "_format_base64url", + "_format_date", + "_format_date_time", + "_format_duration", + "_format_time", + "_parse_base64", + "_parse_base64url", + "_parse_date", + "_parse_date_time", + "_parse_duration", + "_parse_spec_integer", + "_parse_time", + "_quote", + "_transfer_type_convertible", +] + + +@dataclasses.dataclass(frozen=True, slots=True) +class Violation: + """A single constraint failure, located by JSON path.""" + + path: str + reason: str + + +def _quote(value: object) -> str: + """Renders a value in the JSON form every target quotes offending values in.""" + + try: + return json.dumps(value, ensure_ascii=False) + except (TypeError, ValueError): + return repr(value) + + +def _collect( + violations: list[Violation], + path: str, + error: temporalio.exceptions.ApplicationError, +) -> None: + """Re-paths a nested model's violations under `path` and appends them.""" + + if error.type != "PayloadValidationError" or not error.details: + raise error + # Generated failures retain the original list as their first detail. The + # cast is type-checker-only and performs no serialization. + nested_violations = typing.cast(list[Violation], error.details[0]) + for inner in nested_violations: + # A nested violation about the value *itself* carries no path of its own + # (a union branch's own constraint, an element-level check), so the + # prefix is the whole path -- never a dangling separator (P11). + nested = f"{path}.{inner.path}" if inner.path else path + violations.append(Violation(path=nested, reason=inner.reason)) + + +_ModelT = typing.TypeVar("_ModelT") + + +def _transfer_type_convertible( + converter: type[temporalio.converter.TransferTypeConverter[typing.Any, typing.Any]], +) -> collections.abc.Callable[[type[_ModelT]], type[_ModelT]]: + """Registers a transfer type converter on a model class. + + Wraps `temporalio.converter.transfer_type_convertible` to erase the + converter's value-type parameter. Binding it directly on the decorated class + is circular for a static type checker -- the class's type depends on the + decorator, whose value type depends on the class -- which degrades the model + to `Unknown`. Erasing it here keeps the decorator idiomatic at each model and + resolves the cycle. + """ + + return temporalio.converter.transfer_type_convertible(converter) + + +_INTEGER_CAP = (1 << 53) - 1 + + +def _parse_spec_integer( + value: object, path: str, violations: list[Violation] +) -> int | None: + """Parses a JSON number as a spec integer (`1.0` accepted, `1.5` rejected).""" + + # `bool` is a subclass of `int`, so it must be excluded before the int check. + if isinstance(value, bool) or not isinstance(value, (int, float)): + violations.append(Violation(path=path, reason="expected integer")) + return None + if isinstance(value, float): + if not value.is_integer(): + violations.append(Violation(path=path, reason="expected integer")) + return None + out = int(value) + else: + out = value + if abs(out) > _INTEGER_CAP: + violations.append(Violation(path=path, reason="expected integer")) + return None + return out + + +def _binary64(value: float) -> float: + """Narrows a `number` value to the binary64 domain shared by every target.""" + + try: + return float(value) + except OverflowError: + # Past the binary64 range there is nothing to narrow to. The caller's + # finiteness check has already recorded that violation and will raise, so + # hand the value back rather than throwing out of turn and losing the + # other violations collected alongside it. + return value + + +def _json_values_equal(left: typing.Any, right: typing.Any) -> bool: + """Compares JSON values without Python's bool-as-int equality leak.""" + + left_is_number = isinstance(left, (int, float)) and not isinstance(left, bool) + right_is_number = isinstance(right, (int, float)) and not isinstance(right, bool) + if left_is_number or right_is_number: + return left_is_number and right_is_number and left == right + if type(left) is not type(right): + return False + if isinstance(left, list): + left_list = typing.cast("list[typing.Any]", left) + right_list = typing.cast("list[typing.Any]", right) + return len(left_list) == len(right_list) and all( + _json_values_equal(a, b) for a, b in zip(left_list, right_list) + ) + if isinstance(left, dict): + left_dict = typing.cast("dict[typing.Any, typing.Any]", left) + right_dict = typing.cast("dict[typing.Any, typing.Any]", right) + return left_dict.keys() == right_dict.keys() and all( + _json_values_equal(left_dict[key], right_dict[key]) for key in left_dict + ) + return left == right + + +def _check_unique_items( + value: list[typing.Any], path: str, violations: list[Violation] +) -> None: + """Asserts an array's elements are pairwise distinct.""" + + seen: list[typing.Any] = [] + for index, element in enumerate(value): + for earlier, previous in enumerate(seen): + if _json_values_equal(previous, element): + violations.append( + Violation( + path=path, + reason=( + f"duplicate items: element at index {index} " + f"equals index {earlier}" + ), + ) + ) + break + seen.append(element) + + +def _check_contains( + value: list[typing.Any], + matches: typing.Callable[[typing.Any], bool], + min_contains: int, + max_contains: int | None, + bounded_min: bool, + path: str, + violations: list[Violation], +) -> None: + """Asserts how many of an array's elements match the `contains` schema.""" + + match_count = sum(1 for element in value if matches(element)) + if match_count < min_contains: + if bounded_min: + violations.append( + Violation( + path=path, + reason=( + f"too few matching items: at least {min_contains}, " + f"got {match_count}" + ), + ) + ) + else: + violations.append( + Violation(path=path, reason="no element matches the required schema") + ) + if max_contains is not None and match_count > max_contains: + violations.append( + Violation( + path=path, + reason=( + f"too many matching items: at most {max_contains}, " + f"got {match_count}" + ), + ) + ) + + +_TEMPORAL_DATE_TIME_RE = re.compile( + r"^[0-9]{4}-(0[1-9]|1[0-2])-(0[1-9]|[12][0-9]|3[01])[Tt]([01][0-9]|2[0-3]):[0-5][0-9]:[0-5][0-9](\.[0-9]+)?([Zz]|[+-]([01][0-9]|2[0-3]):[0-5][0-9])\Z" +) +_TEMPORAL_DATE_RE = re.compile(r"^[0-9]{4}-(0[1-9]|1[0-2])-(0[1-9]|[12][0-9]|3[01])\Z") +_TEMPORAL_TIME_RE = re.compile( + r"^([01][0-9]|2[0-3]):[0-5][0-9]:[0-5][0-9](\.[0-9]+)?([Zz]|[+-]([01][0-9]|2[0-3]):[0-5][0-9])?\Z" +) +_TEMPORAL_DURATION_RE = re.compile( + r"^PT(?:[0-9]+H(?:[0-9]+M(?:[0-9]+S)?)?|[0-9]+M(?:[0-9]+S)?|[0-9]+S)\Z" +) +_TEMPORAL_MAX_DURATION_SECONDS = ((1 << 63) - 1) // 1_000_000_000 +# A duration component with more digits than the cap itself is over the cap +# whatever those digits are, which is how the magnitude is bounded before `int()` +# sees it: CPython refuses to convert a string of more than 4300 digits. +_TEMPORAL_MAX_DURATION_DIGITS = len(str(_TEMPORAL_MAX_DURATION_SECONDS)) +# `datetime` resolves to microseconds, and `fromisoformat` before Python 3.11 +# parses only the fraction widths `isoformat` writes. +_TEMPORAL_FRACTION_DIGITS = 6 + + +def _days_in_month(year: int, month: int) -> int: + if month in (1, 3, 5, 7, 8, 10, 12): + return 31 + if month in (4, 6, 9, 11): + return 30 + if month == 2: + return 29 if (year % 4 == 0 and year % 100 != 0) or year % 400 == 0 else 28 + return 0 + + +def _valid_temporal_calendar(value: str) -> bool: + if len(value) < 10: + return False + try: + year, month, day = int(value[0:4]), int(value[5:7]), int(value[8:10]) + except ValueError: + return False + # `datetime.MINYEAR` is 1, which is also the shared cross-language floor. + # Year 0000 is rejected rather than shifted into range, and + # `_temporal_reason` says so. + if year < datetime.MINYEAR: + return False + maximum = _days_in_month(year, month) + return maximum > 0 and 1 <= day <= maximum + + +def _temporal_reason(name: str, value: str) -> str: + """The reason a rejected temporal string is reported under. + + Year 0000 earns its own clause so the caller sees the shared calendar floor + rather than only a generic malformed-timestamp reason. + """ + + if value[0:4] == "0000": + return ( + f"must be a valid {name}, got {_quote(value)}: year 0000 is not" + f" representable (datetime.MINYEAR is {datetime.MINYEAR})" + ) + return f"must be a valid {name}, got {_quote(value)}" + + +def _temporal_isoformat(value: str) -> str: + """Rewrites a wire temporal into the spelling `fromisoformat` accepts. + + `Z` becomes `+00:00`, and the fractional second is padded or truncated to + exactly `_TEMPORAL_FRACTION_DIGITS`: before Python 3.11 `fromisoformat` + parses only what `isoformat` writes, so an RFC 3339 `.1` or `.1234567` -- + which every other target accepts -- would otherwise raise. Digits past the + sixth are dropped, the loss at `datetime`'s own resolution that P1 allows; + the canonical output re-trims the padding, so `.1` still writes as `.1`. + """ + + normalized = value.upper() + if normalized.endswith("Z"): + normalized = normalized[:-1] + "+00:00" + dot = normalized.find(".") + if dot < 0: + return normalized + end = dot + 1 + while end < len(normalized) and normalized[end].isdigit(): + end += 1 + fraction = normalized[dot + 1 : end].ljust(_TEMPORAL_FRACTION_DIGITS, "0") + return ( + normalized[: dot + 1] + fraction[:_TEMPORAL_FRACTION_DIGITS] + normalized[end:] + ) + + +def _parse_date_time( + value: str, path: str, violations: list[Violation] +) -> datetime.datetime | None: + if _TEMPORAL_DATE_TIME_RE.match(value) is None or not _valid_temporal_calendar( + value + ): + violations.append( + Violation(path=path, reason=_temporal_reason("date-time", value)) + ) + return None + return datetime.datetime.fromisoformat(_temporal_isoformat(value)) + + +def _parse_date( + value: str, path: str, violations: list[Violation] +) -> datetime.date | None: + if _TEMPORAL_DATE_RE.match(value) is None or not _valid_temporal_calendar(value): + violations.append(Violation(path=path, reason=_temporal_reason("date", value))) + return None + return datetime.date.fromisoformat(value) + + +def _parse_time( + value: str, path: str, violations: list[Violation] +) -> datetime.time | None: + if _TEMPORAL_TIME_RE.match(value) is None: + violations.append(Violation(path=path, reason=_temporal_reason("time", value))) + return None + return datetime.time.fromisoformat(_temporal_isoformat(value)) + + +def _parse_duration( + value: str, path: str, violations: list[Violation] +) -> datetime.timedelta | None: + if _TEMPORAL_DURATION_RE.match(value) is None: + violations.append( + Violation(path=path, reason=_temporal_reason("duration", value)) + ) + return None + total = 0 + number = "" + for char in value[2:]: + if char.isdigit(): + number += char + continue + digits = number.lstrip("0") + number = "" + if len(digits) > _TEMPORAL_MAX_DURATION_DIGITS: + # Over the cap by digit count alone (see the constant), so the + # conversion `int()` would refuse is never attempted. + total = _TEMPORAL_MAX_DURATION_SECONDS + 1 + break + total += int(digits or "0") * {"H": 3600, "M": 60, "S": 1}[char] + if total > _TEMPORAL_MAX_DURATION_SECONDS: + break + if total > _TEMPORAL_MAX_DURATION_SECONDS: + violations.append( + Violation(path=path, reason=_temporal_reason("duration", value)) + ) + return None + return datetime.timedelta(seconds=total) + + +def _check_temporal_offset( + name: str, + value: datetime.datetime | datetime.time, + offset: datetime.timedelta, + path: str, + violations: list[Violation], +) -> None: + """Asserts a UTC offset is a whole number of minutes, the finest the wire + form spells (`tzinfo` allows seconds, which the offset would silently lose). + """ + + if offset % datetime.timedelta(minutes=1): + violations.append( + Violation( + path=path, + reason=( + f"must be a valid {name}, got {_quote(str(value))}: " + f"the UTC offset {offset} is not a whole number of minutes" + ), + ) + ) + + +def _check_date_time( + value: datetime.datetime, path: str, violations: list[Violation] +) -> None: + """Asserts a datetime is writable as a wire date-time (P12). + + A dataclass is constructed unchecked, so a naive datetime -- with no offset + the required wire form could carry -- reaches serialize; without this it + would emit a value this module's own parser rejects. + """ + + offset = value.utcoffset() + if offset is None: + violations.append( + Violation( + path=path, + reason=( + f"must be a valid date-time, got {_quote(str(value))}: " + "a naive datetime carries no UTC offset" + ), + ) + ) + return + _check_temporal_offset("date-time", value, offset, path, violations) + + +def _check_time(value: datetime.time, path: str, violations: list[Violation]) -> None: + """Asserts a time is writable as a wire time (P12). The offset is optional in + the grammar, so only its precision is held to anything.""" + + offset = value.utcoffset() + if offset is not None: + _check_temporal_offset("time", value, offset, path, violations) + + +def _check_duration( + value: datetime.timedelta, path: str, violations: list[Violation] +) -> None: + """Asserts a timedelta is writable as a wire duration (P12): the grammar is + unsigned, whole-second and capped, and a `timedelta` is none of those.""" + + if value < datetime.timedelta(0): + reason = "a duration cannot be negative" + elif value % datetime.timedelta(seconds=1): + reason = "a duration cannot carry a fraction of a second" + elif value.total_seconds() > _TEMPORAL_MAX_DURATION_SECONDS: + reason = f"a duration cannot exceed {_TEMPORAL_MAX_DURATION_SECONDS} seconds" + else: + return + violations.append( + Violation( + path=path, + reason=f"must be a valid duration, got {_quote(str(value))}: {reason}", + ) + ) + + +def _temporal_frac(microsecond: int) -> str: + if microsecond == 0: + return "" + return "." + f"{microsecond:06d}".rstrip("0") + + +def _temporal_offset(value: datetime.datetime | datetime.time) -> str: + offset = value.utcoffset() + if offset is None: + return "" + total = int(offset.total_seconds()) + if total == 0: + return "Z" + sign = "+" if total > 0 else "-" + total = abs(total) + return f"{sign}{total // 3600:02d}:{(total % 3600) // 60:02d}" + + +def _format_date_time(value: datetime.datetime) -> str: + return ( + f"{value.year:04d}-{value.month:02d}-{value.day:02d}" + f"T{value.hour:02d}:{value.minute:02d}:{value.second:02d}" + f"{_temporal_frac(value.microsecond)}{_temporal_offset(value)}" + ) + + +def _format_date(value: datetime.date) -> str: + return f"{value.year:04d}-{value.month:02d}-{value.day:02d}" + + +def _format_time(value: datetime.time) -> str: + return ( + f"{value.hour:02d}:{value.minute:02d}:{value.second:02d}" + f"{_temporal_frac(value.microsecond)}{_temporal_offset(value)}" + ) + + +def _format_duration(value: datetime.timedelta) -> str: + total = int(value.total_seconds()) + if total == 0: + return "PT0S" + hours, remainder = divmod(total, 3600) + minutes, seconds = divmod(remainder, 60) + out = "PT" + if hours: + out += f"{hours}H" + if minutes: + out += f"{minutes}M" + if seconds: + out += f"{seconds}S" + return out + + +_BASE64_RE = re.compile( + "^(?:[A-Za-z0-9+/]{4})*(?:[A-Za-z0-9+/][AQgw]==|[A-Za-z0-9+/]{2}[AEIMQUYcgkosw048]=)?\\Z", + re.ASCII, +) +_BASE64URL_RE = re.compile( + "^(?:[A-Za-z0-9_-]{4})*(?:[A-Za-z0-9_-][AQgw]|[A-Za-z0-9_-]{2}[AEIMQUYcgkosw048])?\\Z", + re.ASCII, +) + + +def _parse_base64(value: str, path: str, violations: list[Violation]) -> bytes | None: + if _BASE64_RE.match(value) is None: + violations.append( + Violation(path=path, reason=f"must be base64-encoded, got {_quote(value)}") + ) + return None + return base64.b64decode(value, validate=True) + + +def _format_base64(value: bytes) -> str: + return base64.b64encode(value).decode("ascii") + + +def _parse_base64url( + value: str, path: str, violations: list[Violation] +) -> bytes | None: + if _BASE64URL_RE.match(value) is None: + violations.append( + Violation( + path=path, reason=f"must be base64url-encoded, got {_quote(value)}" + ) + ) + return None + return base64.urlsafe_b64decode(value + "=" * (-len(value) % 4)) + + +def _format_base64url(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii") diff --git a/temporalio/streams/providers/_nexus_generated/models.py b/temporalio/streams/providers/_nexus_generated/models.py new file mode 100644 index 000000000..fd4c65fb6 --- /dev/null +++ b/temporalio/streams/providers/_nexus_generated/models.py @@ -0,0 +1,1029 @@ +# Generated by nexgen v0.2.4. DO NOT EDIT! + +from __future__ import annotations + +import dataclasses +import typing + +import typing_extensions + +import temporalio.converter +import temporalio.exceptions + +from ._definitions import ( + Violation, + _collect, + _format_base64, + _parse_base64, + _parse_spec_integer, + _quote, + _transfer_type_convertible, +) + +_APPEND_OUTPUT_DECLARED: frozenset[str] = frozenset({"cursor"}) + + +_READ_OUTPUT_DECLARED: frozenset[str] = frozenset({"records", "next_token", "done"}) + + +_RECORD_WIRE_DECLARED: frozenset[str] = frozenset({"token", "record"}) + + +class _AppendInputTransferTypeConverter( + temporalio.converter.TransferTypeConverter["AppendInput", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["AppendInput"] + ) -> "AppendInput": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + stream_value: StreamRef = typing.cast("typing.Any", None) + if "stream" not in raw or raw["stream"] is None: + violations.append(Violation(path="stream", reason="required")) + else: + stream_value_raw = raw["stream"] + try: + stream_value = _StreamRefTransferTypeConverter().from_transfer_type( + stream_value_raw, StreamRef + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, "stream", error) + + producer_id_value: str = typing.cast("typing.Any", None) + if "producer_id" not in raw or raw["producer_id"] is None: + violations.append(Violation(path="producer_id", reason="required")) + else: + producer_id_value_raw = raw["producer_id"] + if not isinstance(producer_id_value_raw, str): + violations.append( + Violation(path="producer_id", reason="expected string") + ) + else: + producer_id_value = producer_id_value_raw + + attempt_value: int = typing.cast("typing.Any", None) + if "attempt" not in raw or raw["attempt"] is None: + violations.append(Violation(path="attempt", reason="required")) + else: + attempt_value_raw = raw["attempt"] + attempt_value_parsed = _parse_spec_integer( + attempt_value_raw, "attempt", violations + ) + if attempt_value_parsed is not None: + attempt_value = attempt_value_parsed + + sequence_value: int = typing.cast("typing.Any", None) + if "sequence" not in raw or raw["sequence"] is None: + violations.append(Violation(path="sequence", reason="required")) + else: + sequence_value_raw = raw["sequence"] + sequence_value_parsed = _parse_spec_integer( + sequence_value_raw, "sequence", violations + ) + if sequence_value_parsed is not None: + sequence_value = sequence_value_parsed + if sequence_value < 0: + violations.append( + Violation( + path="sequence", + reason=f"must be >= 0, got {sequence_value}", + ) + ) + + batch_index_value: int = typing.cast("typing.Any", None) + if "batch_index" not in raw or raw["batch_index"] is None: + violations.append(Violation(path="batch_index", reason="required")) + else: + batch_index_value_raw = raw["batch_index"] + batch_index_value_parsed = _parse_spec_integer( + batch_index_value_raw, "batch_index", violations + ) + if batch_index_value_parsed is not None: + batch_index_value = batch_index_value_parsed + if batch_index_value < 1: + violations.append( + Violation( + path="batch_index", + reason=f"must be >= 1, got {batch_index_value}", + ) + ) + + payloads_value: list[bytes] | None = None + if "payloads" in raw: + payloads_value_raw = raw["payloads"] + if payloads_value_raw is None: + violations.append( + Violation(path="payloads", reason="explicit null not allowed") + ) + else: + if not isinstance(payloads_value_raw, list): + violations.append( + Violation(path="payloads", reason="expected array") + ) + else: + payloads_value_list: list[bytes] = [] + for payloads_value_index, payloads_value_element in enumerate( + typing.cast("list[typing.Any]", payloads_value_raw) + ): + payloads_value_item_path = f"payloads[{payloads_value_index}]" + payloads_value_item_violation_count = len(violations) + payloads_value_item: bytes = typing.cast("typing.Any", None) + if not isinstance(payloads_value_element, str): + violations.append( + Violation( + path=payloads_value_item_path, + reason="expected string", + ) + ) + else: + payloads_value_item_parsed = _parse_base64( + payloads_value_element, + payloads_value_item_path, + violations, + ) + if payloads_value_item_parsed is not None: + payloads_value_item = payloads_value_item_parsed + if len(violations) == payloads_value_item_violation_count: + payloads_value_list.append(payloads_value_item) + payloads_value = payloads_value_list + + finish_value: bool | None = None + if "finish" in raw: + finish_value_raw = raw["finish"] + if finish_value_raw is None: + violations.append( + Violation(path="finish", reason="explicit null not allowed") + ) + else: + if not isinstance(finish_value_raw, bool): + violations.append( + Violation(path="finish", reason="expected boolean") + ) + else: + finish_value = finish_value_raw + + for key in raw: + if ( + key != "stream" + and key != "producer_id" + and key != "attempt" + and key != "sequence" + and key != "batch_index" + and key != "payloads" + and key != "finish" + ): + violations.append(Violation(path=key, reason="unknown field")) + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return AppendInput( + stream=stream_value, + producer_id=producer_id_value, + attempt=attempt_value, + sequence=sequence_value, + batch_index=batch_index_value, + payloads=payloads_value, + finish=finish_value, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "AppendInput") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + try: + out["stream"] = _StreamRefTransferTypeConverter().to_transfer_type( + value.stream + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, "stream", error) + out["producer_id"] = value.producer_id + if abs(value.attempt) > 9007199254740991: + violations.append( + Violation(path="attempt", reason="exceeds ±(2^53-1) integer cap") + ) + out["attempt"] = value.attempt + if abs(value.sequence) > 9007199254740991: + violations.append( + Violation(path="sequence", reason="exceeds ±(2^53-1) integer cap") + ) + if value.sequence < 0: + violations.append( + Violation(path="sequence", reason=f"must be >= 0, got {value.sequence}") + ) + out["sequence"] = value.sequence + if abs(value.batch_index) > 9007199254740991: + violations.append( + Violation(path="batch_index", reason="exceeds ±(2^53-1) integer cap") + ) + if value.batch_index < 1: + violations.append( + Violation( + path="batch_index", reason=f"must be >= 1, got {value.batch_index}" + ) + ) + out["batch_index"] = value.batch_index + if value.payloads is not None: + out["payloads"] = [_format_base64(element) for element in value.payloads] + if value.finish is not None: + out["finish"] = value.finish + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_AppendInputTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class AppendInput: + """One append call: who is writing, where, and what.""" + + stream: StreamRef + + producer_id: str + """Identifies the writer across its retries, so its attempts can be ordered. Never + empty here: the caller resolved it, from the activity context when it was not given. + """ + + attempt: int + """The generation this producer is writing. A later attempt supersedes an earlier one.""" + + sequence: int + """The producer's sequence of the first record in this batch; the batch is numbered + from it and a finish takes the next number. It has to continue where the previous + batch ended, so the handler and the store agree on every record's position within + the attempt. + """ + + batch_index: int + """Counts this producer's batches from 1. The handler answers a repeat of the last + index with the original's position, and refuses one that skips ahead, one already + behind the last, or one that resumes an attempt it never saw start. + """ + + payloads: list[bytes] | None = None + """The record bodies to write, in order, each a serialized + temporal.api.common.v1.Payload. A payload codec configured on the caller has already + run on them. Empty on a call that only finishes. + """ + + finish: bool | None = None + """Write FINISH for this producer after the batch: it will write nothing more on the + topic. Says nothing about the producer's outcome. + """ + + +class _AppendOutputTransferTypeConverter( + temporalio.converter.TransferTypeConverter["AppendOutput", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["AppendOutput"] + ) -> "AppendOutput": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + cursor_value: str | None = None + if "cursor" in raw: + cursor_value_raw = raw["cursor"] + if cursor_value_raw is None: + violations.append( + Violation(path="cursor", reason="explicit null not allowed") + ) + else: + if not isinstance(cursor_value_raw, str): + violations.append( + Violation(path="cursor", reason="expected string") + ) + else: + cursor_value = cursor_value_raw + + additional_properties: dict[str, typing.Any] = {} + for key in raw: + if key not in _APPEND_OUTPUT_DECLARED: + additional_properties[key] = raw[key] + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return AppendOutput( + cursor=cursor_value, + additional_properties=additional_properties, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "AppendOutput") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + if value.cursor is not None: + out["cursor"] = value.cursor + for key, entry in value.additional_properties.items(): + if key in _APPEND_OUTPUT_DECLARED: + violations.append( + Violation( + path=key, + reason="additional property collides with declared property", + ) + ) + else: + out[key] = entry + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_AppendOutputTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class AppendOutput: + """Where the append landed, when the store can say.""" + + cursor: str | None = None + """Opaque token naming the last record written by this call, or by the original when + the call repeated the last batch. Absent when the store learns positions only at + read time; such a caller positions itself with a latest_only read. + """ + + additional_properties: dict[str, typing.Any] = dataclasses.field( + default_factory=dict + ) + + +class _ReadInputTransferTypeConverter( + temporalio.converter.TransferTypeConverter["ReadInput", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["ReadInput"] + ) -> "ReadInput": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + stream_value: StreamRef = typing.cast("typing.Any", None) + if "stream" not in raw or raw["stream"] is None: + violations.append(Violation(path="stream", reason="required")) + else: + stream_value_raw = raw["stream"] + try: + stream_value = _StreamRefTransferTypeConverter().from_transfer_type( + stream_value_raw, StreamRef + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, "stream", error) + + after_token_value: str | None = None + if "after_token" in raw: + after_token_value_raw = raw["after_token"] + if after_token_value_raw is None: + violations.append( + Violation(path="after_token", reason="explicit null not allowed") + ) + else: + if not isinstance(after_token_value_raw, str): + violations.append( + Violation(path="after_token", reason="expected string") + ) + else: + after_token_value = after_token_value_raw + + max_records_value: int | None = None + if "max_records" in raw: + max_records_value_raw = raw["max_records"] + if max_records_value_raw is None: + violations.append( + Violation(path="max_records", reason="explicit null not allowed") + ) + else: + max_records_value_parsed = _parse_spec_integer( + max_records_value_raw, "max_records", violations + ) + if max_records_value_parsed is not None: + max_records_value = max_records_value_parsed + if max_records_value < 1: + violations.append( + Violation( + path="max_records", + reason=f"must be >= 1, got {max_records_value}", + ) + ) + if max_records_value > 1000: + violations.append( + Violation( + path="max_records", + reason=f"must be <= 1000, got {max_records_value}", + ) + ) + + wait_ms_value: int | None = None + if "wait_ms" in raw: + wait_ms_value_raw = raw["wait_ms"] + if wait_ms_value_raw is None: + violations.append( + Violation(path="wait_ms", reason="explicit null not allowed") + ) + else: + wait_ms_value_parsed = _parse_spec_integer( + wait_ms_value_raw, "wait_ms", violations + ) + if wait_ms_value_parsed is not None: + wait_ms_value = wait_ms_value_parsed + if wait_ms_value < 0: + violations.append( + Violation( + path="wait_ms", + reason=f"must be >= 0, got {wait_ms_value}", + ) + ) + if wait_ms_value > 60000: + violations.append( + Violation( + path="wait_ms", + reason=f"must be <= 60000, got {wait_ms_value}", + ) + ) + + latest_only_value: bool | None = None + if "latest_only" in raw: + latest_only_value_raw = raw["latest_only"] + if latest_only_value_raw is None: + violations.append( + Violation(path="latest_only", reason="explicit null not allowed") + ) + else: + if not isinstance(latest_only_value_raw, bool): + violations.append( + Violation(path="latest_only", reason="expected boolean") + ) + else: + latest_only_value = latest_only_value_raw + + last_n_value: int | None = None + if "last_n" in raw: + last_n_value_raw = raw["last_n"] + if last_n_value_raw is None: + violations.append( + Violation(path="last_n", reason="explicit null not allowed") + ) + else: + last_n_value_parsed = _parse_spec_integer( + last_n_value_raw, "last_n", violations + ) + if last_n_value_parsed is not None: + last_n_value = last_n_value_parsed + if last_n_value < 1: + violations.append( + Violation( + path="last_n", + reason=f"must be >= 1, got {last_n_value}", + ) + ) + + for key in raw: + if ( + key != "stream" + and key != "after_token" + and key != "max_records" + and key != "wait_ms" + and key != "latest_only" + and key != "last_n" + ): + violations.append(Violation(path=key, reason="unknown field")) + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return ReadInput( + stream=stream_value, + after_token=after_token_value, + max_records=max_records_value, + wait_ms=wait_ms_value, + latest_only=latest_only_value, + last_n=last_n_value, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "ReadInput") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + try: + out["stream"] = _StreamRefTransferTypeConverter().to_transfer_type( + value.stream + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, "stream", error) + if value.after_token is not None: + out["after_token"] = value.after_token + if value.max_records is not None: + if abs(value.max_records) > 9007199254740991: + violations.append( + Violation( + path="max_records", reason="exceeds ±(2^53-1) integer cap" + ) + ) + if value.max_records < 1: + violations.append( + Violation( + path="max_records", + reason=f"must be >= 1, got {value.max_records}", + ) + ) + if value.max_records > 1000: + violations.append( + Violation( + path="max_records", + reason=f"must be <= 1000, got {value.max_records}", + ) + ) + out["max_records"] = value.max_records + if value.wait_ms is not None: + if abs(value.wait_ms) > 9007199254740991: + violations.append( + Violation(path="wait_ms", reason="exceeds ±(2^53-1) integer cap") + ) + if value.wait_ms < 0: + violations.append( + Violation( + path="wait_ms", reason=f"must be >= 0, got {value.wait_ms}" + ) + ) + if value.wait_ms > 60000: + violations.append( + Violation( + path="wait_ms", reason=f"must be <= 60000, got {value.wait_ms}" + ) + ) + out["wait_ms"] = value.wait_ms + if value.latest_only is not None: + out["latest_only"] = value.latest_only + if value.last_n is not None: + if abs(value.last_n) > 9007199254740991: + violations.append( + Violation(path="last_n", reason="exceeds ±(2^53-1) integer cap") + ) + if value.last_n < 1: + violations.append( + Violation(path="last_n", reason=f"must be >= 1, got {value.last_n}") + ) + out["last_n"] = value.last_n + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_ReadInputTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class ReadInput: + """One read call: where to resume from and how long to wait.""" + + stream: StreamRef + + after_token: str | None = None + """Opaque cursor from an earlier record or append. The read resumes strictly after the + record it names, so a caller never sees that record twice. Empty starts at the + beginning. The token is produced by whichever store sits behind the endpoint, so a + caller cannot tell which one that is, and a token from another store is refused. + """ + + max_records: int | None = None + """Return at most this many records. Defaults to 100 when omitted.""" + + wait_ms: int | None = None + """Wait at most this long for records to arrive before answering. The handler shortens + it to fit the request deadline. A call that collects nothing answers with no records + and the caller's own token. + """ + + latest_only: bool | None = None + """Answer with the newest position and no records, for a reader that wants to follow + from now. + """ + + last_n: int | None = None + """With no after_token, start at the newest this many records, or at all of them when + the stream holds fewer. Records of every kind count. Refused alongside an + after_token, which is how a read resumes. + """ + + +class _ReadOutputTransferTypeConverter( + temporalio.converter.TransferTypeConverter["ReadOutput", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["ReadOutput"] + ) -> "ReadOutput": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + records_value: list[RecordWire] | None = None + if "records" in raw: + records_value_raw = raw["records"] + if records_value_raw is None: + violations.append( + Violation(path="records", reason="explicit null not allowed") + ) + else: + if not isinstance(records_value_raw, list): + violations.append( + Violation(path="records", reason="expected array") + ) + else: + records_value_list: list[RecordWire] = [] + for records_value_index, records_value_element in enumerate( + typing.cast("list[typing.Any]", records_value_raw) + ): + records_value_item_path = f"records[{records_value_index}]" + records_value_item_violation_count = len(violations) + records_value_item: RecordWire = typing.cast("typing.Any", None) + try: + records_value_item = ( + _RecordWireTransferTypeConverter().from_transfer_type( + records_value_element, RecordWire + ) + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, records_value_item_path, error) + if len(violations) == records_value_item_violation_count: + records_value_list.append(records_value_item) + records_value = records_value_list + + next_token_value: str | None = None + if "next_token" in raw: + next_token_value_raw = raw["next_token"] + if next_token_value_raw is None: + violations.append( + Violation(path="next_token", reason="explicit null not allowed") + ) + else: + if not isinstance(next_token_value_raw, str): + violations.append( + Violation(path="next_token", reason="expected string") + ) + else: + next_token_value = next_token_value_raw + + done_value: bool | None = None + if "done" in raw: + done_value_raw = raw["done"] + if done_value_raw is None: + violations.append( + Violation(path="done", reason="explicit null not allowed") + ) + else: + if not isinstance(done_value_raw, bool): + violations.append(Violation(path="done", reason="expected boolean")) + else: + done_value = done_value_raw + + additional_properties: dict[str, typing.Any] = {} + for key in raw: + if key not in _READ_OUTPUT_DECLARED: + additional_properties[key] = raw[key] + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return ReadOutput( + records=records_value, + next_token=next_token_value, + done=done_value, + additional_properties=additional_properties, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "ReadOutput") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + if value.records is not None: + records_out: list[typing.Any] = [] + for records_index, records_element in enumerate(value.records): + try: + records_out.append( + _RecordWireTransferTypeConverter().to_transfer_type( + records_element + ) + ) + except temporalio.exceptions.ApplicationError as error: + _collect(violations, f"records[{records_index}]", error) + out["records"] = records_out + if value.next_token is not None: + out["next_token"] = value.next_token + if value.done is not None: + out["done"] = value.done + for key, entry in value.additional_properties.items(): + if key in _READ_OUTPUT_DECLARED: + violations.append( + Violation( + path=key, + reason="additional property collides with declared property", + ) + ) + else: + out[key] = entry + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_ReadOutputTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class ReadOutput: + """What one read call answered, and where to resume.""" + + records: list[RecordWire] | None = None + """The records after the caller's token, in stream order.""" + + next_token: str | None = None + """Opaque cursor to pass as after_token on the following call. Echoes the caller's own + token when the call collected nothing. + """ + + done: bool | None = None + """The store ended the read: the owning execution, or its chain, is closed and every + retained record after the caller's token has been delivered. Nothing more will + arrive, so the caller stops. + """ + + additional_properties: dict[str, typing.Any] = dataclasses.field( + default_factory=dict + ) + + +class _RecordWireTransferTypeConverter( + temporalio.converter.TransferTypeConverter["RecordWire", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["RecordWire"] + ) -> "RecordWire": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + token_value: str = typing.cast("typing.Any", None) + if "token" not in raw or raw["token"] is None: + violations.append(Violation(path="token", reason="required")) + else: + token_value_raw = raw["token"] + if not isinstance(token_value_raw, str): + violations.append(Violation(path="token", reason="expected string")) + else: + token_value = token_value_raw + + record_value: bytes = typing.cast("typing.Any", None) + if "record" not in raw or raw["record"] is None: + violations.append(Violation(path="record", reason="required")) + else: + record_value_raw = raw["record"] + if not isinstance(record_value_raw, str): + violations.append(Violation(path="record", reason="expected string")) + else: + record_value_parsed = _parse_base64( + record_value_raw, "record", violations + ) + if record_value_parsed is not None: + record_value = record_value_parsed + + additional_properties: dict[str, typing.Any] = {} + for key in raw: + if key not in _RECORD_WIRE_DECLARED: + additional_properties[key] = raw[key] + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return RecordWire( + token=token_value, + record=record_value, + additional_properties=additional_properties, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "RecordWire") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + out["token"] = value.token + out["record"] = _format_base64(value.record) + for key, entry in value.additional_properties.items(): + if key in _RECORD_WIRE_DECLARED: + violations.append( + Violation( + path=key, + reason="additional property collides with declared property", + ) + ) + else: + out[key] = entry + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_RecordWireTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class RecordWire: + """One record on the wire: its cursor and the record itself.""" + + token: str + """Opaque cursor naming this record. Pass it as after_token to resume just past it.""" + + record: bytes + """The serialized temporal.api.stream.v1.StreamRecord: topic, kind, producer_id, + attempt, sequence and body, as the store holds it. + """ + + additional_properties: dict[str, typing.Any] = dataclasses.field( + default_factory=dict + ) + + +class _StreamRefTransferTypeConverter( + temporalio.converter.TransferTypeConverter["StreamRef", typing.Any] +): + @typing_extensions.override + def from_transfer_type( + self, value: typing.Any, type_hint: type["StreamRef"] + ) -> "StreamRef": + violations: list[Violation] = [] + if not isinstance(value, dict): + raise temporalio.converter.create_payload_validation_error( + [Violation(path="", reason="expected object")] + ) + raw = typing.cast("dict[str, typing.Any]", value) + + kind_value: typing.Literal["workflow", "activity", "standalone"] = typing.cast( + "typing.Any", None + ) + if "kind" not in raw or raw["kind"] is None: + violations.append(Violation(path="kind", reason="required")) + else: + kind_value_raw = raw["kind"] + if not isinstance(kind_value_raw, str): + violations.append(Violation(path="kind", reason="expected string")) + elif kind_value_raw not in ("workflow", "activity", "standalone"): + violations.append( + Violation( + path="kind", + reason=f'must be one of ["workflow", "activity", "standalone"], got {_quote(kind_value_raw)}', + ) + ) + else: + kind_value = kind_value_raw + + workflow_id_value: str | None = None + if "workflow_id" in raw: + workflow_id_value_raw = raw["workflow_id"] + if workflow_id_value_raw is None: + workflow_id_value = None + else: + if not isinstance(workflow_id_value_raw, str): + violations.append( + Violation(path="workflow_id", reason="expected string") + ) + else: + workflow_id_value = workflow_id_value_raw + + run_id_value: str | None = None + if "run_id" in raw: + run_id_value_raw = raw["run_id"] + if run_id_value_raw is None: + run_id_value = None + else: + if not isinstance(run_id_value_raw, str): + violations.append( + Violation(path="run_id", reason="expected string") + ) + else: + run_id_value = run_id_value_raw + + activity_id_value: str | None = None + if "activity_id" in raw: + activity_id_value_raw = raw["activity_id"] + if activity_id_value_raw is None: + activity_id_value = None + else: + if not isinstance(activity_id_value_raw, str): + violations.append( + Violation(path="activity_id", reason="expected string") + ) + else: + activity_id_value = activity_id_value_raw + + stream_id_value: str | None = None + if "stream_id" in raw: + stream_id_value_raw = raw["stream_id"] + if stream_id_value_raw is None: + stream_id_value = None + else: + if not isinstance(stream_id_value_raw, str): + violations.append( + Violation(path="stream_id", reason="expected string") + ) + else: + stream_id_value = stream_id_value_raw + + topic_value: str = typing.cast("typing.Any", None) + if "topic" not in raw or raw["topic"] is None: + violations.append(Violation(path="topic", reason="required")) + else: + topic_value_raw = raw["topic"] + if not isinstance(topic_value_raw, str): + violations.append(Violation(path="topic", reason="expected string")) + else: + topic_value = topic_value_raw + + for key in raw: + if ( + key != "kind" + and key != "workflow_id" + and key != "run_id" + and key != "activity_id" + and key != "stream_id" + and key != "topic" + ): + violations.append(Violation(path=key, reason="unknown field")) + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return StreamRef( + kind=kind_value, + workflow_id=workflow_id_value, + run_id=run_id_value, + activity_id=activity_id_value, + stream_id=stream_id_value, + topic=topic_value, + ) + + @typing_extensions.override + def to_transfer_type(self, value: "StreamRef") -> typing.Any: + violations: list[Violation] = [] + out: dict[str, typing.Any] = {} + if typing.cast("object", value.kind) not in ( + "workflow", + "activity", + "standalone", + ): + violations.append( + Violation( + path="kind", + reason=f'must be one of ["workflow", "activity", "standalone"], got {_quote(value.kind)}', + ) + ) + out["kind"] = value.kind + if value.workflow_id is not None: + out["workflow_id"] = value.workflow_id + if value.run_id is not None: + out["run_id"] = value.run_id + if value.activity_id is not None: + out["activity_id"] = value.activity_id + if value.stream_id is not None: + out["stream_id"] = value.stream_id + out["topic"] = value.topic + if violations: + raise temporalio.converter.create_payload_validation_error(violations) + return out + + +@_transfer_type_convertible(_StreamRefTransferTypeConverter) +@dataclasses.dataclass(slots=True, kw_only=True) +class StreamRef: + """A stream, named by its owner and a topic: what an operation returns to hand a stream + to its caller, and what read and append take in place of an owner spelled out. It + names no cursor and no store, so the same reference is good behind any endpoint that + serves the owner. A member the owner kind does not use is absent or null; null is + how the SDK writes an unset member of its own StreamRef. + """ + + kind: typing.Literal["workflow", "activity", "standalone"] + """What owns the stream. A workflow's streams are keyed by workflow_id and, when + pinned, run_id. An activity's own streams are keyed by activity_id and, when a + workflow scheduled it, that workflow's ids. A standalone stream has its own + stream_id and no execution behind it. + """ + + workflow_id: str | None = None + """The owning workflow, or the workflow that scheduled the owning activity. Required + for a workflow owner. + """ + + run_id: str | None = None + """Pin the owner to this run. Absent follows the execution chain across + continue-as-new, which is what a producer normally wants. + """ + + activity_id: str | None = None + """The owning activity. Required for an activity owner.""" + + stream_id: str | None = None + """The stream's own id. Required for a standalone owner.""" + + topic: str + """The topic on the owner's stream.""" diff --git a/temporalio/streams/providers/_nexus_generated/services.py b/temporalio/streams/providers/_nexus_generated/services.py new file mode 100644 index 000000000..591e31046 --- /dev/null +++ b/temporalio/streams/providers/_nexus_generated/services.py @@ -0,0 +1,44 @@ +# Generated by nexgen v0.2.4. DO NOT EDIT! + +from __future__ import annotations + +from nexusrpc import Operation, service + +from .models import ( + AppendInput, + AppendOutput, + ReadInput, + ReadOutput, +) + + +@service +class TemporalStreams: + """The stream endpoint's two operations. One Temporal-authenticated endpoint hides the + store behind it, so an operator switches storage without touching callers. Reads + hand out batches and an append carries one batch per call, because a Nexus operation + per record costs too much for token streams. A record crosses as the serialized + temporal.api.stream.v1.StreamRecord, the same bytes every store keeps, so a caller + in any language decodes it with the api protos alone. + """ + + append: Operation[ + AppendInput, + AppendOutput, + ] = Operation(name="append") + """Append one batch on the caller's account. The handler writes a batch once: a repeat + of the last batch_index answers with where the original landed, and an index that + skips ahead, one already behind the last, a sequence that does not continue, or a + producer attempt the handler has no state for is refused. Supersession records are + not transported: a reader re-synthesizes them from the attempts it observes. + """ + + read: Operation[ + ReadInput, + ReadOutput, + ] = Operation(name="read") + """Answer with the records after the caller's token, or time out. The call parks until + it has max_records records or wait_ms elapses, whichever comes first, and returns + whatever it collected. It says when the store has ended the read, so the caller can + stop. + """ diff --git a/temporalio/streams/providers/nexus.py b/temporalio/streams/providers/nexus.py new file mode 100644 index 000000000..1cc60d236 --- /dev/null +++ b/temporalio/streams/providers/nexus.py @@ -0,0 +1,1242 @@ +"""The Nexus front: the outside surface behind one Temporal-authenticated endpoint. + +Two halves in one module. :class:`NexusStreams` is a provider that implements +the outside half only: it hands out handles that append and read through the +endpoint, and it has no workflow half, because a workflow's publishes and +reads ride the Workflow Task and cannot cross an RPC. +:class:`TemporalStreamsHandler` runs in a worker next to any storage provider +and serves two sync operations, ``append`` and ``read``, by fronting that +provider's own handles, so the store behind the endpoint is invisible to +callers and an operator switches it without touching them. + +The wire types and the service definition come from +``temporal_streams.nexusrpc.yaml``, so any language nexgen targets can be +handed the same contract. A record crosses as the serialized +``temporal.api.stream.v1.StreamRecord``, the same bytes every store keeps, so +a caller in another language decodes it with the api protos alone. What stays +hand-written here is what the generator cannot express yet: the handler's +dedupe and long-poll collect loop, and the caller's batching, cursor handling +and codec. + +Reads hand out batches and an append carries one batch per call, because a +Nexus operation per record costs too much for token streams. Record bodies +cross the handler untouched: it reads the store as raw payloads and forwards +them as they are, so a codec that changes the payload encoding survives the +hop and the handler's worker never needs the key. Supersession records are +not transported: the caller's reader re-synthesizes them from the attempts it +observes, which is the policy module's job on every provider. Cursors pass +through opaque, so the caller cannot tell which store produced them; a +foreign one is refused by the store behind the endpoint and reaches the +caller on the first read. + +Both operations address a stream by a reference, ``StreamRef`` on the wire: +the owner (a workflow, an activity, or a standalone stream), the ids that +name it, and the topic. The handler maps the reference onto the store's own +accessor for that owner and refuses an owner the store cannot host with +:class:`temporalio.streams.StreamUnsupportedError`, so a reference is good +behind any endpoint that serves the owner. The same reference is what an +operation of the application's own returns to hand a stream to its caller; +:meth:`NexusStreamHandle.ref` makes one and ``client.get_stream_handle(ref)`` +opens it. + +The handler keeps one parked read per stream reference and serves +consecutive calls from it, so an idle caller does not leave one abandoned +long poll on the store per call. A call whose token does not match the parked +position replaces the subscription, and an idle one is released after a +minute; that is the residual cost. + +Two stated prototype limits. Append deduplication lives in handler memory, by +batch index per producer attempt, so a handler that has no state for a +producer attempt refuses to continue it rather than starting a fresh delegate +whose numbering the store would drop as a repeat; the caller opens a new +attempt, which readers report as a supersession. And any caller the endpoint +admits may touch any owner's streams in the namespace; the endpoint's own +authorization is the boundary. + +A failure from the endpoint reaches the caller as the +:class:`temporalio.streams.StreamError` the store raised, when the handler +named one, and as :class:`temporalio.service.RPCError` otherwise, never as an +HTTP or urllib exception. +""" + +from __future__ import annotations + +import asyncio +import http.client +import json +import logging +import time +import urllib.error +import urllib.request +import weakref +from collections import OrderedDict +from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Any, Generic, NoReturn, TypeVar, cast + +import nexusrpc +import nexusrpc.handler +from google.protobuf.message import DecodeError + +import temporalio.client +import temporalio.converter +from temporalio.api.common.v1 import Payload +from temporalio.api.operatorservice.v1 import ListNexusEndpointsRequest +from temporalio.client import Client, ClientConfig +from temporalio.common import RawValue +from temporalio.service import ConnectConfig, RPCError, RPCStatusCode, ServiceClient +from temporalio.streams._errors import ( + StreamClosedError, + StreamCursorError, + StreamError, + StreamNotFoundError, + StreamProducerError, + StreamUnsupportedError, +) +from temporalio.streams._provider import StreamHandle, StreamProducer, StreamProvider +from temporalio.streams._record import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamRecord, + check_read_start, +) +from temporalio.streams._ref import StreamOwnerKind, StreamRef +from temporalio.streams._topic import StreamTopic, resolve_topic +from temporalio.streams._wire import ( + RecordDecoder, + WireRecord, + producer_identity, + to_wire, +) +from temporalio.streams.providers._nexus_generated import ( + AppendInput, + AppendOutput, + ReadInput, + ReadOutput, + RecordWire, + TemporalStreams, +) +from temporalio.streams.providers._nexus_generated import StreamRef as WireStreamRef + +__all__ = [ + "NexusProducer", + "NexusStreamHandle", + "NexusStreams", + "TemporalStreamsHandler", + "WireStreamRef", +] + +T = TypeVar("T") + +# The reference's identity on the handler, topic included, so one parked read +# and one producer state key on exactly what the wire names. +_StreamKey = tuple[str, str, str, str, str, str] +_ProducerKey = tuple[_StreamKey, str, int] + +_WORKFLOW_SIDE_ERROR = ( + "the nexus provider is an outside transport; a worker publishes and reads " + "through a storage provider, so give the worker one of those" +) + +# The contract leaves these unset, so the handler is the one place that says +# what an omitted read bound means. The wait is long because the handler +# shortens it to the request deadline anyway, and a shorter default only +# means more round trips on an idle stream. +_DEFAULT_MAX_RECORDS = 100 +_DEFAULT_READ_WAIT = timedelta(seconds=30) +_DEADLINE_MARGIN = timedelta(milliseconds=500) +_DEFAULT_MAX_PRODUCERS = 10_000 +_DEFAULT_SUBSCRIPTION_IDLE = timedelta(seconds=60) +_QUEUE_DEPTH = 1000 +_APPEND_TIMEOUT = timedelta(seconds=30) +_ROUND_TRIP_MARGIN = timedelta(seconds=5) +# What ReadInput accepts in temporal_streams.nexusrpc.yaml. Repeated here so +# a caller's mistake is refused where it was made rather than on the wire. +_MIN_READ_WAIT = timedelta(0) +_MAX_READ_WAIT = timedelta(milliseconds=60_000) +_MIN_MAX_RECORDS = 1 +_MAX_MAX_RECORDS = 1000 + +_definition = nexusrpc.get_service_definition(TemporalStreams) +if _definition is None: # pragma: no cover + raise RuntimeError("the generated stream service carries no nexus definition") +# Both halves read the wire names off the generated definition, so a rename in +# the contract cannot leave the caller posting to a path the handler no longer +# serves. +_SERVICE_NAME = _definition.name +_OPERATION_NAMES = { + operation.method_name: operation.name + for operation in _definition.operation_definitions.values() +} +_APPEND_OPERATION = _OPERATION_NAMES["append"] +_READ_OPERATION = _OPERATION_NAMES["read"] + +# The stream conditions that cross the endpoint under their own name, so the +# caller raises the class the store raised. +_STREAM_ERRORS: dict[str, type[StreamError]] = { + cls.__name__: cls + for cls in ( + StreamError, + StreamNotFoundError, + StreamCursorError, + StreamProducerError, + StreamClosedError, + StreamUnsupportedError, + ) +} +_HTTP_TO_RPC = { + 400: RPCStatusCode.INVALID_ARGUMENT, + 401: RPCStatusCode.UNAUTHENTICATED, + 403: RPCStatusCode.PERMISSION_DENIED, + 404: RPCStatusCode.NOT_FOUND, + 408: RPCStatusCode.DEADLINE_EXCEEDED, + 409: RPCStatusCode.ALREADY_EXISTS, + 412: RPCStatusCode.FAILED_PRECONDITION, + 429: RPCStatusCode.RESOURCE_EXHAUSTED, + 500: RPCStatusCode.INTERNAL, + 501: RPCStatusCode.UNIMPLEMENTED, + 503: RPCStatusCode.UNAVAILABLE, + 504: RPCStatusCode.DEADLINE_EXCEEDED, +} + +_OutputT = TypeVar("_OutputT") + +logger = logging.getLogger(__name__) + + +def _require_topic(topic: str) -> None: + if not topic: + raise ValueError("topic must not be empty") + + +def _check_ref(ref: WireStreamRef) -> None: + """Refuse a reference whose owner lacks the id that names it.""" + _require_topic(ref.topic) + if ref.kind == "workflow" and not ref.workflow_id: + raise ValueError("a stream reference with a workflow owner needs workflow_id") + if ref.kind == "activity" and not ref.activity_id: + raise ValueError("a stream reference with an activity owner needs activity_id") + if ref.kind == "standalone" and not ref.stream_id: + raise ValueError("a stream reference with a standalone owner needs stream_id") + + +def _stream_key(ref: WireStreamRef) -> _StreamKey: + return ( + ref.kind, + ref.workflow_id or "", + ref.run_id or "", + ref.activity_id or "", + ref.stream_id or "", + ref.topic, + ) + + +@dataclass(frozen=True) +class _Address: + """The owner a handle is bound to; a topic completes it into a reference.""" + + kind: StreamOwnerKind + workflow_id: str | None = None + run_id: str | None = None + activity_id: str | None = None + stream_id: str | None = None + + def wire(self, topic: str) -> WireStreamRef: + return WireStreamRef( + kind=self.kind, + workflow_id=self.workflow_id, + run_id=self.run_id, + activity_id=self.activity_id, + stream_id=self.stream_id, + topic=topic, + ) + + def ref(self, topic: str) -> StreamRef: + # The same members under the SDK's own name, so what an operation + # returns is what client.get_stream_handle(ref) opens. + return StreamRef( + self.kind, + topic, + workflow_id=self.workflow_id, + run_id=self.run_id, + activity_id=self.activity_id, + stream_id=self.stream_id, + ) + + +def _handler_error(error: StreamError) -> nexusrpc.HandlerError: + """Carry a stream condition across the endpoint under its own class name.""" + kind = ( + nexusrpc.HandlerErrorType.NOT_FOUND + if isinstance(error, StreamNotFoundError) + else nexusrpc.HandlerErrorType.BAD_REQUEST + ) + return nexusrpc.HandlerError( + f"{type(error).__name__}: {error}", type=kind, retryable_override=False + ) + + +def _bad_request(message: str) -> nexusrpc.HandlerError: + return nexusrpc.HandlerError( + message, type=nexusrpc.HandlerErrorType.BAD_REQUEST, retryable_override=False + ) + + +@dataclass +class _Subscription: + """One parked read on the store, shared by consecutive calls.""" + + records: asyncio.Queue[StreamRecord[Any]] + position: str + last_used: float + last_n: int | None = None + """The newest-N start this subscription was opened with, when it had no token.""" + pump: asyncio.Task[None] | None = None + failure: BaseException | None = None + done: bool = False + + +@dataclass +class _ProducerState: + delegate: StreamProducer + batch_index: int + sequence: int + last_cursor: str | None + finished: bool = False + + +@nexusrpc.handler.service_handler(service=TemporalStreams) +class TemporalStreamsHandler: + """Serves the stream endpoint by fronting one storage provider's handles. + + The provider is an explicit instance the application constructed, so a + caller and a handler can coexist in one process, and so an operator can + run handlers for a new store next to handlers for the old one. + """ + + def __init__( + self, + provider: StreamProvider, + client: Client | None, + *, + max_producers: int = _DEFAULT_MAX_PRODUCERS, + subscription_idle: timedelta = _DEFAULT_SUBSCRIPTION_IDLE, + ) -> None: + """Serve the endpoint out of ``provider``'s handles, opened with ``client``. + + ``max_producers`` bounds the dedupe state kept per producer attempt; + the oldest is evicted, and a producer evicted mid-life gets a clear + refusal on its next batch rather than a silent drop. + ``subscription_idle`` is how long a parked read outlives its last + caller. + """ + self._provider = provider + self._client = client + self._producers: OrderedDict[_ProducerKey, _ProducerState] = OrderedDict() + self._max_producers = max_producers + self._subscriptions: dict[_StreamKey, _Subscription] = {} + # Held weakly, and by the call that is using one for as long as it + # runs. A strong map would keep an entry per address ever read or + # appended to, and nothing would ever reach it again. + self._read_locks: weakref.WeakValueDictionary[_StreamKey, asyncio.Lock] = ( + weakref.WeakValueDictionary() + ) + self._append_locks: weakref.WeakValueDictionary[_ProducerKey, asyncio.Lock] = ( + weakref.WeakValueDictionary() + ) + self._subscription_idle = subscription_idle.total_seconds() + + def _stream(self, ref: WireStreamRef) -> StreamHandle: + """The store's handle for the owner ``ref`` names. + + Raises: + StreamUnsupportedError: The store has no accessor for that owner. + """ + # A storage provider wants a client; the memory provider, which the + # tests front, accepts none, so the cast only lies where it is unread. + client = cast(Client, self._client) + if ref.kind == "workflow": + return self._provider.get_stream_handle( + client, cast(str, ref.workflow_id), run_id=ref.run_id or None + ) + if ref.kind == "activity": + return self._provider.get_activity_stream_handle( + client, + cast(str, ref.activity_id), + workflow_id=ref.workflow_id or None, + run_id=ref.run_id or None, + ) + return self._provider.get_standalone_stream_handle( + client, cast(str, ref.stream_id) + ) + + @nexusrpc.handler.sync_operation + async def append( + self, _ctx: nexusrpc.handler.StartOperationContext, input: AppendInput + ) -> AppendOutput: + """Append on the caller's account, writing a repeated batch once. + + Raises: + nexusrpc.HandlerError: ``BAD_REQUEST`` when the batch index or + sequence does not continue the producer attempt, when the + attempt is one this handler has no state for, when a payload + is not a serialized ``Payload``, or when the store cannot host + the reference's owner; ``NOT_FOUND`` when the store has no + such owner. + """ + try: + return await self._append(input) + except StreamError as error: + raise _handler_error(error) from error + except ValueError as error: + raise _bad_request(str(error)) from error + + async def _append(self, input: AppendInput) -> AppendOutput: + _check_ref(input.stream) + key: _ProducerKey = ( + _stream_key(input.stream), + input.producer_id, + input.attempt, + ) + # The repeat check and the commit have an await between them, so two + # in-flight copies of one batch would both pass the check and both + # reach the store. The read path serialises per address for the same + # reason. + lock = self._append_locks.setdefault(key, asyncio.Lock()) + async with lock: + return await self._append_locked(input, key) + + async def _append_locked( + self, input: AppendInput, key: _ProducerKey + ) -> AppendOutput: + state = self._producers.get(key) + if state is None: + if input.batch_index != 1 or input.sequence != 0: + # A fresh delegate would restart its numbering at zero, and + # the store would drop that as a repeat of the first batch. + # Refusing here turns a silent loss into a failure the caller + # can act on by opening a new attempt. + raise StreamProducerError( + f"producer {input.producer_id!r} attempt {input.attempt} resumed " + f"at batch {input.batch_index} on a handler that has no state for " + "it; open a new attempt" + ) + delegate = self._stream(input.stream).producer( + topic=input.stream.topic, + producer_id=input.producer_id, + attempt=input.attempt, + ) + state = _ProducerState(delegate, 0, 0, None) + elif input.batch_index == state.batch_index: + # The last batch again: the store already holds it, and where it + # landed is the answer the original got. + self._producers.move_to_end(key) + return AppendOutput(cursor=state.last_cursor) + elif input.batch_index < state.batch_index: + raise StreamProducerError( + f"batch {input.batch_index} was written and is no longer the last " + f"for producer {input.producer_id!r} attempt {input.attempt}" + ) + elif input.batch_index != state.batch_index + 1: + raise StreamProducerError( + f"batch {input.batch_index} skips ahead of {state.batch_index} for " + f"producer {input.producer_id!r} attempt {input.attempt}" + ) + elif state.finished: + raise StreamProducerError( + f"producer {input.producer_id!r} attempt {input.attempt} already " + "finished this topic; open a new attempt" + ) + if input.sequence != state.sequence: + raise StreamProducerError( + f"sequence {input.sequence} does not continue at {state.sequence} for " + f"producer {input.producer_id!r} attempt {input.attempt}" + ) + payloads = self._payloads(input.payloads) + cursor: str | None = None + if payloads: + appended = await state.delegate.append(*(RawValue(p) for p in payloads)) + cursor = appended.token if appended is not None else None + if input.finish: + await state.delegate.finish() + # Recorded only once the store accepted the batch, so a failed append + # is not mistaken for a repeat when the caller retries it. A finish is + # recorded the same way rather than dropping the state: a lost + # response would otherwise leave the caller with a batch it can + # neither repeat nor abandon. + state.batch_index = input.batch_index + state.sequence += len(payloads) + state.last_cursor = cursor + state.finished = state.finished or bool(input.finish) + self._producers[key] = state + self._producers.move_to_end(key) + while len(self._producers) > self._max_producers: + self._producers.popitem(last=False) + return AppendOutput(cursor=cursor) + + @staticmethod + def _payloads(encoded: list[bytes] | None) -> list[Payload]: + payloads = [] + for index, raw in enumerate(encoded or []): + try: + payloads.append(Payload.FromString(raw)) + except DecodeError as error: + raise _bad_request( + f"payloads[{index}] is not a serialized Temporal Payload: {error}" + ) from error + return payloads + + @nexusrpc.handler.sync_operation + async def read( + self, ctx: nexusrpc.handler.StartOperationContext, input: ReadInput + ) -> ReadOutput: + """Answer with the records after the caller's token, or time out. + + Raises: + nexusrpc.HandlerError: ``BAD_REQUEST`` when the store refuses the + token or cannot host the reference's owner, ``NOT_FOUND`` when + it has no such owner. + """ + try: + return await self._read(ctx, input) + except StreamError as error: + raise _handler_error(error) from error + except ValueError as error: + raise _bad_request(str(error)) from error + + async def _read( + self, ctx: nexusrpc.handler.StartOperationContext, input: ReadInput + ) -> ReadOutput: + _check_ref(input.stream) + topic = input.stream.topic + stream = self._stream(input.stream) + if input.latest_only: + return ReadOutput(next_token=(await stream.latest(topic=topic)).token) + max_records = ( + _DEFAULT_MAX_RECORDS if input.max_records is None else input.max_records + ) + wait = ( + _DEFAULT_READ_WAIT.total_seconds() + if input.wait_ms is None + else input.wait_ms / 1000 + ) + if ctx.request_deadline is not None: + # The server answers the caller with a timeout at the deadline + # whatever this call does, so collecting past it only holds the + # slot for an answer nobody receives. + deadline = ctx.request_deadline + if deadline.tzinfo is None: + deadline = deadline.replace(tzinfo=timezone.utc) + left = deadline - datetime.now(timezone.utc) - _DEADLINE_MARGIN + wait = max(0.0, min(wait, left.total_seconds())) + after = input.after_token or "" + if after and input.last_n is not None: + raise ValueError( + "pass either after_token or last_n, not both: a token resumes a " + "read and last_n starts one" + ) + key = _stream_key(input.stream) + await self._expire_subscriptions() + lock = self._read_locks.setdefault(key, asyncio.Lock()) + async with lock: + subscription = self._subscriptions.get(key) + if ( + subscription is None + or subscription.position != after + or (not after and subscription.last_n != input.last_n) + ): + if subscription is not None: + await self._drop(key) + subscription = self._subscribe(key, stream, topic, after, input.last_n) + try: + records, next_token = await self._drain(subscription, max_records, wait) + except Exception: + # The pump failed. Keeping the subscription would answer the + # caller's retry from the same token with end of stream, so a + # transient read failure would read as the stream ending. + await self._drop(key) + raise + subscription.position = next_token + subscription.last_used = time.monotonic() + done = ( + subscription.done + and subscription.records.empty() + and subscription.failure is None + ) + if done: + # The store ended the read: the run or chain is closed and the + # tail has been handed over. Nothing more will arrive on it. + await self._drop(key) + return ReadOutput(records=records, next_token=next_token, done=done) + + def _subscribe( + self, + key: _StreamKey, + stream: StreamHandle, + topic: str, + after: str, + last_n: int | None, + ) -> _Subscription: + # Raw payloads: the handler forwards what the store holds without + # decoding it, so an encoding only the caller's codec understands + # passes through untouched. A foreign token is refused right here. + source = ( + stream.read(topic=topic, last=last_n, result_type=RawValue) + if last_n is not None + else stream.read( + topic=topic, + after=Cursor(after) if after else BEGINNING, + result_type=RawValue, + ) + ) + subscription = _Subscription( + records=asyncio.Queue(maxsize=_QUEUE_DEPTH), + position=after, + last_used=time.monotonic(), + last_n=last_n, + ) + + async def pump() -> None: + try: + async for record in source: + if record.kind is RecordKind.SUPERSEDED: + continue + await subscription.records.put(record) + except asyncio.CancelledError: + raise + except Exception as error: + subscription.failure = error + finally: + subscription.done = True + await source.aclose() + + subscription.pump = asyncio.create_task(pump()) + self._subscriptions[key] = subscription + return subscription + + async def _drain( + self, subscription: _Subscription, max_records: int, wait: float + ) -> tuple[list[RecordWire], str]: + records: list[RecordWire] = [] + next_token = subscription.position + queue = subscription.records + deadline = time.monotonic() + wait + while len(records) < max_records: + if queue.empty(): + remaining = deadline - time.monotonic() + if records or remaining <= 0: + break + if subscription.done: + if subscription.failure is None: + break + # Honour the wait even though nothing can arrive, so a + # caller polling a failed stream does not spin on the + # endpoint. + await asyncio.sleep(remaining) + break + try: + record = await asyncio.wait_for(queue.get(), remaining) + except asyncio.TimeoutError: + break + else: + record = queue.get_nowait() + records.append(self._wire(record)) + next_token = record.cursor.token + if not records and subscription.failure is not None: + # Left on the subscription rather than cleared: the caller is + # answered with the failure and the subscription is dropped, so + # there is nothing for a second reader of it to be misled by. + raise subscription.failure + return records, next_token + + @staticmethod + def _wire(record: StreamRecord[Any]) -> RecordWire: + # The record read as RawValue carries the stored body untouched; the + # default converter passes a RawValue through, so no encoding happens + # here. + wire = to_wire( + temporalio.converter.DataConverter.default.payload_converter, + topic=record.topic, + kind=record.kind, + value=record.value, + producer_id=record.producer_id, + attempt=record.attempt, + sequence=record.sequence, + ) + return RecordWire(token=record.cursor.token, record=wire.SerializeToString()) + + async def close(self) -> None: + """Release every parked read. + + Call it when the worker hosting this handler stops, so subscriptions + waiting on the store do not outlive it. + """ + for key in list(self._subscriptions): + await self._drop(key) + + async def _drop(self, key: _StreamKey) -> None: + subscription = self._subscriptions.pop(key, None) + if subscription is None or subscription.pump is None: + return + subscription.pump.cancel() + try: + await subscription.pump + except (asyncio.CancelledError, Exception): + pass + + async def _expire_subscriptions(self) -> None: + cutoff = time.monotonic() - self._subscription_idle + for key in [k for k, s in self._subscriptions.items() if s.last_used < cutoff]: + lock = self._read_locks.get(key) + if lock is not None and lock.locked(): + continue + await self._drop(key) + + +class _EndpointFailure(Exception): + """What the transport saw: the endpoint's answer, or the reason it gave none.""" + + def __init__(self, detail: str, status: int | None) -> None: + super().__init__(detail) + self.detail = detail + self.status = status + + +def _translate(failure: _EndpointFailure) -> Exception: + """The exception a caller raises for an endpoint failure. + + The handler names a stream condition by its class in the failure message, + so the same class is raised here; anything else is an ``RPCError`` with + the status the HTTP answer maps to, or ``UNAVAILABLE`` when there was none. + """ + message = failure.detail + try: + parsed = json.loads(message) + if isinstance(parsed, dict) and isinstance(parsed.get("message"), str): + message = parsed["message"] + except ValueError: + pass + if failure.status is not None: + for name, cls in _STREAM_ERRORS.items(): + prefix = f"{name}: " + if message.startswith(prefix): + return cls(message[len(prefix) :]) + status = _HTTP_TO_RPC.get(failure.status, RPCStatusCode.UNKNOWN) + else: + status = RPCStatusCode.UNAVAILABLE + return RPCError(message, status, b"") + + +def _post( + url: str, body: bytes, headers: Mapping[str, str], timeout: timedelta +) -> bytes: + timeout_ms = int(timeout.total_seconds() * 1000) + # Everything is inside the try, because urllib wraps only the request in + # URLError: a bad url raises from Request(), and a socket timeout or a + # dropped connection raises from getresponse() and read() as + # builtins.TimeoutError or an http.client exception. On 3.11 and later + # TimeoutError is asyncio.TimeoutError, so letting one out would be + # indistinguishable from the caller's own wait_for expiring. + try: + request = urllib.request.Request( + url, + data=body, + headers={ + "Content-Type": "application/json", + # Tells the server how long the handler may park, so it does + # not time the call out ahead of a wait the caller asked for. + "Request-Timeout": f"{timeout_ms}ms", + **headers, + }, + method="POST", + ) + with urllib.request.urlopen( + request, timeout=(timeout + _ROUND_TRIP_MARGIN).total_seconds() + ) as response: + return response.read() + except urllib.error.HTTPError as error: + raise _EndpointFailure( + error.read().decode(errors="replace"), error.code + ) from error + except urllib.error.URLError as error: + raise _EndpointFailure( + f"stream endpoint unreachable at {url}: {error.reason}", None + ) from error + except (TimeoutError, http.client.HTTPException, OSError, ValueError) as error: + raise _EndpointFailure( + f"stream endpoint at {url} did not answer: {error!r}", None + ) from error + + +class _Front: + """Everything a handle needs to reach one endpoint.""" + + def __init__(self, streams: NexusStreams, client: Client | None) -> None: + self._streams = streams + self._client = client + data_converter = streams._data_converter + if data_converter is None: + data_converter = ( + client.data_converter + if client is not None + else temporalio.converter.DataConverter.default + ) + self.converter = data_converter.payload_converter + self.codec = data_converter.payload_codec + self.read_wait = streams._read_wait + self.max_records = streams._max_records + + async def _base_url(self) -> str: + endpoint_id = await self._streams._endpoint_id(self._client) + return ( + f"{self._streams._http_address}/nexus/endpoints/{endpoint_id}" + f"/services/{_SERVICE_NAME}" + ) + + async def invoke( + self, + operation: str, + request: Any, + output: type[_OutputT], + *, + timeout: timedelta, + ) -> _OutputT: + # The contract types carry their own JSON encoding, so the raw caller + # and the worker serving the operation agree on the body without + # either of them spelling the fields out. That encoding is a transfer + # type hook, which only the internal converter applies. + contract = ( + temporalio.converter.DataConverter.default._get_internal_payload_converter() + ) + url = f"{await self._base_url()}/{operation}" + try: + raw = await asyncio.to_thread( + _post, + url, + contract.to_payloads([request])[0].data, + self._streams._headers, + timeout, + ) + except _EndpointFailure as failure: + raise _translate(failure) from failure + payload = Payload(metadata={"encoding": b"json/plain"}, data=raw or b"{}") + return contract.from_payloads([payload], [output])[0] + + +class NexusProducer(Generic[T]): + """Appends through the stream endpoint; the store behind it does the rest. + + Calls are serialized and the batch index is committed only once the + endpoint answered. A batch whose call raised stays pending and is sent + again under the same index, either when the caller retries the same + values or ahead of whatever the caller sends next, so an ambiguous + failure writes the batch once and loses nothing. + """ + + def __init__( + self, + front: _Front, + stream: WireStreamRef, + producer_id: str, + attempt: int, + ) -> None: + """Bind this producer to the stream ``stream`` names behind the endpoint.""" + self._front = front + self._stream = stream + self._producer_id = producer_id + self._attempt = attempt + self._batch_index = 0 + self._sequence = 0 + self._last: Cursor | None = BEGINNING + self._pending: tuple[AppendInput, tuple[bytes, ...], bool] | None = None + self._lock = asyncio.Lock() + + @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 + + async def append(self, *values: T) -> Cursor | None: + """Append ``values`` through the endpoint and return where the last one landed. + + A repeat answers with the original's position and an empty call with + the last one; ``None`` when the store behind the endpoint learns + positions only at read time. + """ + if not values: + return self._last + payloads = self._front.converter.to_payloads(list(values)) + answer = await self._call(payloads, finish=False) + self._last = Cursor(answer.cursor) if answer.cursor else None + return self._last + + async def finish(self) -> None: + """Write ``FINISH`` for this producer on this topic.""" + await self._call([], finish=True) + + async def _call(self, payloads: list[Payload], *, finish: bool) -> AppendOutput: + # Identity is taken before the codec runs: a codec may encrypt with a + # fresh nonce each time, and the retry has to be recognised anyway. + identity = tuple(payload.SerializeToString() for payload in payloads) + if self._front.codec is not None: + # The endpoint is the edge of this process, so a configured codec + # runs here rather than at the handler: whoever hosts the endpoint + # never holds the plaintext. + payloads = await self._front.codec.encode(payloads) + encoded = [payload.SerializeToString() for payload in payloads] + async with self._lock: + if self._pending is not None and self._pending[1:] != (identity, finish): + # The caller moved on from a batch the endpoint may or may + # not have applied. It goes first under its own index, so the + # handler drops it if it landed and writes it if not, and the + # new batch takes the index after it. + await self._send(*self._pending) + elif self._pending is not None: + return await self._send(*self._pending) + request = AppendInput( + stream=self._stream, + producer_id=self._producer_id, + attempt=self._attempt, + sequence=self._sequence, + batch_index=self._batch_index + 1, + payloads=encoded, + finish=finish, + ) + return await self._send(request, identity, finish) + + async def _send( + self, request: AppendInput, identity: tuple[bytes, ...], finish: bool + ) -> AppendOutput: + self._pending = (request, identity, finish) + answer = await self._front.invoke( + _APPEND_OPERATION, request, AppendOutput, timeout=_APPEND_TIMEOUT + ) + self._batch_index = request.batch_index + self._sequence = request.sequence + len(request.payloads or []) + int(finish) + self._pending = None + return answer + + +class NexusStreamHandle: + """One owner's stream through the endpoint, re-synthesizing supersession.""" + + def __init__(self, front: _Front, address: _Address) -> None: + """Address the owner ``address`` names; each call completes it with a topic.""" + self._front = front + self._address = address + + 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, one endpoint batch at a time. + + ``BEGINNING`` and ``last=`` are resolved by the store behind the + endpoint on the first call. ``END`` is the endpoint's newest position, + asked for with a ``latest_only`` call when the read starts, and the + read resumes after it, so it yields only what is appended from then on. + + The token is opaque here, so a cursor from another store is refused by + the store behind the endpoint and raises + :class:`temporalio.streams.StreamCursorError` on the first iteration. + + ``aclose()`` on the result releases nothing at the endpoint at once. + The contract carries no unsubscribe operation, so the handler cannot + be told; it reclaims the parked read when it has gone idle, which is + a minute by default. Until then the subscription and its long poll on + the store stay, and a caller that stops and starts many reads on one + topic should expect that lag rather than an immediate release. + """ + check_read_start(after, last) + topic, result_type = resolve_topic(topic, result_type) + return self._read(topic, after, last, result_type) + + async def _read( + self, topic: str, after: Cursor, last: int | None, result_type: type | None + ) -> AsyncGenerator[StreamRecord[Any], None]: + if after == END: + after = await self.latest(topic=topic) + # Where last= lands is the store's to say, so BEGINNING stands in for + # the position before it: a resume from it may repeat a record, where + # anything later could skip one. + decoder = RecordDecoder( + self._front.converter, result_type, after=after, warn=logger.warning + ) + token = after.token + last_n = last + stream = self._address.wire(topic) + while True: + answer = await self._front.invoke( + _READ_OPERATION, + ReadInput( + stream=stream, + after_token=token, + last_n=last_n, + max_records=self._front.max_records, + wait_ms=int(self._front.read_wait.total_seconds() * 1000), + ), + ReadOutput, + timeout=self._front.read_wait + _ROUND_TRIP_MARGIN, + ) + for wire in answer.records or []: + cursor = Cursor(wire.token) + try: + record = WireRecord.FromString(wire.record) + except DecodeError as error: + # Same answer as every other reader: skip and say so. + logger.warning("skipping stream record at %s: %s", cursor, error) + continue + if self._front.codec is not None and record.HasField("body"): + record.body.CopyFrom( + (await self._front.codec.decode([record.body]))[0] + ) + for out in decoder.decode(cursor, record): + yield out + token = answer.next_token or token + if token: + last_n = None + if answer.done: + return + + async def latest(self, *, topic: str | StreamTopic[Any] | None = None) -> Cursor: + """The newest position on ``topic`` behind the endpoint, for following from now.""" + topic, _ = resolve_topic(topic) + answer = await self._front.invoke( + _READ_OPERATION, + ReadInput(stream=self._address.wire(topic), latest_only=True), + ReadOutput, + timeout=_APPEND_TIMEOUT, + ) + token = answer.next_token or "" + return Cursor(token) if token else BEGINNING + + def producer( + self, + *, + topic: str | StreamTopic[Any] | None = None, + producer_id: str = "", + attempt: int = 0, + ) -> NexusProducer[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 NexusProducer( + self._front, self._address.wire(topic), producer_id, attempt + ) + + def ref(self, *, topic: str | StreamTopic[Any] | None = None) -> StreamRef: + """A :class:`temporalio.streams.StreamRef` to ``topic`` of this owner. + + Plain data naming the owner as this handle addresses it and the + topic, with no cursor and no endpoint, so an operation can return it + and whoever receives it opens it with ``client.get_stream_handle(ref)`` + on the front or on the store itself. + """ + name, _ = resolve_topic(topic) + return self._address.ref(name) + + async def close(self) -> None: + """Refused: the contract carries no operation that seals a stream. + + Raises: + ValueError: The handle is on a workflow's or an activity's stream, + which ends with its owner. + StreamUnsupportedError: The handle is on a standalone stream; seal + it on the store's own handle. + """ + if self._address.kind != "standalone": + raise ValueError( + "only a standalone stream can be closed; this handle is on an owner's " + "stream, which ends when the owner does" + ) + raise StreamUnsupportedError( + "the stream endpoint has no operation that seals a standalone stream; " + "close it on the store's own handle" + ) + + +class NexusStreams(StreamProvider, temporalio.client.Plugin): + """The outside half of a provider, over one Nexus endpoint. + + A client plugin but not a worker plugin: a workflow publishes and reads + through the storage provider its worker was given, and this front only + serves code outside a workflow, so ``Client.connect(plugins=[front])`` + makes ``client.get_stream_handle()`` go through the endpoint while a + worker built from that client is left without a provider. The endpoint + hides which store sits behind it. + """ + + def __init__( + self, + *, + endpoint: str, + http_address: str = "http://127.0.0.1:7243", + headers: Mapping[str, str] | None = None, + read_wait: timedelta = _DEFAULT_READ_WAIT, + max_records: int = _DEFAULT_MAX_RECORDS, + data_converter: temporalio.converter.DataConverter | None = None, + ) -> None: + """Point the front at one endpoint. + + Args: + endpoint: The Nexus endpoint's name, resolved to its id through + the operator service of the client a handle is opened with. + When a handle is opened without a client it has to be the id, + because the HTTP ingress dispatches by id. + http_address: The server's Nexus HTTP ingress. + headers: Sent on every request; where an authorization header + belongs. + read_wait: How long the handler may park a read before answering + with what it has. + max_records: The most records one read answer carries. + data_converter: Overrides the client's converter for record + bodies. A payload codec configured here runs on this side of + the endpoint, so records are encoded before they leave the + process and the handler's worker never holds the key. + """ + if not endpoint: + raise ValueError("endpoint must not be empty") + # Checked here rather than on the first read, because the contract + # refuses an out-of-range value with a payload validation error that + # is neither a StreamError nor an RPCError, a long way from the line + # that got it wrong. + if not _MIN_READ_WAIT <= read_wait <= _MAX_READ_WAIT: + raise ValueError( + f"read_wait must be between {_MIN_READ_WAIT} and {_MAX_READ_WAIT}, " + f"got {read_wait}" + ) + if read_wait.microseconds % 1000: + raise ValueError( + f"read_wait is carried in whole milliseconds, got {read_wait}" + ) + if not _MIN_MAX_RECORDS <= max_records <= _MAX_MAX_RECORDS: + raise ValueError( + f"max_records must be between {_MIN_MAX_RECORDS} and " + f"{_MAX_MAX_RECORDS}, got {max_records}" + ) + self._endpoint = endpoint + self._http_address = http_address.rstrip("/") + self._headers = dict(headers or {}) + self._read_wait = read_wait + self._max_records = max_records + self._data_converter = data_converter + self._resolved: str | None = None + self._resolving = asyncio.Lock() + + def workflow_provider(self) -> NoReturn: + """Raise: this front has no workflow half.""" + raise StreamUnsupportedError(_WORKFLOW_SIDE_ERROR) + + def get_stream_handle( + self, client: Client | None, workflow_id: str, *, run_id: str | None = None + ) -> NexusStreamHandle: + """A handle on ``workflow_id``'s stream through the endpoint. + + ``client`` is not used for transport: it resolves the endpoint's name + and supplies the data converter when none was configured. + """ + return NexusStreamHandle( + _Front(self, client), _Address("workflow", workflow_id, run_id) + ) + + def get_activity_stream_handle( + self, + client: Client | None, + activity_id: str, + *, + workflow_id: str | None = None, + run_id: str | None = None, + ) -> NexusStreamHandle: + """A handle on the streams ``activity_id`` owns, through the endpoint. + + The reference carries the activity owner across; whether the store + behind the endpoint can host it is the store's answer, and a refusal + reaches the caller as :class:`temporalio.streams.StreamUnsupportedError` + on the first call. + """ + return NexusStreamHandle( + _Front(self, client), + _Address("activity", workflow_id, run_id, activity_id), + ) + + def get_standalone_stream_handle( + self, client: Client | None, stream_id: str + ) -> NexusStreamHandle: + """A handle on the standalone stream ``stream_id``, through the endpoint. + + The reference carries the standalone owner across; the store behind + the endpoint answers whether it hosts one, and a refusal reaches the + caller as :class:`temporalio.streams.StreamUnsupportedError` on the + first call. + """ + if not stream_id: + raise ValueError("stream_id must not be empty") + return NexusStreamHandle( + _Front(self, client), _Address("standalone", stream_id=stream_id) + ) + + async def create_standalone_stream( + self, + client: Client | None, + stream_id: str, + *, + retention: timedelta | None = None, + max_records: int | None = None, + max_bytes: int | None = None, + ) -> NoReturn: + """Refused: the contract carries no operation that creates a stream. + + Raises: + StreamUnsupportedError: Always. Create the stream on the store + behind the endpoint and hand its reference to callers. + """ + raise StreamUnsupportedError( + "the stream endpoint has no operation that creates a standalone stream; " + "create it on the store behind the endpoint and pass its StreamRef" + ) + + async def close(self) -> None: + """Nothing to release: each call opens and closes its own connection.""" + + def configure_client(self, config: ClientConfig) -> ClientConfig: + """Set this front as the client's ``stream_provider``.""" + config["stream_provider"] = self + return config + + async def connect_service_client( + self, + config: ConnectConfig, + next: Callable[[ConnectConfig], Awaitable[ServiceClient]], + ) -> ServiceClient: + """Connect unchanged.""" + return await next(config) + + async def _endpoint_id(self, client: Client | None) -> str: + if self._resolved is not None: + return self._resolved + async with self._resolving: + if self._resolved is None: + if client is None: + self._resolved = self._endpoint + else: + found = await client.operator_service.list_nexus_endpoints( + ListNexusEndpointsRequest(name=self._endpoint) + ) + if not found.endpoints: + raise ValueError( + f"no nexus endpoint is named {self._endpoint!r}" + ) + self._resolved = found.endpoints[0].id + return self._resolved diff --git a/temporalio/streams/providers/temporal_streams.nexusrpc.yaml b/temporalio/streams/providers/temporal_streams.nexusrpc.yaml new file mode 100644 index 000000000..1b027cfeb --- /dev/null +++ b/temporalio/streams/providers/temporal_streams.nexusrpc.yaml @@ -0,0 +1,234 @@ +# The stream endpoint's wire contract. Inputs reject unknown fields so a typo +# fails loudly at the handler; outputs accept them so a handler can add a +# field before every caller has been rebuilt from the newer contract. +nexusrpc: "1.0.0" + +services: + TemporalStreams: + fqn: TemporalStreams + description: >- + The stream endpoint's two operations. One Temporal-authenticated endpoint + hides the store behind it, so an operator switches storage without + touching callers. Reads hand out batches and an append carries one batch + per call, because a Nexus operation per record costs too much for token + streams. A record crosses as the serialized + temporal.api.stream.v1.StreamRecord, the same bytes every store keeps, so + a caller in any language decodes it with the api protos alone. + operations: + append: + fqn: append + description: >- + Append one batch on the caller's account. The handler writes a batch + once: a repeat of the last batch_index answers with where the + original landed, and an index that skips ahead, one already behind + the last, a sequence that does not continue, or a producer attempt + the handler has no state for is refused. Supersession records are not + transported: a reader re-synthesizes them from the attempts it + observes. + input: { $ref: "#/$defs/AppendInput" } + output: { $ref: "#/$defs/AppendOutput" } + read: + fqn: read + description: >- + Answer with the records after the caller's token, or time out. The + call parks until it has max_records records or wait_ms elapses, + whichever comes first, and returns whatever it collected. It says + when the store has ended the read, so the caller can stop. + input: { $ref: "#/$defs/ReadInput" } + output: { $ref: "#/$defs/ReadOutput" } + +$defs: + StreamRef: + type: object + description: >- + A stream, named by its owner and a topic: what an operation returns to + hand a stream to its caller, and what read and append take in place of + an owner spelled out. It names no cursor and no store, so the same + reference is good behind any endpoint that serves the owner. A member + the owner kind does not use is absent or null; null is how the SDK + writes an unset member of its own StreamRef. + properties: + kind: + type: string + enum: [workflow, activity, standalone] + description: >- + What owns the stream. A workflow's streams are keyed by workflow_id + and, when pinned, run_id. An activity's own streams are keyed by + activity_id and, when a workflow scheduled it, that workflow's ids. + A standalone stream has its own stream_id and no execution behind + it. + workflow_id: + oneOf: [{ type: string }, { type: "null" }] + description: >- + The owning workflow, or the workflow that scheduled the owning + activity. Required for a workflow owner. + run_id: + oneOf: [{ type: string }, { type: "null" }] + description: >- + Pin the owner to this run. Absent follows the execution chain + across continue-as-new, which is what a producer normally wants. + activity_id: + oneOf: [{ type: string }, { type: "null" }] + description: "The owning activity. Required for an activity owner." + stream_id: + oneOf: [{ type: string }, { type: "null" }] + description: "The stream's own id. Required for a standalone owner." + topic: + type: string + description: "The topic on the owner's stream." + required: + - kind + - topic + additionalProperties: false + + AppendInput: + type: object + description: "One append call: who is writing, where, and what." + properties: + stream: { $ref: "#/$defs/StreamRef" } + producer_id: + type: string + description: >- + Identifies the writer across its retries, so its attempts can be + ordered. Never empty here: the caller resolved it, from the activity + context when it was not given. + attempt: + type: integer + description: >- + The generation this producer is writing. A later attempt supersedes + an earlier one. + sequence: + type: integer + minimum: 0 + description: >- + The producer's sequence of the first record in this batch; the batch + is numbered from it and a finish takes the next number. It has to + continue where the previous batch ended, so the handler and the + store agree on every record's position within the attempt. + batch_index: + type: integer + minimum: 1 + description: >- + Counts this producer's batches from 1. The handler answers a repeat + of the last index with the original's position, and refuses one that + skips ahead, one already behind the last, or one that resumes an + attempt it never saw start. + payloads: + type: array + items: + type: string + contentEncoding: base64 + description: >- + The record bodies to write, in order, each a serialized + temporal.api.common.v1.Payload. A payload codec configured on the + caller has already run on them. Empty on a call that only finishes. + finish: + type: boolean + description: >- + Write FINISH for this producer after the batch: it will write nothing + more on the topic. Says nothing about the producer's outcome. + required: + - stream + - producer_id + - attempt + - sequence + - batch_index + additionalProperties: false + + AppendOutput: + type: object + description: "Where the append landed, when the store can say." + properties: + cursor: + type: string + description: >- + Opaque token naming the last record written by this call, or by the + original when the call repeated the last batch. Absent when the + store learns positions only at read time; such a caller positions + itself with a latest_only read. + additionalProperties: true + + ReadInput: + type: object + description: "One read call: where to resume from and how long to wait." + properties: + stream: { $ref: "#/$defs/StreamRef" } + after_token: + type: string + description: >- + Opaque cursor from an earlier record or append. The read resumes + strictly after the record it names, so a caller never sees that + record twice. Empty starts at the beginning. The token is produced by + whichever store sits behind the endpoint, so a caller cannot tell + which one that is, and a token from another store is refused. + max_records: + type: integer + minimum: 1 + maximum: 1000 + description: "Return at most this many records. Defaults to 100 when omitted." + wait_ms: + type: integer + minimum: 0 + maximum: 60000 + description: >- + Wait at most this long for records to arrive before answering. The + handler shortens it to fit the request deadline. A call that collects + nothing answers with no records and the caller's own token. + latest_only: + type: boolean + description: >- + Answer with the newest position and no records, for a reader that + wants to follow from now. + last_n: + type: integer + minimum: 1 + description: >- + With no after_token, start at the newest this many records, or at + all of them when the stream holds fewer. Records of every kind + count. Refused alongside an after_token, which is how a read + resumes. + required: + - stream + additionalProperties: false + + RecordWire: + type: object + description: "One record on the wire: its cursor and the record itself." + properties: + token: + type: string + description: >- + Opaque cursor naming this record. Pass it as after_token to resume + just past it. + record: + type: string + contentEncoding: base64 + description: >- + The serialized temporal.api.stream.v1.StreamRecord: topic, kind, + producer_id, attempt, sequence and body, as the store holds it. + required: + - token + - record + additionalProperties: true + + ReadOutput: + type: object + description: "What one read call answered, and where to resume." + properties: + records: + type: array + items: + $ref: "#/$defs/RecordWire" + description: "The records after the caller's token, in stream order." + next_token: + type: string + description: >- + Opaque cursor to pass as after_token on the following call. Echoes + the caller's own token when the call collected nothing. + done: + type: boolean + description: >- + The store ended the read: the owning execution, or its chain, is + closed and every retained record after the caller's token has been + delivered. Nothing more will arrive, so the caller stops. + additionalProperties: true diff --git a/tests/streams/test_nexus_provider.py b/tests/streams/test_nexus_provider.py new file mode 100644 index 000000000..efec762c8 --- /dev/null +++ b/tests/streams/test_nexus_provider.py @@ -0,0 +1,1234 @@ +"""Conformance for the Nexus front. + +The caller talks only to the stream endpoint; the handler fronts a storage +provider's own handles, so these tests are the provider-hiding demonstration: +nothing on the caller side names or could name the store. + +The live tests run over the server's Nexus HTTP ingress and are gated behind +``STREAMS_LIVE=nexus`` because they need a dev server with an HTTP port. +Environment: ``TEMPORAL_ADDRESS`` (default ``localhost:7233``) and +``TEMPORAL_HTTP`` (default ``http://127.0.0.1:7243``). Each live test +registers its own Nexus endpoint through the operator service and deletes it +at the end. + +Everything else stands the endpoint up in this process, because what it +checks is the bytes the caller puts on the wire and the handler's own rules. +""" + +from __future__ import annotations + +import asyncio +import base64 +import contextlib +import dataclasses +import gc +import http.client +import json +import os +import uuid +from collections.abc import AsyncIterator, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Any, cast +from unittest import mock + +import nexusrpc +import nexusrpc.handler +import pytest + +import temporalio.converter +from temporalio import nexus as temporal_nexus +from temporalio import streams, workflow +from temporalio.api.common.v1 import Payload +from temporalio.api.nexus.v1 import EndpointSpec, EndpointTarget +from temporalio.api.operatorservice.v1 import ( + CreateNexusEndpointRequest, + DeleteNexusEndpointRequest, +) +from temporalio.client import Client +from temporalio.common import WorkflowIDConflictPolicy +from temporalio.service import RPCError, RPCStatusCode +from temporalio.streams import ( + BEGINNING, + END, + Cursor, + RecordKind, + StreamCursorError, + StreamProducerError, + StreamRef, + StreamUnsupportedError, +) +from temporalio.streams._ref import open_ref +from temporalio.streams._wire import WireRecord +from temporalio.streams.providers import nexus +from temporalio.streams.providers._nexus_generated import AppendInput, ReadInput +from temporalio.streams.providers._nexus_generated import StreamRef as WireStreamRef +from temporalio.streams.providers.memory import MemoryStreams +from temporalio.streams.providers.nexus import ( + NexusStreamHandle, + NexusStreams, + TemporalStreamsHandler, +) +from temporalio.streams.providers.workflow_streams import WorkflowStreamsProvider +from temporalio.worker import Worker +from tests.helpers import new_worker +from tests.streams.test_workflow_streams_provider import EchoLoop, take +from tests.streams.test_workflow_streams_provider import _feed as feed_echo_loop + +live_only = pytest.mark.skipif( + os.environ.get("STREAMS_LIVE") != "nexus", + reason="needs a live server and nexus endpoint; run with STREAMS_LIVE=nexus", +) + +APPEND_OPERATION = nexus._APPEND_OPERATION # pyright: ignore[reportPrivateUsage] +READ_OPERATION = nexus._READ_OPERATION # pyright: ignore[reportPrivateUsage] +INPUTS = "inputs" +DECISIONS = "decisions" + + +async def _live_client() -> Client: + return await Client.connect(os.environ.get("TEMPORAL_ADDRESS", "localhost:7233")) + + +def _live_front(endpoint: str) -> NexusStreams: + return NexusStreams( + endpoint=endpoint, + http_address=os.environ.get("TEMPORAL_HTTP", "http://127.0.0.1:7243"), + read_wait=timedelta(seconds=5), + ) + + +@contextlib.asynccontextmanager +async def _own_endpoint( + client: Client, name: str, task_queue: str +) -> AsyncIterator[str]: + """Register a Nexus endpoint routed to ``task_queue`` for one test's life.""" + created = await client.operator_service.create_nexus_endpoint( + CreateNexusEndpointRequest( + spec=EndpointSpec( + name=name, + target=EndpointTarget( + worker=EndpointTarget.Worker( + namespace=client.namespace, task_queue=task_queue + ) + ), + ) + ) + ) + try: + yield name + finally: + await client.operator_service.delete_nexus_endpoint( + DeleteNexusEndpointRequest( + id=created.endpoint.id, version=created.endpoint.version + ) + ) + + +@live_only +async def test_interface_loop_through_the_nexus_front(): + # The workflow worker and the handler worker both use the storage + # provider; only the caller goes through the front, and it names the + # endpoint the way an operator does, by name. + store = WorkflowStreamsProvider() + client = await _live_client() + workflow_id = f"streams-nexus-live-{uuid.uuid4().hex}" + handler = TemporalStreamsHandler(store, client) + handler_tq = f"handlers-{workflow_id}" + + async with ( + _own_endpoint(client, workflow_id, handler_tq) as endpoint, + Worker(client, task_queue=handler_tq, nexus_service_handlers=[handler]), + ): + front = _live_front(endpoint) + async with Worker( + client, + task_queue=f"tq-{workflow_id}", + workflows=[EchoLoop], + plugins=[store], + ): + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=f"tq-{workflow_id}" + ) + + stream = front.get_stream_handle(client, workflow_id) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}) + await producer.append({"n": 3}) + await producer.finish() + + records = await take(stream.read(topic=DECISIONS, result_type=dict), 4, 60) + assert [r.kind for r in records] == [ + RecordKind.DATA, + RecordKind.DATA, + RecordKind.DATA, + RecordKind.FINISH, + ] + assert [r.value["echo"] for r in records[:3]] == [1, 2, 3] + + # An opaque cursor from behind the front resumes a fresh reader + # just past the record it names. + checkpoint = records[0].cursor + again = await take( + stream.read(topic=DECISIONS, result_type=dict, after=checkpoint), 2, 60 + ) + assert [r.value["echo"] for r in again[:2]] == [2, 3] + + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + await handler.close() + + +class ScrambleCodec(temporalio.converter.PayloadCodec): + """Flips every byte and labels the result with an encoding only it can read.""" + + KEY = 0x5A + + async def encode(self, payloads: Sequence[Payload]) -> list[Payload]: + """Wrap each payload as opaque bytes nothing downstream can read.""" + return [ + Payload( + metadata={"encoding": b"binary/scrambled"}, + data=self._flip(payload.SerializeToString()), + ) + for payload in payloads + ] + + async def decode(self, payloads: Sequence[Payload]) -> list[Payload]: + """Recover the payloads :meth:`encode` wrapped.""" + return [Payload.FromString(self._flip(payload.data)) for payload in payloads] + + @classmethod + def _flip(cls, data: bytes) -> bytes: + return bytes(byte ^ cls.KEY for byte in data) + + +class _NeverCancelled(nexusrpc.handler.OperationTaskCancellation): + def is_cancelled(self) -> bool: + return False + + def cancellation_reason(self) -> str | None: + return None + + def wait_until_cancelled_sync(self, timeout: float | None = None) -> bool: + del timeout + return False + + async def wait_until_cancelled(self) -> None: + await asyncio.Event().wait() + + +def _context( + operation: str, deadline: datetime | None = None +) -> nexusrpc.handler.StartOperationContext: + return nexusrpc.handler.StartOperationContext( + service=nexus._SERVICE_NAME, # pyright: ignore[reportPrivateUsage] + operation=operation, + headers={}, + task_cancellation=_NeverCancelled(), + request_deadline=deadline, + request_id=uuid.uuid4().hex, + ) + + +async def _dispatch( + handler: TemporalStreamsHandler, operation: str, request: Any +) -> Any: + if operation == APPEND_OPERATION: + return await handler.append(_context(operation), request) + return await handler.read(_context(operation), request) + + +def _in_process_endpoint( + handler: TemporalStreamsHandler, + posted: list[bytes], + answered: list[bytes], + *, + fail_after_applying: list[str] | None = None, +) -> Any: + """Serve the caller's posts from ``handler``, recording both directions. + + ``fail_after_applying`` names operations whose first call is applied and + then reported as a transport failure, the ambiguous case a retry has to + survive. + """ + contract = ( + temporalio.converter.DataConverter.default._get_internal_payload_converter() + ) + loop = asyncio.get_running_loop() + failing = set(fail_after_applying or []) + + def post( + url: str, body: bytes, headers: Mapping[str, str], timeout: timedelta + ) -> bytes: + del headers, timeout + posted.append(body) + operation = url.rsplit("/", 1)[1] + request_type = AppendInput if operation == APPEND_OPERATION else ReadInput + request: Any = contract.from_payloads( + [Payload(metadata={"encoding": b"json/plain"}, data=body)], [request_type] + )[0] + try: + answer = asyncio.run_coroutine_threadsafe( + _dispatch(handler, operation, request), loop + ).result() + except nexusrpc.HandlerError as error: + # The ingress answers a handler error with a failure body and the + # status its type maps to. + status = 404 if error.type is nexusrpc.HandlerErrorType.NOT_FOUND else 400 + raise nexus._EndpointFailure( # pyright: ignore[reportPrivateUsage] + json.dumps({"message": str(error)}), status + ) from error + raw = contract.to_payloads([answer])[0].data + answered.append(raw) + if operation in failing: + failing.discard(operation) + raise nexus._EndpointFailure( # pyright: ignore[reportPrivateUsage] + "connection reset after the handler answered", None + ) + return raw + + return post + + +def _front(codec: temporalio.converter.PayloadCodec | None) -> NexusStreams: + return NexusStreams( + endpoint="in-process", + # This endpoint has no other traffic to wait for, so park briefly + # rather than for the contract's default. + read_wait=timedelta(milliseconds=200), + data_converter=dataclasses.replace( + temporalio.converter.DataConverter.default, payload_codec=codec + ), + ) + + +async def _loop_through_an_endpoint( + monkeypatch: pytest.MonkeyPatch, + codec: temporalio.converter.PayloadCodec | None, +) -> tuple[list[Any], list[bytes], list[bytes]]: + store = MemoryStreams() + posted: list[bytes] = [] + answered: list[bytes] = [] + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, posted, answered)) + stream = _front(codec).get_stream_handle(None, "wf-codec") + + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"secret": "tuna"}) + await producer.finish() + + records = await take(stream.read(topic=INPUTS, result_type=dict), 2, timeout=30) + await handler.close() + return records, posted, answered + + +def _appended_bytes(posted: list[bytes]) -> bytes: + out = b"" + for body in posted: + for encoded in json.loads(body).get("payloads", []): + out += base64.b64decode(encoded) + return out + + +def _answered_bodies(answered: list[bytes]) -> bytes: + out = b"" + for body in answered: + for record in json.loads(body).get("records", []): + wire = WireRecord.FromString(base64.b64decode(record["record"])) + out += wire.body.SerializeToString() + return out + + +async def test_the_caller_encodes_records_with_the_configured_codec( + monkeypatch: pytest.MonkeyPatch, +): + records, posted, answered = await _loop_through_an_endpoint( + monkeypatch, ScrambleCodec() + ) + + # The value survives the hop although the handler cannot read the + # encoding, so the body passed through it untouched and the codec ran on + # the caller alone. + assert [record.kind for record in records] == [RecordKind.DATA, RecordKind.FINISH] + assert records[0].value == {"secret": "tuna"} + # Neither direction carried the plaintext. + assert b"tuna" not in _appended_bytes(posted) + assert b"tuna" not in _answered_bodies(answered) + + +async def test_without_a_codec_the_same_records_go_out_in_the_clear( + monkeypatch: pytest.MonkeyPatch, +): + # The contrast is the point: the bytes above differ because of the codec, + # not because the transport obscures them anyway. + records, posted, answered = await _loop_through_an_endpoint(monkeypatch, None) + + assert records[0].value == {"secret": "tuna"} + assert b"tuna" in _appended_bytes(posted) + assert b"tuna" in _answered_bodies(answered) + + +def _ref(workflow_id: str, topic: str = INPUTS) -> WireStreamRef: + return WireStreamRef(kind="workflow", workflow_id=workflow_id, topic=topic) + + +def _append( + workflow_id: str, + batch_index: int, + *values: Any, + sequence: int | None = None, + finish: bool = False, +) -> AppendInput: + converter = temporalio.converter.DataConverter.default.payload_converter + return AppendInput( + stream=_ref(workflow_id), + producer_id="model", + attempt=1, + # Batches of one, numbered in step with the batch index unless a + # case wants them out of step. + sequence=batch_index - 1 if sequence is None else sequence, + batch_index=batch_index, + payloads=[ + converter.to_payloads([value])[0].SerializeToString() for value in values + ], + finish=finish, + ) + + +async def _stored(store: MemoryStreams, workflow_id: str) -> list[Any]: + stream = store.get_stream_handle(None, workflow_id) + end = await stream.latest(topic=INPUTS) + if end == BEGINNING: + return [] + out: list[Any] = [] + async for record in stream.read(topic=INPUTS, result_type=dict): + if record.kind is RecordKind.DATA: + out.append(record.value) + if record.cursor == end: + break + return out + + +async def test_a_handler_without_state_for_a_producer_refuses_to_continue_it(): + # A second handler instance stands in for a restarted or load-balanced + # handler worker: it shares the store but not the dedupe state. + store = MemoryStreams() + first = TemporalStreamsHandler(store, None) + second = TemporalStreamsHandler(store, None) + await _dispatch(first, APPEND_OPERATION, _append("wf", 1, {"n": 1})) + with pytest.raises(nexusrpc.HandlerError) as failed: + await _dispatch(second, APPEND_OPERATION, _append("wf", 2, {"n": 2})) + assert failed.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + assert failed.value.retryable_override is False + # Named so the caller can raise the same class. + assert str(failed.value).startswith("StreamProducerError: ") + # The store holds the first batch and nothing was silently lost or doubled. + assert await _stored(store, "wf") == [{"n": 1}] + + +async def test_the_handler_answers_repeats_and_rejects_gaps(): + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + await _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1})) + second = await _dispatch(handler, APPEND_OPERATION, _append("wf", 2, {"n": 2})) + # A repeat of the last batch is written once and answers with where the + # original landed, so a retrying caller checkpoints the same position. + repeat = await _dispatch(handler, APPEND_OPERATION, _append("wf", 2, {"n": 2})) + assert repeat.cursor == second.cursor + with pytest.raises(nexusrpc.HandlerError) as skipped: + await _dispatch(handler, APPEND_OPERATION, _append("wf", 4, {"n": 4})) + assert skipped.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + with pytest.raises(nexusrpc.HandlerError) as behind: + await _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1})) + assert str(behind.value).startswith("StreamProducerError: ") + with pytest.raises(nexusrpc.HandlerError) as renumbered: + await _dispatch( + handler, APPEND_OPERATION, _append("wf", 3, {"n": 3}, sequence=7) + ) + assert "sequence 7 does not continue at 2" in str(renumbered.value) + assert await _stored(store, "wf") == [{"n": 1}, {"n": 2}] + + +async def test_a_malformed_payload_is_the_callers_fault(): + handler = TemporalStreamsHandler(MemoryStreams(), None) + request = _append("wf", 1) + request.payloads = [b"\xff\xfe not a payload"] + with pytest.raises(nexusrpc.HandlerError) as failed: + await _dispatch(handler, APPEND_OPERATION, request) + assert failed.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + + +async def test_the_caller_retries_an_ambiguous_append_under_the_same_index( + monkeypatch: pytest.MonkeyPatch, +): + posted: list[bytes] = [] + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr( + nexus, + "_post", + _in_process_endpoint( + handler, posted, [], fail_after_applying=[APPEND_OPERATION] + ), + ) + stream = _front(None).get_stream_handle(None, "wf") + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + # A transport failure is an RPCError, never the transport's own type. + with pytest.raises(RPCError) as failed: + await producer.append({"n": 1}) + assert failed.value.status == RPCStatusCode.UNAVAILABLE + landed = await producer.append({"n": 1}) + await producer.append({"n": 2}) + await producer.finish() + indexes = [json.loads(body)["batch_index"] for body in posted] + sequences = [json.loads(body)["sequence"] for body in posted] + # The retry re-sent index 1, which the handler answered as a repeat with + # the original's position, so the store holds each record once. + assert indexes == [1, 1, 2, 3] + assert sequences == [0, 0, 1, 2] + assert landed == await _memory_position(store, "wf", 0) + assert await _stored(store, "wf") == [{"n": 1}, {"n": 2}] + + +async def _memory_position( + store: MemoryStreams, workflow_id: str, index: int +) -> Cursor: + records = await take( + store.get_stream_handle(None, workflow_id).read(topic=INPUTS, result_type=dict), + index + 1, + ) + return records[index].cursor + + +async def test_a_failed_append_goes_out_before_the_next_one( + monkeypatch: pytest.MonkeyPatch, +): + posted: list[bytes] = [] + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr( + nexus, + "_post", + _in_process_endpoint( + handler, posted, [], fail_after_applying=[APPEND_OPERATION] + ), + ) + stream = _front(None).get_stream_handle(None, "wf") + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + with pytest.raises(RPCError): + await producer.append({"n": 1}) + # The caller moves on without retrying; the pending batch is replayed + # first under its own index and the new one takes the next. + await producer.append({"n": 2}) + indexes = [json.loads(body)["batch_index"] for body in posted] + assert indexes == [1, 1, 2] + assert await _stored(store, "wf") == [{"n": 1}, {"n": 2}] + + +async def test_a_stream_condition_crosses_the_endpoint_under_its_own_class( + monkeypatch: pytest.MonkeyPatch, +): + handler = TemporalStreamsHandler(MemoryStreams(), None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + stream = _front(None).get_stream_handle(None, "wf") + # The token is opaque on this side, so the store behind the endpoint is + # what refuses it, and the refusal arrives on the first read as the class + # the store raised. + with pytest.raises(StreamCursorError): + await take(stream.read(topic=INPUTS, after=Cursor("elsewhere:1")), 1) + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}) + # A second handler has no state for the attempt; the caller learns that + # as a producer conflict. + fresh = TemporalStreamsHandler(MemoryStreams(), None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(fresh, [], [])) + with pytest.raises(StreamProducerError): + await producer.append({"n": 2}) + await handler.close() + + +def test_the_front_has_no_workflow_half(): + with pytest.raises(StreamUnsupportedError): + _front(None).workflow_provider() + + +async def test_the_front_registers_on_a_client(client: Client): + # The same accessor as any provider, so code outside a workflow does not + # change when the store moves behind an endpoint. + config = client.config() + config["plugins"] = [_front(None)] + registered = Client(**config) + assert isinstance(registered.get_stream_handle("wf"), NexusStreamHandle) + + +async def test_consecutive_reads_share_one_parked_subscription(): + handler = TemporalStreamsHandler(MemoryStreams(), None) + await _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1}, {"n": 2})) + first = await _dispatch( + handler, + READ_OPERATION, + ReadInput(stream=_ref("wf"), max_records=1, wait_ms=200), + ) + assert len(first.records) == 1 + parked = handler._subscriptions # pyright: ignore[reportPrivateUsage] + pump = next(iter(parked.values())).pump + # A call that resumes where the last one stopped drains the same + # subscription instead of parking a second read on the store. + second = await _dispatch( + handler, + READ_OPERATION, + ReadInput( + stream=_ref("wf"), + after_token=first.next_token, + max_records=1, + wait_ms=200, + ), + ) + assert len(second.records) == 1 + assert second.next_token != first.next_token + assert len(parked) == 1 and next(iter(parked.values())).pump is pump + # A call from somewhere else replaces it. + await _dispatch( + handler, + READ_OPERATION, + ReadInput(stream=_ref("wf"), wait_ms=200), + ) + assert len(parked) == 1 and next(iter(parked.values())).pump is not pump + await handler.close() + assert not parked + + +async def test_the_read_wait_is_cut_to_the_request_deadline(): + handler = TemporalStreamsHandler(MemoryStreams(), None) + started = asyncio.get_running_loop().time() + answer = await handler.read( + _context(READ_OPERATION, datetime.now(timezone.utc) + timedelta(seconds=1)), + ReadInput(stream=_ref("wf"), wait_ms=60000), + ) + assert answer.records == [] and answer.next_token == "" and not answer.done + assert asyncio.get_running_loop().time() - started < 5 + await handler.close() + + +async def test_the_read_ends_when_the_store_ends_it( + client: Client, monkeypatch: pytest.MonkeyPatch +): + # Behind the endpoint sits the Workflow Streams store, whose read ends + # once the run has closed and its tail is served. The front carries that + # end to the caller, whose loop leaves by itself. + store = WorkflowStreamsProvider(poll_cooldown=timedelta(milliseconds=20)) + handler = TemporalStreamsHandler(store, client) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + workflow_id = f"streams-nexus-{uuid.uuid4().hex}" + async with new_worker(client, EchoLoop, plugins=[store]) as worker: + handle = await client.start_workflow( + EchoLoop.run, id=workflow_id, task_queue=worker.task_queue + ) + await feed_echo_loop(store, client, workflow_id) + await handle.signal(EchoLoop.release) + assert await handle.result() == 3 + + stream = _front(None).get_stream_handle(None, 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(), 30) + await handler.close() + assert records == [ + (RecordKind.DATA, {"echo": 1}), + (RecordKind.DATA, {"echo": 2}), + (RecordKind.DATA, {"echo": 3}), + (RecordKind.FINISH, None), + ] + + +class _SlowStore(MemoryStreams): + """The memory store with a real await inside ``append``. + + The handler's repeat check and its commit have an await between them. + Without a gate there, two in-flight copies of one batch both finish the + check before either commits and both reach the store. + """ + + def __init__(self) -> None: + super().__init__() + self.gate = asyncio.Event() + self.appends = 0 + + def get_stream_handle( + self, client: Any, workflow_id: str, *, run_id: str | None = None + ) -> Any: + handle = super().get_stream_handle(client, workflow_id, run_id=run_id) + outer = self + make_producer = handle.producer + + def producer(**kwargs: Any) -> Any: + delegate = make_producer(**kwargs) + inner_append = delegate.append + + async def append(*values: Any) -> Any: + outer.appends += 1 + await outer.gate.wait() + return await inner_append(*values) + + delegate.append = append # type: ignore[method-assign] + return delegate + + handle.producer = producer # type: ignore[method-assign] + return handle + + +async def test_two_copies_of_one_batch_reach_the_store_once(): + store = _SlowStore() + handler = TemporalStreamsHandler(store, None) + both = asyncio.gather( + _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1})), + _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1})), + ) + await asyncio.sleep(0.1) + store.gate.set() + first, second = await asyncio.wait_for(both, 10) + # One of them wrote and the other was answered from the state it left. + # Counted at the store because a store that dedupes by itself would hide + # a second commit the handler should never have made. + assert store.appends == 1 + assert first.cursor == second.cursor + assert await _stored(store, "wf") == [{"n": 1}] + + +async def test_a_finished_batch_can_be_retried(): + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + await _dispatch(handler, APPEND_OPERATION, _append("wf", 1, {"n": 1})) + finished = await _dispatch( + handler, APPEND_OPERATION, _append("wf", 2, {"n": 2}, finish=True) + ) + # The caller never saw the answer. The byte-identical retry has to be + # answerable, because a finish the caller cannot repeat is a batch it can + # neither complete nor abandon. + again = await _dispatch( + handler, APPEND_OPERATION, _append("wf", 2, {"n": 2}, finish=True) + ) + assert again.cursor == finished.cursor + assert await _stored(store, "wf") == [{"n": 1}, {"n": 2}] + # And the attempt is over: a batch after the finish is refused rather + # than landing behind the marker. + with pytest.raises(nexusrpc.HandlerError) as after: + await _dispatch(handler, APPEND_OPERATION, _append("wf", 3, {"n": 3})) + assert "already finished" in str(after.value) + + +class _FailingReadStore(MemoryStreams): + """A store whose read fails once, the way a transient store failure does.""" + + def __init__(self) -> None: + super().__init__() + self.reads = 0 + + def get_stream_handle( + self, client: Any, workflow_id: str, *, run_id: str | None = None + ) -> Any: + handle = super().get_stream_handle(client, workflow_id, run_id=run_id) + outer = self + + def read(**kwargs: Any) -> Any: + outer.reads += 1 + if outer.reads == 1: + + async def failing() -> AsyncIterator[Any]: + for record in cast("list[Any]", []): + yield record + raise RuntimeError("the store went away") + + return failing() + return MemoryStreams.get_stream_handle( + outer, client, workflow_id, run_id=run_id + ).read(**kwargs) + + handle.read = read # type: ignore[method-assign] + return handle + + +async def test_a_failed_read_is_not_reported_as_the_end_of_the_stream(): + store = _FailingReadStore() + handler = TemporalStreamsHandler(store, None) + await ( + store.get_stream_handle(None, "wf") + .producer(topic=INPUTS, producer_id="model", attempt=1) + .append({"n": 1}) + ) + + request = ReadInput(stream=_ref("wf"), wait_ms=200, max_records=10) + with pytest.raises(RuntimeError, match="the store went away"): + await _dispatch(handler, READ_OPERATION, request) + # The retry from the same token has to re-subscribe rather than be + # answered done=True off the subscription the failed pump left behind. + answer = await _dispatch(handler, READ_OPERATION, request) + assert [WireRecord.FromString(r.record).sequence for r in answer.records] == [1] + assert answer.done is False + await handler.close() + + +async def test_the_handler_does_not_grow_a_lock_per_address(): + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + for index in range(25): + workflow_id = f"wf-{index}" + await _dispatch( + handler, APPEND_OPERATION, _append(workflow_id, 1, {"n": index}) + ) + await _dispatch( + handler, + READ_OPERATION, + ReadInput(stream=_ref(workflow_id), wait_ms=0, max_records=1), + ) + gc.collect() + # Held weakly, so an address nobody is reading or appending to leaves + # nothing behind. A strong map would hold one entry per address forever. + assert len(handler._read_locks) == 0 # pyright: ignore[reportPrivateUsage] + assert len(handler._append_locks) == 0 # pyright: ignore[reportPrivateUsage] + await handler.close() + + +def test_the_read_bounds_are_checked_where_they_are_set(): + # The contract accepts a wait up to a minute and a batch up to a + # thousand; past that the endpoint answers with a payload validation + # error, which is neither a stream condition nor an RPC failure. + with pytest.raises(ValueError, match="read_wait"): + NexusStreams(endpoint="e", read_wait=timedelta(minutes=2)) + with pytest.raises(ValueError, match="read_wait"): + NexusStreams(endpoint="e", read_wait=timedelta(seconds=-1)) + with pytest.raises(ValueError, match="max_records"): + NexusStreams(endpoint="e", max_records=0) + with pytest.raises(ValueError, match="max_records"): + NexusStreams(endpoint="e", max_records=1001) + NexusStreams(endpoint="e", read_wait=timedelta(seconds=60), max_records=1000) + + +def test_a_socket_failure_never_reaches_the_caller_as_a_timeout(): + # urllib wraps only the request in URLError, so a read timeout comes out + # of getresponse() as builtins.TimeoutError, which on 3.11 and later is + # asyncio.TimeoutError and would be taken for the caller's own deadline. + def _raise(*args: Any, **kwargs: Any) -> Any: + del args, kwargs + raise TimeoutError("the socket read timed out") + + with mock.patch("urllib.request.urlopen", _raise): + with pytest.raises(nexus._EndpointFailure): # pyright: ignore[reportPrivateUsage] + nexus._post( # pyright: ignore[reportPrivateUsage] + "http://127.0.0.1:1/x", b"{}", {}, timedelta(seconds=1) + ) + + def _disconnect(*args: Any, **kwargs: Any) -> Any: + del args, kwargs + raise http.client.RemoteDisconnected("the endpoint hung up") + + with mock.patch("urllib.request.urlopen", _disconnect): + with pytest.raises(nexus._EndpointFailure): # pyright: ignore[reportPrivateUsage] + nexus._post( # pyright: ignore[reportPrivateUsage] + "http://127.0.0.1:1/x", b"{}", {}, timedelta(seconds=1) + ) + + # A url urllib cannot even build a request from does not escape either. + with pytest.raises(nexus._EndpointFailure): # pyright: ignore[reportPrivateUsage] + nexus._post("not a url", b"{}", {}, timedelta(seconds=1)) # pyright: ignore[reportPrivateUsage] + + +async def test_the_front_carries_every_read_start(monkeypatch: pytest.MonkeyPatch): + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + stream = _front(None).get_stream_handle(None, "wf-starts") + producer = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await producer.append({"n": 1}, {"n": 2}, {"n": 3}) + + # last= rides the contract to the store behind the endpoint. + newest = await take( + stream.read(topic=INPUTS, result_type=dict, last=2), 2, timeout=30 + ) + assert [record.value for record in newest] == [{"n": 2}, {"n": 3}] + + # END is the endpoint's newest position when the read starts. + at_end = stream.read(topic=INPUTS, result_type=dict, after=END) + first = asyncio.ensure_future(at_end.__anext__()) + try: + for _ in range(100): + await producer.append({"n": "new"}) + done, _ = await asyncio.wait({first}, timeout=0.1) + if done: + break + record = await asyncio.wait_for(first, 30) + finally: + await at_end.aclose() + assert record.value == {"n": "new"} + + # A token resumes a read and last_n starts one, so the two are refused. + with pytest.raises(nexusrpc.HandlerError) as both: + await _dispatch( + handler, + READ_OPERATION, + ReadInput( + stream=_ref("wf-starts"), + after_token=newest[0].cursor.token, + last_n=1, + ), + ) + assert both.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + await handler.close() + + +async def test_an_activity_owner_crosses_the_endpoint(monkeypatch: pytest.MonkeyPatch): + # The reference carries the owner, so an activity's own stream is reached + # through the same two operations and lands on the store's activity + # accessor, apart from the workflow's topics of the same name. + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + front = _front(None) + own = front.get_activity_stream_handle(None, "act", workflow_id="wf") + await own.producer(topic=INPUTS, producer_id="model", attempt=1).append({"n": 1}) + records = await take(own.read(topic=INPUTS, result_type=dict), 1, timeout=30) + assert records[0].value == {"n": 1} + assert await _stored(store, "wf") == [] + stored = store.get_activity_stream_handle(None, "act", workflow_id="wf") + assert [r.value async for r in _take_data(stored, 1)] == [{"n": 1}] + await handler.close() + + +async def _take_data(stream: Any, count: int) -> AsyncIterator[Any]: + async for record in stream.read(topic=INPUTS, result_type=dict): + if record.kind is RecordKind.DATA: + yield record + count -= 1 + if count == 0: + return + + +async def test_an_owner_the_store_cannot_host_is_refused_under_its_own_class( + client: Client, monkeypatch: pytest.MonkeyPatch +): + # Whether an owner is supported is the store's answer, not the front's: + # the Workflow Streams store hosts only the streams of a workflow's + # activities, so a standalone activity's are refused, and its refusal + # reaches the caller as the class it raised. + handler = TemporalStreamsHandler(WorkflowStreamsProvider(), client) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + own = _front(None).get_activity_stream_handle(None, "act") + with pytest.raises(StreamUnsupportedError, match="activity"): + await own.producer(topic=INPUTS, producer_id="model", attempt=1).append( + {"n": 1} + ) + with pytest.raises(StreamUnsupportedError, match="activity"): + await take(own.read(topic=INPUTS), 1) + # No provider addresses a standalone stream yet, so that owner is refused + # for every store. + standalone = WireStreamRef(kind="standalone", stream_id="s-1", topic=INPUTS) + with pytest.raises(nexusrpc.HandlerError) as refused: + await _dispatch(handler, READ_OPERATION, ReadInput(stream=standalone)) + assert str(refused.value).startswith("StreamUnsupportedError: ") + await handler.close() + + +async def test_a_reference_missing_its_owner_id_is_the_callers_fault(): + handler = TemporalStreamsHandler(MemoryStreams(), None) + for ref in ( + WireStreamRef(kind="workflow", topic=INPUTS), + WireStreamRef(kind="activity", workflow_id="wf", topic=INPUTS), + WireStreamRef(kind="standalone", topic=INPUTS), + WireStreamRef(kind="workflow", workflow_id="wf", topic=""), + ): + with pytest.raises(nexusrpc.HandlerError) as failed: + await _dispatch(handler, READ_OPERATION, ReadInput(stream=ref)) + assert failed.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + with pytest.raises(nexusrpc.HandlerError) as failed: + request = _append("wf", 1, {"n": 1}) + request.stream = ref + await _dispatch(handler, APPEND_OPERATION, request) + assert failed.value.type is nexusrpc.HandlerErrorType.BAD_REQUEST + + +async def test_a_handle_names_its_stream_as_a_ref_that_opens_on_the_front( + monkeypatch: pytest.MonkeyPatch, +): + store = MemoryStreams() + handler = TemporalStreamsHandler(store, None) + monkeypatch.setattr(nexus, "_post", _in_process_endpoint(handler, [], [])) + front = _front(None) + handle = front.get_stream_handle(None, "wf", run_id="r1") + ref = handle.ref(topic=INPUTS) + assert ref == StreamRef.for_workflow("wf", run_id="r1", topic=INPUTS) + # Plain data: the default converter carries it the way an operation + # result or a workflow argument is carried. + converter = temporalio.converter.DataConverter.default.payload_converter + carried = converter.from_payloads([converter.to_payloads([ref])[0]], [StreamRef]) + assert carried == [ref] + # Opened on the front, the ref reaches the same stream and its topic is + # the handle's default; the client accessor goes through open_ref. The + # in-process endpoint is reached by id, so no client is needed to resolve it. + no_client: Any = None + opened = open_ref(front, no_client, ref) + await opened.producer(producer_id="model", attempt=1).append({"n": 1}) + records = await take(handle.read(topic=INPUTS, result_type=dict), 1, timeout=30) + assert records[0].value == {"n": 1} + assert opened.ref() == ref + # Every owner kind names itself. + assert front.get_activity_stream_handle( + None, "act", workflow_id="wf" + ).ref() == StreamRef.for_activity("act", workflow_id="wf") + assert front.get_standalone_stream_handle(None, "s-1").ref( + topic="t" + ) == StreamRef.for_standalone("s-1", topic="t") + await handler.close() + + +def test_an_sdk_stream_ref_crosses_the_wire_model_with_its_nulls(): + # The SDK's default converter writes null for every member a ref does not + # use. The wire model declares those members nullable, so a ref an + # operation returned as plain data parses into the same wire reference + # the front builds, and the wire form of that reference reads back as + # the SDK type with nothing lost. + converter = temporalio.converter.DataConverter.default.payload_converter + contract = ( + temporalio.converter.DataConverter.default._get_internal_payload_converter() + ) + for ref in ( + StreamRef.for_workflow("wf", topic=INPUTS), + StreamRef.for_workflow("wf", run_id="r1", topic=INPUTS), + StreamRef.for_activity("act", workflow_id="wf", topic=INPUTS), + StreamRef.for_standalone("s-1", topic=INPUTS), + ): + as_sdk = converter.to_payloads([ref])[0] + assert b"null" in as_sdk.data, as_sdk.data + wire = contract.from_payloads([as_sdk], [WireStreamRef])[0] + assert wire == WireStreamRef( + kind=ref.kind, + topic=ref.topic, + workflow_id=ref.workflow_id, + run_id=ref.run_id, + activity_id=ref.activity_id, + stream_id=ref.stream_id, + ) + back = converter.from_payloads([contract.to_payloads([wire])[0]], [StreamRef])[ + 0 + ] + assert back == ref + + +async def test_the_front_refuses_what_the_contract_cannot_carry(): + # No create and no seal operation: an owned stream cannot be closed by + # anyone, and a standalone stream is created and sealed on the store. + front = _front(None) + with pytest.raises(ValueError, match="only a standalone stream"): + await front.get_stream_handle(None, "wf").close() + with pytest.raises(StreamUnsupportedError, match="seals"): + await front.get_standalone_stream_handle(None, "s-1").close() + with pytest.raises(StreamUnsupportedError, match="creates"): + await front.create_standalone_stream(None, "s-1") + with pytest.raises(ValueError, match="stream_id"): + front.get_standalone_stream_handle(None, "") + + +# --------------------------------------------------------------------------- +# A stream as an operation result and as an operation input (June s8 b). +# --------------------------------------------------------------------------- + + +@dataclass +class ScoreUpdate: + """One score change.""" + + home: int + away: int + + +@dataclass +class GameRequest: + """Which game to start, and where its workflow runs.""" + + game_id: str + task_queue: str + + +@dataclass +class ScorePost: + """A score for the stream given, and whether it is the last one.""" + + stream: StreamRef + score: ScoreUpdate + last: bool = False + + +@dataclass +class CallerInput: + """The endpoint to call, the game to ask for, and the scores to post.""" + + endpoint: str + game_id: str + task_queue: str + scores: list[ScoreUpdate] + + +SCORES = streams.topic("scores", ScoreUpdate) +COMMANDS = streams.topic("commands", ScoreUpdate) + + +@workflow.defn +class ScoreboardGame: + """Reads posted scores on ``commands`` and publishes each on ``scores``.""" + + @workflow.run + async def run(self) -> int: + """Relay every posted score until the poster finishes, then finish.""" + posted = 0 + scores = workflow.stream_writer(SCORES) + async for record in workflow.stream_reader(COMMANDS): + if record.kind is RecordKind.FINISH: + break + if record.kind is RecordKind.DATA: + assert record.value is not None + scores.publish(record.value) + posted += 1 + scores.finish() + return posted + + +@nexusrpc.service +class ScoreboardService: + """Starts a game and hands back its stream; takes a stream to post to.""" + + start_game: nexusrpc.Operation[GameRequest, StreamRef] + post_score: nexusrpc.Operation[ScorePost, None] + + +@nexusrpc.handler.service_handler(service=ScoreboardService) +class ScoreboardHandler: + """Serves the two operations as a client of the stream front.""" + + @nexusrpc.handler.sync_operation + async def start_game( + self, _ctx: nexusrpc.handler.StartOperationContext, input: GameRequest + ) -> StreamRef: + """Start the game workflow and return a ref to its ``scores`` topic.""" + client = temporal_nexus.client() + # Reusing a running game makes a retried start hand back the same ref. + await client.start_workflow( + ScoreboardGame.run, + id=input.game_id, + task_queue=input.task_queue, + id_conflict_policy=WorkflowIDConflictPolicy.USE_EXISTING, + ) + return client.get_stream_handle(input.game_id).ref(topic=SCORES) + + @nexusrpc.handler.sync_operation + async def post_score( + self, ctx: nexusrpc.handler.StartOperationContext, input: ScorePost + ) -> None: + """Append the score to the stream the caller named.""" + # The request id is the producer, so a retried post writes once. + producer = ( + temporal_nexus.client() + .get_stream_handle(input.stream) + .producer(producer_id=f"post-{ctx.request_id}", attempt=1) + ) + await producer.append(input.score) + if input.last: + await producer.finish() + + +@workflow.defn +class ScoreboardCaller: + """Calls the operations and passes the ref on; it never reads the stream.""" + + @workflow.run + async def run(self, input: CallerInput) -> StreamRef: + """Start a game, post the scores to it, and return where to read them.""" + scoreboard = workflow.create_nexus_client( + service=ScoreboardService, endpoint=input.endpoint + ) + ref = await scoreboard.execute_operation( + ScoreboardService.start_game, + GameRequest(input.game_id, input.task_queue), + ) + for index, score in enumerate(input.scores): + await scoreboard.execute_operation( + ScoreboardService.post_score, + ScorePost( + ref.with_topic(COMMANDS), score, last=index == len(input.scores) - 1 + ), + ) + return ref + + +@live_only +async def test_an_operation_returns_a_stream_ref_and_another_takes_one(): + # June s8 (b) without the emulation: the operation's result type is the + # SDK's StreamRef, the calling workflow passes it on as data, the client + # opens it on the front it carries and long-polls it, and a second + # operation appends to a stream it was given as input. The handlers are + # clients of the front like any process; only the game workflow's worker + # holds the store. + store = WorkflowStreamsProvider() + plain = await _live_client() + game_id = f"streams-nexus-ref-{uuid.uuid4().hex}" + handler_tq = f"handlers-{game_id}" + stream_handler = TemporalStreamsHandler(store, plain) + + async with ( + _own_endpoint(plain, game_id, handler_tq) as endpoint, + Worker( + plain, + task_queue=handler_tq, + nexus_service_handlers=[stream_handler], + ), + ): + front = _live_front(endpoint) + config = plain.config() + config["plugins"] = [front] + fronted = Client(**config) + async with ( + Worker( + fronted, + task_queue=f"scoreboard-{game_id}", + nexus_service_handlers=[ScoreboardHandler()], + ), + Worker( + plain, + task_queue=f"tq-{game_id}", + workflows=[ScoreboardGame, ScoreboardCaller], + plugins=[store], + ), + ): + # The scoreboard service sits behind its own endpoint. + async with _own_endpoint( + plain, f"scoreboard-{game_id}", f"scoreboard-{game_id}" + ) as scoreboard_endpoint: + posted = [ScoreUpdate(1, 0), ScoreUpdate(1, 1), ScoreUpdate(2, 1)] + ref = await plain.execute_workflow( + ScoreboardCaller.run, + CallerInput(scoreboard_endpoint, game_id, f"tq-{game_id}", posted), + id=f"{game_id}-caller", + task_queue=f"tq-{game_id}", + ) + assert ref == StreamRef.for_workflow(game_id, topic=SCORES) + + # The ref is plain data the client opens on its own provider, + # here the front, and reads with the endpoint's long poll. It + # names the topic, not the record type, so the reader says + # what to decode as, the way a runtime topic name does. + records = await take( + fronted.get_stream_handle(ref).read(result_type=ScoreUpdate), 4, 60 + ) + assert [record.kind for record in records] == [ + RecordKind.DATA, + RecordKind.DATA, + RecordKind.DATA, + RecordKind.FINISH, + ] + assert [record.value for record in records[:3]] == posted + assert await plain.get_workflow_handle(game_id).result() == 3 + await stream_handler.close() diff --git a/tests/streams/test_streams_conformance.py b/tests/streams/test_streams_conformance.py index 993cba69b..eceb5af2b 100644 --- a/tests/streams/test_streams_conformance.py +++ b/tests/streams/test_streams_conformance.py @@ -11,6 +11,11 @@ skipped with a reason on a provider whose ``append()`` learns positions at read time. +``NexusStreams`` is deliberately absent from ``SETUPS``. It is a front with no +workflow half (its ``workflow_provider()`` raises), so there is no store for a +host workflow to own; ``test_nexus_provider`` covers it through an endpoint +that fronts a storage provider. + What this file pins down is what a provider owes: producer identity, retry deduplication, positions, supersession, topic addressing, cursor resumption, cursor ownership, releasing a read the caller stopped early, naming a stream