#!/usr/bin/env python3
r"""
check_svcert_anc.py -- stand-alone verifier for the ancillary file sv_certificates_w19-21.json.

Uses ONLY the three ancillary JSON files (laurent_C4_w7-12.json, zsv_depth3_w11-21.json,
sv_certificates_w19-21.json) and exact rational arithmetic (Python fractions); no other code or data.
Optionally compares the C_{8,1,1,1} certificate with the printed equation of the paper (its LaTeX source).

Checks
  (a1) certificate -> Laurent coefficient, by substitution: for each of the 26 deep coefficients, replacing every
       zsv(a,b,c) of the decomposition by its entry in zsv_reduced (the value written in the reduced basis of the
       Laurent file) and adding the odd-zeta monomials reproduces, monomial by monomial, the coefficient c_{3-w}
       of laurent_C4_w7-12.json; and the 'certifies' field is that coefficient verbatim.
       What 'reduced' means: both sides are compared in the basis in which the Laurent file is ALREADY written
       (monomials in zeta(2), odd zeta values and the listed depth-two/three generators); no reduction code runs here.
  (a2) formula -> reduced value, by substitution: for each of the 49 weight-19/21 formulas of zsv_depth3_w11-21.json
       (2*zeta(a,b,c) + correction), replacing every monomial that is not a basis monomial by its entry in
       'reductions' reproduces zsv_reduced exactly. The reductions (identities among multiple zeta values implied by
       the regularized double-shuffle relations in depth <= 3) are data of the certificate file: they are the one
       ingredient this script takes as given rather than re-derives. Every non-basis monomial of those formulas must
       have a reduction, and every reduction must be written in basis monomials only.
  (b)  counts: 26 certificates = 11 at transcendental weight 19 (w = 11) + 15 at 21 (w = 12); eleven terms, all zsv
       values, in each weight-19 decomposition; fourteen, eleven zsv values and three odd-zeta monomials, in each
       weight-21 decomposition; every zsv key is a weight-(2w-3) entry of zsv_depth3_w11-21.json; all rationals exact.
  (c)  the C_{8,1,1,1} decomposition equals the printed equation (gen/svcert8111.tex of the paper source) term by
       term: same zsv(a,b,c), same rational, same order, same count. Run when the .tex file is given or found.
  (d)  phi table: 32 generators = the basis generators (12 of depth two, 20 of depth three), every word of the
       weight of its generator, no pure power of f2 in an even-weight generator (the stated normalization).
  (e)  phi, sv and the values are consistent: with phi extended multiplicatively (phi(zeta(n)) = f_n, phi(zeta(2)) =
       f2, shuffle product) and sv(w) = sum_{w=uv} reverse(u) shuffle v, sv(f2) = 0 (the orientation stated in the
       file), phi(zsv_reduced(a,b,c)) = sv(phi(zeta(a,b,c))) word by word for all 49 weight-19/21 triples, zeta(a,b,c)
       being a basis generator or replaced by its reduction.
Usage
  python3 check_svcert_anc.py [--anc DIR] [--svcert-tex PATH]
  DIR defaults to the directory of this script if it holds the JSON files, else ../anc; PATH defaults to
  ../gen/svcert8111.tex relative to DIR when that file exists (check (c) is reported as not run otherwise).
Exit status 0 and last line 'ALL CHECKS PASS' iff every check that ran passed (and (c) ran when --svcert-tex is given).
"""
import os, re, sys, json, argparse
from fractions import Fraction as Fr
from collections import Counter

CERT = "sv_certificates_w19-21.json"; LAUR = "laurent_C4_w7-12.json"; ZSV = "zsv_depth3_w11-21.json"

