88import pytest
99
1010from durable_workflow import Client , Worker , serializer , workflow
11+ from durable_workflow .workflow import commands_to_server_commands , replay
1112
1213MEMO_ENTRIES : dict [str , Any ] = {
1314 "binary" : b"same" ,
2627 "CGxvbmcEDgxuZXN0ZWQOBAphbHBoYQQCCGJldGEEBAAIdGV4dAoIc2FtZQA="
2728)
2829
30+ CONTINUE_INITIAL_MEMO = {
31+ "existing" : "retained" ,
32+ "overwritten" : "before" ,
33+ }
34+
35+ CONTINUE_MEMO_PATCH = {
36+ "added" : "from-upsert" ,
37+ "overwritten" : "after" ,
38+ }
39+
40+ CONTINUE_MERGED_MEMO = {
41+ "added" : "from-upsert" ,
42+ "existing" : "retained" ,
43+ "overwritten" : "after" ,
44+ }
45+
2946
3047@workflow .defn (name = "tests.memo-restart-python" )
3148class MemoRestartWorkflow :
@@ -46,6 +63,15 @@ def run(self, ctx): # type: ignore[no-untyped-def]
4663 return "python-replayed-memo"
4764
4865
66+ @workflow .defn (name = "tests.memo-yielded-continue-python" )
67+ class MemoYieldedContinueWorkflow :
68+ def run (self , ctx , generation : int ): # type: ignore[no-untyped-def]
69+ if generation == 0 :
70+ yield ctx .upsert_memo (CONTINUE_MEMO_PATCH )
71+ yield ctx .continue_as_new (1 )
72+ return "python-continued-with-memo"
73+
74+
4975async def _wait_until_waiting_with_memo (handle : Any ) -> Any :
5076 deadline = asyncio .get_running_loop ().time () + 30
5177 last_description = None
@@ -94,6 +120,7 @@ async def _seed_worker_contract(client: Client, worker: Worker) -> None:
94120 workflow_definition_fingerprints = worker .workflow_definition_fingerprints ,
95121 workflow_command_contracts = worker .workflow_command_contracts ,
96122 supported_activity_types = [],
123+ capabilities = ["memo_upserts" ],
97124 )
98125
99126
@@ -175,3 +202,105 @@ async def test_fresh_python_worker_replays_persisted_server_memo_history(
175202 final_memo_events = _memo_events (final_history )
176203 assert len (final_memo_events ) == 1
177204 _assert_typed_memo_event (final_memo_events [0 ])
205+
206+
207+ @pytest .mark .asyncio
208+ async def test_yielded_continue_inherits_merged_memo_after_worker_restart (
209+ server_url : str ,
210+ server_token : str ,
211+ ) -> None :
212+ suffix = uuid .uuid4 ().hex [:8 ]
213+ task_queue = f"memo-continue-python-{ suffix } "
214+ workflow_id = f"memo-continue-python-{ suffix } "
215+
216+ async with Client (server_url , token = server_token , namespace = "default" ) as client :
217+ first_worker = Worker (
218+ client ,
219+ task_queue = task_queue ,
220+ workflows = [MemoYieldedContinueWorkflow ],
221+ activities = [],
222+ worker_id = f"memo-continue-before-{ suffix } " ,
223+ )
224+ await _seed_worker_contract (client , first_worker )
225+
226+ handle = await client .start_workflow (
227+ workflow_type = "tests.memo-yielded-continue-python" ,
228+ task_queue = task_queue ,
229+ workflow_id = workflow_id ,
230+ input = [0 ],
231+ memo = CONTINUE_INITIAL_MEMO ,
232+ )
233+ assert handle .run_id is not None
234+ first_run_id = handle .run_id
235+
236+ try :
237+ first_task = await client .poll_workflow_task (
238+ worker_id = first_worker .worker_id ,
239+ task_queue = task_queue ,
240+ timeout = 10.0 ,
241+ )
242+ assert first_task is not None , "expected the initial workflow task"
243+
244+ payload_codec = first_task .get ("payload_codec" ) or serializer .AVRO_CODEC
245+ decoded = serializer .decode_envelope (
246+ first_task .get ("arguments" ),
247+ codec = payload_codec ,
248+ )
249+ start_input = decoded if isinstance (decoded , list ) else [decoded ]
250+ outcome = replay (
251+ MemoYieldedContinueWorkflow ,
252+ first_task .get ("history_events" , []),
253+ start_input ,
254+ run_id = first_task .get ("run_id" , "" ),
255+ )
256+ commands = commands_to_server_commands (
257+ outcome .commands ,
258+ task_queue ,
259+ payload_codec = payload_codec ,
260+ )
261+ assert [command ["type" ] for command in commands ] == [
262+ "upsert_memo" ,
263+ "continue_as_new" ,
264+ ]
265+
266+ await client .complete_workflow_task (
267+ task_id = first_task ["task_id" ],
268+ lease_owner = first_worker .worker_id ,
269+ workflow_task_attempt = first_task .get ("workflow_task_attempt" , 1 ),
270+ commands = commands ,
271+ )
272+ finally :
273+ await client .deregister_worker_registration (first_worker .worker_id )
274+
275+ replacement_worker = Worker (
276+ client ,
277+ task_queue = task_queue ,
278+ workflows = [MemoYieldedContinueWorkflow ],
279+ activities = [],
280+ worker_id = f"memo-continue-after-{ suffix } " ,
281+ poll_timeout = 1.0 ,
282+ shutdown_timeout = 5.0 ,
283+ )
284+ completed = await replacement_worker .run_until (
285+ workflow_id = workflow_id ,
286+ timeout = 30.0 ,
287+ poll_interval = 0.1 ,
288+ )
289+
290+ assert (completed .status or "" ).lower () == "completed"
291+ assert completed .output == "python-continued-with-memo"
292+ assert completed .run_id != first_run_id
293+ assert completed .memo == CONTINUE_MERGED_MEMO
294+
295+ runs = await handle .list_runs ()
296+ assert runs .run_count == 2
297+ assert [run .run_id for run in runs .runs ] == [first_run_id , completed .run_id ]
298+ assert runs .runs [1 ].memo == CONTINUE_MERGED_MEMO
299+
300+ first_history = await client .get_history (workflow_id , first_run_id )
301+ successor_history = await client .get_history (workflow_id , completed .run_id or "" )
302+ first_memo_events = _memo_events (first_history )
303+ assert len (first_memo_events ) == 1
304+ assert serializer .decode_envelope (first_memo_events [0 ]["payload" ]["entries" ]) == CONTINUE_MEMO_PATCH
305+ assert serializer .decode_envelope (first_memo_events [0 ]["payload" ]["merged" ]) == CONTINUE_MERGED_MEMO
306+ assert _memo_events (successor_history ) == []
0 commit comments