22import uuid
33from dataclasses import dataclass
44from datetime import timedelta
5+ from typing import Any , cast
56
67import nexusrpc
78import pytest
2425 NexusOperationFailureError ,
2526 WorkflowExecutionStatus ,
2627 WorkflowFailureError ,
28+ WorkflowUpdateStage ,
2729)
2830from 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+
752812async 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
16531726class UpdatableWorkflow :
16541727 """Workflow that accepts updates and exits when it receives a specific status"""
0 commit comments