diff --git a/changes/vercel-workflow/20260925-disposed-hook-replay.bugfix.md b/changes/vercel-workflow/20260925-disposed-hook-replay.bugfix.md new file mode 100644 index 00000000..cf813ca9 --- /dev/null +++ b/changes/vercel-workflow/20260925-disposed-hook-replay.bugfix.md @@ -0,0 +1 @@ +Fix workflow replay failures when a hook is disposed before its recorded registration is replayed. diff --git a/src/vercel-workflow/tests/unit/test_workflow_disposed_hook_replay.py b/src/vercel-workflow/tests/unit/test_workflow_disposed_hook_replay.py new file mode 100644 index 00000000..b9314db1 --- /dev/null +++ b/src/vercel-workflow/tests/unit/test_workflow_disposed_hook_replay.py @@ -0,0 +1,155 @@ +import asyncio +import dataclasses +from datetime import datetime, timezone + +import pytest + +from tests.payloads import PLAIN_ENCODER +from vercel.workflow._internal import core, errors, runtime, world as w + +TOKEN = "disposed-hook" +registry = core.Workflows(as_vercel_job=False) + + +@dataclasses.dataclass +class Payload(core.BaseHook): + value: str + + +def _context() -> runtime.WorkflowOrchestratorContext: + return runtime.WorkflowOrchestratorContext( + [], + run_id="wrun_test", + seed="seed", + started_at=0, + registry=registry, + ) + + +@pytest.mark.parametrize("outcome", ["created", "conflict", "legacy_conflict"]) +async def test_disposal_does_not_abandon_a_pending_registration(outcome: str) -> None: + context = _context() + hook_id = context.create_hook(TOKEN, Payload)._correlation_id + hook = context.hooks[hook_id] + registration = asyncio.create_task(context.run_hook_conflict(correlation_id=hook_id)) + await asyncio.sleep(0) + assert not registration.done() + + try: + context.dispose_hook(correlation_id=hook_id) + assert not registration.done() + + context.events.append( + w.HookCreatedEventData(token=TOKEN).into_event(hook_id) + if outcome == "created" + else w.HookConflictEvent( + correlation_id=hook_id, + event_data=w.HookConflictEventData( + token=TOKEN, + conflicting_run_id="wrun_owner" if outcome == "conflict" else None, + ), + ) + ) + context.resume() + await asyncio.sleep(0) + assert registration.done() + if outcome == "created": + assert await registration is None + assert hook.has_created_event + assert hook.conflict_error is None + elif outcome == "conflict": + owner = await registration + assert owner is not None + assert owner.run_id == "wrun_owner" + assert isinstance(hook.conflict_error, errors.HookConflictError) + else: + with pytest.raises(errors.HookConflictError): + await registration + + assert hook.disposed + assert hook_id not in context.suspensions + assert context.replay_index == len(context.events) + finally: + registration.cancel() + await asyncio.gather(registration, return_exceptions=True) + + +async def test_disposed_hook_still_discards_recorded_payloads() -> None: + context = _context() + hook_id = context.create_hook(TOKEN, Payload)._correlation_id + context.dispose_hook(correlation_id=hook_id) + context.events.extend( + [ + w.HookCreatedEventData(token=TOKEN).into_event(hook_id), + w.HookReceivedEventData( + token=TOKEN, payload=PLAIN_ENCODER.encode({"value": "discarded"}) + ).into_event(hook_id), + w.HookDisposedEvent(correlation_id=hook_id), + ] + ) + + for _ in context.events: + context.resume() + + assert not context.hooks[hook_id].buffered_results + assert context.hooks[hook_id].has_dispose_event + assert hook_id not in context.suspensions + with pytest.raises(StopAsyncIteration): + await context.run_hook(correlation_id=hook_id) + + +@pytest.mark.parametrize("recorded_kind", ["step", "wait"]) +async def test_disposed_hook_still_occupies_its_recorded_position(recorded_kind: str) -> None: + context = _context() + hook_id = context.create_hook(TOKEN, Payload)._correlation_id + position = hook_id.split("_", 1)[1] + event = ( + w.StepCreatedEventData( + step_name="previous_step", input=PLAIN_ENCODER.encode([]) + ).into_event(f"step_{position}") + if recorded_kind == "step" + else w.WaitCreatedEventData(resume_at=datetime(2026, 1, 1, tzinfo=timezone.utc)).into_event( + f"wait_{position}" + ) + ) + context.events.append(event) + context.dispose_hook(correlation_id=hook_id) + + async def replay() -> None: + context.resume() + + try: + runtime._run_isolated(replay(), loop_factory=asyncio.new_event_loop) + except asyncio.CancelledError: + pass + + assert isinstance(context.resume_exception, runtime.NondeterminismError) + assert str(context.resume_exception) == ( + f"workflow replay diverged at position {position}: recorded a {recorded_kind!r} call, " + "but the body now issues a 'hook' call. The workflow body is non-deterministic." + ) + assert context.suspended + + +async def test_cancellation_hook_replays_without_a_user_hook() -> None: + context = _context() + cancellation = runtime.Cancellation( + correlation_id="hook_1", token="abrt_1", step_id="step_1", requested=True + ) + context.suspensions[cancellation.correlation_id] = cancellation + context.events.extend( + [ + w.HookCreatedEventData(token=cancellation.token).into_event( + cancellation.correlation_id + ), + w.HookReceivedEventData(payload=PLAIN_ENCODER.encode(None)).into_event( + cancellation.correlation_id + ), + ] + ) + + context.resume() + assert cancellation.has_created_event + assert not context.hooks + context.resume() + assert not context.suspensions diff --git a/src/vercel-workflow/tests/unit/test_workflow_hook_registration.py b/src/vercel-workflow/tests/unit/test_workflow_hook_registration.py index 3d0bf78e..39840bf1 100644 --- a/src/vercel-workflow/tests/unit/test_workflow_hook_registration.py +++ b/src/vercel-workflow/tests/unit/test_workflow_hook_registration.py @@ -75,6 +75,21 @@ async def pending_step() -> str: return "done" +@registry.step +async def gate_step() -> str: + return "ready" + + +@registry.workflow +async def dispose_hook_after_parallel_step() -> str: + pending = asyncio.create_task(pending_step()) + await gate_step() + approval = Approval.wait(token=TOKEN) + await pending + approval.dispose() + return await pending_step() + + @registry.workflow async def create_then_step_then_await() -> dict[str, object]: approval = Approval.wait(token=TOKEN) @@ -156,6 +171,84 @@ async def test_hook_is_registered_when_the_run_first_suspends(world) -> None: assert (await world.hooks_get_by_token(TOKEN)).run_id == run_id +@pytest.mark.parametrize("conflicted", [False, True]) +async def test_step_completion_before_hook_registration_can_replay( + world, monkeypatch, conflicted: bool +) -> None: + if conflicted: + owner_id = await _create_run(world, create_then_step_then_await.workflow_id) + await _invoke(owner_id, create_then_step_then_await.workflow_id) + + run_id = await _create_run(world, dispose_hook_after_parallel_step.workflow_id) + await _invoke(run_id, dispose_hook_after_parallel_step.workflow_id) + events = (await world.events_list(run_id)).data + steps = { + event.event_data.step_name: event.correlation_id + for event in events + if isinstance(event, w.StepCreatedEvent) + } + pending_id = steps[pending_step.name] + + async def finish_step(step_id: str, step_name: str) -> None: + await runtime.workflow_handler( + w.WorkflowInvokePayload(run_id=run_id, step_id=step_id, step_name=step_name).model_dump( + by_alias=True + ), + attempt=1, + queue_name=w.get_queue_name(dispose_hook_after_parallel_step.workflow_id), + message_id="msg_step", + registry=registry, + ) + + await finish_step(steps[gate_step.name], gate_step.name) + + create_event = world.events_create + + async def complete_step_before_registering_hook(run_id, event): + if isinstance(event, w.HookCreatedEvent): + # The earlier step finishes after the replay snapshot was loaded, + # but before this invocation persists its newly created hook. + await finish_step(pending_id, pending_step.name) + return await create_event(run_id, event) + + with monkeypatch.context() as patch: + patch.setattr(world, "events_create", complete_step_before_registering_hook) + await _invoke(run_id, dispose_hook_after_parallel_step.workflow_id) + + events = (await world.events_list(run_id)).data + completion_index = next( + index + for index, event in enumerate(events) + if isinstance(event, w.StepCompletedEvent) and event.correlation_id == pending_id + ) + registration_index = next( + index + for index, event in enumerate(events) + if isinstance(event, w.HookConflictEvent if conflicted else w.HookCreatedEvent) + ) + assert completion_index < registration_index + + await _invoke(run_id, dispose_hook_after_parallel_step.workflow_id) + assert (await world.runs_get(run_id)).status == "running" + events = (await world.events_list(run_id)).data + final_step_id = next( + event.correlation_id + for event in events + if isinstance(event, w.StepCreatedEvent) + and event.event_data.step_name == pending_step.name + and event.correlation_id != pending_id + ) + await finish_step(final_step_id, pending_step.name) + await _invoke(run_id, dispose_hook_after_parallel_step.workflow_id) + + assert await Run(run_id).return_value() == "done" + if conflicted: + assert (await world.hooks_get_by_token(TOKEN)).run_id == owner_id + else: + with pytest.raises(w.HookNotFoundError): + await world.hooks_get_by_token(TOKEN) + + @pytest.mark.parametrize( "workflow", [create_then_step_then_await, confirm_then_step_then_await], diff --git a/src/vercel-workflow/vercel/workflow/_internal/runtime.py b/src/vercel-workflow/vercel/workflow/_internal/runtime.py index 89f338ba..87b9c2c8 100644 --- a/src/vercel-workflow/vercel/workflow/_internal/runtime.py +++ b/src/vercel-workflow/vercel/workflow/_internal/runtime.py @@ -24,6 +24,7 @@ Sequence, ) from datetime import datetime +from itertools import chain from typing import Any, Generic, ParamSpec, TypeVar, overload from urllib.parse import parse_qsl, urlsplit @@ -1170,17 +1171,13 @@ def resume(self) -> None: ), ) return - if event.correlation_id not in self.suspensions: + elif event.correlation_id not in self.suspensions: match event: case ( # A step's attribute write. It answers no call in # this body, so consume it and move on. w.AttrSetEvent(correlation_id=None) | w.AttrSetEvent(event_data=w.AttrSetEventData(writer=w.StepAttributeWriter())) - # A hook received without being registered - # means it has been disposed of. Nothing to do - # but drop it. - | w.HookReceivedEvent() ): return case ( @@ -1198,7 +1195,7 @@ def resume(self) -> None: # instead of yielding forever (the matching ID will never # appear, so plain `return` would deadlock the run). pos = _correlation_ulid(slot_id) - for sus in self.suspensions.values(): + for sus in chain(self.suspensions.values(), self.hooks.values()): if _correlation_ulid(sus.correlation_id) == pos: self._fail_nondeterminism( sus, @@ -1219,6 +1216,7 @@ def resume(self) -> None: ) return + hook: BaseSuspension | None match event: case w.StepCreatedEvent( event_data=w.StepCreatedEventData(step_name=name, input=recorded_input) @@ -1243,7 +1241,10 @@ def resume(self) -> None: sus.has_created_event = True case w.HookCreatedEvent(): - hook = self.suspensions[event.correlation_id] + hook = self.hooks.get(event.correlation_id) + if hook is None: + # Internal step-cancellation hooks only live in suspensions. + hook = self.suspensions[event.correlation_id] hook.has_created_event = True if isinstance(hook, Hook): while hook.conflict_futures: @@ -1321,10 +1322,9 @@ def resume(self) -> None: conflicting_run_id=conflicting_run_id, ) ): - conflicting_hook = self.suspensions.get(event.correlation_id) + conflicting_hook = self.hooks.get(event.correlation_id) if conflicting_hook is not None: - self.suspensions.pop(event.correlation_id) - assert isinstance(conflicting_hook, Hook) + self.suspensions.pop(event.correlation_id, None) conflict_error = errors.HookConflictError(token, conflicting_run_id) conflicting_hook.conflict_error = conflict_error conflicting_hook.conflicting_run = ( @@ -1346,7 +1346,10 @@ def resume(self) -> None: future.set_exception(conflict_error) case w.HookReceivedEvent(event_data=w.HookReceivedEventData(payload=data)): - hook = self.suspensions[event.correlation_id] + hook = self.suspensions.get(event.correlation_id) + if hook is None: + # Disposed hooks no longer receive payloads. + return if isinstance(hook, Cancellation): # A step cancellation already recorded: deregister it so # the flush does not send it again.