# ---------------- monomial syntax of the ancillary files ----------------
def parse_mono(s):
    """'zeta(2)^2*zeta(3)*zeta(3,13)' -> (z2 power, sorted tuple of odd arguments, sorted tuple of index tuples)"""
    s = s.strip()
    if s in ("1", ""): return (0, (), ())
    z2 = 0; odd = []; mz = []
    for f in s.split("*"):
        m = re.fullmatch(r"zeta\(([\d,]+)\)(?:\^(\d+))?", f.strip())
        if not m: raise ValueError("bad factor %r in %r" % (f, s))
        idx = tuple(int(x) for x in m.group(1).split(",")); e = int(m.group(2) or 1)
        if len(idx) == 1:
            if idx[0] == 2: z2 += e
            else:
                if idx[0] % 2 == 0: raise ValueError("even single zeta other than zeta(2) in %r" % s)
                odd += [idx[0]] * e
        else:
            if e != 1: raise ValueError("power of a multiple zeta value in %r" % s)
            mz.append(idx)
    return (z2, tuple(sorted(odd)), tuple(sorted(mz)))

def weight(key): return 2 * key[0] + sum(key[1]) + sum(sum(t) for t in key[2])

def slot_dict(slot):
    """{'den': D, 'num': {monomial: N}} -> {key: Fraction}"""
    out = {}
    for m, n in slot["num"].items():
        k = parse_mono(m); out[k] = out.get(k, Fr(0)) + Fr(int(n), int(slot["den"]))
    return {k: v for k, v in out.items() if v != 0}

def add_into(acc, d, q):
    for k, c in d.items(): acc[k] = acc.get(k, Fr(0)) + q * c

def clean(d): return {k: v for k, v in d.items() if v != 0}

def formula_terms(entry):
    """zsv_depth3 entry -> {key: Fraction} of 2*zeta(a,b,c) + correction (unreduced)"""
    m = re.fullmatch(r"2\*zeta\((\d+),(\d+),(\d+)\)", entry["leading"]); assert m, entry["leading"]
    tri = tuple(int(x) for x in m.groups())
    d = {(0, (), (tri,)): Fr(2)}
    for mono, q in entry["correction"].items():
        k = parse_mono(mono); d[k] = d.get(k, Fr(0)) + Fr(q)
    return tri, clean(d)

# ---------------- f-alphabet arithmetic ----------------
_SH = {}
def shuffle(u, v):
    """shuffle product of two words (tuples) -> Counter{word: multiplicity}"""
    if not u: return Counter({v: 1})
    if not v: return Counter({u: 1})
    key = (u, v)
    if key in _SH: return _SH[key]
    out = Counter()
    for w, c in shuffle(u[1:], v).items(): out[(u[0],) + w] += c
    for w, c in shuffle(u, v[1:]).items(): out[(v[0],) + w] += c
    _SH[key] = out
    return out

def fmul(A, B):
    """product in U tensor Q[f2]: {(j, word): Fraction}"""
    out = {}
    for (ja, ua), qa in A.items():
        for (jb, ub), qb in B.items():
            for w, c in shuffle(ua, ub).items():
                k = (ja + jb, w); out[k] = out.get(k, Fr(0)) + qa * qb * c
    return clean(out)

def parse_word(s):
    """'f2^2 f3 f11' -> (2, (3, 11)); 'f3 f5' -> (0, (3, 5)); 'f2 f9' -> (1, (9,))"""
    toks = s.split(); j = 0
    if toks and re.fullmatch(r"f2(\^\d+)?", toks[0]):
        j = int(toks[0][3:]) if "^" in toks[0] else 1; toks = toks[1:]
    u = []
    for t in toks:
        m = re.fullmatch(r"f(\d+)", t)
        if not m or int(m.group(1)) % 2 == 0 or int(m.group(1)) < 3: raise ValueError("bad letter %r in %r" % (t, s))
        u.append(int(m.group(1)))
    return (j, tuple(u))

def sv(A):
    """single-valued map in the file's word orientation: sv(f2) = 0; sv(w) = sum_{w = uv} reverse(u) shuffle v"""
    out = {}
    for (j, u), q in A.items():
        if j: continue
        for k in range(len(u) + 1):
            for w, c in shuffle(tuple(reversed(u[:k])), u[k:]).items():
                out[(0, w)] = out.get((0, w), Fr(0)) + q * c
    return clean(out)

