#!/usr/bin/env python3
"""cegm-verify.py -- standalone verifier for the CEGM X(3,6) companion bundle.

Run from the bundle root (the directory containing this file):

    python3 cegm-verify.py             # full verification, 80 working digits
    python3 cegm-verify.py --selftest  # in-memory mutation control
    python3 cegm-verify.py --dps 120   # crank working precision

Requires: Python 3 stdlib + mpmath. No other imports, no network, no
absolute paths; all files are read relative to this script's directory.

What it checks (sources: the paper "Grassmannian string integrals", Section 3
and Appendices A-B, and the ancillary artifacts; nothing is asserted beyond
what those documents state):
  1. SHA-256 pins of all 8 ancillary files, re-checked on every load. The
     five JSON files are pinned byte-exact to the files of record; the three
     markdown records are pinned as the public renders shipped here.
  2. Every exact fraction parses, is in lowest terms (from the raw string),
     with positive denominator.
  3. Cross-artifact identities documented in the sources:
     - t2_GATES.json w4.Q4 == w4.coords.z4 (same quantity, same file);
     - w5_collapse_verdict.json Q5,R5 == the fractions displayed in
       ancillary/THEOREM_W5.md Sec 4 ("byte-identical"), as Fractions;
     - w5_collapse_verdict.json R5_zero consistent with R5.
  4. Recomputed closed forms (mpmath zetas under mp.workdps; every
     reference value is read from the pinned files, none hard-coded here):
     - c5(s*) = Q5*zeta(5) + R5*zeta(2)*zeta(3) vs c5_exact_numeric
       [paper Sec. 3.3, Eq. (15); w5_collapse_verdict.json]
     - c3(s*) = Q3*zeta(3) vs c3           [paper Sec. 3.2 table; t2_c3_final.json]
     - c4(s*) = Q4*zeta(4) vs w4.c4        [paper Sec. 3.2 table; t2_GATES.json w4]
     - second point (t2_exact_s2.json): c0_float = c0,
       c2_float = R2*zeta(2), c3_float = Q3*zeta(3)
       [paper Sec. 3.2, "The same statements hold at the second kinematic point"]
  Fields with no documented decomposition (R3b, R3d, E_table entries) get
  the fraction-hygiene checks only; this is stated in the output.

Exit codes: 0 full PASS; 2 pin mismatch; 3 value mismatch; 4 missing file.
Mutation selftest (--selftest) perturbs copies IN MEMORY ONLY; the shipped
files are never modified.
"""

import argparse
import copy
import hashlib
import json
import math
import os
import re
import sys
from fractions import Fraction

from mpmath import mp, mpf

# SHA-256 pins of the ancillary files (integrity pins, not result values;
# every result value used in a pass criterion is read from these files).
PINS = {
    "ancillary/APPENDIX_ARTIFACT.md":
        "46a4bb4a7f71003711a56dd75021d7a9cef245f6c37a3e431c776dd240c05aa8",
    "ancillary/PROOF_W4W3.md":
        "7244e5f67a8da74dfbe002756ccceaeb2904d47264672671f8c8e79435e6b144",
    "ancillary/THEOREM_W5.md":
        "4dfc15db7b53b3987affc234c68427029876f4bafbad25bea9a8c504ccd48e4c",
    "ancillary/t2_GATES.json":
        "5a76cb4a8a4c936a5a72053d80bd4d6d7814c92f7ab23152871bb3be3587fdba",
    "ancillary/t2_c3_final.json":
        "b816e109cafa807a96885fed88fe007cb1888c2b2ffef3666be77d1f7e451bbb",
    "ancillary/t2_exact_rationals.json":
        "591917ab4ae302223629180b74954459952447604dfbcb5a18dc5cced051884b",
    "ancillary/t2_exact_s2.json":
        "6281d01fbabb07df7e899ed3ab6720208ba1dda5d9a1834f418db6b894dca264",
    "ancillary/w5_collapse_verdict.json":
        "7711d574ee5ffcad4052dbcab183b096f13b207f9a771f1b848cbe9d883484d4",
}


