#!/usr/bin/env python3
"""lbl3x-k3letters-census.py -- the all-orders graded census of the rotated K3 corner of the row-21 (LBL3X, crossed
light-by-light box) connection: the enumeration arm beside lbl3x-evaluate.py.

WHAT IT ENUMERATES.  The K3 corner {13, 14, 15} of the on-curve connection, rotated to the intrinsic frame by the exact
eps-rational intertwiner and written in the Sym^2 S.D gauge over Q(s)[y]{w^pm, w'} (y^2 = q4*q16), is generated by an
operator with four coefficients (a2, a1, a0, b) exact in (eps, s) -- vendor_row21_k3letters/rot_dial1__LEPS_OPERATOR.json.
This script expands the operator in eps to order K (--korder), forms at each order k the rotated block
G_k = S^-1 . Theta_c^-1 . B_k . Theta_c . S, and records every non-zero slot G_k[i,j] with its grade k + j - i
(the slot census); it then ENUMERATES the distinct integrands (the letters) as the dimension over Q of the span of the
non-zero slot entries, computed exactly (fmpq_mat rank) from the coefficient functions sampled at --npts rational
s-points, at two disjoint sample sets (both must agree), with the cumulative rank by order (the saturation order), a
planted spurious letter that must be REFUSED by name (rank + 1 at both sets; --plant w1 or w0) and a duplicate member
that must be absorbed (rank unchanged).  The eps-structure of the four coefficients is read off exactly (numerator and
denominator degrees in eps, whether the eps-dependence of the denominator separates from s): a polynomial operator
terminates the expansion at its top degree, so a census to that order is the all-orders census (the run of record:
degrees 1, 2, 3, 3, eps-denominator degree 0 -- the census terminates at eps^3; orders 4-6 zero slots).

WHAT IT RE-ASSERTS ON EVERY RUN (the identities of the record, each one a hard assertion): the Sym^2 certificate of the
pulled-back Picard-Fuchs companion, the rho ODE, S^-1 S = I, the gauged connection M against its target, the
eps-series of the operator against the vendored IIB_CORE.json at orders 0..3, and the intertwining identity of
Theta_c = U_L . Theta.

VENDORED DATA (vendor_row21_k3letters/ beside this script; every file sha256-pinned in PINS, its own
RECORD_COPIES.sha256 carrying the same rows): the ring machinery rot_probe_i__ratlib.py (exact univariate rational
functions and the quadratic extension), rot_dial1__r2lib.py (bivariate exact rational functions in (eps, s)),
rot_probe_iib__iib_lib.py (the Sym^2 S-gauge), and the objects rot_dial1__LEPS_OPERATOR.json (the operator),
rot_probe_iib__IIB_CORE.json (the eps-series of record at orders 0..3), rot_probe_i__block1315_raw.json (the raw K3
corner) and rot_probe_i__THETA.json (the intertwiner).  The modules are loaded from those byte copies by file (no
bytecode written); a pin mismatch exits 3, a missing file exits 4.

HOW TO RUN (from this directory; python-flint, sympy and the standard library):
  python3 lbl3x-k3letters-census.py --korder 3 --out CENSUS_K3.json                       # the run of record (about 1 min)
  python3 lbl3x-k3letters-census.py --korder 6 --seed 2 --npts 48 --plant w0 --out CENSUS_K6.json   # the second grading (about 3 min)
  python3 lbl3x-k3letters-census.py --korder 1 --npts 8 --out CENSUS_pilot.json           # the smallest form (seconds)
Options: --korder K (the eps order, required), --out FILE (never overwritten; required), --npts N (rational s-points
per sample set, default 40), --seed S (the sample generator's seed, default 1), --plant w1|w0 (the planted spurious
letter's key: a rational multiple of w' or of w, default w1).

WHAT IT PRINTS: one timed line per stage (the frame identities, the eps-structure of the operator, the per-order slot
counts, the census totals) and the rank line
  ranks: all R1/R2 pos P1/P2 sat_order K planted_refused True dup_absorbed True
(R1/R2 = the letter rank over Q at the two sample sets, P1/P2 the same over the grade >= 1 slots), then WROTE <out>.
The run of record (--korder 3): 27 non-zero slots (9 per order), 4 at grade <= 0, ranks 25/25 and 21/21, saturation
order 3, the planted letter REFUSED, the duplicate absorbed; the JSON carries the slot census, the grades, the ranks,
the cumulative rank by order, the planted / duplicate controls, every slot integrand as an exact string, the identities
and the eps-structure, the walls per order and the stamp (date -u).
"""
import json, time, sys, os, argparse, random, itertools, subprocess, hashlib
from fractions import Fraction

