#!/usr/bin/env python3
"""IBP-Lapper page evaluator — standalone, no dependencies (python3 stdlib only).

The point of a lambda-witness is that ANYONE can check it with an independent
parser that imports none of the solver. This script is that parser: ~30 lines
of exact integer arithmetic mod p. It verifies the shipped example witness
against the shipped system and self-tests the certificate layer:

  [1] verify a true lambda-witness against its system      -> must PASS
  [2] verify the shipped MUTATED witness (bad coeffs)      -> must FAIL
  [3] table-mode: assert a WRONG claimed table row         -> must FAIL
                  (the confident-wrong-table / COPAIR class)
  [4] if the bench bundle is unpacked alongside, run a T1 instance

The identity checked (see the page's Certificates section):
      sum_i lam[i] * R_i  ==  e_target - sum_m c[m] * e_m   (mod p)
over the system's ordered sparse rows R_i. Exact match, no tolerances: one
wrong coefficient, one wrong lambda entry, one stray residual column -> FAIL.
The witness's recorded system fingerprint is the library's binding of witness
to system file; this independent checker instead verifies the identity itself
against the rows on disk, which is the stronger statement.

Deep numbers on the page are artifact-cited, not recomputed here: 0.69 ms/row
verify -> runs/laportaml/receipt_verify_timing_jul7/VERIFY_TIMING.json; bench
walls -> lapper_bench PAGE_TABLE.md; the wrong-table -> runs/laportaml/copair_day1/.

Run:  python3 winnow-evaluate.py        (exit 0 = all checks pass)
"""
import glob
import json
import os
import sys


def load_system(path):
    """system.jsonl: one sparse row per line, {col: val} with values in F_p."""
    return [{int(c): int(v) for c, v in json.loads(line).items()}
            for line in open(path) if line.strip()]


def verify(rows, p, target, c, lam_idx, lam_val, table_row=None):
    """Independent lambda-witness check. Returns (ok, detail).

    Accumulate residual = sum_i lam[i]*R_i (mod p) sparsely, then require it
    to equal {target: 1} u {m: -c[m]} EXACTLY. table_row, if given, replaces
    c as the claimed reduction row (table-mode: does the claimed table match
    what the witness actually proves?).
    """
    claimed = table_row if table_row is not None else c
    residual = {}
    for i, lv in zip(lam_idx, lam_val):
        for col, val in rows[i].items():
            residual[col] = (residual.get(col, 0) + lv * val) % p
    expected = {int(m): (-int(v)) % p for m, v in claimed.items()}
    expected[int(target)] = 1
    residual = {k: v for k, v in residual.items() if v % p}
    expected = {k: v for k, v in expected.items() if v % p}
    if residual == expected:
        return True, "exact match"
    extra = sorted(set(residual) - set(expected))
    missing = sorted(set(expected) - set(residual))
    diff = sorted(k for k in set(residual) & set(expected) if residual[k] != expected[k])
    return False, f"mismatch (extra cols {extra}, missing {missing}, wrong-value {diff})"


def load_witness(path):
    w = json.load(open(path))
    assert w.get("kind") == "lambda-witness", f"not a lambda-witness: {path}"
    return (int(w["p"]), int(w["target"]["col"]), w["c"],
            [int(i) for i in w["lam"]["idx"]], [int(v) for v in w["lam"]["val"]],
            int(w["system"]["n_rows"]))


def main():
    here = os.path.dirname(os.path.abspath(__file__))
    rows = load_system(os.path.join(here, "system.jsonl"))
    print(f"[0] loaded system: {len(rows)} rows")

    p, tgt, c, lidx, lval, n_rows = load_witness(os.path.join(here, "w_202.json"))
    assert n_rows == len(rows), f"witness expects {n_rows} rows, system has {len(rows)}"
    ok, det = verify(rows, p, tgt, c, lidx, lval)
    assert ok, f"true witness failed to verify: {det}"
    print(f"[1] verify(true witness w_202): PASS ({det}; target e_{tgt} = "
          + " + ".join(f"{v}*e_{m}" for m, v in sorted(c.items())) + f" mod {p})")

    mp_, mtgt, mc, mlidx, mlval, _ = load_witness(os.path.join(here, "w_mut.json"))
    ok_mut, det_mut = verify(rows, mp_, mtgt, mc, mlidx, mlval)
    assert not ok_mut, "MUTATION NOT CAUGHT — certificate layer broken"
    print(f"[2] verify(mutated witness w_mut): FAIL as required ({det_mut})")

    wrong = dict(c)
    k = next(iter(wrong))
    wrong[k] = int(wrong[k]) + 1
    ok_bad, det_bad = verify(rows, p, tgt, c, lidx, lval, table_row=wrong)
    assert not ok_bad, "WRONG claimed table row passed — table-mode broken"
    ok_true, _ = verify(rows, p, tgt, c, lidx, lval, table_row=c)
    assert ok_true, "true claimed table row failed"
    print(f"[3] table-mode: wrong claimed row FAILS ({det_bad}); true claimed row PASSES")

    h = glob.glob(os.path.join(here, "lapper_bench*", "harness", "run_bench.sh"))
    if h:
        import subprocess
        print(f"[4] T1 bench: running {h[0]} T1 (zero-mismatch gate; needs the winnow library + bash>=4.4) ...")
        subprocess.run(["bash", h[0], "T1"], check=False)
    else:
        print("[4] T1 bench: optional — `tar xzf lapper_bench_v0.1.tar.gz` here to enable; skipped")

    print("\nALL CHECKS PASS — witness verifies, mutation caught, wrong table row caught.")


if __name__ == "__main__":
    main()
