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
32 changes: 20 additions & 12 deletions temporalio/contrib/external_workflow_streams/_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@
and diverge from the live run.
"""

#: Where per-Run subscription state hangs off the Workflow instance. Reserved,
#: Where per-Run subscription state hangs off the SDK's Run object. Reserved,
#: and deliberately not in the ``__temporal_workflow_stream*`` namespace the
#: shipped contrib feature already owns.
_RUN_STATE_ATTR = "__temporal_external_stream_state"
Expand Down Expand Up @@ -175,9 +175,11 @@ def discard_pending(self, wait_id: int) -> None:
class _RunState:
"""Per-Run subscription bookkeeping.

Lives on the Workflow instance rather than in a module global, so it shares
the instance's lifetime exactly -- a module global would outlive an evicted
Run and hand its wait ids to the next one.
Lives on the SDK's object for the Run rather than in a module global, so it
shares the Run's lifetime exactly -- a module global would outlive an evicted
Run and hand its wait ids to the next one. Not on the user's Workflow object
either: that object does not exist while its ``@workflow.init`` constructor
runs, and a constructor may publish.
"""

runtime: ExternalStreamRuntime | None = None
Expand All @@ -186,26 +188,32 @@ class _RunState:
pending: dict[int, Any] = field(default_factory=dict)


def _run_holder() -> Any:
"""The object per-Run state hangs off: the SDK's runtime for the current Run."""
return temporalio.workflow._Runtime.current()


def _run_state() -> _RunState:
instance = temporalio.workflow.instance()
state = getattr(instance, _RUN_STATE_ATTR, None)
holder = _run_holder()
state = getattr(holder, _RUN_STATE_ATTR, None)
if state is None:
state = _RunState()
setattr(instance, _RUN_STATE_ATTR, state)
setattr(holder, _RUN_STATE_ATTR, state)
return state


def _install_runtime( # pyright: ignore[reportUnusedFunction]
instance: Any, runtime: ExternalStreamRuntime
holder: Any, runtime: ExternalStreamRuntime
) -> None:
"""Gives a Workflow instance its handle to the Worker's manager.
"""Gives a Run its handle to the Worker's manager.

Called by the Worker when it creates the instance.
Called by the Worker with its object for the Run, before the Workflow's
constructor runs, so a constructor can subscribe and publish.
"""
state = getattr(instance, _RUN_STATE_ATTR, None)
state = getattr(holder, _RUN_STATE_ATTR, None)
if state is None:
state = _RunState()
setattr(instance, _RUN_STATE_ATTR, state)
setattr(holder, _RUN_STATE_ATTR, state)
state.runtime = runtime


Expand Down
27 changes: 13 additions & 14 deletions temporalio/worker/_workflow_instance.py
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,19 @@ def __init__(self, det: WorkflowInstanceDetails) -> None:
self._extern_functions = det.extern_functions
self._external_stream_runtime = det.external_stream_runtime
self._external_streams_configured = det.external_streams_configured
# Installed on this object, before the Workflow's constructor runs, so a
# ``@workflow.init`` constructor can publish and subscribe. Per-Run state
# must share the Run's lifetime, which this object has and a module global
# would not.
if (
self._external_stream_runtime is not None
and self._external_streams_configured
):
from temporalio.contrib.external_workflow_streams._api import (
_install_runtime,
)

_install_runtime(self, self._external_stream_runtime)
self._pending_replay_finish: Any = None
"""A marker replay whose last segment the activation's own drain serves.

Expand Down Expand Up @@ -3211,20 +3224,6 @@ def _instantiate_workflow_object(self) -> Any:
else:
workflow_instance = self._defn.cls()

# Hand the workflow object its External Stream handle. It goes on the
# object rather than in a module global because per-Run state must share
# the Run's lifetime exactly -- a global would outlive an evicted Run and
# hand its wait ids to the next one.
if (
self._external_stream_runtime is not None
and self._external_streams_configured
):
from temporalio.contrib.external_workflow_streams._api import (
_install_runtime,
)

_install_runtime(workflow_instance, self._external_stream_runtime)

if self._defn.versioning_behavior:
self._versioning_behavior = self._defn.versioning_behavior
# If there's a dynamic config function, call it now after we've instantiated the object
Expand Down
21 changes: 16 additions & 5 deletions tests/contrib/external_workflow_streams/test_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@
import pytest