class VerifyError(Exception):
    rc = 1


class PinMismatchError(VerifyError):
    rc = 2


class ValueMismatchError(VerifyError):
    rc = 3

    def __init__(self, msg, achieved=None):
        super().__init__(msg)
        self.achieved = achieved


class MissingFileError(VerifyError):
    rc = 4


BASE = os.path.dirname(os.path.abspath(__file__))


def load_bytes(rel):
    """Read a pinned file; the pin is checked on EVERY load (no TOFU)."""
    path = os.path.join(BASE, rel)
    if not os.path.isfile(path):
        raise MissingFileError("missing file: %s" % rel)
    with open(path, "rb") as fh:
        data = fh.read()
    got = hashlib.sha256(data).hexdigest()
    want = PINS[rel]
    if got != want:
        raise PinMismatchError(
            "pin mismatch: %s sha256=%s... pinned=%s..." % (rel, got[:16], want[:16]))
    return data


def load_all():
    """Load and pin-check all 8 ancillary files; parse the JSONs."""
    for rel in sorted(PINS):
        load_bytes(rel)  # pin-check even the files we do not parse
    return {
        "gates": json.loads(load_bytes("ancillary/t2_GATES.json").decode()),
        "c3": json.loads(load_bytes("ancillary/t2_c3_final.json").decode()),
        "rat": json.loads(load_bytes("ancillary/t2_exact_rationals.json").decode()),
        "s2": json.loads(load_bytes("ancillary/t2_exact_s2.json").decode()),
        "verdict": json.loads(load_bytes("ancillary/w5_collapse_verdict.json").decode()),
        "theorem_w5": load_bytes("ancillary/THEOREM_W5.md").decode(),
    }


def check_fraction(s, where):
    """Parse an exact rational string; enforce lowest terms on the RAW string."""
    if not re.fullmatch(r"-?\d+(/\d+)?", s):
        raise ValueMismatchError("%s: not an exact rational string" % where)
    if "/" in s:
        num_s, den_s = s.split("/")
        num, den = int(num_s), int(den_s)
        if den <= 0:
            raise ValueMismatchError("%s: nonpositive denominator" % where)
        if math.gcd(abs(num), den) != 1:
            raise ValueMismatchError("%s: not in lowest terms" % where)
        return Fraction(num, den)
    return Fraction(int(s))


def sig_digits(s):
    """Number of significant digits printed in a decimal string."""
    mant = s.strip().lstrip("+-").split("e")[0].split("E")[0]
    return len(mant.replace(".", "").lstrip("0"))


def fr_to_mpf(fr):
    return mpf(fr.numerator) / mpf(fr.denominator)


def check_value(name, combo, recorded_str, dps, source, quiet=False):
    """Recompute combo() and compare against the recorded decimal string.

    Bar: |diff| <= 1 ulp of the recorded print at min(printed-1, dps-5)
    significant digits (tolerates a rounded-vs-truncated last printed
    digit, and caps at working precision). Returns floored verified-digit
    count, capped at both the printed digit count and dps.
    """
    printed = sig_digits(recorded_str)
    with mp.workdps(dps + 35):
        computed = combo()
        recorded = mpf(recorded_str)
        diff = abs(computed - recorded)
        lead = int(mp.floor(mp.log10(abs(recorded))))
        d_eff = min(printed - 1, dps - 5)
        tol = mpf(10) ** (lead - d_eff + 1)
        if diff == 0:
            achieved = min(printed, dps)
        else:
            achieved = int(mp.floor(-mp.log10(diff / abs(recorded))))
            achieved = min(achieved, printed, dps)
        if diff > tol:
            raise ValueMismatchError(
                "%s: recomputed %s vs recorded %s -- agreement %d digits, "
                "below the %d-digit bar" % (
                    name, mp.nstr(computed, 10), mp.nstr(recorded, 10),
                    max(achieved, 0), d_eff),
                achieved=max(achieved, 0))
        if not quiet:
            print("[PASS] %s = %s : verified to %d digits (recorded prints %d; "
                  "|diff| %s)  [%s]" % (
                      name, mp.nstr(computed, 10), achieved, printed,
                      mp.nstr(diff, 2), source))
    return achieved


