#!/usr/bin/env python3
"""Consistency checks on the files of this package (run from the package root: python3 tools/check_package.py).
 1. sha256 of every file against README.md's manifest (if README.md exists).
 2. v^T A(eps) = 0 for every left null vector in response_matrices/null_vectors_S*.json against A_S*.json, all eps.
 3. Eq. (36): I_0 + (16/3) I_2 - 6 I_3 = 0 on the values in boundary_constants_PX2.json (pole exactly, finite part numerically).
 4. Static-slice closed forms vs the numerical values (I_r5, I_r6, I_r7).
 5. c_M reassembled from cM_laurent.json (cores x prefactors x identity coefficients): eps^-4, eps^-3 cancellation and c_M = 1;
    Feynman leading terms: eps^-4 coefficient from the closed-form cores = 1/(6144 pi^6).
 6. operators.txt/values: the K3' Wronskian at gamma = 1/2 has determinant consistent between the two K3-type files' structure (light format check),
    and every POINT block has r*r entries."""
import json, re, os, sys, hashlib, glob
import mpmath as mp
mp.mp.dps = 70
root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
ok = True
def rep(name, cond, detail=''):
    global ok
    ok &= bool(cond)
    print(f"[{'PASS' if cond else 'FAIL'}] {name}" + (f": {detail}" if detail else ''))
# 1
readme = os.path.join(root, 'README.md')
if os.path.exists(readme):
    rows = re.findall(r'^\| `([^`]+)` \| (\d+) \| ([0-9a-f]{64}) \|', open(readme).read(), flags=re.M)
    bad = []
    for path, size, sha in rows:
        p = os.path.join(root, path)
        h = hashlib.sha256(open(p, 'rb').read()).hexdigest() if os.path.exists(p) else 'missing'
        if h != sha: bad.append(path)
    rep('manifest sha256', not bad and rows, f'{len(rows)} files' + (f'; mismatches {bad}' if bad else ''))
# 2
def cnum(e):
    return mp.mpc(0) if e == 0 else mp.mpc(mp.mpf(e['re']), mp.mpf(e['im']))
worst = mp.mpf(0); ncheck = 0
for S in ('S1', 'S2', 'S3'):
    A = json.load(open(os.path.join(root, 'response_matrices', f'A_{S}.json')))
    N = json.load(open(os.path.join(root, 'response_matrices', f'null_vectors_{S}.json')))
    for k in range(len(A['eps'])):
        Ak = [[cnum(x) for x in row] for row in A['A'][k]]
        scale = max(abs(x) for row in Ak for x in row)
        for v in N['left_null_vectors']:
            if 'weights_exact' in v:
                w = {r: (mp.mpf(x.split('/')[0]) / mp.mpf(x.split('/')[1]) if '/' in x else mp.mpf(x)) for r, x in zip(v['rows'], v['weights_exact'])}
            else:
                pe = v['per_eps'][k]
                assert mp.almosteq(mp.mpf(pe['eps']), mp.mpf(A['eps'][k]), rel_eps=mp.mpf('1e-30'))
                names = v['support']; rows = v['rows']
                w = {r: mp.mpc(mp.mpf(pe['weights'][n]['re']), mp.mpf(pe['weights'][n]['im'])) for r, n in zip(rows, names)}
            for j in range(A['n_columns']):
                s = mp.fsum(w[r]*Ak[r][j] for r in w)
                worst = max(worst, abs(s)/scale); ncheck += 1
rep('v^T A = 0 (all sectors, all eps, all columns)', worst < mp.mpf('1e-45'), f'{ncheck} column sums, largest |v^T A|/max|A| = {mp.nstr(worst, 3)}')
# 3, 4
B = json.load(open(os.path.join(root, 'boundary_constants_PX2.json')))
M = {m['name']: m for m in B['masters']}
pi = mp.pi; ln2 = mp.log(2); gE = mp.euler
def ev(expr):
    e = expr.replace('^', '**').replace('log(2)', 'ln2').replace('gammaE', 'gE')
    e = re.sub(r'(?<![A-Za-z0-9_])i(?![A-Za-z0-9_])', 'I', e)      # imaginary unit written as a bare i
    e = re.sub(r'(?<![A-Za-z0-9_.])(\d+)/(\d+)(?![0-9.])', r'mp.mpf(\1)/\2', e)   # exact rationals
    return eval(e, {'pi': pi, 'ln2': ln2, 'gE': gE, 'mp': mp, 'I': mp.mpc(0, 1)})