T0 = time.time()
def TT(m): print('%-60s %8.1fs' % (m, time.time() - T0), flush=True)
def utc(): return subprocess.run(['date', '-u', '+%Y-%m-%dT%H:%M:%SZ'], capture_output=True, text=True).stdout.strip()

HERE = os.path.dirname(os.path.abspath(__file__))
VENDOR_DIR = 'vendor_row21_k3letters'
RC = os.path.join(HERE, VENDOR_DIR)
# sha256 of every vendored file this script reads (the rows of vendor_row21_k3letters/RECORD_COPIES.sha256; exit 3 on a mismatch, exit 4 when missing)
PINS = {
    'rot_probe_i__ratlib.py': 'c28337116aa90f9ed39f3e2f3167fee8f007f216a3d1cfeb6a8521db7428b612',
    'rot_dial1__r2lib.py': '24ac4536265258c1e04890ecb340b7d386ca85db78f33b88ddb5112615617b56',
    'rot_probe_iib__iib_lib.py': '7468bcfe08851a6b750771382a87e0f73c4d66fcb5906b06338300674d3ef97d',
    'rot_dial1__LEPS_OPERATOR.json': '67ff5e6d38a1675518537a2243cd92d4d24f9f9ae512c6f5a32db38a59808f7e',
    'rot_probe_iib__IIB_CORE.json': '008ff805f85df658fcf1cc34222becab7a471c174a3ea0a3af6525453ad729a7',
    'rot_probe_i__block1315_raw.json': '908a9f3b144e6054d252d22c084f25fa61019fa340fc7e1b36ad84a1c7cafb5c',
    'rot_probe_i__THETA.json': 'baf13ffc170f3361f82d6a5d9e3a9e74dd2376c3a012263a1ab7c6d789df956a',
}
def check_pins():
    for rel in sorted(PINS):
        p = os.path.join(RC, rel)
        if not os.path.exists(p):
            print('MISSING data file %s/%s (expected beside the script, sha256 %s...): exit 4' % (VENDOR_DIR, rel, PINS[rel][:16]))
            sys.exit(4)
        got = hashlib.sha256(open(p, 'rb').read()).hexdigest()
        if got != PINS[rel]:
            print('PIN MISMATCH %s/%s: sha256 %s... != pinned %s...: exit 3' % (VENDOR_DIR, rel, got[:16], PINS[rel][:16]))
            sys.exit(3)
    print('[pins] %s/: %d files sha256 OK' % (VENDOR_DIR, len(PINS)), flush=True)

# the record's own ring machinery, imported from the byte copies (module names as the record expects)
import importlib.util
def load(name, fn):
    spec = importlib.util.spec_from_file_location(name, os.path.join(RC, fn))
    m = importlib.util.module_from_spec(spec); sys.modules[name] = m; spec.loader.exec_module(m); return m

