mirror of
https://github.com/run-llama/workflows-py.git
synced 2026-08-24 20:01:34 -04:00
fix: rebuild_state_from_ticks clears in_progress before replaying (#282)
This commit is contained in:
committed by
GitHub
parent
bfbfba4780
commit
85f948e504
@@ -0,0 +1,7 @@
|
||||
---
|
||||
"llama-index-workflows": patch
|
||||
---
|
||||
|
||||
fix: rebuild_state_from_ticks clears in_progress before replaying
|
||||
|
||||
Fixed ctx.to_dict() failing with "Worker X not found in in_progress" when checkpointing resumed workflows. The function now also rewinds in progress when recreating from ticks, to match the actual behavior when resuming a workflow.
|
||||
@@ -319,7 +319,23 @@ def rebuild_state_from_ticks(
|
||||
state: BrokerState,
|
||||
ticks: list[WorkflowTick],
|
||||
) -> BrokerState:
|
||||
"""Rebuild the state from a list of ticks"""
|
||||
"""Rebuild the state from a list of ticks.
|
||||
|
||||
When reconstructing state (e.g., for checkpointing), we must first apply
|
||||
rewind_in_progress() to match what happens at runtime when resuming a workflow.
|
||||
This clears in_progress, moves events back to the queue, and then re-assigns
|
||||
new worker IDs starting from 0.
|
||||
|
||||
Without this, resuming a workflow and then checkpointing again would fail
|
||||
because the original in_progress worker IDs don't match the new worker IDs
|
||||
assigned after rewind.
|
||||
"""
|
||||
# Apply rewind_in_progress to match what happens at runtime when resuming.
|
||||
# This re-assigns worker IDs so they align with the ticks that were recorded
|
||||
# after the workflow was resumed.
|
||||
state, _ = rewind_in_progress(state, time.time())
|
||||
|
||||
# Replay ticks to rebuild state
|
||||
for tick in ticks:
|
||||
state, _ = _reduce_tick(
|
||||
tick, state, time.time()
|
||||
|
||||
@@ -34,6 +34,7 @@ from workflows.runtime.control_loop import (
|
||||
_process_publish_event_tick,
|
||||
_process_step_result_tick,
|
||||
_process_timeout_tick,
|
||||
rebuild_state_from_ticks,
|
||||
rewind_in_progress,
|
||||
)
|
||||
from workflows.runtime.types.commands import (
|
||||
@@ -69,6 +70,7 @@ from workflows.runtime.types.ticks import (
|
||||
TickPublishEvent,
|
||||
TickStepResult,
|
||||
TickTimeout,
|
||||
WorkflowTick,
|
||||
)
|
||||
|
||||
|
||||
@@ -904,3 +906,128 @@ def test_no_idle_event_when_workflow_completes(base_state: BrokerState) -> None:
|
||||
if isinstance(c, CommandPublishEvent) and isinstance(c.event, WorkflowIdleEvent)
|
||||
]
|
||||
assert len(idle_commands) == 0
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Tests for rebuild_state_from_ticks
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_rebuild_state_from_ticks_clears_in_progress(base_state: BrokerState) -> None:
|
||||
"""
|
||||
Test that rebuild_state_from_ticks clears in_progress before replaying ticks.
|
||||
|
||||
This is critical for checkpointing resumed workflows. When a workflow is resumed:
|
||||
1. The checkpoint has in_progress workers with IDs like [1, 2, 3]
|
||||
2. rewind_in_progress() clears in_progress and assigns new IDs [0, 1, 2]
|
||||
3. New ticks reference the new worker IDs [0, 1, 2]
|
||||
4. When checkpointing again, rebuild_state_from_ticks must also clear in_progress
|
||||
before replaying ticks, otherwise worker IDs won't match.
|
||||
|
||||
Without the fix, this would raise: "Worker 0 not found in in_progress"
|
||||
"""
|
||||
event1 = MyTestEvent(value=1)
|
||||
event2 = MyTestEvent(value=2)
|
||||
|
||||
# Simulate checkpoint state with in_progress workers using IDs 1, 2
|
||||
# (as if they were mid-execution when checkpoint was taken)
|
||||
shared_state = StepWorkerState(
|
||||
step_name="test_step",
|
||||
collected_events={},
|
||||
collected_waiters=[],
|
||||
)
|
||||
base_state.workers["test_step"].in_progress = [
|
||||
InProgressState(
|
||||
event=event1,
|
||||
worker_id=1, # Original worker ID from checkpoint
|
||||
shared_state=shared_state,
|
||||
attempts=0,
|
||||
first_attempt_at=100.0,
|
||||
),
|
||||
InProgressState(
|
||||
event=event2,
|
||||
worker_id=2, # Original worker ID from checkpoint
|
||||
shared_state=shared_state,
|
||||
attempts=0,
|
||||
first_attempt_at=100.0,
|
||||
),
|
||||
]
|
||||
|
||||
# Simulate ticks from a resumed run where rewind_in_progress assigned new IDs
|
||||
# These ticks reference worker IDs 0 and 1 (not 1 and 2 from checkpoint)
|
||||
ticks: list[WorkflowTick] = [
|
||||
# Worker 0 starts (after rewind assigned new ID)
|
||||
TickAddEvent(event=event1),
|
||||
# Worker 0 completes
|
||||
TickStepResult(
|
||||
step_name="test_step",
|
||||
worker_id=0, # New ID assigned after rewind
|
||||
event=event1,
|
||||
result=[StepWorkerResult(result=OtherEvent(data="done1"))],
|
||||
),
|
||||
# Worker 1 starts (after rewind assigned new ID)
|
||||
TickAddEvent(event=event2),
|
||||
# Worker 1 completes
|
||||
TickStepResult(
|
||||
step_name="test_step",
|
||||
worker_id=0, # Reuses ID 0 since previous worker completed
|
||||
event=event2,
|
||||
result=[StepWorkerResult(result=StopEvent(result="done2"))],
|
||||
),
|
||||
]
|
||||
|
||||
# This should NOT raise "Worker 0 not found in in_progress"
|
||||
# because rebuild_state_from_ticks now clears in_progress before replaying
|
||||
final_state = rebuild_state_from_ticks(base_state, ticks)
|
||||
|
||||
# Verify the workflow completed
|
||||
assert final_state.is_running is False
|
||||
assert len(final_state.workers["test_step"].in_progress) == 0
|
||||
|
||||
|
||||
def test_rebuild_state_from_ticks_preserves_queue_order(
|
||||
base_state: BrokerState,
|
||||
) -> None:
|
||||
"""
|
||||
Test that rebuild_state_from_ticks applies rewind_in_progress which moves
|
||||
in_progress events to the front of the queue and then re-starts them.
|
||||
|
||||
Since the base fixture has num_workers=1, only one event can be in_progress
|
||||
at a time. The originally in_progress event (event1) should be re-started
|
||||
with worker_id=0, and event2 should remain in the queue.
|
||||
"""
|
||||
event1 = MyTestEvent(value=1)
|
||||
event2 = MyTestEvent(value=2)
|
||||
|
||||
# State with in_progress worker
|
||||
shared_state = StepWorkerState(
|
||||
step_name="test_step",
|
||||
collected_events={},
|
||||
collected_waiters=[],
|
||||
)
|
||||
base_state.workers["test_step"].in_progress = [
|
||||
InProgressState(
|
||||
event=event1,
|
||||
worker_id=0,
|
||||
shared_state=shared_state,
|
||||
attempts=2, # Already retried twice
|
||||
first_attempt_at=100.0,
|
||||
),
|
||||
]
|
||||
# Also has queued event
|
||||
base_state.workers["test_step"].queue = [
|
||||
EventAttempt(event=event2, attempts=0, first_attempt_at=None)
|
||||
]
|
||||
|
||||
# No ticks - test that rebuild applies rewind_in_progress
|
||||
result = rebuild_state_from_ticks(base_state, [])
|
||||
|
||||
# rewind_in_progress re-starts workers, so event1 should be back in in_progress
|
||||
# with worker_id=0 (reassigned) and retry info preserved
|
||||
assert len(result.workers["test_step"].in_progress) == 1
|
||||
assert result.workers["test_step"].in_progress[0].event == event1
|
||||
assert result.workers["test_step"].in_progress[0].worker_id == 0
|
||||
assert result.workers["test_step"].in_progress[0].attempts == 2
|
||||
# Queue should have event2 (since num_workers=1, only 1 can be in_progress)
|
||||
assert len(result.workers["test_step"].queue) == 1
|
||||
assert result.workers["test_step"].queue[0].event == event2
|
||||
|
||||
@@ -9,7 +9,7 @@ import logging
|
||||
import pickle
|
||||
import threading
|
||||
import weakref
|
||||
from typing import Any, Callable, Union, cast
|
||||
from typing import Any, Callable, Optional, Union, cast
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
@@ -1046,3 +1046,61 @@ async def test_inner_step_can_access_run_id_from_instrument_tags() -> None:
|
||||
assert handler.run_id is not None
|
||||
assert run_id["run_id"] is not None
|
||||
assert run_id["run_id"] == handler.run_id
|
||||
|
||||
|
||||
class Par(Event):
|
||||
id: int
|
||||
|
||||
|
||||
class ParDone(Event):
|
||||
id: int
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_workflow_parallel_resume() -> None:
|
||||
allowed_done = asyncio.Event()
|
||||
resume_event = asyncio.Event()
|
||||
allowed_index = 0
|
||||
|
||||
class ParallelResumeWorkflow(Workflow):
|
||||
@step
|
||||
async def step1(self, ev: StartEvent, ctx: Context) -> Optional[Par]: # noqa - python 3.9 struggles here with | None
|
||||
for i in range(4):
|
||||
ctx.send_event(Par(id=i))
|
||||
return None
|
||||
|
||||
@step(num_workers=4)
|
||||
async def par(self, ev: Par) -> ParDone:
|
||||
if ev.id != allowed_index:
|
||||
await resume_event.wait()
|
||||
return ParDone(id=ev.id)
|
||||
|
||||
@step
|
||||
async def step3(self, ev: ParDone, ctx: Context) -> Optional[StopEvent]: # noqa - python 3.9 struggles here with | None
|
||||
if ev.id == allowed_index:
|
||||
allowed_done.set()
|
||||
if ctx.collect_events(ev, [ParDone] * 4) is None:
|
||||
return None
|
||||
return StopEvent(result="Done")
|
||||
|
||||
wf = ParallelResumeWorkflow(timeout=10)
|
||||
handler = wf.run()
|
||||
await allowed_done.wait()
|
||||
serialized_ctx = handler.ctx.to_dict()
|
||||
try:
|
||||
handler.cancel()
|
||||
await handler.cancel_run()
|
||||
except Exception:
|
||||
pass
|
||||
# immediately resume the workflow
|
||||
allowed_index = 3
|
||||
allowed_done.clear()
|
||||
new_handler = wf.run(ctx=Context.from_dict(wf, serialized_ctx))
|
||||
await allowed_done.wait()
|
||||
# serialize again to detect inconsistencies
|
||||
serialized_ctx = new_handler.ctx.to_dict()
|
||||
|
||||
# finally resume the workflow, and complete
|
||||
new_handler = wf.run(ctx=Context.from_dict(wf, serialized_ctx))
|
||||
resume_event.set()
|
||||
await new_handler
|
||||
|
||||
Reference in New Issue
Block a user