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
36 changes: 34 additions & 2 deletions python/formats/solver_trace_snapshot.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,12 +191,17 @@ def trace_from_document(document: object) -> SolverTrace:
_require_literal(root["geometry"], GEOMETRY, "$.geometry")
_require_sha256(root["source_formula_sha256"], "$.source_formula_sha256")
_require_sha256(root["region_sha256"], "$.region_sha256")
if root["solution_sha256"] is not None:
_require_sha256(root["solution_sha256"], "$.solution_sha256")
solution_digest = root["solution_sha256"]
if solution_digest is not None:
_require_sha256(solution_digest, "$.solution_sha256")
solver = _require_string(root["solver"], "$.solver")
if solver not in SOLVERS:
raise PipelineSnapshotError("$.solver: is not a supported native solver")
status_value = _require_status(root["status"], "$.status")
if (status_value == "sat") != (solution_digest is not None):
raise PipelineSnapshotError(
"$.solution_sha256: must be present exactly for SAT traces"
)

layout = _require_object(root["layout"], "$.layout")
_require_exact_fields(
Expand Down Expand Up @@ -443,6 +448,32 @@ def _validate_trace_region_state(
)


def _validate_solution_bundle_identity(
documents: dict[str, dict[str, object]],
) -> None:
solution = documents["solution"]
region = documents["region"]
tileset = documents["tileset"]
if solution["bounds"] != region["bounds"]:
raise PipelineSnapshotError(
"solution.bounds: does not match the referenced region"
)
if solution["tile_table"] != tileset["tiles"]:
raise PipelineSnapshotError(
"solution.tile_table: does not match the referenced tileset"
)
cells = _require_array(solution["cells"], "solution.cells")
active = _require_array(region["active"], "region.active")
if tuple(cell is not None for cell in cells) != tuple(active):
raise PipelineSnapshotError(
"solution.cells: active map does not match the referenced region"
)
if solution["boundary"] != region["boundary"]:
raise PipelineSnapshotError(
"solution.boundary: does not match the referenced region"
)


def validate_solver_trace_manifest(document: object) -> None:
"""Validate the v3 manifest without reading referenced artifacts."""
root = _require_object(document, "$")
Expand Down Expand Up @@ -592,6 +623,7 @@ def load_solver_trace_bundle(
raise PipelineSnapshotError(
"trace.solution_sha256: does not match solution artifact"
)
_validate_solution_bundle_identity(documents)
solution = documents["solution"]
bounds = _require_object(solution["bounds"], "solution.bounds")
width = bounds["max_x_inclusive"] - bounds["min_x_inclusive"] + 1
Expand Down
29 changes: 27 additions & 2 deletions python/model/solver_trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -221,12 +221,20 @@ def replay_solver_trace(
"""
domains = list(trace.initial_domains)
changes: list[tuple[int, int]] = []
search_change_floor: int | None = None
states: list[tuple[int, ...]] = []
checkpoints = {item.event_sequence: item for item in trace.checkpoints}
if len(checkpoints) != len(trace.checkpoints):
raise ValueError("checkpoint event sequences must be unique")

for index, event in enumerate(trace.events):
if event.cell is not None and event.cell >= len(domains):
raise ValueError("trace event cell lies outside the trace")
if event.phase == TRACE_SEARCH:
if search_change_floor is None:
search_change_floor = len(changes)
elif event.phase == TRACE_INITIAL and search_change_floor is not None:
raise ValueError("initial phase cannot follow search")
if event.kind != TRACE_RESULT and event.status is not None:
raise ValueError("only the result event may publish a status")
if event.kind != TRACE_DOMAIN_REDUCTION and event.reason is not None:
Expand All @@ -243,7 +251,7 @@ def replay_solver_trace(
elif event.kind == TRACE_DOMAIN_REDUCTION:
if event.phase not in TRACE_PHASES or event.reason not in TRACE_REASONS:
raise ValueError("domain reduction requires phase and reason")
if event.cell is None or event.cell >= len(domains):
if event.cell is None:
raise ValueError("domain reduction cell lies outside the trace")
if event.old_domain is None or event.new_domain is None:
raise ValueError("domain reduction requires old and new domains")
Expand All @@ -258,7 +266,11 @@ def replay_solver_trace(
if event.change_mark != len(changes):
raise ValueError("domain reduction change_mark is inconsistent")
elif event.kind == TRACE_BACKTRACK:
if event.phase != TRACE_SEARCH or event.change_mark > len(changes):
if (
event.phase != TRACE_SEARCH
or search_change_floor is None
or not search_change_floor <= event.change_mark <= len(changes)
):
raise ValueError("backtrack change_mark is inconsistent")
while len(changes) > event.change_mark:
cell, old_domain = changes.pop()
Expand All @@ -282,6 +294,19 @@ def replay_solver_trace(
raise ValueError("decision new_domain must be a singleton")
if event.old_domain is None or event.new_domain & ~event.old_domain:
raise ValueError("decision must select from its old domain")
next_event = trace.events[index + 1]
if next_event.sequence == event.sequence + 1 and not (
next_event.kind == TRACE_DOMAIN_REDUCTION
and next_event.phase == event.phase
and next_event.reason == TRACE_REASON_DECISION
and next_event.depth == event.depth
and next_event.cell == event.cell
and next_event.old_domain == event.old_domain
and next_event.new_domain == event.new_domain
):
raise ValueError(
"decision does not match its following domain reduction"
)
elif event.kind == TRACE_PROPAGATION:
if event.phase not in TRACE_PHASES:
raise ValueError("propagation requires a phase")
Expand Down
141 changes: 140 additions & 1 deletion renderer/test_wang_trace.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

from dataclasses import replace
import hashlib
import json
from pathlib import Path
Expand All @@ -11,7 +12,7 @@
import pytest

from wang_hex_port import WangSquareRenderError
from wang_trace import load_trace_bundle
from wang_trace import TraceEvent, TraceSnapshot, load_trace_bundle, replay_trace
from wang_trace_render import render_trace_assets


Expand All @@ -30,6 +31,59 @@ def _tree_bytes(directory: Path) -> dict[str, bytes]:
}


def _rewrite_artifact(copied, manifest, name, document):
encoded = (json.dumps(document, ensure_ascii=False, indent=2) + "\n").encode(
"utf-8"
)
digest = hashlib.sha256(encoded).hexdigest()
artifact_name = f"{name}-{digest}.json"
(copied / artifact_name).write_bytes(encoded)
manifest["artifacts"][name]["path"] = artifact_name
manifest["artifacts"][name]["sha256"] = digest
return digest


def _small_trace():
events = (
TraceEvent(0, "root", "initial", None, 0, None, 0, None, None, None),
TraceEvent(
1,
"domain_reduction",
"initial",
"propagation",
0,
0,
1,
(1 << 23) - 1,
3,
None,
),
TraceEvent(2, "propagation", "initial", None, 0, None, 1, None, None, None),
TraceEvent(3, "decision", "search", None, 1, 0, 1, 3, 1, None),
TraceEvent(4, "domain_reduction", "search", "decision", 1, 0, 2, 3, 1, None),
TraceEvent(5, "backtrack", "search", None, 0, 0, 1, None, None, None),
TraceEvent(6, "result", None, None, 0, None, 1, None, None, "sat"),
)
return TraceSnapshot(
solver="reference",
status="sat",
source_formula_sha256="1" * 64,
region_sha256="2" * 64,
solution_sha256="3" * 64,
width=1,
height=1,
event_capacity=len(events),
observed_event_count=len(events),
truncated=False,
checkpoint_interval=0,
checkpoint_capacity=0,
checkpoints_truncated=False,
initial_domains=((1 << 23) - 1,),
events=events,
checkpoints=(),
)


def test_loads_and_replays_hash_bound_trace_without_solver_imports():
bundle = load_trace_bundle(MANIFEST)

Expand Down Expand Up @@ -93,6 +147,91 @@ def test_rejects_semantic_delta_tampering_even_with_updated_hash(tmp_path):
load_trace_bundle(manifest_path)


def test_replay_rejects_false_decisions_backtracks_and_cells():
trace = _small_trace()
assert replay_trace(trace)[-1] == (3,)

events = list(trace.events)
events[3] = replace(events[3], new_domain=2)
with pytest.raises(WangSquareRenderError, match="following domain reduction"):
replay_trace(replace(trace, events=tuple(events)))

events = list(trace.events)
events[5] = replace(events[5], change_mark=0)
with pytest.raises(WangSquareRenderError, match="backtrack mark"):
replay_trace(replace(trace, events=tuple(events)))

for index in (3, 2):
events = list(trace.events)
events[index] = replace(events[index], cell=1)
with pytest.raises(WangSquareRenderError, match="outside layout"):
replay_trace(replace(trace, events=tuple(events)))


@pytest.mark.parametrize(
("mutation", "message"),
(
("tile_table", "tile_table"),
("active", "active map"),
("boundary", "solution.*boundary"),
("bounds", "solution.*bounds"),
),
)
def test_rejects_solution_identity_drift(tmp_path, mutation, message):
copied = tmp_path / "bundle"
shutil.copytree(FIXTURE_DIRECTORY, copied)
manifest_path = copied / "manifest.json"
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
solution_reference = manifest["artifacts"]["solution"]
solution = json.loads(
(copied / solution_reference["path"]).read_text(encoding="utf-8")
)
if mutation == "tile_table":
for tile in solution["tile_table"]:
for direction in ("N", "E", "S", "W"):
tile["edges"][direction] += 100
for sides in solution["boundary"]:
if sides is None:
continue
for direction in ("N", "E", "S", "W"):
if sides[direction] is not None:
sides[direction] += 100
elif mutation == "active":
index = next(
index
for index, tile_id in enumerate(solution["cells"])
if tile_id is not None
)
solution["cells"][index] = None
solution["boundary"][index] = None
elif mutation == "boundary":
sides = next(
sides
for sides in solution["boundary"]
if sides is not None
and any(value is not None for value in sides.values())
)
direction = next(
direction
for direction, value in sides.items()
if value is not None
)
sides[direction] = None
else:
for coordinate in ("min_x_inclusive", "max_x_inclusive"):
solution["bounds"][coordinate] += 1

solution_digest = _rewrite_artifact(copied, manifest, "solution", solution)
trace_reference = manifest["artifacts"]["trace"]
trace = json.loads((copied / trace_reference["path"]).read_text(encoding="utf-8"))
trace["solution_sha256"] = solution_digest
_rewrite_artifact(copied, manifest, "trace", trace)
manifest_path.write_text(json.dumps(manifest), encoding="utf-8")

with pytest.raises(WangSquareRenderError, match=message):
load_trace_bundle(manifest_path)


def test_rejects_inactive_cell_state_even_with_updated_hash(tmp_path):
copied = tmp_path / "bundle"
shutil.copytree(FIXTURE_DIRECTORY, copied)
Expand Down
50 changes: 46 additions & 4 deletions renderer/wang_trace.py
Original file line number Diff line number Diff line change
Expand Up @@ -422,6 +422,7 @@ def replay_trace(trace: TraceSnapshot) -> tuple[tuple[int, ...], ...]:

domains = list(trace.initial_domains)
changes: list[tuple[int, int]] = []
search_change_floor: int | None = None
states: list[tuple[int, ...]] = []
checkpoints = {item.event_sequence: item for item in trace.checkpoints}
if len(checkpoints) != len(trace.checkpoints):
Expand Down Expand Up @@ -453,6 +454,13 @@ def replay_trace(trace: TraceSnapshot) -> tuple[tuple[int, ...], ...]:
_fail("trace $.capacity.checkpoints_truncated", "is inconsistent")

for index, event in enumerate(trace.events):
if event.cell is not None and event.cell >= area:
_fail(f"trace $.events[{index}].cell", "lies outside layout")
if event.phase == "search":
if search_change_floor is None:
search_change_floor = len(changes)
elif event.phase == "initial" and search_change_floor is not None:
_fail(f"trace $.events[{index}].phase", "returns to initial after search")
if event.kind != "result" and event.status is not None:
_fail(f"trace $.events[{index}].status", "is reserved for result")
if event.kind != "domain_reduction" and event.reason is not None:
Expand All @@ -469,7 +477,7 @@ def replay_trace(trace: TraceSnapshot) -> tuple[tuple[int, ...], ...]:
elif event.kind == "domain_reduction":
if event.phase not in _PHASES or event.reason not in _REASONS:
_fail(f"trace $.events[{index}]", "reduction needs phase and reason")
if event.cell is None or event.cell >= area:
if event.cell is None:
_fail(f"trace $.events[{index}].cell", "lies outside layout")
if event.old_domain is None or event.new_domain is None:
_fail(f"trace $.events[{index}]", "reduction needs both domains")
Expand All @@ -482,7 +490,11 @@ def replay_trace(trace: TraceSnapshot) -> tuple[tuple[int, ...], ...]:
if event.change_mark != len(changes):
_fail(f"trace $.events[{index}].change_mark", "is inconsistent")
elif event.kind == "backtrack":
if event.phase != "search" or event.change_mark > len(changes):
if (
event.phase != "search"
or search_change_floor is None
or not search_change_floor <= event.change_mark <= len(changes)
):
_fail(f"trace $.events[{index}]", "backtrack mark is inconsistent")
while len(changes) > event.change_mark:
cell, previous = changes.pop()
Expand All @@ -503,6 +515,20 @@ def replay_trace(trace: TraceSnapshot) -> tuple[tuple[int, ...], ...]:
_fail(f"trace $.events[{index}].new_domain", "is not a singleton")
if event.old_domain is None or event.new_domain & ~event.old_domain:
_fail(f"trace $.events[{index}]", "decision lies outside old domain")
next_event = trace.events[index + 1]
if next_event.sequence == event.sequence + 1 and not (
next_event.kind == "domain_reduction"
and next_event.phase == event.phase
and next_event.reason == "decision"
and next_event.depth == event.depth
and next_event.cell == event.cell
and next_event.old_domain == event.old_domain
and next_event.new_domain == event.new_domain
):
_fail(
f"trace $.events[{index}]",
"decision does not match its following domain reduction",
)
elif event.kind == "propagation":
if event.phase not in _PHASES:
_fail(f"trace $.events[{index}].phase", "is required")
Expand Down Expand Up @@ -648,8 +674,24 @@ def load_trace_bundle(path: str | Path) -> TraceBundle:
solution = _project_wang_presentation(documents["solution"])
if trace.solution_sha256 != digests["solution"]:
_fail("trace $.solution_sha256", "does not match solution")
if (solution.width, solution.height) != (trace.width, trace.height):
_fail("solution $.bounds", "does not match trace layout")
if (
solution.min_x,
solution.min_y,
solution.max_x,
solution.max_y,
) != (
region.min_x,
region.min_y,
region.max_x,
region.max_y,
):
_fail("solution $.bounds", "does not match referenced region")
if solution.tile_edges != tileset.tile_edges:
_fail("solution $.tile_table", "does not match referenced tileset")
if tuple(cell is not None for cell in solution.cells) != region.active:
_fail("solution $.cells", "active map does not match referenced region")
if solution.boundary != region.boundary:
_fail("solution $.boundary", "does not match referenced region")
if not trace.truncated:
final_domains = replay_trace(trace)[-1]
expected_domains = tuple(
Expand Down
Loading