import temporalio.converter
import temporalio.workflow
from temporalio.contrib.external_workflow_streams import _api
from temporalio.contrib.external_workflow_streams._api import (
DEFAULT_IDLE_TIMEOUT,
Expand Down Expand Up @@ -129,7 +128,10 @@ class FakeInstance:
@pytest.fixture
def workflow_instance(monkeypatch: pytest.MonkeyPatch) -> FakeInstance:
instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
return instance


Expand Down Expand Up @@ -223,7 +225,10 @@ def test_wait_ids_reproduce_across_two_runs_of_the_same_code(

def run_once() -> list[int]:
instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
_install_runtime(instance, FakeRuntime())
return [
external_stream.topic(name).subscribe().wait_id
Expand Down Expand Up @@ -257,12 +262,18 @@ def test_a_second_run_restarts_the_counter(monkeypatch: pytest.MonkeyPatch) -> N
A module global would outlive the Run and hand its wait ids to the next one.
"""
first_instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: first_instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: first_instance,
)
_install_runtime(first_instance, FakeRuntime())
external_stream.topic("tokens").subscribe()

second_instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: second_instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: second_instance,
)
_install_runtime(second_instance, FakeRuntime())

assert external_stream.topic("tokens").subscribe().wait_id == 1
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,6 @@
import pytest

import temporalio.converter
import temporalio.workflow
from temporalio.contrib.external_workflow_streams import _runtime as _runtime_module
from temporalio.contrib.external_workflow_streams._annotation import (
MAX_ANNOTATION_BYTES,
Expand Down Expand Up @@ -73,7 +72,10 @@ class FakeInstance:
@pytest.fixture
def workflow_instance(monkeypatch: pytest.MonkeyPatch) -> FakeInstance:
instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
return instance


Expand Down
16 changes: 12 additions & 4 deletions tests/contrib/external_workflow_streams/test_multi_stream.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,6 @@
import pytest_asyncio

import temporalio.converter
import temporalio.workflow
from temporalio import workflow
from temporalio.client import Client
from temporalio.contrib.external_workflow_streams._annotation import decode_annotation
Expand Down Expand Up @@ -215,7 +214,10 @@ class Instance:
"""Stands in for the Workflow object the per-Run state hangs off."""

instance = Instance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
runtime = make_runtime(manager, backend, idle=timedelta(seconds=1))
_install_runtime(instance, runtime) # type: ignore[arg-type]

Expand Down Expand Up @@ -518,7 +520,10 @@ class FakeInstance:
@pytest.fixture
def fake_runtime(monkeypatch: pytest.MonkeyPatch) -> FakeRuntime:
instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
runtime = FakeRuntime()
_install_runtime(instance, runtime) # type: ignore[arg-type]
return runtime
Expand Down Expand Up @@ -802,7 +807,10 @@ class Instance:
pass

instance = Instance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
runtime = make_runtime(manager, backend)
_install_runtime(instance, runtime) # type: ignore[arg-type]
return runtime
Expand Down
5 changes: 4 additions & 1 deletion tests/contrib/external_workflow_streams/test_replay.py
Original file line number Diff line number Diff line change
Expand Up @@ -2103,7 +2103,10 @@ async def test_a_segment_replays_in_its_recorded_cross_stream_order(

runtime = make_runtime(manager, backend)
instance = FakeInstance()
monkeypatch.setattr(temporalio.workflow, "instance", lambda: instance)
monkeypatch.setattr(
"temporalio.contrib.external_workflow_streams._api._run_holder",
lambda: instance,
)
_install_runtime(instance, runtime)

runtime.begin_replay(plan.annotation.header.streams)
Expand Down
54 changes: 54 additions & 0 deletions tests/contrib/external_workflow_streams/test_worker_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,32 @@ async def run(self) -> None:
timer.cancel()


@workflow.defn
class SubscribeInInitWorkflow:
"""Subscribes from its ``@workflow.init`` constructor and consumes in ``run``.

The user's object does not exist yet while the constructor runs, so per-Run
stream state cannot hang off it.
"""

@workflow.init
def __init__(self, expected: int) -> None:
self._tokens = (
external_stream.with_options(idle_timeout=timedelta(seconds=30))
.topic("tokens", type=str)
.subscribe()
)

@workflow.run
async def run(self, expected: int) -> list[str]:
seen: list[str] = []
async for token in self._tokens:
seen.append(token)
if len(seen) >= expected:
break
return seen


@pytest.fixture
def backend() -> MemoryStreamBackend:
return MemoryStreamBackend()
Expand Down Expand Up @@ -210,6 +236,34 @@ async def test_a_workflow_consumes_records_it_never_read_itself(
]


async def test_a_constructor_can_subscribe(
client: Client, backend: MemoryStreamBackend
) -> None:
task_queue = f"tq-{uuid.uuid4()}"
async with Worker(
client,
task_queue=task_queue,
workflows=[SubscribeInInitWorkflow],
external_stream_backend=backend,
):
handle = await client.start_workflow(
SubscribeInInitWorkflow.run,
2,
id=f"wf-{uuid.uuid4()}",
task_queue=task_queue,
)
description = await handle.describe()
key = StreamKey(
client.namespace,
handle.id,
description.raw_description.workflow_execution_info.first_run_id,
"tokens",
)
await publish(backend, key, ["alpha", "beta"])

assert await asyncio.wait_for(handle.result(), 30) == ["alpha", "beta"]


async def test_no_stream_payload_reaches_history(
client: Client, backend: MemoryStreamBackend
) -> None:
Expand Down
Loading