From 0ccbcbee177335cf6e658f5d67a6f0a0f2f076c6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Tue, 29 Sep 2026 22:54:57 +0000 Subject: [PATCH] Avoid canonical re-encoding for freshly generated planner reports --- src/art/trainer_rank/_planner_misses.py | 40 +++++- tests/unit/test_planner_generated_reports.py | 144 +++++++++++++++++++ 2 files changed, 180 insertions(+), 4 deletions(-) create mode 100644 tests/unit/test_planner_generated_reports.py diff --git a/src/art/trainer_rank/_planner_misses.py b/src/art/trainer_rank/_planner_misses.py index 0fbed41b1..889d90e60 100644 --- a/src/art/trainer_rank/_planner_misses.py +++ b/src/art/trainer_rank/_planner_misses.py @@ -199,17 +199,23 @@ def _source_files() -> dict[str, dict[str, str | int]]: return result -def validate_report(raw: bytes) -> dict[str, Any]: - """Validate the bounded, canonical transport without loading code/tensors.""" +def _decode_report(raw: bytes, *, generated: bool = False) -> dict[str, Any]: + """Parse and check schema; generated bytes must come directly from _encode.""" if len(raw) > MAX_REPORT_BYTES: raise ValueError("report exceeds byte limit") def pairs(items: list[tuple[str, Any]]) -> dict[str, Any]: result = {} + previous: str | None = None for key, value in items: if key in result: raise ValueError("duplicate report key") + # Encoder keys may be nonstrings/custom strings: their encoded or + # Unicode-normalized order can differ from the decoded string order. + if generated and previous is not None and key < previous: + raise ValueError("report must be canonical JSON with a final newline") result[key] = value + previous = key return result record = json.loads(raw, object_pairs_hook=pairs) @@ -287,6 +293,12 @@ def pairs(items: list[tuple[str, Any]]) -> dict[str, Any]: or abs(observed - predicted) * 100 <= record["threshold_pct"] * predicted ): raise ValueError("report percentage/threshold does not match measurements") + return record + + +def validate_report(raw: bytes) -> dict[str, Any]: + """Validate the bounded, canonical transport without loading code/tensors.""" + record = _decode_report(raw) if _encode(record) != raw: raise ValueError("report must be canonical JSON with a final newline") return record @@ -300,7 +312,23 @@ def persist_report( retention: RetentionLimits | None = None, ) -> Path: """Durably retain exact bytes; duplicate delivery is safe, conflicts refuse.""" - record = validate_report(raw) + return _persist_report( + raw, + validate_report(raw), + spool_dir, + planning_budget=planning_budget, + retention=retention, + ) + + +def _persist_report( + raw: bytes, + record: dict[str, Any], + spool_dir: Path, + *, + planning_budget: bool, + retention: RetentionLimits | None, +) -> Path: if retention is not None and spool_dir != retention.spool_dir: raise ValueError("planner report spool differs from assigned allowance") charged = ( @@ -520,8 +548,12 @@ def report( f"replay unavailable: {'ValueError' if isinstance(exc, _ReportTooLarge) else type(exc).__name__}" ] raw = _encode(record) - path = persist_report( + # Only this producer bypasses the canonical re-encode: _encode + # already fixed scalar spelling/escaping/whitespace and the newline. + # Decode still checks key normalization, duplicates and full schema. + path = _persist_report( raw, + _decode_report(raw, generated=True), self.spool_dir if retention is None else retention.spool_dir, planning_budget=planning_budget, retention=retention, diff --git a/tests/unit/test_planner_generated_reports.py b/tests/unit/test_planner_generated_reports.py new file mode 100644 index 000000000..f43118613 --- /dev/null +++ b/tests/unit/test_planner_generated_reports.py @@ -0,0 +1,144 @@ +"""Generated reports keep the external transport contract without re-encoding.""" + +import json + +import pytest +from test_planner_retention_budget import emit, ledger, limits + +from art.trainer_rank import _planner_misses as reports + + +class UnorderedString(str): + def __lt__(self, other): + return False + + +class DuplicateItems(dict): + def items(self): + return [("same", 1), ("same", 2)] + + +@pytest.mark.parametrize( + "payload,retained", + [ + ({"a": [1, 2.0, False, None]}, True), + ({"a": (1, 2, 3)}, True), + ({1: "one", 2: "two"}, True), + ({2: "two", 10: "ten"}, False), + ({False: 1, True: 2}, True), + ({None: 1}, True), + ({1.0: 1, 2.5: 2}, True), + ({"\ud800\udc00": 1, "\ue000": 2}, False), + ({"\ud800\udc00": 1, "\U00010000": 2}, False), + ({UnorderedString("z"): 1, UnorderedString("a"): 2}, False), + (DuplicateItems(seed=True), False), + ({"a": "\ud800\udc00/\n\u0000\udfff\u2028"}, True), + ({"a": [0.0, -0.0, 5e-324, -5e-324, 1.7976931348623157e308]}, True), + ], +) +def test_factory_outputs_keep_public_canonical_acceptance( + tmp_path, monkeypatch, payload, retained +): + monkeypatch.setattr(reports, "_source_files", lambda: {}) + encoded = [] + encode = reports._encode + + def capture(record, **kwargs): + raw = encode(record, **kwargs) + encoded.append(raw) + return raw + + monkeypatch.setattr(reports, "_encode", capture) + bound = limits(tmp_path) + reporter = reports.Reporter(5) + with reports.report_retention_scope(bound): + path = emit(reporter, replay_factory=lambda: {"nested": payload}) + assert (path is not None) == retained + # The first encoding is the actual producer output, before any validation. + raw = encoded[0] + if retained: + reports.validate_report(raw) + assert path.read_bytes() == raw + assert sum(x[1] for x in ledger(bound)["charges"].values()) == len(raw) + assert ledger(bound)["omitted"] == reporter.failures == 0 + else: + with pytest.raises(ValueError): + reports.validate_report(raw) + assert not list(bound.spool_dir.glob("*.json")) + assert reporter.failures == 1 + + +@pytest.mark.parametrize("payload", [{1: 1, "1": 2}, {"a": float("nan")}]) +def test_factory_encoding_failure_still_retains_original_fallback( + tmp_path, monkeypatch, payload +): + monkeypatch.setattr(reports, "_source_files", lambda: {}) + path = emit( + reports.Reporter(5, spool_dir=tmp_path / "reports"), + replay_factory=lambda: payload, + oom=True, + observed_peak_bytes=None, + partial_peak_bytes=99, + ) + record = reports.validate_report(path.read_bytes()) + assert record["oom"] and record["partial_peak_bytes"] == 99 + assert record["replay"] is None and not record["replay_complete"] + assert record["incomplete_reasons"][0].startswith("replay unavailable: ") + + +def test_generated_report_encodes_and_parses_once(tmp_path, monkeypatch): + monkeypatch.setattr(reports, "_source_files", lambda: {}) + calls = {"encode": 0, "decode": 0} + encode, decode = reports._encode, json.loads + + def counted_encode(*args, **kwargs): + calls["encode"] += 1 + return encode(*args, **kwargs) + + def counted_decode(*args, **kwargs): + calls["decode"] += 1 + return decode(*args, **kwargs) + + monkeypatch.setattr(reports, "_encode", counted_encode) + monkeypatch.setattr(json, "loads", counted_decode) + path = emit(reports.Reporter(5, spool_dir=tmp_path / "reports")) + assert path is not None + assert calls == {"encode": 1, "decode": 1} + # Public delivery remains a separate trust boundary, even for the same bytes. + reports.persist_report(path.read_bytes(), tmp_path / "external") + assert calls == {"encode": 2, "decode": 2} + + +@pytest.mark.parametrize( + "mutation", ["space", "newline", "escape", "number", "nested_nan"] +) +def test_external_transport_still_requires_full_canonical_validation( + tmp_path, monkeypatch, mutation +): + monkeypatch.setattr(reports, "_source_files", lambda: {}) + raw = emit(reports.Reporter(5, spool_dir=tmp_path / "seed")).read_bytes() + if mutation == "space": + raw = b" " + raw + elif mutation == "newline": + raw = raw[:-1] + elif mutation == "escape": + raw = raw.replace(b'"forward"', b'"for\\u0077ard"') + elif mutation == "number": + raw = raw.replace(b'"threshold_pct":5', b'"threshold_pct":5e0') + else: + raw = raw.replace(b'"replay":{', b'"replay":{"arbitrary":NaN,') + for function in ( + lambda: reports.validate_report(raw), + lambda: reports.persist_report(raw, tmp_path / "external"), + ): + with pytest.raises(ValueError): + function() + assert not (tmp_path / "external").exists() + + +@pytest.mark.parametrize("name", ["validate_report", "persist_report"]) +def test_generated_shortcut_is_not_a_public_flag(tmp_path, name): + function = getattr(reports, name) + args = (b"{}\n", tmp_path) if name == "persist_report" else (b"{}\n",) + with pytest.raises(TypeError): + function(*args, generated=True)