class Phi:
    def __init__(self, table):
        self.gen = {}
        for name, e in table.items():
            k = parse_mono(name); assert k[0] == 0 and not k[1] and len(k[2]) == 1, name
            self.gen[k[2][0]] = {parse_word(w): Fr(q) for w, q in e["image"].items()}
    def of_key(self, key):
        z2, odd, mz = key
        v = {(z2, ()): Fr(1)}
        for n in odd: v = fmul(v, {(0, (n,)): Fr(1)})
        for t in mz:
            if t not in self.gen: raise KeyError("no phi image for zeta%s (not a basis generator)" % (t,))
            v = fmul(v, self.gen[t])
        return v
    def of_dict(self, d):
        out = {}
        for k, q in d.items(): add_into(out, self.of_key(k), q)
        return clean(out)

# ---------------- the printed equation ----------------
TERM_RE = re.compile(r"([+-]?)\\tfrac\{(\d+)\}\{(\d+)\}\\,\\zsv\((\d+),(\d+),(\d+)\)")
def parse_svcert_tex(tex):
    body = "\n".join(l for l in tex.split("\n") if not l.lstrip().startswith("%"))
    assert "\\label{svcert8111}" in body and "c_{-8}(C_{8,1,1,1})" in body, "not the svcert8111 display"
    body = body.replace("&{}", "").replace("&", "").replace(" ", "").replace("\n", "")
    out = []
    for s, num, den, a, b, c in TERM_RE.findall(body):
        fr = Fr(int(num), int(den)); out.append(((int(a), int(b), int(c)), -fr if s == "-" else fr))
    return out

