Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -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-*/
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
110 changes: 110 additions & 0 deletions streams_demo/agent_loop.py
Original file line number Diff line number Diff line change
@@ -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
36 changes: 36 additions & 0 deletions streams_demo/provider_setup.py
Original file line number Diff line number Diff line change
@@ -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()
174 changes: 174 additions & 0 deletions streams_demo/run_demo.py
Original file line number Diff line number Diff line change
@@ -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()))
Loading