ap = argparse.ArgumentParser(description='the all-orders graded census of the rotated K3 corner of the row-21 connection (see the module docstring)')
ap.add_argument('--korder', type=int, required=True, help='the eps order K of the census (3 = the run of record; the operator is polynomial in eps, so orders above its degree are zero)')
ap.add_argument('--out', required=True, help='the census JSON to write (never overwritten)')
ap.add_argument('--npts', type=int, default=40, help='rational s-points per sample set for the exact rank (default 40)')
ap.add_argument('--seed', type=int, default=1, help='the sample generator seed (default 1)')
ap.add_argument('--plant', default='w1', help='planted spurious letter key: w1 (varpi prime times rational) or w0')
a = ap.parse_args()
K = a.korder
if os.path.exists(a.out): raise SystemExit('refusing to overwrite %s' % a.out)
sys.dont_write_bytecode = True
check_pins()
ratlib = load('ratlib', 'rot_probe_i__ratlib.py')
r2lib = load('r2lib', 'rot_dial1__r2lib.py')
L = load('iib_lib', 'rot_probe_iib__iib_lib.py')
from ratlib import R, R0, R1, _poly, A2, mmul, inv3, solve3
from iib_lib import parse_r, radd, rscale, rC, MM, Mconst, Mdiffr, Msub, ZR
from r2lib import R2, CTX
from flint import fmpq, fmpq_poly, fmpq_mat

GLOG = {}
# ---------------- pulled banana PF companion (the record's frame; z = +1/u branch) ----------------
pc = [[0, -4, 64], [0, 1, -68, 448], [0, 0, 3, -90, 384], [0, 0, 0, 1, -20, 64]]
def peval(row, zr):
    acc, pw = R0, R1
    for c in row: acc = acc + pw * R(c); pw = pw * zr
    return acc
zs = R(_poly([-8, -3]), _poly([0, 0, 3])); zp = zs.diff()
p3 = peval(pc[3], zs)
red = [R0 - peval(pc[k], zs) / p3 for k in range(3)]
rows = [[R1, R0, R0]]
for k in range(3):
    r = rows[k]; nr = [r[b].diff() for b in range(3)]
    for b in range(3):
        if b + 1 < 3: nr[b + 1] = nr[b + 1] + r[b] * zp
        else:
            for c in range(3): nr[c] = nr[c] + r[b] * zp * red[c]
    rows.append(nr)
mu = solve3(rows[:3], rows[3])
aPF = [R0 - mu[0], R0 - mu[1], R0 - mu[2]]
Cmp_PF = [[R0, R1, R0], [R0, R0, R1], [mu[0], mu[1], mu[2]]]
b1 = aPF[2] / R(3); b0 = (aPF[1] - R(2) * b1 * b1 - b1.diff()) / R(4)
GLOG['sym2_certificate_PF'] = (aPF[0] - (R(4) * b0 * b1 + R(2) * b0.diff())).is_zero()
q4, q16 = _poly([32, 12, 3]), _poly([128, 48, 3])
A2.P = q4 * q16
tpow, rr, rho = L.find_rho(b1, A2.P)
GLOG['rho'] = {'t': tpow, 'rho_r': repr(rr)}
GLOG['rho_ode'] = (rho.diff() / rho + A2(b1)).is_zero()
S, S_inv = L.build(b0, b1, rho)
Iden = MM(S_inv, S); okI = True
for i in range(3):
    for j in range(3):
        d = dict(Iden[i][j])
        if i == j: d[(0, 0)] = d.get((0, 0), L.AZ) - L.AO
        okI &= all(v.is_zero() for v in d.values())
GLOG['S_identity'] = okI
Mg = Msub(MM(S_inv, MM(Mconst(Cmp_PF), S)), MM(S_inv, Mdiffr(S)))
target = [[ZR, {(-1, 0): rho}, ZR], [ZR, ZR, {(-1, 0): A2(R(2)) * rho}], [ZR, ZR, ZR]]
GLOG['M_identity_PF'] = all(not radd(Mg[i][j], rscale(target[i][j], A2(R(-1)))) for i in range(3) for j in range(3))
TT('PF frame + S gauge: sym2 %s rho_t %d S %s M %s' % (GLOG['sym2_certificate_PF'], tpow, okI, GLOG['M_identity_PF']))
assert GLOG['sym2_certificate_PF'] and GLOG['rho_ode'] and okI and GLOG['M_identity_PF']

