Skip to content

Commit d67e609

Browse files
authored
Require wait_for_stage parameter for Nexus Workflow Updates (#1883)
1 parent 3757b91 commit d67e609

4 files changed

Lines changed: 86 additions & 1 deletion

File tree

‎CHANGELOG.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@ to include examples, links to docs, or any other relevant information.
3737
remain available at runtime and retain their static type information.
3838
New code should depend on `temporalio-openai-agents` directly and import
3939
`temporalio.openai_agents`.
40+
- **Experimental**: Nexus Workflow Updates now require `wait_for_stage` to be explicitly set to `ACCEPTED`.
4041

4142
### Fixed
4243

‎temporalio/nexus/_operation_context.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -711,13 +711,16 @@ async def _start_nexus_operation_workflow_update( # pyright: ignore[reportUnuse
711711
update: str | Callable,
712712
arg: Any = temporalio.common._arg_unset,
713713
args: Sequence[Any] = [],
714+
wait_for_stage: temporalio.client.WorkflowUpdateStage,
714715
update_id: str | None = None,
715716
result_type: type | None = None,
716717
rpc_metadata: Mapping[str, str | bytes] = {},
717718
rpc_timeout: timedelta | None = None,
718719
run_id: str | None = None,
719720
first_execution_run_id: str | None = None,
720721
) -> temporalio.client.WorkflowUpdateHandle[Any]:
722+
if wait_for_stage != temporalio.client.WorkflowUpdateStage.ACCEPTED:
723+
raise ValueError("Only ACCEPTED wait stage is supported")
721724
# Default update ID to the Nexus request ID for retry-safety (matches sdk-go).
722725
update_id = update_id or temporal_context.nexus_context.request_id
723726
workflow_handle = temporal_context.client.get_workflow_handle(
@@ -728,7 +731,7 @@ async def _start_nexus_operation_workflow_update( # pyright: ignore[reportUnuse
728731
update,
729732
arg,
730733
args=args,
731-
wait_for_stage=temporalio.client.WorkflowUpdateStage.ACCEPTED, # hardcoded as nexus only supports async updates
734+
wait_for_stage=wait_for_stage,
732735
id=update_id,
733736
result_type=result_type,
734737
rpc_metadata=rpc_metadata,

‎temporalio/nexus/_temporal_client.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
Any,
1111
Concatenate,
1212
Generic,
13+
Literal,
1314
TypeVar,
1415
cast,
1516
overload,
@@ -294,6 +295,7 @@ async def start_workflow_update(
294295
workflow_id: str,
295296
update: temporalio.workflow.UpdateMethodMultiParam[[Any], ReturnType],
296297
*,
298+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
297299
update_id: str | None = None,
298300
rpc_metadata: Mapping[str, str | bytes] = {},
299301
rpc_timeout: timedelta | None = None,
@@ -311,6 +313,7 @@ async def start_workflow_update(
311313
],
312314
arg: ParamType,
313315
*,
316+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
314317
update_id: str | None = None,
315318
rpc_metadata: Mapping[str, str | bytes] = {},
316319
rpc_timeout: timedelta | None = None,
@@ -326,6 +329,7 @@ async def start_workflow_update(
326329
update: temporalio.workflow.UpdateMethodMultiParam[MultiParamSpec, ReturnType],
327330
*,
328331
args: MultiParamSpec.args, # type: ignore
332+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
329333
update_id: str | None = None,
330334
rpc_metadata: Mapping[str, str | bytes] = {},
331335
rpc_timeout: timedelta | None = None,
@@ -342,6 +346,7 @@ async def start_workflow_update(
342346
arg: Any = temporalio.common._arg_unset,
343347
*,
344348
args: Sequence[Any] = [],
349+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
345350
update_id: str | None = None,
346351
result_type: type[ReturnType] | None = None,
347352
rpc_metadata: Mapping[str, str | bytes] = {},
@@ -358,6 +363,7 @@ async def start_workflow_update(
358363
arg: Any = temporalio.common._arg_unset,
359364
*,
360365
args: Sequence[Any] = [],
366+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
361367
update_id: str | None = None,
362368
result_type: type | None = None,
363369
rpc_metadata: Mapping[str, str | bytes] = {},
@@ -679,6 +685,7 @@ async def start_workflow_update(
679685
arg: Any = temporalio.common._arg_unset,
680686
*,
681687
args: Sequence[Any] = [],
688+
wait_for_stage: Literal[temporalio.client.WorkflowUpdateStage.ACCEPTED],
682689
update_id: str | None = None,
683690
result_type: type | None = None,
684691
rpc_metadata: Mapping[str, str | bytes] = {},
@@ -699,6 +706,7 @@ async def start_workflow_update(
699706
update=update,
700707
arg=arg,
701708
args=args,
709+
wait_for_stage=wait_for_stage,
702710
update_id=update_id,
703711
result_type=result_type,
704712
rpc_metadata=rpc_metadata,

‎tests/nexus/test_temporal_operation.py‎

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
import uuid
33
from dataclasses import dataclass
44
from datetime import timedelta
5+
from typing import Any, cast
56

67
import nexusrpc
78
import pytest
@@ -24,6 +25,7 @@
2425
NexusOperationFailureError,
2526
WorkflowExecutionStatus,
2627
WorkflowFailureError,
28+
WorkflowUpdateStage,
2729
)
2830
from temporalio.common import (
2931
NexusOperationExecutionStatus,
@@ -116,6 +118,7 @@ class TestService:
116118
sync_result: Operation[Input, str]
117119
custom_cancel: Operation[str, None]
118120
update_op: Operation[Input, str]
121+
bad_update_stage_op: Operation[Input, str]
119122
query_op: Operation[str, bool]
120123
echo_activity: Operation[Input, str]
121124
error_activity: Operation[Input, None]
@@ -134,6 +137,7 @@ def __init__(self) -> None:
134137
self.started_custom_cancel_workflow = asyncio.Event()
135138
self.started_custom_cancel_activity = asyncio.Event()
136139
self.custom_cancel_activity_called = asyncio.Event()
140+
self.bad_update_stage_error: ValueError | None = None
137141

138142
@nexus.temporal_operation
139143
async def echo(
@@ -290,9 +294,30 @@ async def update_op(
290294
input.value,
291295
UpdatableWorkflow.do_update,
292296
input.update_value,
297+
wait_for_stage=WorkflowUpdateStage.ACCEPTED,
293298
update_id=input.update_id,
294299
)
295300

301+
@nexus.temporal_operation
302+
async def bad_update_stage_op(
303+
self,
304+
_ctx: nexus.TemporalStartOperationContext,
305+
client: nexus.TemporalNexusClient,
306+
input: Input,
307+
) -> nexus.TemporalOperationResult[str]:
308+
try:
309+
return await client.start_workflow_update(
310+
input.value,
311+
UpdatableWorkflow.do_update,
312+
input.update_value,
313+
# cast to bypass type checker
314+
wait_for_stage=cast(Any, WorkflowUpdateStage.COMPLETED),
315+
update_id=input.update_id,
316+
)
317+
except ValueError as err:
318+
self.bad_update_stage_error = err
319+
return nexus.TemporalOperationResult.sync(str(err))
320+
296321
@nexus.temporal_operation
297322
async def query_op(
298323
self,
@@ -749,6 +774,41 @@ async def test_temporal_operation_update_workflow_delayed(
749774
assert expected_backward_link in handler_links
750775

751776

777+
async def test_start_workflow_update_rejects_non_accepted_wait_for_stage(
778+
client: Client, env: WorkflowEnvironment
779+
) -> None:
780+
if env.supports_time_skipping:
781+
pytest.skip("Update workflow tests don't work with time-skipping server")
782+
task_queue = str(uuid.uuid4())
783+
endpoint_name = make_nexus_endpoint_name(task_queue)
784+
await env.create_nexus_endpoint(endpoint_name, task_queue)
785+
service_handler = TestServiceHandler()
786+
async with Worker(
787+
env.client,
788+
task_queue=task_queue,
789+
nexus_service_handlers=[service_handler],
790+
workflows=[UpdatableWorkflow, BadUpdateStageCaller],
791+
):
792+
update_workflow_id = f"updatable-workflow-{uuid.uuid4()}"
793+
await client.start_workflow(
794+
UpdatableWorkflow.run, id=update_workflow_id, task_queue=task_queue
795+
)
796+
result = await client.execute_workflow(
797+
BadUpdateStageCaller.run,
798+
Input(
799+
value=update_workflow_id,
800+
task_queue=task_queue,
801+
update_value="Created",
802+
),
803+
task_queue=task_queue,
804+
id=f"bad-update-stage-caller-{uuid.uuid4()}",
805+
)
806+
807+
assert isinstance(service_handler.bad_update_stage_error, ValueError)
808+
assert result == str(service_handler.bad_update_stage_error)
809+
assert result == "Only ACCEPTED wait stage is supported"
810+
811+
752812
async def test_temporal_operation_cancel_rejects_unknown_tokens():
753813
class FakeNexusTaskCancellation(OperationTaskCancellation):
754814
def is_cancelled(self) -> bool:
@@ -1649,6 +1709,19 @@ async def run(self, input: Input) -> str:
16491709
return await op_handle
16501710

16511711

1712+
@workflow.defn
1713+
class BadUpdateStageCaller:
1714+
"""Caller workflow for an update op that requests an unsupported update stage."""
1715+
1716+
@workflow.run
1717+
async def run(self, input: Input) -> str:
1718+
client = workflow.create_nexus_client(
1719+
service=TestService,
1720+
endpoint=make_nexus_endpoint_name(input.task_queue),
1721+
)
1722+
return await client.execute_operation(TestService.bad_update_stage_op, input)
1723+
1724+
16521725
@workflow.defn
16531726
class UpdatableWorkflow:
16541727
"""Workflow that accepts updates and exits when it receives a specific status"""

0 commit comments

Comments
 (0)