f = {n: M[n]['values']['graviton_i0']['fey'] for n in ('I_0', 'I_2', 'I_3')}
pole = ev(f['I_0']['-1']['exact_re']) + 1j*ev(f['I_0']['-1']['exact_im']) + mp.mpf(16)/3*(ev(f['I_2']['-1']['exact_re']) + 1j*ev(f['I_2']['-1']['exact_im'])) - 6*(ev(f['I_3']['-1']['exact_re']) + 1j*ev(f['I_3']['-1']['exact_im']))
fin = sum(c*mp.mpc(mp.mpf(f[n]['0']['re']), mp.mpf(f[n]['0']['im'])) for c, n in ((1, 'I_0'), (mp.mpf(16)/3, 'I_2'), (-6, 'I_3')))
numpole = sum(c*mp.mpc(mp.mpf(f[n]['-1']['re']), mp.mpf(f[n]['-1']['im'])) for c, n in ((1, 'I_0'), (mp.mpf(16)/3, 'I_2'), (-6, 'I_3')))
rep('Eq. (36) at the pole, exact forms', abs(pole) < mp.mpf('1e-60'), f'|combination| = {mp.nstr(abs(pole), 3)}')
rep('Eq. (36) at the pole, numerical strings', abs(numpole)/abs(mp.mpf(f["I_0"]["-1"]["re"])) < mp.mpf('1e-28'), f'relative {mp.nstr(abs(numpole)/26.9, 3)}')
rep('Eq. (36) at eps^0 (26-28 correct digits per entry)', abs(fin)/388 < mp.mpf('1e-25'), f'|combination|/|I_0| = {mp.nstr(abs(fin)/388, 3)}')
for n in ('I_r5', 'I_r6', 'I_r7'):
    v = M[n]['values']; worst = 0
    for k in ('0', '1'):
        ex = ev(v['exact'][k]); nu = mp.mpc(mp.mpf(v['numerical_sigma_plus'][k]['re']), mp.mpf(v['numerical_sigma_plus'][k]['im']))
        worst = max(worst, abs(ex - nu)/abs(ex))
    rep(f'{n} closed forms vs numerical values (eps^0, eps^1)', worst < mp.mpf('1e-35'), f'max relative difference {mp.nstr(worst, 3)}')
# 5
C = json.load(open(os.path.join(root, 'cM_laurent.json')))
IA = C['retarded']['implementation_A']
def L(d):  # dict k-> str  to dict int->mpf
    return {int(k): mp.mpf(v) for k, v in d.items()}
def mul(a, b, kmax):
    out = {}
    for i, x in a.items():
        for j, y in b.items():
            if i + j <= kmax: out[i+j] = out.get(i+j, 0) + x*y
    return out
j1 = {k: mp.mpf(IA['cores']['values_25_digits'][f'j_1^({k})']) for k in (0, 1, 2)}
j3 = {k: mp.mpf(IA['cores']['values_25_digits'][f'j_3^({k})']) for k in (0, 1, 2)}
N1, N3 = L(IA['prefactor_laurent']['N_1']), L(IA['prefactor_laurent']['N_3'])
r2, r1 = L(IA['ibp_coefficient_laurent']['r_2 (multiplies I_1)']), L(IA['ibp_coefficient_laurent']['r_1 (multiplies I_3)'])
I1 = mul(N1, j1, 1); I3 = mul(N3, j3, 1)
I2 = {k: mul(r2, I1, -2).get(k, 0) + mul(r1, I3, -2).get(k, 0) for k in (-4, -3, -2)}
cM = -(6*(8*pi)**4/5)*I2[-2]
rep('c_M reassembled from cores, prefactors and r_1, r_2 (18-digit prefactor strings)', abs(cM - 1) < mp.mpf('1e-15') and abs(I2[-4]/I2[-2]) < 1e-15 and abs(I2[-3]/I2[-2]) < 1e-15,
    f'c_M = {mp.nstr(cM, 20)}; I2[eps^-4]/I2[eps^-2] = {mp.nstr(I2[-4]/I2[-2], 3)}, I2[eps^-3]/I2[eps^-2] = {mp.nstr(I2[-3]/I2[-2], 3)}')
FP = C['feynman_propagators']
j1F = -mp.mpf(1)/90 - 1/(3*pi**2); j3F = mp.mpf(8)/315 + 4/(9*pi**2)
I2F4 = r2[-3]*N1[-1]*j1F + r1[-3]*N3[-1]*j3F
rep('Feynman eps^-4 coefficient from the closed-form cores = 1/(6144 pi^6)', abs(I2F4*6144*pi**6 - 1) < mp.mpf('1e-16'),
    f'{mp.nstr(I2F4, 20)} vs {mp.nstr(1/(6144*pi**6), 20)}; quoted {FP["I2_eps^-4"]["value_route_1"]}')
Ss = mp.mpf(FP['quadrant_moments']['values']['S_same']['numerical_25_digits']); So = mp.mpf(FP['quadrant_moments']['values']['S_opp']['numerical_25_digits'])
rep('retarded leading moment -2 S_same + 2 S_opp = 2 pi^2/15', abs((-2*Ss + 2*So)/(2*pi**2/15) - 1) < mp.mpf('1e-24'), mp.nstr(-2*Ss + 2*So, 25))
# 6
npt = 0; good = True
for vf in sorted(glob.glob(os.path.join(root, 'bound_frobenius', 'values_*.txt'))):
    r = 3 if ('K30' in vf or 'K3P' in vf) else 4
    blocks = open(vf).read().split('POINT ')[1:]
    for b in blocks:
        npt += 1
        nW = len(re.findall(r'^\s+W\[\d,\d\] = \S+\s+\S+\s*$', b, flags=re.M))
        good &= (nW == r*r) and ('z_re=' in b and 'z_im=' in b)
rep('value files: every POINT block complete', good and npt == 48, f'{npt} points')
print('ALL CHECKS PASSED' if ok else 'SOME CHECKS FAILED')
sys.exit(0 if ok else 1)
