diff --git a/temporalio/contrib/external_workflow_streams/_api.py b/temporalio/contrib/external_workflow_streams/_api.py index c89986954..feca4ca7d 100644 --- a/temporalio/contrib/external_workflow_streams/_api.py +++ b/temporalio/contrib/external_workflow_streams/_api.py @@ -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" @@ -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 @@ -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 diff --git a/temporalio/worker/_workflow_instance.py b/temporalio/worker/_workflow_instance.py index c7a8ae7ca..8dcd507ff 100644 --- a/temporalio/worker/_workflow_instance.py +++ b/temporalio/worker/_workflow_instance.py @@ -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. @@ -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 diff --git a/tests/contrib/external_workflow_streams/test_api.py b/tests/contrib/external_workflow_streams/test_api.py index 08d725761..86f36bf33 100644 --- a/tests/contrib/external_workflow_streams/test_api.py +++ b/tests/contrib/external_workflow_streams/test_api.py @@ -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, @@ -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 @@ -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 @@ -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 diff --git a/tests/contrib/external_workflow_streams/test_delivery_budget.py b/tests/contrib/external_workflow_streams/test_delivery_budget.py index 79386d6fe..3c766417b 100644 --- a/tests/contrib/external_workflow_streams/test_delivery_budget.py +++ b/tests/contrib/external_workflow_streams/test_delivery_budget.py @@ -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, @@ -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 diff --git a/tests/contrib/external_workflow_streams/test_multi_stream.py b/tests/contrib/external_workflow_streams/test_multi_stream.py index c90acb21c..2bdb5c90b 100644 --- a/tests/contrib/external_workflow_streams/test_multi_stream.py +++ b/tests/contrib/external_workflow_streams/test_multi_stream.py @@ -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 @@ -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] @@ -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 @@ -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 diff --git a/tests/contrib/external_workflow_streams/test_replay.py b/tests/contrib/external_workflow_streams/test_replay.py index bc9ee978c..d0d15490b 100644 --- a/tests/contrib/external_workflow_streams/test_replay.py +++ b/tests/contrib/external_workflow_streams/test_replay.py @@ -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) diff --git a/tests/contrib/external_workflow_streams/test_worker_integration.py b/tests/contrib/external_workflow_streams/test_worker_integration.py index cf797719f..f2ed9d9fb 100644 --- a/tests/contrib/external_workflow_streams/test_worker_integration.py +++ b/tests/contrib/external_workflow_streams/test_worker_integration.py @@ -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() @@ -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: