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
9 changes: 6 additions & 3 deletions temporalio/streams/_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,10 +322,13 @@ def open_writer(self, topic: str) -> WriteSink:
...

def on_workflow_start(self) -> None:
"""Called before the workflow function runs.
"""Called before the workflow function runs, and before the first task's handlers.

A provider that serves outside readers through handlers on the
workflow registers them here, before the first task completes.
After the workflow's own ``__init__`` and before any Signal or Update
of the first task is handled, which the SDK does ahead of the
workflow function. A provider that serves outside readers through
handlers on the workflow registers them here, so an Update that
arrives with the first task finds them.
"""
...

Expand Down
1 change: 1 addition & 0 deletions temporalio/worker/_replayer.py
Original file line number Diff line number Diff line change
Expand Up @@ -290,6 +290,7 @@ def on_eviction_hook(
default_workflow_logic_flags=frozenset(
self._default_workflow_logic_flags
),
stream_provider=self._config.get("stream_provider"),
)
external_storage = data_converter.external_storage
storage_driver_types = (
Expand Down
6 changes: 6 additions & 0 deletions temporalio/worker/_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -468,6 +468,11 @@ def _init_from_config(self, client: temporalio.client.Client, config: WorkerConf

# Prepend applicable client interceptors to the given ones
client_config = config["client"].config(active_config=True) # type: ignore[reportTypedDictNotRequiredAccess]
# A provider registered on the client serves its workers too, so one
# registration covers every context that asks for a stream.
stream_provider = config.get("stream_provider") or client_config.get(
"stream_provider"
)
interceptors_from_client = cast(
list[Interceptor],
[i for i in client_config["interceptors"] if isinstance(i, Interceptor)],
Expand Down Expand Up @@ -572,6 +577,7 @@ def check_activity(activity: str):
encode_headers=client_config["header_codec_behavior"]
!= HeaderCodecBehavior.NO_CODEC,
max_workflow_task_external_storage_concurrency=max_workflow_task_external_storage_concurrency,
stream_provider=stream_provider,
)

tuner = config.get("tuner")
Expand Down
60 changes: 60 additions & 0 deletions temporalio/worker/_workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from dataclasses import dataclass
from datetime import timezone
from types import TracebackType
from typing import Any, cast

import temporalio.api.common.v1
import temporalio.bridge.proto.common
Expand All @@ -24,6 +25,7 @@
import temporalio.converter
import temporalio.converter._extstore
import temporalio.exceptions
import temporalio.streams
import temporalio.workflow
from temporalio.bridge.worker import PollShutdownError
from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo
Expand All @@ -35,9 +37,11 @@
_relax_sandbox_for_debugger,
)
from ._interceptor import (
ExecuteWorkflowInput,
Interceptor,
WorkflowInboundInterceptor,
WorkflowInterceptorClassInput,
WorkflowOutboundInterceptor,
)
from ._workflow_instance import (
_DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS,
Expand All @@ -55,6 +59,55 @@
LOG_PROTOS = False


class _StreamHooksInterceptor(WorkflowInboundInterceptor):
"""Brackets the workflow function with the stream provider's lifecycle hooks.

Installed by the worker when it has a stream provider, so no workflow
code has to call anything before it runs or before it returns. The start
hook runs when the instance's loop first turns, after the workflow's own
``__init__`` and before the first task's Signals and Updates are handled,
because the SDK handles those ahead of the workflow function and a handler
registered any later would be missed by an Update that arrives with that
task. The finish hook runs when the function returns, raises or continues as new,
because a provider that parked a reader against the run has to let go
either way.
It does not run when the run is being evicted from the cache or when the
abandoned coroutine is collected: neither is the workflow ending, the
instance's state is not to be touched during eviction, and at collection
time the runtime on the thread belongs to whichever workflow happens to
be running, so the hook would act on that one.
"""

def init(self, outbound: WorkflowOutboundInterceptor) -> None:
super().init(outbound)
# The hook has to run after the workflow's own __init__, which may
# register handlers the provider adopts, and before the first task's
# Signals and Updates are handled, which the SDK does ahead of the
# workflow function. The instance is its own event loop and nothing
# is queued on it yet, so a callback queued now runs first when that
# loop first turns, which is after every job of the activation has
# been applied and before any task they created takes a step.
runtime = temporalio.workflow._Runtime.current()
loop = cast(asyncio.AbstractEventLoop, cast(object, runtime))
loop.call_soon(lambda: runtime.workflow_streams().provider.on_workflow_start())

async def execute_workflow(self, input: ExecuteWorkflowInput) -> Any:
runtime = temporalio.workflow._Runtime.current()
provider = runtime.workflow_streams().provider
try:
result = await self.next.execute_workflow(input)
except GeneratorExit:
raise
except BaseException:
# Eviction cancels the primary task the same way a workflow
# cancellation does, and only the cancellation is a run ending.
if not runtime.workflow_is_evicting():
await provider.on_workflow_finish()
raise
await provider.on_workflow_finish()
return result


# Value was chosen abitrarily as a small number that allows some concurrency and prevents
# large numbers of concurrent external storage operations causing resource contention.
# This default limit is per workflow task activation and does not limit the total number
Expand Down Expand Up @@ -104,6 +157,7 @@ def __init__(
encode_headers: bool,
max_workflow_task_external_storage_concurrency: int,
default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] | None = None,
stream_provider: temporalio.streams.StreamProvider | None = None,
) -> None:
# Debug mode is enabled if specified or if the TEMPORAL_DEBUG env var is truthy
debug_mode = debug_mode or bool(os.environ.get("TEMPORAL_DEBUG"))
Expand Down Expand Up @@ -162,6 +216,11 @@ def __init__(
__temporal_assert_local_activity_valid=assert_local_activity_valid,
)
)
self._stream_provider = stream_provider
if stream_provider is not None:
# Innermost, so the lifecycle hooks bracket the workflow function
# itself, after every user interceptor has done its own setup.
self._interceptor_classes.append(_StreamHooksInterceptor)

self._workflow_failure_exception_types = workflow_failure_exception_types
self._patch_activation_callback = patch_activation_callback
Expand Down Expand Up @@ -743,6 +802,7 @@ def _create_workflow_instance(
last_completion_result=init.last_completion_result,
last_failure=last_failure,
default_workflow_logic_flags=frozenset(self._default_workflow_logic_flags),
stream_provider=self._stream_provider,
)
if defn.sandboxed:
return self._workflow_runner.create_instance(det)
Expand Down
22 changes: 22 additions & 0 deletions temporalio/worker/_workflow_instance.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,9 @@
import temporalio.converter
import temporalio.exceptions
import temporalio.nexus.system
import temporalio.streams
import temporalio.workflow
import temporalio.workflow._streams
from temporalio.converter import StorageDriverStoreContext, StorageDriverWorkflowInfo
from temporalio.converter._payload_converter import (
_TemporalTransferTypePayloadConverter,
Expand Down Expand Up @@ -184,6 +186,7 @@ class WorkflowInstanceDetails:
default_workflow_logic_flags: frozenset[_WorkflowLogicFlag] = field(
default_factory=lambda: _DEFAULT_ENABLED_WORKFLOW_LOGIC_FLAGS
)
stream_provider: temporalio.streams.StreamProvider | None = None


class WorkflowInstance(ABC):
Expand Down Expand Up @@ -303,6 +306,8 @@ def __init__(self, det: WorkflowInstanceDetails) -> None:
det.worker_level_failure_exception_types
)
self._patch_activation_callback = det.patch_activation_callback
self._stream_provider = det.stream_provider
self._streams: temporalio.workflow._streams._WorkflowStreams | None = None
# Keyed by channel name; one subscription per channel per run
self._channel_subscriptions: dict[
str, temporalio.workflow.ChannelSubscription
Expand Down Expand Up @@ -1392,6 +1397,9 @@ def workflow_instance(self) -> Any:
def workflow_is_continue_as_new_suggested(self) -> bool:
return self._continue_as_new_suggested

def workflow_is_evicting(self) -> bool:
return self._deleting

def workflow_is_target_worker_deployment_version_changed(self) -> bool:
return self._target_worker_deployment_version_changed

Expand Down Expand Up @@ -1844,6 +1852,20 @@ async def workflow_start_nexus_operation(
)
)

def workflow_streams(self) -> temporalio.workflow._streams._WorkflowStreams:
if self._streams is None:
if self._stream_provider is None:
raise RuntimeError(
"no stream provider is configured on this worker; pass one with "
"Worker(plugins=[provider]) or stream_provider="
)
# The workflow half is made per instance, so whatever it keeps
# dies with the instance the way handlers do.
self._streams = temporalio.workflow._streams._WorkflowStreams(
self._stream_provider.workflow_provider()
)
return self._streams

def workflow_subscribe_channel(
self, channel: str
) -> temporalio.workflow.ChannelSubscription:
Expand Down
10 changes: 10 additions & 0 deletions temporalio/workflow/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -157,6 +157,12 @@
logger,
unsafe,
)
from ._streams import (
StreamReader,
StreamWriter,
stream_reader,
stream_writer,
)
from ._workflow_ops import (
ChildWorkflowCancellationType,
ChildWorkflowConfig,
Expand Down Expand Up @@ -268,6 +274,10 @@
"Notification",
"linked_channel",
"subscribe_channel",
"StreamReader",
"StreamWriter",
"stream_reader",
"stream_writer",
"ChildWorkflowCancellationType",
"ChildWorkflowConfig",
"ChildWorkflowHandle",
Expand Down
14 changes: 14 additions & 0 deletions temporalio/workflow/_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from ._event_groups import EventGroup
from ._exceptions import ContinueAsNewVersioningBehavior, VersioningIntent
from ._nexus import NexusOperationCancellationType, NexusOperationHandle
from ._streams import _WorkflowStreams
from ._workflow_ops import (
ChildWorkflowCancellationType,
ChildWorkflowHandle,
Expand Down Expand Up @@ -354,6 +355,16 @@ def workflow_instance(self) -> Any: ...
@abstractmethod
def workflow_is_continue_as_new_suggested(self) -> bool: ...

@abstractmethod
def workflow_is_evicting(self) -> bool:
"""Whether this instance is being dropped from the cache rather than ending.

Eviction cancels the primary task the way a workflow cancellation
does, so anything that runs on the way out has to be able to tell the
two apart. Instance state must not be touched while this is true.
"""
...

@abstractmethod
def workflow_is_target_worker_deployment_version_changed(self) -> bool: ...

Expand Down Expand Up @@ -500,6 +511,9 @@ async def workflow_start_nexus_operation(
event_groups: Sequence[EventGroup] | None = None,
) -> NexusOperationHandle[OutputT]: ...

@abstractmethod
def workflow_streams(self) -> _WorkflowStreams: ...

@abstractmethod
def workflow_subscribe_channel(self, channel: str) -> ChannelSubscription: ...

Expand Down
Loading
Loading