Skip to content
Draft
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
Expand Up @@ -148,6 +148,8 @@ Resume from memory:
result = pipeline.run(data={}, pipeline_snapshot=snapshot)
```

The saved component visit counts stay unchanged after a resumed run, including when a downstream component fails. Resuming the same snapshot again therefore uses the original visit counts for `max_runs_per_component`.

Resume from disk:

```python
Expand Down
3 changes: 2 additions & 1 deletion haystack/core/pipeline/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -376,7 +376,8 @@ def run( # noqa: PLR0915, PLR0912, C901
_validate_pipeline_snapshot_against_pipeline(pipeline_snapshot, self.graph)

# Handle resuming the pipeline from a snapshot
component_visits = pipeline_snapshot.pipeline_state.component_visits
# Advancing or failing a resumed run must not consume the snapshot's saved run budget.
component_visits = pipeline_snapshot.pipeline_state.component_visits.copy()
ordered_component_names = pipeline_snapshot.ordered_component_names
data = _deserialize_value_with_schema(pipeline_snapshot.original_input_data)

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
---
fixes:
- |
``Pipeline.run()`` no longer changes the visit counts stored in a
``PipelineSnapshot`` when resuming it. Replaying the same snapshot or
retrying after a downstream failure now preserves the saved state and
does not incorrectly exhaust ``max_runs_per_component``.
43 changes: 40 additions & 3 deletions test/core/pipeline/test_breakpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
from haystack.components.joiners import BranchJoiner, ListJoiner
from haystack.components.routers import ConditionalRouter
from haystack.components.routers.conditional_router import Route
from haystack.core.errors import BreakpointException, PipelineInvalidPipelineSnapshotError
from haystack.core.errors import BreakpointException, PipelineInvalidPipelineSnapshotError, PipelineRuntimeError
from haystack.core.pipeline import Pipeline
from haystack.core.pipeline.breakpoint import (
HAYSTACK_PIPELINE_SNAPSHOT_SAVE_ENABLED,
Expand Down Expand Up @@ -171,8 +171,8 @@ def run(self, value: int) -> dict[str, Any]:
return {"done": f"finished at {value}"}


def _three_component_pipeline() -> Pipeline:
pipeline = Pipeline()
def _three_component_pipeline(*, max_runs_per_component: int = 100) -> Pipeline:
pipeline = Pipeline(max_runs_per_component=max_runs_per_component)
pipeline.add_component("comp1", _AppendingComponent())
pipeline.add_component("comp2", _AppendingComponent())
pipeline.add_component("comp3", _AppendingComponent())
Expand All @@ -191,6 +191,42 @@ def _looping_pipeline() -> Pipeline:


class TestResumeFromPipelineSnapshot:
def test_resuming_snapshot_preserves_saved_state_and_run_budget(self):
pipeline = _three_component_pipeline(max_runs_per_component=1)

with pytest.raises(BreakpointException) as exc_info:
pipeline.run(data={"comp1": {"input_value": "test"}}, break_point=Breakpoint(component_name="comp2"))
snapshot = exc_info.value.pipeline_snapshot
assert snapshot is not None
saved_state = snapshot.to_dict()
expected = {"comp3": {"result": "test_processed_processed_processed"}}

assert pipeline.run(data={}, pipeline_snapshot=snapshot) == expected
assert pipeline.run(data={}, pipeline_snapshot=snapshot) == expected
assert snapshot.to_dict() == saved_state

def test_failed_resume_does_not_modify_snapshot_or_prevent_retry(self, monkeypatch):
pipeline = _three_component_pipeline(max_runs_per_component=1)

with pytest.raises(BreakpointException) as exc_info:
pipeline.run(data={"comp1": {"input_value": "test"}}, break_point=Breakpoint(component_name="comp2"))
snapshot = exc_info.value.pipeline_snapshot
assert snapshot is not None
saved_state = snapshot.to_dict()

def fail(input_value: str) -> dict[str, str]:
raise RuntimeError(f"injected downstream failure for {input_value}")

with monkeypatch.context() as patch:
patch.setattr(pipeline.get_component("comp3"), "run", fail)
with pytest.raises(PipelineRuntimeError, match="injected downstream failure"):
pipeline.run(data={}, pipeline_snapshot=snapshot)

assert pipeline.run(data={}, pipeline_snapshot=snapshot) == {
"comp3": {"result": "test_processed_processed_processed"}
}
assert snapshot.to_dict() == saved_state

def test_break_point_with_pipeline_snapshot_steps_through_pipeline(self):
pipeline = _three_component_pipeline()

Expand All @@ -207,6 +243,7 @@ def test_break_point_with_pipeline_snapshot_steps_through_pipeline(self):
second_snapshot = exc_info.value.pipeline_snapshot
assert second_snapshot is not None
assert second_snapshot.pipeline_state.component_visits == {"comp1": 1, "comp2": 1, "comp3": 0}
assert first_snapshot.pipeline_state.component_visits == {"comp1": 1, "comp2": 0, "comp3": 0}

# resume from the second snapshot and run to completion
result = pipeline.run(data={}, pipeline_snapshot=second_snapshot)
Expand Down