# ---------------- exact-eps operator (LEPS_OPERATOR.json) -> eps-Taylor series to order K ----------------
LEPS = json.load(open(os.path.join(RC, 'rot_dial1__LEPS_OPERATOR.json')))
import sympy as sp
def pr2(t): return r2lib.from_sympy(sp.sympify(t.replace('^', '**')))
OPS = {k: pr2(LEPS[k]) for k in ('a2', 'a1', 'a0', 'b')}
TT('LEPS parsed (sympy -> flint fmpq_mpoly)')
def upoly(dct):
    if not dct: return fmpq_poly([])
    return fmpq_poly([dct.get(m, fmpq(0)) for m in range(max(dct) + 1)])
def group(p):
    G = {}
    for (i, m), c in zip(p.monoms(), p.coeffs()): G.setdefault(int(i), {})[int(m)] = c
    return G
def series(x, KO):
    NK, DK = group(x.n), group(x.d)
    assert DK.get(0), 'ep-pole at ep=0'
    d = [R(upoly(DK.get(k, {})), 1) for k in range(KO)]
    out = []
    for k in range(KO):
        t = R(upoly(NK.get(k, {})), 1)
        for m in range(1, k + 1): t = t - d[m] * out[k - m]
        out.append(t / d[0])
    return out
# eps-structure: degrees and whether the eps-denominator is s-independent (=> finite recurrence => all-orders finite)
EPS = {}
for k, x in OPS.items():
    NK, DK = group(x.n), group(x.d)
    dden = max(DK); dnum = max(NK)
    # s-independence of the eps-dependence of the denominator: every eps-power slice of D is a rational multiple of D_0(s)?
    D0 = upoly(DK[0]); sep = True
    for i, dct in DK.items():
        Di = upoly(dct)
        q = R(Di, 1) / R(D0, 1)
        if q.n.degree() > 0 or q.d.degree() > 0: sep = False
    EPS[k] = {'deg_eps_num': dnum, 'deg_eps_den': dden, 'den_eps_separable_from_s': sep}
TT('eps-structure: %s' % json.dumps(EPS))
aser = {k: series(OPS[k], K + 1) for k in ('a2', 'a1', 'a0', 'b')}
TT('eps-series to order %d' % K)
# cross-check vs the record IIB_CORE.json aser (orders 0..3) by repr
CORE = json.load(open(os.path.join(RC, 'rot_probe_iib__IIB_CORE.json')))
GLOG['aser_match_IIB_CORE_orders_0_3'] = all(repr(aser[k][m]) == CORE['aser'][k][m] for k in ('a2', 'a1', 'a0', 'b') for m in range(min(4, K + 1)))
TT('aser vs IIB_CORE.json (orders 0..3): %s' % GLOG['aser_match_IIB_CORE_orders_0_3'])
assert GLOG['aser_match_IIB_CORE_orders_0_3']

# ---------------- Theta_c = U_L . Theta, verified by the intertwining identity ----------------
RAWB = json.load(open(os.path.join(RC, 'rot_probe_i__block1315_raw.json')))
def entryk(i, j, k):
    acc = R0
    for part in range(4):
        e = RAWB.get(f'{i},{j},{part}')
        if e is None: continue
        if k < len(e['P']): acc = acc + R(_poly(e['P'][k]), _poly(e['Q'][0]))
    return acc