# ---------------- the checks ----------------
def run_checks(L, Z, C, svcert_tex=None, require_c=False):
    """L, Z, C: parsed laurent / zsv / certificate JSON objects; svcert_tex: text of gen/svcert8111.tex or None.
    Returns (ok, report_lines)."""
    rep = []; ok = True
    gens2 = [parse_mono(g)[2][0] for g in C["basis_generators"]["depth2"]]
    gens3 = [parse_mono(g)[2][0] for g in C["basis_generators"]["depth3"]]
    basis = set(gens2) | set(gens3)
    is_basis = lambda key: all(t in basis for t in key[2])
    zred = {k: slot_dict(e["value"]) for k, e in C["zsv_reduced"].items()}
    red = {parse_mono(m): slot_dict(e["value"]) for m, e in C["reductions"].items()}
    certs = C["certificates"]
    # (b) counts
    n19 = [k for k, e in certs.items() if e["transcendental_weight"] == 19]; n21 = [k for k, e in certs.items() if e["transcendental_weight"] == 21]
    comp = {19: set(), 21: set()}; bad = []
    for fk, e in certs.items():
        a4 = [int(x) for x in fk.split(",")]; w = sum(a4); tw = 2 * w - 3
        if not (e["weight"] == w and e["k"] == 3 - w and e["transcendental_weight"] == tw and len(a4) == 4): bad.append((fk, "header"))
        nz = nodd = 0
        for m, q in e["decomposition"].items():
            fr = Fr(q)
            if not re.fullmatch(r"-?\d+(/\d+)?", q) or (("%d/%d" % (fr.numerator, fr.denominator)) if fr.denominator != 1 else "%d" % fr.numerator) != q: bad.append((fk, "rational " + q))
            mm = re.fullmatch(r"zsv\((\d+),(\d+),(\d+)\)", m)
            if mm:
                t = tuple(int(x) for x in mm.groups()); zk = "%d,%d,%d" % t
                if zk not in Z["formulas"] or int(Z["formulas"][zk]["weight"]) != tw or zk not in C["zsv_reduced"] or sum(t) != tw or any(x % 2 == 0 or x < 3 for x in t): bad.append((fk, m))
                nz += 1
            else:
                key = parse_mono(m)
                if key[0] != 0 or key[2] or weight(key) != tw: bad.append((fk, m))
                nodd += 1
        if (nz + nodd, nz, nodd) != (e["n_terms"], e["n_zsv_terms"], e["n_odd_zeta_monomial_terms"]): bad.append((fk, "term counts"))
        comp[tw].add((nz, nodd))
    g = (len(certs) == 26 and len(n19) == 11 and len(n21) == 15 and comp[19] == {(11, 0)} and comp[21] == {(11, 3)} and not bad
         and all(sum(int(x) for x in k.split(",")) == 11 for k in n19) and all(sum(int(x) for x in k.split(",")) == 12 for k in n21))
    ok &= g
    rep.append("(b)  counts: %d certificates = %d (tw 19, w=11) + %d (tw 21, w=12); terms per decomposition tw19 %s, tw21 %s as (zsv, odd monomials); keys and rationals well-formed: %s%s"
               % (len(certs), len(n19), len(n21), sorted(comp[19]), sorted(comp[21]), "PASS" if g else "FAIL", "" if g else " %s" % bad[:6]))
    # (a1) substitution certificate -> Laurent slot
    nbad = 0; nverb = 0; msgs = []
    for fk, e in certs.items():
        slot = L["functions"][fk]["coefficients"][str(e["k"])]
        if slot == e["certifies"]: nverb += 1
        else: msgs.append("%s: 'certifies' != Laurent file slot" % fk)
        acc = {}
        for m, q in e["decomposition"].items():
            mm = re.fullmatch(r"zsv\((\d+),(\d+),(\d+)\)", m)
            if mm: add_into(acc, zred["%s,%s,%s" % mm.groups()], Fr(q))
            else: add_into(acc, {parse_mono(m): Fr(1)}, Fr(q))
        if clean(acc) != slot_dict(slot): nbad += 1; msgs.append("%s: substitution != c_{%d}" % (fk, e["k"]))
    g = nbad == 0 and nverb == len(certs) == 26; ok &= g
    rep.append("(a1) substitution zsv -> zsv_reduced in each decomposition reproduces c_{3-w} of %s monomial by monomial: %d/%d; 'certifies' = that slot verbatim: %d/%d: %s%s"
               % (LAUR, len(certs) - nbad, len(certs), nverb, len(certs), "PASS" if g else "FAIL", "" if g else " " + "; ".join(msgs[:4])))
    # (a2) substitution formula -> reduced value; coverage of the reductions
    hi = {k: e for k, e in Z["formulas"].items() if int(e["weight"]) in (19, 21)}
    nb = 0; need = set(); msgs = []
    for k, e in hi.items():
        tri, terms = formula_terms(e); acc = {}
        for key, q in terms.items():
            if is_basis(key): add_into(acc, {key: Fr(1)}, q)
            elif key in red: add_into(acc, red[key], q); need.add(key)
            else: msgs.append("%s: monomial %r has no reduction" % (k, key)); nb += 1
        if k not in zred: msgs.append("%s: no zsv_reduced entry" % k); nb += 1; continue
        if clean(acc) != zred[k]: nb += 1; msgs.append("%s: substituted formula != zsv_reduced" % k)
        if int(C["zsv_reduced"][k]["weight"]) != sum(tri) or C["zsv_reduced"][k]["label"] != e["label"]: nb += 1; msgs.append("%s: weight/label" % k)
    red_in_basis = all(all(is_basis(x) for x in d) and all(weight(x) == weight(kk) for x in d) for kk, d in red.items())
    zred_in_basis = all(all(is_basis(x) for x in d) for d in zred.values())
    unused = set(red) - need
    g = nb == 0 and len(hi) == 49 and len(zred) == 49 and red_in_basis and zred_in_basis and not unused; ok &= g
    rep.append("(a2) substitution of the %d reductions into the %d weight-19/21 formulas of %s reproduces zsv_reduced exactly: %d/%d; reductions and reduced values written in basis monomials only: %s; every reduction used: %s: %s%s"
               % (len(red), len(hi), ZSV, len(hi) - nb, len(hi), red_in_basis and zred_in_basis, not unused, "PASS" if g else "FAIL", "" if g else " " + "; ".join(msgs[:4])))
    # (c) printed equation
    if svcert_tex is not None:
        printed = parse_svcert_tex(svcert_tex)
        mine = [((int(a), int(b), int(c)), Fr(q)) for m, q in certs["8,1,1,1"]["decomposition"].items() for a, b, c in [re.fullmatch(r"zsv\((\d+),(\d+),(\d+)\)", m).groups()]]
        g = printed == mine and len(printed) == 11; ok &= g
        rep.append("(c)  C_{8,1,1,1}: the %d terms of the printed equation (svcert8111) = the certificate term by term (triple, rational, order): %s" % (len(printed), "PASS" if g else "FAIL"))
    else:
        rep.append("(c)  C_{8,1,1,1} vs the printed equation: NOT RUN (paper source gen/svcert8111.tex not given)")
        if require_c: ok = False
    # (d) phi table structure
    P = C["phi"]; msgs = []
    names2 = set(C["basis_generators"]["depth2"]); names3 = set(C["basis_generators"]["depth3"])
    if set(P) != names2 | names3: msgs.append("generator set != basis_generators")
    for name, e in P.items():
        t = parse_mono(name)[2][0]; w = sum(t)
        if e["weight"] != w or e["depth"] != len(t): msgs.append("%s header" % name)
        for wd, q in e["image"].items():
            j, u = parse_word(wd); Fr(q)
            if 2 * j + sum(u) != w: msgs.append("%s: word %r has the wrong weight" % (name, wd))
            if len(t) == 2 and not u: msgs.append("%s: pure f2 power present" % name)
    g = not msgs and len(names2) == 12 and len(names3) == 20 and len(P) == 32; ok &= g
    rep.append("(d)  phi table: %d generators (%d depth two, %d depth three), words homogeneous of the generator's weight, no pure f2-power word in an even-weight generator: %s%s"
               % (len(P), len(names2), len(names3), "PASS" if g else "FAIL", "" if g else " " + "; ".join(msgs[:4])))
    # (e) phi(zsv_reduced) == sv(phi(zeta(a,b,c)))
    phi = Phi(P); ne = 0; msgs = []
    for k, d in zred.items():
        tri = tuple(int(x) for x in k.split(","))
        lhs = phi.of_dict(d)
        src = {(0, (), (tri,)): Fr(1)} if tri in basis else red[(0, (), (tri,))]
        rhs = sv(phi.of_dict(src))
        if lhs == rhs and lhs: ne += 1
        else: msgs.append(k)
    g = ne == len(zred) == 49; ok &= g
    rep.append("(e)  phi(zsv_reduced(a,b,c)) = sv(phi(zeta(a,b,c))) word by word, f2-free, for %d/%d weight-19/21 triples: %s%s"
               % (ne, len(zred), "PASS" if g else "FAIL", "" if g else " " + ", ".join(msgs[:6])))
    return ok, rep

