diff --git a/.gitignore b/.gitignore index 8cd439e05..15aa7014a 100644 --- a/.gitignore +++ b/.gitignore @@ -15,3 +15,6 @@ temporalio/bridge/temporal_sdk_bridge* tags /.claude tmpclaude-* + +# Output of a streams_demo run, written next to the script. +streams_demo/results-*/ diff --git a/pyproject.toml b/pyproject.toml index a83b34cd7..cacec1ad0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -267,7 +267,7 @@ reportUnnecessaryIsInstance = "none" reportUnnecessaryTypeIgnoreComment = "none" reportUnusedCallResult = "none" reportUnknownLambdaType = "none" -include = ["temporalio", "tests"] +include = ["temporalio", "tests", "streams_demo"] exclude = [ # Exclude auto generated files "temporalio/api", diff --git a/streams_demo/agent_loop.py b/streams_demo/agent_loop.py new file mode 100644 index 000000000..df5eac4d6 --- /dev/null +++ b/streams_demo/agent_loop.py @@ -0,0 +1,110 @@ +"""One agent loop, written once, run on every stream provider. + +Reads input, decides on it, publishes the decision, and runs an ordinary +Activity in the same workflow task; the Activity reports on its own +workflow's stream in turn. Also handles the two control records the contract +defines, so a retried producer and a finished topic are exercised rather than +described. + +This file is byte-identical in the server-side tree and the client-side tree. +Nothing in it names a provider: the workflow asks its runtime, the Activity +asks its context, and the process that runs them registered the provider once. +""" + +from __future__ import annotations + +from datetime import timedelta +from typing import Any + +from temporalio import activity, streams, workflow +from temporalio.streams import RecordKind + +# Defined once and shared by the workflow, the Activity and the demo's +# reader, so the type each topic carries is stated in one place. +DECISIONS = streams.topic("decisions", dict[str, Any]) +INPUTS = streams.topic("inputs", dict[str, Any]) +RECEIPTS = streams.topic("receipts", dict[str, Any]) + + +@activity.defn(name="RecordDecision") +async def record_decision(decision: dict[str, Any]) -> str: + """An ordinary command in the same task as the publish. + + Appends its receipt onto its own workflow's stream too, under the + Activity's own identity, so a reader outside sees the decision and the + record of it side by side. + """ + receipt = f"recorded:{decision['source']}:{decision['branch']}" + await activity.stream_handle().producer(topic=RECEIPTS).append({"receipt": receipt}) + return receipt + + +def decide(token: dict[str, Any]) -> dict[str, Any]: + """The decision the workflow is here to make.""" + if token["value"] % 2 == 0: + return { + "source": token["id"], + "branch": "even", + "computed": token["value"] * 10, + } + return {"source": token["id"], "branch": "odd", "computed": token["value"] + 100} + + +@workflow.defn(name="StreamContractDemo", sandboxed=False) +class AgentLoop: + """Read, decide, write, until the producer says it has finished.""" + + @workflow.run + async def run(self, limit: int) -> list[dict[str, Any]]: + """Decide on at most ``limit`` inputs, then return the trace.""" + inputs = workflow.stream_reader(INPUTS) + decisions = workflow.stream_writer(DECISIONS) + trace: list[dict[str, Any]] = [] + accepted = 0 + try: + async for record in inputs: + if record.kind is RecordKind.SUPERSEDED: + # A newer attempt of the same producer started writing. The + # decisions already published stand, so the workflow says so + # rather than pretending they can be withdrawn. + assert record.supersession is not None + trace.append( + { + "kind": "superseded", + "producer": record.producer_id, + "replaced": record.supersession.previous_attempt, + "attempt": record.supersession.attempt, + } + ) + decisions.publish( + {"retracting_attempt": record.supersession.previous_attempt} + ) + continue + if record.kind is RecordKind.FINISH: + # The producer says it is done, which is what ends the loop. + # Counting decisions instead would leave the terminal record + # unread and let the workflow finish while its producer is + # still writing. + trace.append({"kind": "finish", "producer": record.producer_id}) + break + assert record.value is not None + decision = decide(record.value) + decisions.publish(decision) + receipt = await workflow.execute_activity( + record_decision, + decision, + activity_id=f"decision-{decision['source']}", + start_to_close_timeout=timedelta(seconds=10), + ) + trace.append( + {"kind": "decision", "value": decision, "receipt": receipt} + ) + accepted += 1 + if accepted >= limit: + # A bound so a stuck producer cannot run this forever. The + # terminal record above is the ordinary way out. + break + finally: + inputs.close() + decisions.finish() + return trace diff --git a/streams_demo/provider_setup.py b/streams_demo/provider_setup.py new file mode 100644 index 000000000..2e4bcc350 --- /dev/null +++ b/streams_demo/provider_setup.py @@ -0,0 +1,36 @@ +"""Pick the provider for a demo run from the environment. + +``STREAMS_PROVIDER`` names a provider; this base tree carries only ``memory``, +and each provider branch adds its own name here. The demo needs a Temporal +server to run the workflow either way; ``TEMPORAL_ADDRESS`` points at it. +""" + +from __future__ import annotations + +import os + +from temporalio.streams.providers import ProviderPlugin +from temporalio.streams.providers.memory import MemoryStreams + +NAME = os.environ.get("STREAMS_PROVIDER", "memory") + +# The memory provider is not replay-safe, so its demo keeps the cache warm. +# Storage providers run with the smallest cache they support instead. +WORKFLOW_CACHE = int(os.environ.get("STREAMS_WORKFLOW_CACHE", "512")) + + +async def open() -> tuple[str, ProviderPlugin]: + """The server to connect to and the provider the worker and the client share.""" + if NAME != "memory": + raise SystemExit(f"this tree carries no stream provider named {NAME!r}") + return os.environ.get("TEMPORAL_ADDRESS", "localhost:7233"), MemoryStreams() + + +async def close(provider: ProviderPlugin) -> None: + """Let go of whatever :func:`open` acquired. + + The memory provider holds no connection, so this is its ``close()`` and + nothing more. A provider branch that opens one closes it the same way, so + the demo's teardown reads the same on every provider. + """ + await provider.close() diff --git a/streams_demo/run_demo.py b/streams_demo/run_demo.py new file mode 100644 index 000000000..769d7bf18 --- /dev/null +++ b/streams_demo/run_demo.py @@ -0,0 +1,174 @@ +"""Run the shared agent loop against whichever provider is configured. + +Three cases, the same on every provider: + +- read, decide, publish and an ordinary Activity in the same workflow task, + with the smallest workflow cache the provider supports, so that as much of + the run as it allows is rebuilt rather than remembered; +- a producer whose second attempt supersedes its first, which the reader has + to report and the workflow has to act on; +- an outside reader following what the workflow published, and the receipts + the Activity appended on the workflow's stream from inside its own context. + +The provider is registered once, on the client; the worker inherits it and +every context asks for its stream without naming it. Byte-identical in every +tree. ``provider_setup`` is what differs, and it is the only import here that +names a provider. +""" + +from __future__ import annotations + +import asyncio +import json +import sys +import time +import uuid +from pathlib import Path +from typing import Any + +from temporalio.api.enums.v1 import EventType +from temporalio.client import Client +from temporalio.streams import RecordKind, StreamHandle +from temporalio.worker import Worker + +# Run as a script rather than imported as a module, so the two siblings are +# reached by name off this directory rather than through a package path. +sys.path.insert(0, str(Path(__file__).resolve().parent)) +import provider_setup # noqa: E402 # pyright: ignore[reportImplicitRelativeImport] +from agent_loop import ( # noqa: E402 # pyright: ignore[reportImplicitRelativeImport] + DECISIONS, + INPUTS, + RECEIPTS, + AgentLoop, + record_decision, +) + +DECISION_LIMIT = 8 +EXPECTED_OUTPUT = 6 + + +async def collect_output(stream: StreamHandle, want: int) -> list[dict]: + """Read ``want`` decisions off the workflow's stream from outside it.""" + seen: list[dict] = [] + async for record in stream.read(topic=DECISIONS): + seen.append({"kind": record.kind.name, "value": record.value}) + if len(seen) >= want: + break + return seen + + +async def main() -> int: + """Run the demo once and write what happened next to this file.""" + out = Path(__file__).resolve().parent / f"results-{provider_setup.NAME}" + out.mkdir(exist_ok=True) + target, provider = await provider_setup.open() + client = await Client.connect(target, namespace="default", plugins=[provider]) + + uid = f"ai198-contract-{provider_setup.NAME}-" + uuid.uuid4().hex + record: dict[str, Any] = { + "provider": provider_setup.NAME, + "workflow_id": uid, + "target": target, + "max_cached_workflows": provider_setup.WORKFLOW_CACHE, + } + + async with Worker( + client, + task_queue=uid, + workflows=[AgentLoop], + activities=[record_decision], + max_cached_workflows=provider_setup.WORKFLOW_CACHE, + ): + handle = await client.start_workflow( + AgentLoop.run, DECISION_LIMIT, id=uid, task_queue=uid + ) + stream = client.get_stream_handle(uid) + output = asyncio.create_task(collect_output(stream, EXPECTED_OUTPUT)) + + # The first attempt writes two records and then stops, as a failed + # activity would. The second writes different inputs under the same + # logical producer, which is what the reader has to report. + first = stream.producer(topic=INPUTS, producer_id="model", attempt=1) + await first.append({"id": "r1", "value": 1}, {"id": "r2", "value": 2}) + await asyncio.sleep(0.5) + second = stream.producer(topic=INPUTS, producer_id="model", attempt=2) + await second.append({"id": "r3", "value": 3}, {"id": "r4", "value": 4}) + await second.finish() + + # A failed Workflow Task is not an outcome. The server rejects a + # completion that raced newly buffered events, and the retry usually + # gets through, so reading the first failure as the result reports a + # working run as a broken one. Wait for a terminal event, then report + # the retries separately so they are neither the headline nor hidden. + deadline = time.monotonic() + 120 + # Taken from the enum rather than written out, because guessing these + # numbers is how a run that completed gets reported as terminated. + terminal = { + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_COMPLETED: "completed", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_FAILED: "workflow_failed", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_TIMED_OUT: "execution_timed_out", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_CANCELED: "canceled", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_TERMINATED: "terminated", + EventType.EVENT_TYPE_WORKFLOW_EXECUTION_CONTINUED_AS_NEW: "continued_as_new", + } + while time.monotonic() < deadline: + history = await handle.fetch_history() + reached = [ + terminal[e.event_type] + for e in history.events + if e.event_type in terminal + ] + if reached: + record["outcome"] = reached[-1] + if record["outcome"] == "completed": + record["trace"] = await handle.result() + break + await asyncio.sleep(0.1) + else: + record["outcome"] = "timed_out_waiting" + history = await handle.fetch_history() + + retried = [ + e + for e in history.events + if e.event_type == EventType.EVENT_TYPE_WORKFLOW_TASK_FAILED + ] + record["task_failures"] = [ + { + "event_id": e.event_id, + "cause": int(e.workflow_task_failed_event_attributes.cause), + "message": e.workflow_task_failed_event_attributes.failure.message, + } + for e in retried + ] + + try: + record["observed_output"] = await asyncio.wait_for(output, timeout=20) + except asyncio.TimeoutError: + output.cancel() + record["observed_output"] = "timed_out" + + async def receipts() -> list[Any]: + # Ends by itself once the workflow is closed and the tail served. + return [ + r.value + async for r in stream.read(topic=RECEIPTS) + if r.kind is RecordKind.DATA + ] + + try: + record["receipts"] = await asyncio.wait_for(receipts(), timeout=20) + except asyncio.TimeoutError: + record["receipts"] = "timed_out" + + history = await handle.fetch_history() + record["history_events"] = len(history.events) + (out / "history.json").write_text(history.to_json()) + await provider_setup.close(provider) + (out / "results.json").write_text(json.dumps(record, indent=2) + "\n") + print(json.dumps(record, indent=2)) + return 0 if record["outcome"] == "completed" else 1 + + +if __name__ == "__main__": + sys.exit(asyncio.run(main()))