Skip to content
Merged
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
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Fix workflow replay failures when a hook is disposed before its recorded registration is replayed.
155 changes: 155 additions & 0 deletions src/vercel-workflow/tests/unit/test_workflow_disposed_hook_replay.py
Original file line number Diff line number Diff line change
@@ -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
93 changes: 93 additions & 0 deletions src/vercel-workflow/tests/unit/test_workflow_hook_registration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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],
Expand Down
25 changes: 14 additions & 11 deletions src/vercel-workflow/vercel/workflow/_internal/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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 (
Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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]
Comment on lines +1245 to +1247

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should we put the step-cancellation hooks in hooks instead/also?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes! Let me do that in a new PR.

hook.has_created_event = True
if isinstance(hook, Hook):
while hook.conflict_futures:
Expand Down Expand Up @@ -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 = (
Expand All @@ -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.
Expand Down
Loading