def extract_theorem_w5_fractions(md_text):
    """Pull the Q5, R5 fractions displayed in THEOREM_W5.md Sec 4."""
    m = re.search(
        r"Q5\s*=\s*([0-9\s]+)/\s*([0-9\s]+)\n\s*R5\s*=\s*(-?[0-9\s]+)/\s*([0-9\s]+)",
        md_text)
    if not m:
        raise ValueMismatchError("THEOREM_W5.md: Q5/R5 display not found")
    strip = lambda g: int(re.sub(r"\s", "", g))
    return (Fraction(strip(m.group(1)), strip(m.group(2))),
            Fraction(strip(m.group(3)), strip(m.group(4))))


def suite(data, dps, quiet=False):
    """Run every check. Raises on the first failure. Returns digit summary."""
    say = (lambda *a: None) if quiet else print

    # --- fraction hygiene (parse + lowest terms + positive denominator) ---
    fracs = {}
    registry = [
        ("gates", ("w0", "c0"), "t2_GATES.json w0.c0"),
        ("gates", ("w4", "Q4"), "t2_GATES.json w4.Q4"),
        ("gates", ("w4", "coords", "z4"), "t2_GATES.json w4.coords.z4"),
        ("c3", ("Q3",), "t2_c3_final.json Q3"),
        ("rat", ("R2",), "t2_exact_rationals.json R2"),
        ("rat", ("R3b",), "t2_exact_rationals.json R3b"),
        ("rat", ("R3d",), "t2_exact_rationals.json R3d"),
        ("s2", ("c0",), "t2_exact_s2.json c0"),
        ("s2", ("R2",), "t2_exact_s2.json R2"),
        ("s2", ("Q3",), "t2_exact_s2.json Q3"),
        ("verdict", ("Q5",), "w5_collapse_verdict.json Q5"),
        ("verdict", ("R5",), "w5_collapse_verdict.json R5"),
    ]
    for src, path, label in registry:
        node = data[src]
        for k in path:
            node = node[k]
        fracs[label] = check_fraction(node, label)
    n_etab = 0
    for key, val in data["s2"]["E_table"].items():
        check_fraction(val, "t2_exact_s2.json E_table[%s]" % key)
        n_etab += 1
    say("[PASS] fraction hygiene: %d named fractions + %d E_table entries "
        "parse, lowest terms, positive denominators" % (len(registry), n_etab))
    say("       (R3b, R3d and the E_table entries have no decomposition "
        "documented in the paper; hygiene checks only)")

    # --- cross-artifact identities (documented) ---
    if fracs["t2_GATES.json w4.Q4"] != fracs["t2_GATES.json w4.coords.z4"]:
        raise ValueMismatchError("t2_GATES.json: w4.Q4 != w4.coords.z4")
    say("[PASS] t2_GATES.json w4.Q4 == w4.coords.z4 (exact)")

    q5_md, r5_md = extract_theorem_w5_fractions(data["theorem_w5"])
    if fracs["w5_collapse_verdict.json Q5"] != q5_md:
        raise ValueMismatchError(
            "Q5 in w5_collapse_verdict.json != Q5 displayed in THEOREM_W5.md")
    if fracs["w5_collapse_verdict.json R5"] != r5_md:
        raise ValueMismatchError(
            "R5 in w5_collapse_verdict.json != R5 displayed in THEOREM_W5.md")
    say("[PASS] Q5, R5 identical between w5_collapse_verdict.json and "
        "THEOREM_W5.md Sec 4 (exact Fractions; the documented "
        "two-derivation byte-identity)")

    r5 = fracs["w5_collapse_verdict.json R5"]
    if (r5 == 0) != bool(data["verdict"]["R5_zero"]):
        raise ValueMismatchError("R5_zero flag inconsistent with R5 value")
    say("[PASS] R5 != 0, consistent with the recorded R5_zero flag")

    # --- recomputed closed forms (references all read from pinned files) ---
    digits = {}
    q5 = fracs["w5_collapse_verdict.json Q5"]
    digits["c5"] = check_value(
        "c5(s*) = Q5*zeta(5) + R5*zeta(2)*zeta(3)",
        lambda: fr_to_mpf(q5) * mp.zeta(5) + fr_to_mpf(r5) * mp.zeta(2) * mp.zeta(3),
        data["verdict"]["c5_exact_numeric"], dps,
        "paper Sec. 3.3 Eq. (15); w5_collapse_verdict.json", quiet)

    q3 = fracs["t2_c3_final.json Q3"]
    digits["c3"] = check_value(
        "c3(s*) = Q3*zeta(3)",
        lambda: fr_to_mpf(q3) * mp.zeta(3),
        data["c3"]["c3"], dps,
        "paper Sec. 3.2 table; t2_c3_final.json", quiet)

    q4 = fracs["t2_GATES.json w4.Q4"]
    digits["c4"] = check_value(
        "c4(s*) = Q4*zeta(4)",
        lambda: fr_to_mpf(q4) * mp.zeta(4),
        data["gates"]["w4"]["c4"], dps,
        "paper Sec. 3.2 table; t2_GATES.json w4 closed_form", quiet)

    c0_s2 = fracs["t2_exact_s2.json c0"]
    digits["s2_c0"] = check_value(
        "second point: c0 (exact rational)",
        lambda: fr_to_mpf(c0_s2),
        data["s2"]["c0_float"], dps,
        "paper Sec. 3.2 second-point paragraph; t2_exact_s2.json", quiet)

    r2_s2 = fracs["t2_exact_s2.json R2"]
    digits["s2_c2"] = check_value(
        "second point: c2 = R2*zeta(2)",
        lambda: fr_to_mpf(r2_s2) * mp.zeta(2),
        data["s2"]["c2_float"], dps,
        "paper Sec. 3.2; t2_GATES.json w2 closed_form; t2_exact_s2.json", quiet)

    q3_s2 = fracs["t2_exact_s2.json Q3"]
    digits["s2_c3"] = check_value(
        "second point: c3 = Q3*zeta(3)",
        lambda: fr_to_mpf(q3_s2) * mp.zeta(3),
        data["s2"]["c3_float"], dps,
        "paper Sec. 3.2; t2_GATES.json w3 closed_form; t2_exact_s2.json", quiet)

    return digits


