From 344902c17ef9ea4770db7fb35d390eac631855c0 Mon Sep 17 00:00:00 2001 From: lda Date: Sat, 5 Sep 2026 20:25:11 +0700 Subject: [PATCH] fix: harden run step budget checkpoint validation --- src/wf_core/run_codec.py | 91 ++++++++++++++++++---- tests/core/test_run_step_budget.py | 120 +++++++++++++++++++++++++++++ 2 files changed, 196 insertions(+), 15 deletions(-) diff --git a/src/wf_core/run_codec.py b/src/wf_core/run_codec.py index c96b1e98..fc0b45cb 100644 --- a/src/wf_core/run_codec.py +++ b/src/wf_core/run_codec.py @@ -1,7 +1,7 @@ from __future__ import annotations from copy import deepcopy -from typing import Any, Literal +from typing import Any, Literal, cast from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError @@ -37,45 +37,91 @@ def dump_run_state(run: RunState) -> dict[str, object]: ).model_dump(mode="json") +def _is_strict_int(value: object) -> bool: + """Return True only for actual ints, excluding bools. + + Pydantic's lax mode coerces ``True``/``"10"`` to ``1``/``10`` before + ``RunLimits.__post_init__`` runs, which would defeat the positive + non-boolean integer contract. Raw-envelope checks must therefore use + exact-type comparison instead of ``isinstance``. + """ + return type(value) is int + + +def _check_step_number(value: object, *, steps_executed: int, what: str) -> None: + """Reject corrupt step numbers that are not None or coherent history. + + ``None`` survives for unadmitted pending frames and upgraded pre-budget + history. Any assigned number is one-based and can never exceed the + authoritative ``steps_executed`` counter. + """ + if value is None: + return + if not _is_strict_int(value) or not 1 <= value <= steps_executed: # type: ignore[operator] + raise ValueError( + "invalid persisted workflow run state: " + f"{what} has incoherent step number {value!r}" + ) + + def _inject_v1_budget_defaults(state: dict[str, Any]) -> dict[str, Any]: """Copy a v1 state dict with the one-time step budget defaults applied. Version-1 envelopes predate step budgets, so they receive the default limit, a zeroed counter, and an unassigned number per frame exactly once - at load time. Pre-budget trace entries and any outstanding interrupt keep - an unassigned (``None``) number: attempts made before the upgrade are - outside the new budget. + at load time. Assignment overwrites rather than preserves: any budget + fields present in a v1 envelope are smuggled (v1 writers never emitted + them) and must not survive the upgrade. Pre-budget trace entries and any + outstanding interrupt keep an unassigned (``None``) number: attempts made + before the upgrade are outside the new budget. """ upgraded = deepcopy(state) - upgraded.setdefault("limits", {"max_steps": RunLimits().max_steps}) - upgraded.setdefault("steps_executed", 0) + upgraded["limits"] = {"max_steps": RunLimits().max_steps} + upgraded["steps_executed"] = 0 frames = upgraded.get("frames") if isinstance(frames, dict): for frame in frames.values(): if isinstance(frame, dict): - frame.setdefault("step_number", None) + frame["step_number"] = None trace = upgraded.get("trace") if isinstance(trace, list): for entry in trace: if isinstance(entry, dict): - entry.setdefault("step_number", None) + entry["step_number"] = None interrupt = upgraded.get("interrupt") if isinstance(interrupt, dict): - interrupt.setdefault("step_number", None) + interrupt["step_number"] = None return upgraded def _require_v2_budget_fields(state: dict[str, Any]) -> None: - """Reject v2 payloads missing budget fields as corrupt state. + """Reject v2 payloads with missing or incoherent budget fields as corrupt. Unlike v1, a v2 envelope promises budget fields; a missing counter is - corruption, not another request for defaults. Trace and interrupt entries - always carry the key (``None`` only for upgraded pre-budget history), so a - missing key is likewise corrupt even though the dataclass default would - otherwise mask it. + corruption, not another request for defaults. Values are validated + strictly on the raw envelope (exact ints, ranges, coherence) because lax + coercion would otherwise accept bools, numeric strings, negatives, or + future frame numbers and silently inflate or distort the budget. Trace + and interrupt entries always carry the key (``None`` only for unadmitted + or upgraded pre-budget history), so a missing key is likewise corrupt + even though the dataclass default would otherwise mask it. """ - if "limits" not in state or "steps_executed" not in state: + limits = state.get("limits") + if not isinstance(limits, dict) or "max_steps" not in limits: raise ValueError("invalid persisted workflow run state: missing step budget") + max_steps = limits["max_steps"] + if not _is_strict_int(max_steps) or max_steps < 1: # type: ignore[operator] + raise ValueError( + "invalid persisted workflow run state: corrupt step budget limit" + ) + steps_executed = state.get("steps_executed") + if not _is_strict_int(steps_executed): + raise ValueError("invalid persisted workflow run state: missing step budget") + if not 0 <= steps_executed <= max_steps: # type: ignore[operator] + raise ValueError( + "invalid persisted workflow run state: corrupt step counter" + ) + exec_count = cast(int, steps_executed) frames = state.get("frames") if not isinstance(frames, dict): raise ValueError("invalid persisted workflow run state: missing frames") @@ -85,6 +131,11 @@ def _require_v2_budget_fields(state: dict[str, Any]) -> None: "invalid persisted workflow run state: " f"frame {frame_id!r} is missing its step number" ) + _check_step_number( + frame["step_number"], + steps_executed=exec_count, + what=f"frame {frame_id!r}", + ) trace = state.get("trace") if isinstance(trace, list): for position, entry in enumerate(trace): @@ -93,6 +144,11 @@ def _require_v2_budget_fields(state: dict[str, Any]) -> None: "invalid persisted workflow run state: " f"trace entry {position!r} is missing its step number" ) + _check_step_number( + entry["step_number"], + steps_executed=exec_count, + what=f"trace entry {position!r}", + ) interrupt = state.get("interrupt") if interrupt is not None: if not isinstance(interrupt, dict) or "step_number" not in interrupt: @@ -100,6 +156,11 @@ def _require_v2_budget_fields(state: dict[str, Any]) -> None: "invalid persisted workflow run state: " "interrupt is missing its step number" ) + _check_step_number( + interrupt["step_number"], + steps_executed=exec_count, + what="interrupt", + ) def _restore_root_alias(run: RunState) -> RunState: diff --git a/tests/core/test_run_step_budget.py b/tests/core/test_run_step_budget.py index f4e79a60..96f83d95 100644 --- a/tests/core/test_run_step_budget.py +++ b/tests/core/test_run_step_budget.py @@ -263,6 +263,126 @@ def test_v2_missing_frame_step_number_is_corrupt() -> None: load_run_state_with_upgrade(stored) +def test_v1_smuggled_budget_fields_receive_defaults() -> None: + """V1 envelopes predate budgets; smuggled fields must not survive.""" + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=7)) + stored = _strip_to_v1(dump_run_state(run)) + state = cast(dict[str, Any], stored["state"]) + state["limits"] = {"max_steps": 999_999} + state["steps_executed"] = 999 + + restored, upgraded = load_run_state_with_upgrade(stored) + + assert upgraded is True + assert restored.limits.max_steps == 10_000 + assert restored.steps_executed == 0 + + +def test_v2_bool_max_steps_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + stored = dump_run_state(run) + cast(dict[str, Any], stored["state"])["limits"] = {"max_steps": True} + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_str_max_steps_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + stored = dump_run_state(run) + cast(dict[str, Any], stored["state"])["limits"] = {"max_steps": "10"} + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_negative_steps_executed_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + stored = dump_run_state(run) + cast(dict[str, Any], stored["state"])["steps_executed"] = -100 + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_exceeding_steps_executed_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + stored = dump_run_state(run) + cast(dict[str, Any], stored["state"])["steps_executed"] = 5 + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_bool_steps_executed_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + stored = dump_run_state(run) + cast(dict[str, Any], stored["state"])["steps_executed"] = True + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_incoherent_frame_step_number_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + admit_step_attempt(run, run.current_frame(), workflow.start) + stored = dump_run_state(run) + state = cast(dict[str, Any], stored["state"]) + frames = cast(dict[str, Any], state["frames"]) + cast(dict[str, Any], frames["root"])["step_number"] = 999 + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_bool_frame_step_number_is_corrupt() -> None: + from typing import Any, cast + + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + admit_step_attempt(run, run.current_frame(), workflow.start) + stored = dump_run_state(run) + state = cast(dict[str, Any], stored["state"]) + frames = cast(dict[str, Any], state["frames"]) + cast(dict[str, Any], frames["root"])["step_number"] = True + + with pytest.raises(ValueError): + load_run_state_with_upgrade(stored) + + +def test_v2_pending_frame_none_number_round_trips() -> None: + """Fresh V2 runs legitimately persist unadmitted (None) frame numbers.""" + workflow = _minimal_workflow() + run = create_run_state(workflow, {}, limits=RunLimits(max_steps=2)) + + restored, upgraded = load_run_state_with_upgrade(dump_run_state(run)) + + assert upgraded is False + assert restored.steps_executed == 0 + assert restored.frames["root"].step_number is None + + # --- Task 2: sync dispatch and trace numbering ---