idx = [13, 14, 15]
K0 = [[entryk(x_, y_, 0) for y_ in idx] for x_ in idx]
u1 = K0[0]
u2 = [u1[j].diff() + sum((u1[kk] * K0[kk][j] for kk in range(3)), R0) for j in range(3)]
U_L = [[R1, R0, R0], u1, u2]
TH = json.load(open(os.path.join(RC, 'rot_probe_i__THETA.json')))
Theta = [[parse_r(TH['Theta'][i][j]) for j in range(3)] for i in range(3)]
Tc = mmul(U_L, Theta)
Cmp_I0 = [[R0, R1, R0], [R0, R0, R1], [R0 - aser['a0'][0], R0 - aser['a1'][0], R0 - aser['a2'][0]]]
dev = [[Tc[i][j].diff() - sum((Cmp_I0[i][kk] * Tc[kk][j] for kk in range(3)), R0) + sum((Tc[i][kk] * Cmp_PF[kk][j] for kk in range(3)), R0) for j in range(3)] for i in range(3)]
GLOG['Theta_c_intertwine'] = all(dev[i][j].is_zero() for i in range(3) for j in range(3))
TT('Theta_c intertwine: %s' % GLOG['Theta_c_intertwine'])
assert GLOG['Theta_c_intertwine']
Tc_inv = inv3(Tc)

# ---------------- the graded census at eps^1..eps^K ----------------
census, slots, walls = {}, {}, {}
def serd(d): return {'%d,%d' % kk: repr(v) for kk, v in d.items()}
for k in range(1, K + 1):
    tk = time.time()
    Bk = [[R0] * 3, [R0] * 3, [R0 - aser['a0'][k], R0 - aser['a1'][k], R0 - aser['a2'][k]]]
    if all(x.is_zero() for x in Bk[2]):
        walls[k] = time.time() - tk; continue
    Gk = MM(S_inv, MM(Mconst(mmul(Tc_inv, mmul(Bk, Tc))), S))
    for i in range(3):
        for j in range(3):
            if Gk[i][j]:
                name = 'G%d[%d,%d]' % (k, i + 1, j + 1)
                census[name] = k + j - i
                slots[name] = Gk[i][j]
    walls[k] = time.time() - tk
    TT('order %d: %d nonzero slots' % (k, sum(1 for n in census if n.startswith('G%d[' % k))))
grade_le0 = {n: g for n, g in census.items() if g <= 0}
TT('census done: %d slots, %d with grade<=0' % (len(census), len(grade_le0)))

# ---------------- letter enumeration: exact Q-rank of the integrand set at two disjoint rational s-sets ----------------
keys = sorted({kk for d in slots.values() for kk in d})
def vec(d, pts):
    v = []
    for kk in keys:
        el = d.get(kk)
        for p in pts:
            if el is None: v += [fmpq(0), fmpq(0)]
            else: v += [el.a(p), el.b(p)]
    return v
def rank_of(vecs):
    if not vecs: return 0
    M = fmpq_mat(len(vecs), len(vecs[0]), [x for v in vecs for x in v])
    return M.rank()
rng = random.Random(a.seed)
def sample(n, avoid):
    out = set()
    while len(out) < n:
        p = Fraction(rng.randint(-400, 400), rng.randint(1, 97))
        if p in avoid or p == 0 or 3 * p + 8 == 0 or 3 * p + 16 == 0: continue
        if (3 * p * p + 12 * p + 32) == 0 or (3 * p * p + 48 * p + 128) == 0: continue
        out.add(p)
    return sorted(out)
pts1 = sample(a.npts, set()); pts2 = sample(a.npts, set(pts1))
names = sorted(slots)
def enum(names_sub, pts):
    return rank_of([vec(slots[n], pts) for n in names_sub])
rank1_all = enum(names, pts1); rank2_all = enum(names, pts2)
pos = [n for n in names if census[n] >= 1]
rank1_pos = enum(pos, pts1); rank2_pos = enum(pos, pts2)
# per-order cumulative rank (saturation order)
cum = []
for k in range(1, K + 1):
    sub = [n for n in names if int(n[1:n.index('[')]) <= k]
    cum.append({'order_le': k, 'rank_set1': enum(sub, pts1), 'rank_set2': enum(sub, pts2), 'n_slots': len(sub)})