def mutate_last_digit(frac_str):
    """Return frac_str with the numerator's last digit changed (mod 10)."""
    num_s, den_s = frac_str.split("/")
    new_last = str((int(num_s[-1]) + 1) % 10)
    return num_s[:-1] + new_last + "/" + den_s


def selftest(dps):
    print("== selftest: in-memory mutation control (shipped files untouched) ==")
    print("-- baseline: full suite must PASS --")
    suite(load_all(), dps, quiet=True)
    print("baseline: FULL PASS")

    # Mutation A: Q5 numerator, last digit.
    data_a = load_all()
    orig = data_a["verdict"]["Q5"]
    data_a["verdict"]["Q5"] = mutate_last_digit(orig)
    print("-- mutation A: Q5 numerator last digit %s -> %s (in memory) --" % (
        orig.split("/")[0][-1], data_a["verdict"]["Q5"].split("/")[0][-1]))
    q5_mut = Fraction(data_a["verdict"]["Q5"])
    r5 = Fraction(data_a["verdict"]["R5"])
    numeric_tripped = True
    try:
        check_value(
            "c5 (mutated Q5, numeric check alone)",
            lambda: fr_to_mpf(q5_mut) * mp.zeta(5)
            + fr_to_mpf(r5) * mp.zeta(2) * mp.zeta(3),
            data_a["verdict"]["c5_exact_numeric"], dps,
            "selftest", quiet=True)
        numeric_tripped = False
    except ValueMismatchError:
        pass
    if numeric_tripped:
        print("   numeric c5 check alone: tripped")
    else:
        print("   numeric c5 check alone: NOT tripped -- a last-digit "
              "perturbation of the 63-digit numerator shifts c5 by about "
              "1 part in 10^62, below the 50-digit recorded print; only an "
              "exact check can see it (reported honestly)")
    try:
        suite(data_a, dps, quiet=True)
        raise VerifyError("selftest FAILED: mutation A was NOT detected")
    except ValueMismatchError as exc:
        depth = ("exact (rational-arithmetic) level -- the suite's exact "
                 "checks (lowest-terms, THEOREM_W5.md byte-identity) see "
                 "what the 50-digit recorded print cannot"
                 if not numeric_tripped else "%s digits" % exc.achieved)
        print("   DETECTED: %s" % exc)
        print("   detection depth: %s" % depth)

    # Mutation B: a t2 fraction -- Q3 of t2_c3_final.json, numerator last digit.
    data_b = load_all()
    orig_b = data_b["c3"]["Q3"]
    data_b["c3"]["Q3"] = mutate_last_digit(orig_b)
    print("-- mutation B: t2_c3_final.json Q3 numerator last digit %s -> %s "
          "(in memory) --" % (orig_b.split("/")[0][-1],
                              data_b["c3"]["Q3"].split("/")[0][-1]))
    q3_mut = Fraction(data_b["c3"]["Q3"])
    try:
        check_value(
            "c3 (mutated Q3, value check)",
            lambda: fr_to_mpf(q3_mut) * mp.zeta(3),
            data_b["c3"]["c3"], dps, "selftest", quiet=True)
        raise VerifyError("selftest FAILED: mutation B value check PASSED")
    except ValueMismatchError as exc:
        print("   DETECTED by the c3 value check: %s" % exc)
        print("   detection depth: agreement fell to %d digits" % exc.achieved)
    try:
        suite(data_b, dps, quiet=True)
        raise VerifyError("selftest FAILED: mutation B not detected by suite")
    except ValueMismatchError:
        pass

    print("selftest PASS: both in-memory mutations detected; "
          "shipped files were never modified")
    return 0


def main():
    ap = argparse.ArgumentParser(description="CEGM X(3,6) bundle verifier")
    ap.add_argument("--dps", type=int, default=80,
                    help="working decimal precision (default 80)")
    ap.add_argument("--selftest", action="store_true",
                    help="in-memory mutation control")
    args = ap.parse_args()

    try:
        if args.selftest:
            return selftest(args.dps)
        print("== CEGM X(3,6) bundle verification (dps %d) ==" % args.dps)
        data = load_all()
        print("[PASS] pins: all %d ancillary files match their SHA-256 pins "
              "(re-checked on every load)" % len(PINS))
        digits = suite(data, args.dps)
        print("== FULL PASS ==")
        print("headline: the recomputed c5(s*) = Q5*zeta(5) + "
              "R5*zeta(2)*zeta(3) is verified to %d digits against the "
              "recorded value" % digits["c5"])
        return 0
    except VerifyError as exc:
        print("REFUSAL (%s, rc=%d): %s" % (type(exc).__name__,
                                           getattr(exc, "rc", 1), exc))
        return getattr(exc, "rc", 1)


if __name__ == "__main__":
    sys.exit(main())
