diff --git a/docs/design/meta-model-traces.md b/docs/design/meta-model-traces.md index 805db55d..55e00fd0 100644 --- a/docs/design/meta-model-traces.md +++ b/docs/design/meta-model-traces.md @@ -22,6 +22,24 @@ telemetry), and — for decode rows — policy checkpoint SHA, decode-config hash (`decode_config_hash`), tokenizer/grammar versions, and seed. Rollouts from different checkpoints are never mixed unlabeled. +### Solver-transition events (VSS1-04 / SLM-64) + +When the verified solver runs during decode (`verified_solver_decode`, default +off), each decision emits typed events into the same decode row's `events` +list — `solver_state`, `support_result`, `certified_deduction`, `decision`, +`backtrack`, `nogood`, `solver_terminal` — built by +[`dsl/solver/replay.py`](../../src/slm_training/dsl/solver/replay.py). Exact-closure +decode emits the `solver_state` / `support_result` / `certified_deduction` / +`solver_terminal` subset; the reversible search controller additionally emits +`decision` / `backtrack` / `nogood`. The row also carries a bounded `solver` +sidecar — `{schema_version, certificate_mode, certificates, counters}` — where +`certificate_mode` (`none`/`summary`/`full`) gates certificate detail: `none` +keeps only counters + honest status, `summary` compact descriptors, `full` the +replay material (each certificate's `to_dict()`, whose recomputed digest must +equal its id). Events and certificates carry only token/path ids and SHA-256 +digests — never raw region/user text. The schema bumps the decode trace to +`version = 3`; v1/v2 rows load and replay unchanged. + ## Replayability `replay_violations(trace)` certifies the decode stream is self-consistent: @@ -31,6 +49,19 @@ trace with steps must carry a final canvas. Empty list = replayable; the fixture test proves both directions (a clean fixture decode passes, a corrupted canvas is caught). +For solver-transition events it additionally invokes +`solver_replay_violations` ([`dsl/solver/replay.py`](../../src/slm_training/dsl/solver/replay.py)), +which checks the ten VSS1-04 invariants: fingerprint lineage (each transition's +`before_fingerprint` matches the active replay state), certified deductions +remove only live values and — in `full` mode — cite a present, digest-consistent +certificate (tamper detection), `unknown` support never removes, decisions select +exactly one live value and record the rest, backtracks restore a recorded +state/level, a `nogood` is never a certified deduction, a `solved` terminal +carries a verifier report, `certified_unsat` is impossible once any +`unknown`/budget/truncation appears, event counts match the sidecar counters, and +a truncated snapshot is reported non-replayable rather than accepted as an +exhaustive proof. Violations are human-readable strings, never assertions. + ## Retention and bucket layout Local stores live where run evidence already lives diff --git a/docs/design/telemetry.md b/docs/design/telemetry.md index 8f456aae..d13ac6ec 100644 --- a/docs/design/telemetry.md +++ b/docs/design/telemetry.md @@ -41,6 +41,32 @@ the remote endpoint is missing or unavailable. **Generate:** `generate_batch` → `generate_once` / `best_of_n_rank`, plus `context_encode` inside the model. +## Decode-stats solver work metrics (VSS1-04 / SLM-64) + +The verified solver's per-decode work is measured on the existing +[`DecodeStats`](../../src/slm_training/models/decode_stats.py) envelope (not a new +owner). All fields default to zero on every historical/default path (solver +disabled), and solver wall time is separated from `denoiser_ms` / `projection_ms`. +Stable names: + +| Field | Meaning | +| --- | --- | +| `solver_ms` | Solver wall time (`timed_ms`), separate from denoiser/projection. | +| `solver_enabled` | `1` when the solver ran on a decision, else `0`. | +| `solver_closure_passes` | Exact-closure fixed-point passes. | +| `solver_support_queries` / `solver_support_cache_hits` | Support-oracle queries and request-local cache hits. | +| `solver_supported` / `solver_unsupported` / `solver_unknown` | Tri-state support verdict counts. | +| `solver_certified_removed` | Candidates removed by replay-valid certificates. | +| `solver_decisions` / `solver_backtracks` / `solver_nogoods` | Reversible-search work (controller path). | +| `solver_expanded_nodes` / `solver_verifier_calls` | Enumeration nodes and verifier calls. | +| `solver_certificate_replay_failures` | Certificate replays that failed (0 at decode — closure never removes on a failed replay; populated by offline trace audits). | +| `solver_terminal_status` | Honest terminal: `unknown` / `certified_unsat` / `budget_exhausted` (closure never claims `solved`). | + +They surface **only** under `metrics["decode_stats"]` in `eval_.json` (and, +transitively, `scoreboard.json`) via `aggregate_stats`; no new top-level metric +keys or files. They do not overload the existing grammar/lattice candidate +counters. + ## How to use ```bash diff --git a/docs/design/verified-scope-solver.md b/docs/design/verified-scope-solver.md index 1e0e71a5..f8645606 100644 --- a/docs/design/verified-scope-solver.md +++ b/docs/design/verified-scope-solver.md @@ -349,6 +349,40 @@ removal/keep/bottom/coverage semantics are pinned Torch-free by or experiment ran; the enabled path is unmeasured and this makes no correctness, readiness, speed, or ship claim.** +## Implemented trace + replay (VSS1-04 / SLM-64) + +[`dsl/solver/replay.py`](../../src/slm_training/dsl/solver/replay.py) turns the +solver artifacts into typed, replayable events recorded inside the existing +decode trace (`DecodeTraceRecorder.events`, schema bumped to `version = 3`), and +`solver_replay_violations` validates them. Event kinds: `solver_state`, +`support_result`, `certified_deduction`, `decision`, `backtrack`, `nogood`, +`solver_terminal`. Exact-closure decode (`models/twotower.py` +`_record_solver_metrics`, gated on an attached recorder) emits the closure subset; +the reversible controller additionally emits decisions/backtracks/nogoods. A +bounded `solver` sidecar carries `{schema_version, certificate_mode, certificates, +counters}`; `solver_certificate_mode` gates certificate detail (`none` counters + +honest status only, `summary` compact descriptors, `full` replay material). Only +token/path ids and SHA-256 digests are stored — never raw region/user text. + +The replay validator (invoked by `harnesses/distill/trace_store.py` +`replay_violations`) checks ten invariants: fingerprint lineage; certified +deductions remove only live values and — in `full` mode — cite a present, +digest-consistent certificate (tamper detection, using the same canonical-JSON +digest as `SupportCertificate.digest`); `unknown` never removes; decisions select +one live value and record the rest; backtracks restore a recorded state/level; +a `nogood` is never a certified deduction; a `solved` terminal carries a verifier +report; `certified_unsat` is impossible once any `unknown`/budget/truncation +appears; event counts match the sidecar counters; truncated snapshots are +reported non-replayable. Solver work metrics (`solver_ms` separated from +denoiser/projection, plus tri-state and certified-removal counters) ride the +existing `DecodeStats` → `metrics["decode_stats"]` envelope, zero on the default +path. Closure never reports `solved` (it prunes; it does not materialize a +verifier-accepted terminal). Pinned by `tests/test_dsl/test_solver_replay.py` +(validator + closure round-trip + tamper) and +`tests/test_harnesses/distill/test_solver_trace.py` (recorder capture, clean +replay, counters, historical compat). **No train/eval/benchmark/checkpoint ran; +this makes no solver-quality, correctness, speed, or ship claim.** + ## Reference support semantics | Verdict | Requirement | Removal permitted? | diff --git a/docs/design/vss1-04-solver-trace-replay-20260718.md b/docs/design/vss1-04-solver-trace-replay-20260718.md new file mode 100644 index 00000000..fd0a2ffa --- /dev/null +++ b/docs/design/vss1-04-solver-trace-replay-20260718.md @@ -0,0 +1,65 @@ +# VSS1-04 solver trace + replay — fixture evidence (2026-07-18) + +Fixture-grade wiring evidence for SLM-64 (VSS1-04): making certified-solver +transitions replayable and measured on the **existing** decode trace + telemetry +owners. No second trace store, no new output root, no custom binary format. + +## What was implemented + +- [`dsl/solver/replay.py`](../../src/slm_training/dsl/solver/replay.py) — Torch-free + typed event builders (`solver_state`, `support_result`, `certified_deduction`, + `decision`, `backtrack`, `nogood`, `solver_terminal`), mode-gated certificate + serialization (`none` / `summary` / `full`), and `solver_replay_violations` + (ten invariants, human-readable strings). +- [`harnesses/distill/trace_store.py`](../../src/slm_training/harnesses/distill/trace_store.py) + — decode trace **schema `version = 3`** (backward-compatible v1/v2 readers), + `DecodeTraceRecorder.record_solver` sidecar, and `replay_violations` extended to + run the solver invariants on any solver events present. +- [`models/decode_stats.py`](../../src/slm_training/models/decode_stats.py) — solver + work-metric counters + `solver_ms` (separated from denoiser/projection), zero on + the default path, surfaced only under `metrics["decode_stats"]` via + `aggregate_stats`. +- [`models/twotower.py`](../../src/slm_training/models/twotower.py) + `_record_solver_metrics` — folds closure counters into `DecodeStats` and, when a + `DecodeTraceRecorder` is attached, emits the closure-subset events + a bounded + certificate/counter sidecar. No-op when neither stats nor a recorder is active. + +## Schema version + +- Decode trace: `TRACE_VERSION = 3` (was 2). v1/v2 rows load and replay unchanged. +- Solver event stream: `SOLVER_TRACE_SCHEMA_VERSION = 1`. +- Certificate schema: `CERTIFICATE_SCHEMA_VERSION = 1` (unchanged; VSS0-04). + +## Privacy / boundedness + +Events and certificates carry only token/path ids and SHA-256 digests — never raw +region/user text (the terminal verifier-report summarizer drops non-allowlisted +strings). The `solver_state` domain snapshot is bounded; a truncated snapshot sets +`trace_truncated=true` and the validator reports the trace non-replayable rather +than accepting bounded evidence as an exhaustive proof. + +## Test command + +```bash +python -m pytest \ + tests/test_dsl/test_solver_replay.py \ + tests/test_harnesses/distill/test_solver_trace.py \ + tests/test_models/test_decode_stats.py \ + tests/test_models/test_trace_store.py \ + tests/test_harnesses/distill/test_meta_traces.py -q +python -m scripts.repo_policy +``` + +Result: solver-replay + trace + decode-stats + historical-compat suites pass; +`repo_policy` ok. (The issue's suggested `tests/test_runtime` path does not hold +trace tests — the runtime-trace tests are `tests/test_runtime_trace.py`; the +decode-trace/replay tests are the paths above.) + +## Honesty + +Fixture-grade wiring only: the event schema, the ten replay invariants (including +`full`-mode certificate tamper detection), the decode-stats counters, and +historical-trace compatibility are tested on tiny closed fixtures. No model, +checkpoint, training corpus, or eval run is produced or claimed. Closure never +reports `solved`. **No solver-quality, correctness, speed, or ship claim is +made.** diff --git a/src/slm_training/dsl/solver/replay.py b/src/slm_training/dsl/solver/replay.py new file mode 100644 index 00000000..885313e2 --- /dev/null +++ b/src/slm_training/dsl/solver/replay.py @@ -0,0 +1,575 @@ +"""Replayable solver-transition trace events + validator (VSS1-04 / SLM-64). + +Torch-free. Turns the certified-solver artifacts (exact-closure / search results +and their certificates) into typed, replayable events recorded inside the +existing ``DecodeTraceRecorder.events`` stream, and validates that every +destructive transition is consistent and — in ``full`` certificate mode — +certificate-checked. + +Honesty invariants (owned here; see ``docs/design/verified-scope-solver.md``): + +* ``unknown`` support never removes a candidate; +* a ``certified_deduction`` removes only currently-live values and cites a + certificate; in ``full`` mode that certificate must be present and its + recomputed digest must equal its id (tamper detection); +* a ``nogood`` is never a certified deduction — a "deduction" with no certificate + is a nogood masquerading as proof and is a violation; +* ``certified_unsat`` is impossible once any ``unknown`` / budget / truncation + appears on the path; +* a ``solved`` terminal must carry a final verifier report; +* a truncated (bounded) snapshot is reported as **non-replayable**, never + accepted as an exhaustive proof. + +The event stream is a subset of the schema per producer: exact-closure decode +emits ``solver_state`` / ``support_result`` / ``certified_deduction`` / +``solver_terminal``; the reversible search controller additionally emits +``decision`` / ``backtrack`` / ``nogood``. The validator handles the full schema. + +This module writes no Torch, runs no model, and makes no quality/ship claim. +""" + +from __future__ import annotations + +import hashlib +import json +from typing import Any + +SOLVER_TRACE_SCHEMA_VERSION = 1 + +SOLVER_EVENT_KINDS = frozenset( + { + "solver_state", + "support_result", + "certified_deduction", + "decision", + "backtrack", + "nogood", + "solver_terminal", + } +) + +CERTIFICATE_MODES = ("none", "summary", "full") + +# Bounded live-value snapshot per ``solver_state`` (privacy + boundedness). +_MAX_DOMAIN_SNAPSHOT = 512 + +# Verifier-report keys whose string values are provenance labels (never user +# text); every other string is dropped so no raw region text can leak. +_REPORT_STR_ALLOW = frozenset({"name", "profile", "status", "verifier", "verdict"}) + + +def _canonical(obj: Any) -> str: + return json.dumps(obj, sort_keys=True, separators=(",", ":"), default=str) + + +def _digest(obj: Any) -> str: + return hashlib.sha256(_canonical(obj).encode()).hexdigest() + + +def _value_key(value_dict: dict) -> str: + return _canonical(value_dict) + + +def _hole_key(hole_dict: dict) -> str: + return _canonical(hole_dict) + + +# --------------------------------------------------------------------------- # +# Event builders +# --------------------------------------------------------------------------- # + + +def solver_state_event(state, *, max_snapshot: int = _MAX_DOMAIN_SNAPSHOT) -> dict: + """A ``solver_state`` event with a bounded live-domain snapshot. + + ``domain`` maps each hole to its live value keys so the validator can verify + deductions/decisions remove/select only live values. If the snapshot exceeds + ``max_snapshot`` it is truncated and ``trace_truncated`` is set — the validator + then refuses to treat the trace as a replayable exhaustive proof. + """ + domain: dict[str, list[str]] = {} + total = 0 + truncated = False + for hole in state.holes: + hole_key = _hole_key(hole.hole_id.to_dict()) + values: list[str] = [] + for value in hole.values: + if total >= max_snapshot: + truncated = True + break + values.append(_value_key(value.to_dict())) + total += 1 + domain[hole_key] = values + if truncated: + break + return { + "kind": "solver_state", + "state_fingerprint": state.fingerprint, + "problem_id": state.problem_id, + "pack_id": state.pack_id, + "constraint_version": state.constraint_version, + "bounds": state.bounds.to_dict(), + "decision_level": state.decision_level, + "domain_summary": state.summary(), + "domain": domain, + "trace_truncated": truncated, + } + + +def support_result_event( + *, + state_fingerprint: str, + hole_id_dict: dict, + candidate_dict: dict, + verdict: str, + certificate_id: str | None = None, + witness_digest: str | None = None, + stop_reason: str | None = None, + coverage: tuple[str, ...] = (), + counters: dict | None = None, +) -> dict: + return { + "kind": "support_result", + "state_fingerprint": state_fingerprint, + "hole_id": hole_id_dict, + "candidate": candidate_dict, + "verdict": verdict, + "certificate_id": certificate_id, + "witness_digest": witness_digest, + "stop_reason": stop_reason, + "coverage": list(coverage), + "counters": counters or {}, + } + + +def certified_deduction_event(deduction) -> dict: + return {"kind": "certified_deduction", **deduction.to_dict()} + + +def decision_event(decision) -> dict: + return {"kind": "decision", **decision.to_dict()} + + +def backtrack_event( + *, + from_fingerprint: str, + to_fingerprint: str, + from_level: int, + to_level: int, + decision_id: str, + conflict_kind: str, +) -> dict: + return { + "kind": "backtrack", + "from_fingerprint": from_fingerprint, + "to_fingerprint": to_fingerprint, + "from_level": from_level, + "to_level": to_level, + "decision_id": decision_id, + "conflict_kind": conflict_kind, + } + + +def nogood_event(nogood) -> dict: + return {"kind": "nogood", **nogood.to_dict()} + + +def _summarize_report(report: Any) -> Any: + if report is None: + return None + if isinstance(report, bool): + return report + if isinstance(report, (int, float)): + return report + if isinstance(report, str): + return report[:64] + if isinstance(report, dict): + summary: dict[str, Any] = {} + for key, value in report.items(): + if isinstance(value, bool) or isinstance(value, (int, float)): + summary[key] = value + elif isinstance(value, str) and key in _REPORT_STR_ALLOW: + summary[key] = value[:64] + return summary + return None + + +def solver_terminal_event( + *, + status: str, + source_digest: str | None = None, + verifier_report: Any = None, + certificate_mode: str = "full", + trace_truncated: bool = False, +) -> dict: + return { + "kind": "solver_terminal", + "status": str(status), + "source_digest": source_digest, + "verifier_report": _summarize_report(verifier_report), + "certificate_mode": certificate_mode, + "trace_truncated": bool(trace_truncated), + } + + +# --------------------------------------------------------------------------- # +# Certificate serialization by mode +# --------------------------------------------------------------------------- # + + +def serialize_certificates(certificate_store: dict, mode: str) -> dict: + """Bounded, mode-gated certificate artifacts keyed by certificate id. + + ``none`` → ``{}`` (aggregate counters/status only); ``summary`` → compact, + non-replayable descriptors; ``full`` → the replay material (each cert's + ``to_dict()``, whose recomputed digest must equal its id). + """ + if mode not in CERTIFICATE_MODES: + raise ValueError(f"unsupported solver_certificate_mode: {mode!r}") + if mode == "none": + return {} + out: dict[str, dict] = {} + for cid, cert in certificate_store.items(): + payload = cert.to_dict() + if mode == "full": + out[cid] = payload + else: # summary + out[cid] = { + "schema_version": payload.get("schema_version"), + "verdict": payload.get("verdict"), + "exhausted": payload.get("exhausted"), + "coverage_observations": payload.get("coverage_observations"), + "witness_digest": payload.get("witness_digest"), + } + return out + + +# --------------------------------------------------------------------------- # +# Producers: solver result -> ordered event stream +# --------------------------------------------------------------------------- # + + +def solver_events_from_closure( + result, + root_state, + *, + certificate_mode: str = "full", +) -> list[dict]: + """Event stream for one exact-closure decode prune. + + Ordered: root ``solver_state``, ``support_result`` (supported witnesses then + unknown queries), ``certified_deduction`` (closure application order — passes + chain by fingerprint), and a ``solver_terminal``. + """ + events: list[dict] = [solver_state_event(root_state)] + for witness in result.witnesses: + events.append( + support_result_event( + state_fingerprint=root_state.fingerprint, + hole_id_dict=witness.hole_id.to_dict(), + candidate_dict=witness.value.to_dict(), + verdict="supported", + certificate_id=witness.certificate_id, + witness_digest=witness.witness_digest, + coverage=("complete",), + ) + ) + for query in result.unknown_queries: + events.append( + support_result_event( + state_fingerprint=query.state_fingerprint, + hole_id_dict=query.hole_id.to_dict(), + candidate_dict=query.candidate.to_dict(), + verdict="unknown", + stop_reason=result.stop_reason, + ) + ) + for deduction in result.deductions: + events.append(certified_deduction_event(deduction)) + truncated = any(e.get("trace_truncated") for e in events) + events.append( + solver_terminal_event( + status=closure_status(result), + certificate_mode=certificate_mode, + trace_truncated=truncated, + ) + ) + return events + + +def closure_status(result) -> str: + """Honest terminal status for a closure *prune*. + + Closure never claims ``solved``: it prunes to a live subset but does not + itself materialize a verifier-accepted terminal (that is the controller's job, + and the decode's own final validate). It reports ``certified_unsat`` only on a + certified bottom, ``budget_exhausted`` on a budget stop, else ``unknown``. + """ + if result.state.is_bottom: + return "certified_unsat" + if result.stop_reason and result.stop_reason.startswith("budget"): + return "budget_exhausted" + return "unknown" + + +def solver_events_from_search( + result, + root_state, + *, + certificate_mode: str = "full", + verifier_report: Any = None, +) -> list[dict]: + """Event stream for a reversible search-controller run (decisions/nogoods).""" + events: list[dict] = [solver_state_event(root_state)] + for deduction in result.deductions: + events.append(certified_deduction_event(deduction)) + for decision in result.decisions: + events.append(decision_event(decision)) + for nogood in result.nogoods: + events.append(nogood_event(nogood)) + events.append( + solver_terminal_event( + status=getattr(result.status, "value", str(result.status)), + verifier_report=verifier_report + if verifier_report is not None + else result.verifier_report, + certificate_mode=certificate_mode, + ) + ) + return events + + +# --------------------------------------------------------------------------- # +# Aggregate counters +# --------------------------------------------------------------------------- # + + +def solver_trace_counters(events: list[dict]) -> dict[str, int]: + """Per-kind aggregate counts derived from an event stream (for invariant 9).""" + counts = { + "solver_states": 0, + "support_supported": 0, + "support_unsupported": 0, + "support_unknown": 0, + "certified_deductions": 0, + "certified_removed": 0, + "decisions": 0, + "backtracks": 0, + "nogoods": 0, + } + for event in events: + kind = event.get("kind") + if kind == "solver_state": + counts["solver_states"] += 1 + elif kind == "support_result": + verdict = event.get("verdict") + if verdict == "supported": + counts["support_supported"] += 1 + elif verdict == "unsupported": + counts["support_unsupported"] += 1 + elif verdict == "unknown": + counts["support_unknown"] += 1 + elif kind == "certified_deduction": + counts["certified_deductions"] += 1 + counts["certified_removed"] += len(event.get("removed", [])) + elif kind == "decision": + counts["decisions"] += 1 + elif kind == "backtrack": + counts["backtracks"] += 1 + elif kind == "nogood": + counts["nogoods"] += 1 + return counts + + +# --------------------------------------------------------------------------- # +# Replay validator +# --------------------------------------------------------------------------- # + + +def solver_replay_violations( + events: list[dict], + *, + certificates: dict | None = None, + certificate_mode: str = "full", + counters: dict | None = None, +) -> list[str]: + """Validate one solver event stream; empty list ⇒ replayable. + + Returns human-readable violation strings (never raises). Checks the ten + VSS1-04 invariants: fingerprint lineage, live-only removals + certificate + replay (full mode), unknown-never-removes, single-live decisions, backtrack + lineage, nogood-not-a-deduction, solved-has-report, certified-unsat purity, + counter agreement, and truncation honesty. + """ + certificates = certificates or {} + violations: list[str] = [] + + active_fp: str | None = None + live: dict[str, set[str]] = {} + # backtrack targets: fingerprint -> (level, live snapshot) + recorded: dict[str, tuple[int, dict[str, set[str]]]] = {} + pending_before: str | None = None + pending_after: str | None = None + saw_unknown = False + saw_budget = False + saw_truncation = False + + def commit_pending() -> None: + nonlocal active_fp, pending_before, pending_after + if pending_before is not None and pending_after is not None: + active_fp = pending_after + pending_before = None + pending_after = None + + for index, event in enumerate(events): + kind = event.get("kind") + if kind not in SOLVER_EVENT_KINDS: + violations.append(f"event {index}: unknown solver event kind {kind!r}") + continue + if event.get("trace_truncated"): + saw_truncation = True + + if kind == "solver_state": + commit_pending() + active_fp = event.get("state_fingerprint") + live = { + hole_key: set(values) + for hole_key, values in (event.get("domain") or {}).items() + } + recorded[active_fp] = ( + int(event.get("decision_level", 0)), + {h: set(v) for h, v in live.items()}, + ) + + elif kind == "support_result": + verdict = event.get("verdict") + if verdict == "unknown": + saw_unknown = True + if event.get("stop_reason", "") and str( + event.get("stop_reason") + ).startswith("budget"): + saw_budget = True + + elif kind == "certified_deduction": + before = event.get("before_fingerprint") + after = event.get("after_fingerprint") + # Passes share a before/after; a new before commits the prior pass. + if before != pending_before: + commit_pending() + if active_fp is not None and before != active_fp: + violations.append( + f"event {index}: deduction before_fingerprint " + f"{before!r} != active state {active_fp!r}" + ) + pending_before = before + pending_after = after + hole_key = _hole_key(event.get("hole_id", {})) + removed = [_value_key(v) for v in event.get("removed", [])] + cert_ids = event.get("certificate_ids", []) + if not cert_ids: + violations.append( + f"event {index}: certified_deduction cites no certificate " + "(a nogood must not be relabeled a certified deduction)" + ) + live_here = live.get(hole_key, set()) + for value_key in removed: + if value_key not in live_here: + violations.append( + f"event {index}: deduction removes non-live value at hole" + ) + else: + live_here.discard(value_key) + live[hole_key] = live_here + if certificate_mode == "full": + for cid in cert_ids: + if cid not in certificates: + violations.append( + f"event {index}: certificate {cid[:12]}… missing in full mode" + ) + elif _digest(certificates[cid]) != cid: + violations.append( + f"event {index}: certificate {cid[:12]}… digest mismatch (tampered)" + ) + + elif kind == "decision": + commit_pending() + before = event.get("before_fingerprint") + if active_fp is not None and before != active_fp: + violations.append( + f"event {index}: decision before_fingerprint " + f"{before!r} != active state {active_fp!r}" + ) + hole_key = _hole_key(event.get("hole_id", {})) + chosen = _value_key(event.get("chosen", {})) + live_here = live.get(hole_key, set()) + if chosen not in live_here: + violations.append( + f"event {index}: decision selects a non-live value" + ) + alternatives = {_value_key(v) for v in event.get("alternatives", [])} + expected_alts = live_here - {chosen} + if alternatives != expected_alts: + violations.append( + f"event {index}: decision alternatives do not match remaining live values" + ) + after = event.get("after_fingerprint") + recorded[before] = ( + int(event.get("level", 0)), + {h: set(v) for h, v in live.items()}, + ) + live = {h: set(v) for h, v in live.items()} + live[hole_key] = {chosen} + active_fp = after + + elif kind == "backtrack": + commit_pending() + to_fp = event.get("to_fingerprint") + to_level = int(event.get("to_level", 0)) + if to_fp not in recorded: + violations.append( + f"event {index}: backtrack to unrecorded state {to_fp!r}" + ) + else: + level, snapshot = recorded[to_fp] + if level != to_level: + violations.append( + f"event {index}: backtrack to_level {to_level} != recorded level {level}" + ) + active_fp = to_fp + live = {h: set(v) for h, v in snapshot.items()} + + elif kind == "nogood": + if not event.get("provenance"): + violations.append( + f"event {index}: nogood missing provenance" + ) + + elif kind == "solver_terminal": + commit_pending() + status = event.get("status") + if status == "solved" and event.get("verifier_report") is None: + violations.append( + f"event {index}: solved terminal without a verifier report" + ) + if status == "certified_unsat" and ( + saw_unknown or saw_budget or saw_truncation + ): + violations.append( + f"event {index}: certified_unsat with unknown/budget/truncation on the path" + ) + + if saw_truncation: + violations.append( + "solver trace is truncated: bounded evidence is not a replayable " + "exhaustive proof" + ) + + if counters is not None: + derived = solver_trace_counters(events) + for key, value in derived.items(): + if key in counters and int(counters[key]) != int(value): + violations.append( + f"counter {key} mismatch: trace {counters[key]} != events {value}" + ) + + return violations diff --git a/src/slm_training/harnesses/distill/trace_store.py b/src/slm_training/harnesses/distill/trace_store.py index eed05f13..829b200a 100644 --- a/src/slm_training/harnesses/distill/trace_store.py +++ b/src/slm_training/harnesses/distill/trace_store.py @@ -21,7 +21,7 @@ from pathlib import Path from typing import Any, Iterator -TRACE_VERSION = 2 +TRACE_VERSION = 3 # Config keys that change decode behavior (used for the decode-config hash). _DECODE_KEYS = ( @@ -109,6 +109,9 @@ def __init__( self.repair_commit_count = 0 self.remask_count = 0 self._depth = 0 + # VSS1-04: mode-serialized solver certificates + aggregate counters the + # replay validator needs (None until record_solver is called). + self.solver: dict[str, Any] | None = None # ── model-side hooks ───────────────────────────────────────────────── @@ -155,6 +158,16 @@ def step( def event(self, kind: str, **payload: Any) -> None: self.events.append({"kind": kind, "depth": self._depth, **payload}) + def record_solver(self, solver_block: dict[str, Any]) -> None: + """Attach the VSS1-04 solver certificate/counter block for replay. + + The typed solver transition events themselves are appended via + ``event("solver_state"|...)``; this stores the mode-serialized (bounded) + certificates and aggregate counters the replay validator cross-checks. + Overwrites any prior block (one solver run per decode trace). + """ + self.solver = dict(solver_block) + def end(self, *, canvas: list[int] | None = None, text: str | None = None) -> None: self._depth = max(0, self._depth - 1) if self._depth > 0: @@ -178,7 +191,7 @@ def finalize( final = dict(self.final or {}) if final_text is not None: final["text"] = final_text - return { + trace = { "version": TRACE_VERSION, "meta": {**self.meta, **meta}, "steps": self.steps, @@ -195,6 +208,11 @@ def finalize( "labels": labels or {}, "recorded_at": datetime.now(timezone.utc).isoformat(), } + if self.solver is not None: + # VSS1-04 solver certificate/counter sidecar; absent on non-solver + # traces so historical v1/v2 readers are unaffected. + trace["solver"] = self.solver + return trace class TraceStore: @@ -383,6 +401,28 @@ def replay_violations(trace: dict[str, Any]) -> list[str]: final = (trace.get("final") or {}).get("canvas") if steps and final is None: violations.append("trace has steps but no final canvas") + + # VSS1-04: validate any solver-transition events via the solver replay + # checker. Decode-only traces carry no such events, so this is a no-op there + # and historical v1/v2 traces are unaffected. + from slm_training.dsl.solver.replay import ( + SOLVER_EVENT_KINDS, + solver_replay_violations, + ) + + solver_events = [ + e for e in trace.get("events", []) if e.get("kind") in SOLVER_EVENT_KINDS + ] + if solver_events: + solver = trace.get("solver") or {} + violations.extend( + solver_replay_violations( + solver_events, + certificates=solver.get("certificates"), + certificate_mode=solver.get("certificate_mode", "summary"), + counters=solver.get("counters"), + ) + ) return violations diff --git a/src/slm_training/models/decode_stats.py b/src/slm_training/models/decode_stats.py index a9f13207..19f563f3 100644 --- a/src/slm_training/models/decode_stats.py +++ b/src/slm_training/models/decode_stats.py @@ -93,6 +93,26 @@ class DecodeStats: constraint_graph_edges: int = 0 completion_bound_known: int = 0 completion_bound_unknown: int = 0 + # VSS1-04 (SLM-64): verified-solver decode work metrics. Zero on every + # historical/default path (solver disabled); solver wall time is separated + # from denoiser_ms/projection_ms. Names are stable and documented in + # docs/design/telemetry.md. + solver_ms: float = 0.0 + solver_enabled: int = 0 + solver_closure_passes: int = 0 + solver_support_queries: int = 0 + solver_support_cache_hits: int = 0 + solver_supported: int = 0 + solver_unsupported: int = 0 + solver_unknown: int = 0 + solver_certified_removed: int = 0 + solver_decisions: int = 0 + solver_backtracks: int = 0 + solver_nogoods: int = 0 + solver_expanded_nodes: int = 0 + solver_verifier_calls: int = 0 + solver_certificate_replay_failures: int = 0 + solver_terminal_status: str = "" constrained_dead_ends: int = 0 constrained_dead_end_last_position: int = -1 constrained_dead_end_forced_rank: int = -1 @@ -247,6 +267,21 @@ def aggregate_stats(rows: list[DecodeStats]) -> dict[str, Any]: "constraint_graph_edges", "completion_bound_known", "completion_bound_unknown", + "solver_ms", + "solver_enabled", + "solver_closure_passes", + "solver_support_queries", + "solver_support_cache_hits", + "solver_supported", + "solver_unsupported", + "solver_unknown", + "solver_certified_removed", + "solver_decisions", + "solver_backtracks", + "solver_nogoods", + "solver_expanded_nodes", + "solver_verifier_calls", + "solver_certificate_replay_failures", ] out: dict[str, Any] = {"n": len(rows)} for key in keys: diff --git a/src/slm_training/models/twotower.py b/src/slm_training/models/twotower.py index acc96be8..f27a2249 100644 --- a/src/slm_training/models/twotower.py +++ b/src/slm_training/models/twotower.py @@ -3744,6 +3744,7 @@ def _solver_prune_forest(self, forest, prefix): OpenUIWellFormedVerifier, ) from slm_training.dsl.solver.state import SolverBounds + from slm_training.models.decode_stats import get_active_stats, timed_ms from slm_training.models.dsl_tokenizer import is_dsl_native_tokenizer if not is_dsl_native_tokenizer(self.tokenizer): @@ -3772,12 +3773,70 @@ def _solver_prune_forest(self, forest, prefix): ) provider = EnumerativeSupportProvider(expander, OpenUIWellFormedVerifier()) policy = str(getattr(self.config, "solver_unknown_policy", "keep_and_rank")) - pruned, _result = solver_prune( - forest, prefix, provider, pack_id="openui", constraint_version=cv, - bounds=bounds, unknown_policy=policy, state=expander.root_state(), cache={}, - ) + root_state = expander.root_state() + certificate_store: dict = {} + stats = get_active_stats() + # Solver wall time is separated from denoiser_ms/projection_ms (VSS1-04). + with timed_ms(stats, "solver_ms"): + pruned, result = solver_prune( + forest, prefix, provider, pack_id="openui", constraint_version=cv, + bounds=bounds, unknown_policy=policy, state=root_state, cache={}, + certificate_store=certificate_store, + ) + if result is not None: + self._record_solver_metrics( + result, root_state, certificate_store, stats + ) return pruned + def _record_solver_metrics(self, result, root_state, certificate_store, stats): + """VSS1-04: fold solver work into decode stats and, when a trace recorder + is attached, emit replayable solver-transition events + a bounded + certificate/counter sidecar. Counters ride the existing DecodeStats + envelope; nothing is emitted when neither stats nor a recorder is active. + """ + from slm_training.dsl.solver.replay import ( + SOLVER_TRACE_SCHEMA_VERSION, + closure_status, + serialize_certificates, + solver_events_from_closure, + solver_trace_counters, + ) + + if stats is not None: + counters = result.counters + stats.solver_enabled = 1 + stats.solver_closure_passes += counters.passes + stats.solver_support_queries += counters.support_queries + stats.solver_support_cache_hits += counters.cache_hits + stats.solver_supported += counters.supported + stats.solver_unsupported += counters.unsupported + stats.solver_unknown += counters.unknown + stats.solver_certified_removed += counters.candidates_removed + stats.solver_expanded_nodes += counters.expanded_nodes + stats.solver_verifier_calls += counters.verifier_calls + stats.solver_terminal_status = closure_status(result) + + recorder = getattr(self, "trace_recorder", None) + if recorder is None: + return + mode = str(getattr(self.config, "solver_certificate_mode", "summary")) + events = solver_events_from_closure( + result, root_state, certificate_mode=mode + ) + for event in events: + kind = event["kind"] + payload = {key: value for key, value in event.items() if key != "kind"} + recorder.event(kind, **payload) + recorder.record_solver( + { + "schema_version": SOLVER_TRACE_SCHEMA_VERSION, + "certificate_mode": mode, + "certificates": serialize_certificates(certificate_store, mode), + "counters": solver_trace_counters(events), + } + ) + def _compiler_ltr_decode_one( self, ctx: torch.Tensor, diff --git a/tests/test_dsl/test_solver_replay.py b/tests/test_dsl/test_solver_replay.py new file mode 100644 index 00000000..89352ad1 --- /dev/null +++ b/tests/test_dsl/test_solver_replay.py @@ -0,0 +1,339 @@ +"""VSS1-04 (SLM-64): solver-transition replay events + validator — core logic. + +Torch-free tests for `dsl/solver/replay.py`: a clean full-mode stream replays +with zero violations, and every honesty invariant (fingerprint lineage, +live-only removals, certificate tamper detection, unknown-never-removes, +single-live decisions, backtrack lineage, nogood-not-a-deduction, +solved-has-report, certified-unsat purity, counter agreement, truncation +honesty) is independently detected. Model wiring/parity is covered under +tests/test_models/ and tests/test_harnesses/. +""" + +from __future__ import annotations + +import copy + +import pytest + +from slm_training.dsl.solver.replay import ( + CERTIFICATE_MODES, + SOLVER_EVENT_KINDS, + _digest, + serialize_certificates, + solver_replay_violations, + solver_trace_counters, +) + +_ROOT = "fp_root" +_S1 = "fp_s1" + + +def _cert_payload(tag: str) -> dict: + # An opaque certificate dict whose id is its own sha256 digest (as in the + # real store: certificate_store[cert.digest] = cert). + return { + "schema_version": 1, + "verdict": "unsupported", + "exhausted": True, + "coverage_observations": ["complete"], + "tag": tag, + } + + +def _hole(name: str = "h0") -> dict: + return {"namespace": "ns", "path": [name], "kind": "component"} + + +def _val(token: int) -> dict: + return {"tag": "path", "value": f'{{"token_ids":[{token}]}}'} + + +def _clean_full_trace(): + """A valid full-mode stream: root state, one certified removal, solved.""" + cert = _cert_payload("c1") + cid = _digest(cert) + events = [ + { + "kind": "solver_state", + "state_fingerprint": _ROOT, + "problem_id": "p", + "pack_id": "openui", + "constraint_version": "cv", + "bounds": {}, + "decision_level": 0, + "domain_summary": {"hole_count": 1}, + # value keys are canonical-JSON of the value dict (validator convention) + "domain": {_hole_str(): [_valkey(10), _valkey(20), _valkey(30)]}, + "trace_truncated": False, + }, + { + "kind": "certified_deduction", + "before_fingerprint": _ROOT, + "after_fingerprint": _S1, + "hole_id": _hole(), + "removed": [_val(10)], + "certificate_ids": [cid], + "reason": "certified_unsupported", + }, + { + "kind": "solver_terminal", + "status": "solved", + "source_digest": "src", + "verifier_report": {"name": "OpenUIWellFormed", "accepted": True}, + "certificate_mode": "full", + "trace_truncated": False, + }, + ] + certificates = {cid: cert} + return events, certificates + + +def _hole_str() -> str: + from slm_training.dsl.solver.replay import _hole_key + + return _hole_key(_hole()) + + +def _valkey(token: int) -> str: + from slm_training.dsl.solver.replay import _value_key + + return _value_key(_val(token)) + + +def test_clean_full_trace_replays_without_violations(): + events, certs = _clean_full_trace() + assert solver_replay_violations(events, certificates=certs, certificate_mode="full") == [] + + +def test_deduction_removing_non_live_value_is_detected(): + # "unknown-preservation": a value kept (never in the live snapshot) cannot be + # certified-removed. + events, certs = _clean_full_trace() + events[1]["removed"] = [_val(99)] # 99 was never live + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("non-live" in v for v in violations) + + +def test_missing_certificate_is_detected(): + events, _certs = _clean_full_trace() + violations = solver_replay_violations(events, certificates={}, certificate_mode="full") + assert any("missing" in v for v in violations) + + +def test_tampered_certificate_is_detected(): + events, certs = _clean_full_trace() + cid = next(iter(certs)) + certs[cid] = {**certs[cid], "tag": "TAMPERED"} # digest no longer matches id + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("digest mismatch" in v or "tampered" in v.lower() for v in violations) + + +def test_nogood_relabeled_as_deduction_is_detected(): + events, certs = _clean_full_trace() + events[1]["certificate_ids"] = [] # a deduction with no certificate is a nogood + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("nogood" in v.lower() for v in violations) + + +def test_solved_without_verifier_report_is_detected(): + events, certs = _clean_full_trace() + events[-1]["verifier_report"] = None + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("without a verifier report" in v for v in violations) + + +def test_certified_unsat_with_unknown_is_detected(): + events, certs = _clean_full_trace() + events.insert(1, { + "kind": "support_result", + "state_fingerprint": _ROOT, + "hole_id": _hole(), + "candidate": _val(20), + "verdict": "unknown", + "certificate_id": None, + "witness_digest": None, + "stop_reason": None, + "coverage": [], + "counters": {}, + }) + events[-1]["status"] = "certified_unsat" + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("certified_unsat" in v for v in violations) + + +def test_bad_before_fingerprint_lineage_is_detected(): + events, certs = _clean_full_trace() + events[1]["before_fingerprint"] = "fp_wrong" + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("!= active state" in v for v in violations) + + +def test_backtrack_to_unrecorded_state_is_detected(): + events, certs = _clean_full_trace() + events.insert(2, { + "kind": "backtrack", + "from_fingerprint": _S1, + "to_fingerprint": "fp_never_recorded", + "from_level": 1, + "to_level": 0, + "decision_id": "d0", + "conflict_kind": "certified_bottom", + }) + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("unrecorded state" in v for v in violations) + + +def test_truncated_trace_is_non_replayable(): + events, certs = _clean_full_trace() + events[0]["trace_truncated"] = True + violations = solver_replay_violations(events, certificates=certs, certificate_mode="full") + assert any("truncated" in v for v in violations) + + +def test_counter_mismatch_is_detected(): + events, certs = _clean_full_trace() + counters = solver_trace_counters(events) + counters["certified_deductions"] += 5 # lie about the count + violations = solver_replay_violations( + events, certificates=certs, certificate_mode="full", counters=counters + ) + assert any("counter" in v for v in violations) + + +def test_matching_counters_pass(): + events, certs = _clean_full_trace() + counters = solver_trace_counters(events) + assert solver_replay_violations( + events, certificates=certs, certificate_mode="full", counters=counters + ) == [] + + +def test_summary_and_none_modes_are_honest_about_limits(): + events, certs = _clean_full_trace() + cid = next(iter(certs)) + + class _Cert: + def __init__(self, payload): + self._p = payload + + def to_dict(self): + return self._p + + store = {cid: _Cert(certs[cid])} + full = serialize_certificates(store, "full") + summary = serialize_certificates(store, "summary") + none = serialize_certificates(store, "none") + assert none == {} + # summary drops the replay material (no 'tag'); full keeps it. + assert "tag" in full[cid] and "tag" not in summary[cid] + # In summary mode a tampered cert is NOT caught (summary is not a replay + # guarantee) — honest limitation, not a false pass in full mode. + tampered = copy.deepcopy(events) + tampered_certs = {cid: {**certs[cid], "tag": "X"}} + assert solver_replay_violations( + tampered, certificates=tampered_certs, certificate_mode="summary" + ) == [] + assert solver_replay_violations( + tampered, certificates=tampered_certs, certificate_mode="full" + ) != [] + + +def test_no_raw_text_leaks_into_terminal_report(): + from slm_training.dsl.solver.replay import solver_terminal_event + + event = solver_terminal_event( + status="solved", + verifier_report={ + "name": "OpenUIWellFormed", + "accepted": True, + "secret_note": "user typed their password here", + }, + ) + report = event["verifier_report"] + assert report["name"] == "OpenUIWellFormed" + assert report["accepted"] is True + assert "secret_note" not in report # non-allowlisted string dropped + + +def test_bad_mode_rejected(): + with pytest.raises(ValueError, match="solver_certificate_mode"): + serialize_certificates({}, "drop") + assert set(CERTIFICATE_MODES) == {"none", "summary", "full"} + assert "solver_state" in SOLVER_EVENT_KINDS + + +def test_closure_events_round_trip_replays_and_detects_tamper(): + """A real ClosureResult (with a real SupportCertificate) round-trips through + the producer + serializer and replays clean in full mode; tampering the + stored certificate breaks the digest check.""" + from slm_training.dsl.solver.closure import ( + CertifiedDeduction, + ClosureCounters, + ClosureResult, + ) + from slm_training.dsl.solver.replay import ( + serialize_certificates, + solver_events_from_closure, + ) + from slm_training.dsl.solver.state import ( + DomainValue, + FiniteDomainState, + HoleDomain, + HoleId, + SolverBounds, + SupportVerdict, + ) + from slm_training.dsl.solver.support import ( + SEARCH_ORDER, + SupportCertificate, + SupportQuery, + ) + + bounds = SolverBounds( + max_tokens=100, max_nodes=100, max_depth=8, max_backtracks=8, + max_verifier_calls=100, + ) + hole = HoleId(namespace="ns", path=("h0",), kind="component") + v_keep = DomainValue(tag="path", payload_json='{"token_ids":[20]}') + v_drop = DomainValue(tag="path", payload_json='{"token_ids":[10]}') + root = FiniteDomainState( + problem_id="p", pack_id="openui", constraint_version="cv", bounds=bounds, + holes=(HoleDomain(hole_id=hole, values=(v_drop, v_keep), metadata={}),), + ) + refined = root.refine(hole, (v_keep,)) + cert = SupportCertificate( + schema_version=1, + query=SupportQuery( + state_fingerprint=root.fingerprint, hole_id=hole, candidate=v_drop + ), + verdict=SupportVerdict.UNSUPPORTED, + problem_id="p", pack_id="openui", constraint_version="cv", bounds=bounds, + search_order=SEARCH_ORDER, explored_state_fingerprints=(), + coverage_observations=("complete",), verifier_profile="stub", exhausted=True, + ) + cid = cert.digest + deduction = CertifiedDeduction( + before_fingerprint=root.fingerprint, after_fingerprint=refined.fingerprint, + hole_id=hole, removed=(v_drop,), certificate_ids=(cid,), + reason="certified_unsupported", + ) + result = ClosureResult( + state=refined, deductions=(deduction,), unknown_queries=(), witnesses=(), + counters=ClosureCounters( + passes=1, support_queries=2, unsupported=1, candidates_removed=1 + ), + reached_fixed_point=True, + ) + events = solver_events_from_closure(result, root, certificate_mode="full") + certs = serialize_certificates({cid: cert}, "full") + assert solver_replay_violations( + events, certificates=certs, certificate_mode="full" + ) == [] + # The producer never claims a bare closure prune is "solved". + terminal = next(e for e in events if e["kind"] == "solver_terminal") + assert terminal["status"] == "unknown" + # Tamper the stored certificate -> digest no longer matches its id. + certs[cid] = {**certs[cid], "verifier_profile": "TAMPERED"} + assert solver_replay_violations( + events, certificates=certs, certificate_mode="full" + ) != [] diff --git a/tests/test_harnesses/distill/test_solver_trace.py b/tests/test_harnesses/distill/test_solver_trace.py new file mode 100644 index 00000000..e6aca6bd --- /dev/null +++ b/tests/test_harnesses/distill/test_solver_trace.py @@ -0,0 +1,120 @@ +"""VSS1-04 (SLM-64): model-level solver trace + replay + decode-stats wiring. + +Attaches a `DecodeTraceRecorder` to a solver-enabled TwoTower and drives one +`_solver_prune_forest` decision (fast — no full decode loop). Asserts the +recorder captures replayable solver-transition events + a bounded certificate +sidecar (`replay_violations` clean), the solver work-metric counters land in the +`DecodeStats` envelope (zero on the default path), and historical decode-only +traces still replay. Core validator semantics live in +tests/test_dsl/test_solver_replay.py. +""" + +from __future__ import annotations + +from slm_training.dsl.grammar.fastpath.compiler_draft import build_completion_forest +from slm_training.dsl.schema import ExampleRecord +from slm_training.harnesses.distill.trace_store import ( + DecodeTraceRecorder, + replay_violations, +) +from slm_training.models.decode_stats import DecodeStats, aggregate_stats, collect_decode_stats +from slm_training.models.twotower import TwoTowerConfig, TwoTowerModel + + +def _solver_model(): + record = ExampleRecord( + id="compiler", + prompt="card", + openui='root = Card([title])\ntitle = TextContent(":hero.title")\n', + placeholders=[":hero.title"], + split="train", + source="fixture", + ) + config = TwoTowerConfig( + context_backend="scratch", + output_tokenizer="lexer", + d_model=32, + n_heads=2, + context_layers=1, + denoiser_layers=1, + max_prompt_len=32, + max_target_len=32, + grammar_ltr_max_tokens=32, + gen_steps=1, + seed=0, + verified_solver_decode=True, + solver_max_nodes=4, + solver_certificate_mode="full", + ) + model = TwoTowerModel.from_records([record], config=config, device="cpu") + model.eval() + return model + + +def test_recorder_captures_replayable_solver_events(): + model = _solver_model() + prefix = [model.tokenizer.bos_id] + forest = build_completion_forest(model.tokenizer, prefix) + recorder = DecodeTraceRecorder() + model.trace_recorder = recorder + + model._solver_prune_forest(forest, prefix) + + from slm_training.dsl.solver.replay import SOLVER_EVENT_KINDS + + solver_events = [e for e in recorder.events if e.get("kind") in SOLVER_EVENT_KINDS] + assert solver_events, "expected solver-transition events on the recorder" + assert any(e["kind"] == "solver_state" for e in solver_events) + assert any(e["kind"] == "solver_terminal" for e in solver_events) + assert recorder.solver is not None + assert recorder.solver["certificate_mode"] == "full" + + trace = recorder.finalize() + assert trace["version"] == 3 + assert "solver" in trace + # The captured solver stream replays with zero violations. + assert replay_violations(trace) == [] + + +def test_solver_counters_land_in_decode_stats_envelope(): + model = _solver_model() + prefix = [model.tokenizer.bos_id] + forest = build_completion_forest(model.tokenizer, prefix) + + with collect_decode_stats() as stats: + model._solver_prune_forest(forest, prefix) + + assert stats.solver_enabled == 1 + assert stats.solver_terminal_status in {"unknown", "certified_unsat", "budget_exhausted"} + # Solver time is tracked separately from denoiser/projection. + assert stats.solver_ms >= 0.0 + # Counters surface (only) under metrics["decode_stats"] via aggregate_stats. + agg = aggregate_stats([stats]) + assert "solver_enabled_sum" in agg + assert "solver_support_queries_sum" in agg + + +def test_solver_counters_default_zero_when_disabled(): + stats = DecodeStats() + assert stats.solver_enabled == 0 + assert stats.solver_ms == 0.0 + assert stats.solver_certified_removed == 0 + assert stats.solver_terminal_status == "" + agg = aggregate_stats([stats]) + assert agg["solver_enabled_sum"] == 0.0 + assert agg["solver_certified_removed_mean"] == 0.0 + + +def test_historical_decode_only_trace_still_replays(): + # A v2-style decode trace with no solver events / no solver block is + # unaffected by the VSS1-04 replay extension. + trace = { + "version": 2, + "meta": {}, + "steps": [ + {"step": 0, "canvas": [5, 0], "commits": [{"t": 0, "id": 5}], "remasks": []} + ], + "events": [], + "final": {"canvas": [5, 2], "text": "x"}, + } + assert replay_violations(trace) == []