sat = next((c['order_le'] for c in cum if c['rank_set1'] == rank1_all), None)
# planted spurious letter: a rational-function multiple of varpi' (key (0,1)) with a foreign denominator (s^2+1)
plant = {(0, 1): A2(R(_poly([1, 2, 3]), _poly([1, 0, 1])))} if a.plant == 'w1' else {(1, 0): A2(R(_poly([1, 2, 3]), _poly([1, 0, 1])))}
rank_plant1 = rank_of([vec(slots[n], pts1) for n in names] + [vec(plant, pts1)])
rank_plant2 = rank_of([vec(slots[n], pts2) for n in names] + [vec(plant, pts2)])
planted_refused = (rank_plant1 == rank1_all + 1) and (rank_plant2 == rank2_all + 1)
# negative control of the rank machinery: a genuine member (2 x an existing slot) must NOT raise the rank
dup = rscale(slots[names[0]], A2(R(2)))
rank_dup = rank_of([vec(slots[n], pts1) for n in names] + [vec(dup, pts1)])
dup_absorbed = (rank_dup == rank1_all)
TT('ranks: all %d/%d pos %d/%d sat_order %s planted_refused %s dup_absorbed %s' % (rank1_all, rank2_all, rank1_pos, rank2_pos, sat, planted_refused, dup_absorbed))

out = {'receipt': 'graded census of blockdiag(1,S).D on the rotated K3 corner of the row-21 connection, eps^1..eps^%d' % K,
       'stamp_utc': utc(), 'korder': K, 'npts': a.npts, 'seed': a.seed,
       'frame': 'intrinsic cyclic (I,I\',I\'\') -> Theta_c=U_L.Theta -> pulled banana PF companion (z=+1/u) -> Sym2 S.D gauge over Q(s)[y]{w^pm,w\'}, y^2=q4*q16 (the vendored ring machinery, rot_probe_iib__iib_lib.py)',
       'identities': GLOG, 'eps_structure_of_LEPS': EPS,
       'slot_census': census, 'n_slots': len(census), 'grade_le0_slots': grade_le0, 'n_grade_le0': len(grade_le0),
       'letters': {'definition': 'distinct integrands = dim_Q of the span of the nonzero slot entries (ring elements of Q(s)[y]{w^pm,w\'}), rank computed exactly (fmpq_mat) from the coefficient functions sampled at npts rational s-points; two disjoint sample sets must agree',
                   'rank_all_slots_set1': rank1_all, 'rank_all_slots_set2': rank2_all, 'rank_grade_ge1_slots_set1': rank1_pos, 'rank_grade_ge1_slots_set2': rank2_pos,
                   'sets_agree': (rank1_all == rank2_all) and (rank1_pos == rank2_pos), 'cumulative_by_order': cum, 'saturation_order': sat,
                   'monomial_keys_w0pow_w1pow': ['%d,%d' % kk for kk in keys],
                   'planted_spurious_letter': {'key': a.plant, 'element': serd(plant), 'rank_with_plant_set1': rank_plant1, 'rank_with_plant_set2': rank_plant2, 'REFUSED': planted_refused},
                   'duplicate_member_control': {'element': '2 x %s' % names[0], 'rank_with_dup': rank_dup, 'absorbed': dup_absorbed}},
       'walls_s_per_order': walls, 'wall_total_s': time.time() - T0,
       'slot_integrands': {n: serd(slots[n]) for n in names},
       'PRODUCER': {'script': os.path.basename(__file__), 'script_sha256': hashlib.sha256(open(os.path.abspath(__file__), 'rb').read()).hexdigest(),
                    'vendored_pins': dict(PINS), 'stamp_source': 'date -u'}}
json.dump(out, open(a.out, 'w'), indent=1)
print('WROTE', a.out, flush=True)