def main():
    here = os.path.dirname(os.path.abspath(__file__))
    ap = argparse.ArgumentParser(description="verify sv_certificates_w19-21.json against the two ancillary data files")
    ap.add_argument("--anc", default=None, help="directory holding the three JSON files")
    ap.add_argument("--svcert-tex", default=None, help="path to gen/svcert8111.tex of the paper source (check (c))")
    a = ap.parse_args()
    anc = a.anc or (here if os.path.exists(os.path.join(here, CERT)) else os.path.normpath(os.path.join(here, "..", "anc")))
    tex = a.svcert_tex
    if tex is None:
        cand = os.path.normpath(os.path.join(anc, "..", "gen", "svcert8111.tex"))
        if os.path.exists(cand): tex = cand
    L = json.load(open(os.path.join(anc, LAUR), encoding="utf-8")); Z = json.load(open(os.path.join(anc, ZSV), encoding="utf-8"))
    C = json.load(open(os.path.join(anc, CERT), encoding="utf-8"))
    text = open(tex, encoding="utf-8").read() if tex else None
    print("check_svcert_anc.py on %s (%s generated %s; %s %s; %s %s)%s" % (anc, CERT, C.get("generated"), LAUR, L.get("generated"), ZSV, Z.get("generated"), ("; (c) against %s" % tex) if tex else ""))
    ok, rep = run_checks(L, Z, C, text, require_c=a.svcert_tex is not None)
    print("\n".join(rep))
    print("ALL CHECKS PASS" if ok else "CHECK FAILED")
    sys.exit(0 if ok else 1)

if __name__ == "__main__":
    main()
