#!/usr/bin/env python3
r"""eq_box_evaluate_standalone.py -- C4(t) and psi4(t) of the conformally coupled scalar one-loop box on the equilateral line
(x_v = 1 + t, regular tetrahedron P_v = S = T = 1), standalone release.

    C4(t)  = u(t) + Omega * v(t)          (u, v canonical solutions of the exact 28-dim differential system for the correlator C4, Sec. 6.5-6.6 and App. B.6 of the paper; Omega a parameter)
    psi4(t) = c(t) . w(t)                  (exact 34-dim system for psi4 = the wavefunction coefficient V_4 of the paper; boundary vector fitted to the co-author's table inside the 8-dim physical space)

Modes
  explicit : large-X series  sum_n X^-n (alpha_n + beta_n log X), X = 1+t, n <= 44  (use for X >= 14; instant)
  transport: Taylor-series integration of the exact system d/dt w = G(t) w from the base point t0 = 3/4 to t along the real axis, with the
             the boundary-vector convention (detours through the UPPER half plane around every real singular point in between); minutes in pure python.
  auto     : explicit if X >= 14 and its tail bound meets --need, else transport.
Fail-closed: exit 3 if the certified digits are below --need; exit 4 for t outside the supported range (t <= 0 -- negative t is not supported by this release -- or |t| < 1e-6, or t exactly
a singular point of the system: t = 0 is a regular singular point of the system -- the physical value there is finite but needs the local
Frobenius connection, which this tool does not ship; use t >= 1e-6).
Dependencies: python >= 3.8, mpmath.  Optional: none.  No network, no other tools.  Data files (same directory or --datadir):
  system_C4_28.json, system_psi4_34.json (the exact rational systems, converted from the authors' pipeline export), model_C4.json, model_psi4.json,
  selftest_reference.json.  If system_*.json are absent but the pipeline's original export files of the two systems (kept with the derivation records in the
  paper's repository; not shipped here) are in --datadir, they are parsed on the fly (slower start; singular points then read from singular_points.json or found numerically).
usage:
  python3 eq_box_evaluate_standalone.py --which C4 --t 15                      # auto mode
  python3 eq_box_evaluate_standalone.py --which C4 --t 1/10 --mode transport --dps 28
  python3 eq_box_evaluate_standalone.py --which psi4 --t 7/13 --mode transport
  python3 eq_box_evaluate_standalone.py --selftest                             # 5 checks, exit 0 iff all pass
  python3 eq_box_evaluate_standalone.py --which C4 --t 20 --omega 21.44442951393577451950316676534   # vary Omega
"""
import sys, os, json, time, argparse
from fractions import Fraction
try:
    import mpmath as mp
except ImportError:
    print('this tool needs mpmath (pip install mpmath)'); sys.exit(2)
try: sys.set_int_max_str_digits(0)
except Exception: pass
HERE = os.path.dirname(os.path.abspath(__file__))

# ----------------------------------------------------------------------------------------------------------------- data loading
def load_json(datadir, name):
    fn = os.path.join(datadir, name)
    if not os.path.exists(fn): return None
    with open(fn) as f: return json.load(f)

def load_system(datadir, which):
    tag = 'C4_28' if which == 'C4' else 'psi4_34'
    d = load_json(datadir, 'system_%s.json' % tag)
    if d is not None: return d
    # fallback: parse the pipeline's original export of the system on the fly (its file-name pattern is the glob below; those files live in the paper's repository, not in anc/)
    import glob
    cands = sorted(glob.glob(os.path.join(datadir, 'MPHYS_%s_GEN_*.json' % which)))
    if not cands: raise SystemExit('missing system_%s.json (and no MPHYS_%s_GEN_*.json in %s)' % (tag, which, datadir))
    sys.path.insert(0, HERE)
    from _ratparse import parse_ratfun
    src = json.load(open(cands[-1])); n = src['dim_M_phys']
    G = {'%d,%d' % (i, j): parse_ratfun(src['G_phys'][i][j]) for i in range(n) for j in range(n) if src['G_phys'][i][j] not in ('0', '', 0)}
    c = [parse_ratfun(x) for x in src['c_phys']]
    spf = load_json(datadir, 'singular_points.json')
    if spf is not None and which in spf:
        sings = spf[which]; print('parsed %s on the fly (%d entries); singular points from singular_points.json' % (os.path.basename(cands[-1]), len(G)))
    else:
        print('parsed %s on the fly (%d entries); singular_points.json absent -> locating singular points numerically with mpmath polyroots (SLOW: can take > 30 min) ...' % (os.path.basename(cands[-1]), len(G))); sys.stdout.flush()
        sings = numeric_singular_points(G, c)
    return dict(fam=which, dim=n, G=G, c=c, singular_points=sings, base='3/4')

def numeric_singular_points(G, c):
    seen = {}; pts = []
    with mp.workdps(40):
        for num, den in list(G.values()) + [x for x in c if x is not None]:
            key = tuple(den)
            if key in seen or len(den) <= 1: continue
            seen[key] = 1
            coeffs = [mp.mpf(Fraction(x).numerator) / Fraction(x).denominator for x in reversed(den)]          # highest first for polyroots
            try: roots = mp.polyroots(coeffs, maxsteps=200, extraprec=200)
            except Exception: roots = mp.polyroots(coeffs, maxsteps=600, extraprec=600)
            for r in roots: pts.append([mp.nstr(mp.re(r), 30), mp.nstr(mp.im(r), 30)])
    return pts

# ----------------------------------------------------------------------------------------------------------------- fixed-point transport core
# All transport arithmetic is done in binary FIXED POINT with Python integers (value = n / 2^B), complex numbers as (re, im) int pairs.
# This is ~50x faster than mpmath objects in pure python and exact in its bookkeeping; B carries generous guard bits.
def to_fix(fr, B):
    """Fraction -> fixed-point int (round to nearest)"""
    n, d = fr.numerator, fr.denominator
    return (n << B) // d if n >= 0 else -((-n << B) // d)
def taylor_shift_fix(coeffs, zr, zi, M, B):
    """first M+1 Taylor coefficients (complex fixed pairs) of p(z+u); coeffs: list of real fixed ints ascending; z = (zr, zi) fixed"""
    ar = list(coeffs); ai = [0] * len(coeffs); out = []
    for k in range(M + 1):
        L = len(ar)
        if L == 0: out.append((0, 0)); continue
        qr = [0] * (L - 1); qi = [0] * (L - 1); cr = 0; ci = 0
        for i in range(L - 1, -1, -1):
            # acc = acc*z + a_i
            tr = (cr * zr - ci * zi) >> B; ti = (cr * zi + ci * zr) >> B
            cr = tr + ar[i]; ci = ti + ai[i]
            if i > 0: qr[i - 1] = cr; qi[i - 1] = ci
        out.append((cr, ci)); ar, ai = qr, qi
    return out
def series_div_fix(nser, dser, M, B):
    """complex series division to order M (fixed point)"""
    d0r, d0i = dser[0]; den2 = (d0r * d0r + d0i * d0i) >> B          # |d0|^2 fixed
    out = []
    for k in range(M + 1):
        sr, si = nser[k] if k < len(nser) else (0, 0)
        for j in range(1, min(k, len(dser) - 1) + 1):
            djr, dji = dser[j]; orr, oi = out[k - j]
            sr -= (djr * orr - dji * oi) >> B; si -= (djr * oi + dji * orr) >> B
        # (s / d0) = s * conj(d0) / |d0|^2
        nr = (sr * d0r + si * d0i) >> B; ni = (si * d0r - sr * d0i) >> B
        out.append(((nr << B) // den2, (ni << B) // den2))
    return out
class System:
    def __init__(self, d, B):
        self.n = d['dim']; self.B = B
        self.G = {}; self.dens = {}
        for key, (num, den) in d['G'].items():
            i, j = (int(x) for x in key.split(','))
            dk = tuple(den)
            if dk not in self.dens: self.dens[dk] = [to_fix(Fraction(x), B) for x in den]
            self.G[(i, j)] = ([to_fix(Fraction(x), B) for x in num], dk)
        self.c = [None if e is None else ([Fraction(x) for x in e[0]], [Fraction(x) for x in e[1]]) for e in d['c']]
        self.sings = [complex(float(s[0]), float(s[1])) for s in d['singular_points']]
        self.real_sings = sorted(set(s.real for s in self.sings if abs(s.imag) < 1e-20))
    def dist(self, z): return min(abs(z - s) for s in self.sings)
    def step_series(self, w, zr, zi, M):
        """Taylor coefficients W_k (complex fixed vectors) of the solution at centre z, W_0 = w:  (k+1) W_{k+1} = sum_{m<=k} G_m W_{k-m}"""
        B = self.B; n = self.n
        dser = {dk: taylor_shift_fix(dc, zr, zi, M, B) for dk, dc in self.dens.items()}
        Gs = {}
        for (i, j), (num, dk) in self.G.items():
            Gs[(i, j)] = series_div_fix(taylor_shift_fix(num, zr, zi, M, B), dser[dk], M, B)
        W = [list(w)]
        for k in range(M):
            accr = [0] * n; acci = [0] * n
            for (i, j), gser in Gs.items():
                sr = 0; si = 0
                for m in range(k + 1):
                    gr, gi = gser[m]; wr, wi = W[k - m][j]
                    sr += gr * wr - gi * wi; si += gr * wi + gi * wr
                accr[i] += sr; acci[i] += si
            W.append([((accr[i] >> B) // (k + 1), (acci[i] >> B) // (k + 1)) for i in range(n)])
        return W
def eval_series_fix(W, ur, ui, B):
    n = len(W[0]); M = len(W) - 1; out = []
    for i in range(n):
        cr, ci = W[M][i]
        for k in range(M - 1, -1, -1):
            tr = (cr * ur - ci * ui) >> B; ti = (cr * ui + ci * ur) >> B
            wr, wi = W[k][i]; cr = tr + wr; ci = ti + wi
        out.append((cr, ci))
    return out
def bitmag(pair): return max(abs(pair[0]).bit_length(), abs(pair[1]).bit_length())
def hop_path(x0, x1, real_sings):
    """convention: if a real singular point lies strictly between x0 and x1, go x0 -> x0+ih -> x1+ih -> x1 (UPPER half plane)"""
    lo, hi = min(x0, x1), max(x0, x1)
    if not [s for s in real_sings if lo < s < hi]: return [complex(x1, 0)]
    h = min(0.125, max(abs(x1 - x0) / 2, 1.0 / 64))
    return [complex(x0, h), complex(x1, h), complex(x1, 0)]
def transport(sysd, w0, t0, t1, dps, ratio=0.2, verbose=False):
    """march the complex vector w0 (mpc list) from t0 to t1 (Fractions) with the fixed-point Taylor method; returns (I(t1) as mpc, w(t1), steps).
    Steps land exactly on binary-representable points; step size = ratio x distance to the nearest singular point of G; order M with ratio^M ~ 2^-B."""
    wp = dps + 24                                  # guard digits (the solution vectors have large cancelling components in this basis)
    B = int(3.3219 * wp) + 16                      # fixed-point fraction bits
    M = int(B / (-__import__('math').log2(ratio))) + 2
    S_ = System(sysd, B)
    x0 = float(t0); x1 = float(t1)
    legs = hop_path(x0, x1, S_.real_sings)
    with mp.workdps(wp + 10):
        w = [(int(mp.floor(mp.re(x) * 2 ** B + mp.mpf(1) / 2)), int(mp.floor(mp.im(x) * 2 ** B + mp.mpf(1) / 2))) for x in w0]
    # exact fixed representation of the endpoints: t0, t1 rationals -> use exact Fraction arithmetic for the LAST step so that we arrive exactly at t1
    pos = complex(x0, 0); posr = to_fix(Fraction(t0), B); posi = 0
    steps = 0; T0 = time.time()
    targets = []
    for z in legs:
        zr = to_fix(Fraction(t1), B) if z.real == x1 else to_fix(Fraction(t0), B) if z.real == x0 else int(z.real * 2 ** 52) * (1 << B) // (1 << 52)
        zi = int(round(z.imag * 2 ** 52)) * (1 << B) // (1 << 52)
        targets.append((z, zr, zi))
    for (z, zr, zi) in targets:
        while True:
            remr, remi = zr - posr, zi - posi
            remf = complex(remr / 2 ** B, remi / 2 ** B); rem_abs = abs(remf)
            if rem_abs < 2.0 ** (-(B - 24)): posr, posi, pos = zr, zi, z; break
            d = S_.dist(pos)
            if d < 1e-12: raise SystemExit('path hits a singular point of the system near t=%s' % pos)
            h = min(rem_abs, ratio * d)
            if h >= rem_abs: ur, ui = remr, remi                     # land exactly on the target
            else:
                f = h / rem_abs; ur = int(remr * f); ui = int(remi * f)
            W = S_.step_series(w, posr, posi, M)
            while True:
                wn = eval_series_fix(W, ur, ui, B)
                sc = max(bitmag(x) for x in wn); ub = max(bitmag((ur, ui)) - B, -10 ** 9)     # log2|u| (u fixed)
                tail = max(max(bitmag(x) for x in W[M]) + M * ub, max(bitmag(x) for x in W[M - 1]) + (M - 1) * ub)
                if tail < sc - (B - 8): break
                ur //= 2; ui //= 2
                if bitmag((ur, ui)) < B - 70: raise SystemExit('step underflow near t=%s' % pos)
            w = wn; posr += ur; posi += ui; pos = complex(posr / 2 ** B, posi / 2 ** B); steps += 1
            if verbose and steps % 10 == 0: print('  step %d at t=%.6f%+.6fi (%.0fs)' % (steps, pos.real, pos.imag, time.time() - T0)); sys.stdout.flush()
    with mp.workdps(wp):
        wm = [mp.mpc(mp.mpf(wr) / 2 ** B, mp.mpf(wi) / 2 ** B) for (wr, wi) in w]
        tt = mp.mpf(t1.numerator) / t1.denominator; tot = mp.mpc(0)
        for j, e in enumerate(S_.c):
            if e is None: continue
            num = mp.mpf(0)
            for c_ in reversed(e[0]): num = num * tt + mp.mpf(c_.numerator) / c_.denominator
            den = mp.mpf(0)
            for c_ in reversed(e[1]): den = den * tt + mp.mpf(c_.numerator) / c_.denominator
            tot += num / den * wm[j]
    return tot, wm, steps

# ----------------------------------------------------------------------------------------------------------------- explicit series
def explicit_value(model, which, tq, omega, dps):
    ser = model['series']; nmin, nmax = ser['nmin'], ser['nmax']
    with mp.workdps(dps + 30):
        X = mp.mpf(tq.numerator) / tq.denominator + 1
        if X <= 2: return None
        L = mp.log(X); tot = mp.mpf(0); terms = []
        for nn in range(nmin, nmax + 1):
            if which == 'C4':
                au, bu = (mp.mpf(x) for x in ser['u'][str(nn)]); av = mp.mpf(ser['v'][str(nn)])
                al, be = au + omega * av, bu
            else:
                al, be = (mp.mpf(x) for x in ser['coefficients'][str(nn)])
            term = X ** (-nn) * (al + be * L); tot += term; terms.append(abs(term))
        r = mp.mpf(2) / X; tail = max(terms[-3:]) * r / (1 - r)
        cert = float(-mp.log10(tail / abs(tot))) if tail > 0 else 99.0
    return tot, cert

def model_vector(model, which, omega):
    if which == 'C4':
        u = [mp.mpc(mp.mpf(x[0]), mp.mpf(x[1])) for x in model['u_t0']]; v = [mp.mpc(mp.mpf(x[0]), mp.mpf(x[1])) for x in model['v_t0']]
        return [u[i] + omega * v[i] for i in range(len(u))]
    return [mp.mpc(mp.mpf(x[0]), mp.mpf(x[1])) for x in model['vector_t0']]

MODEL_DIGITS = {'C4': 28.0, 'psi4': 27.0}     # accuracy of the boundary data (see model_*.json 'accuracy')

def evaluate(which, tq, mode, dps, need, datadir, omega=None, verbose=False):
    mp.mp.dps = max(mp.mp.dps, dps + 60)          # global working precision for all parsing of model data (components up to 1e9 with heavy cancellation)
    model = load_json(datadir, 'model_%s.json' % which)
    if model is None: raise SystemExit('missing model_%s.json in %s' % (which, datadir))
    with mp.workdps(dps + 40):
        omega_v = mp.mpf(omega if omega is not None else model.get('Omega_default', '0'))
    X = float(tq) + 1
    if tq <= Fraction(0, 1) and abs(tq) >= Fraction(1,10**6): raise SystemExit(4 if not print('t <= 0 is outside the supported range of this release (physical region t > 0; the space was built with single-valuedness through t = -1/2 only)') else 4)
    if abs(tq) < Fraction(1, 10**6): 
        print('t = %s: t = 0 is a regular singular point of the system; the physical value is finite but needs the local Frobenius connection (not shipped). Use |t| >= 1e-6.' % tq); sys.exit(4)
    res = {}
    if mode in ('explicit', 'auto', 'both'):
        e = explicit_value(model, which, tq, omega_v, dps) if X >= 2.5 else None
        if e is not None:
            val, cert = e; cert = min(cert, MODEL_DIGITS[which]); res['explicit'] = (val, cert)
            if mode == 'explicit' or (mode == 'auto' and X >= 14 and cert >= need): return val, cert, 'explicit', res
        elif mode == 'explicit': raise SystemExit('explicit mode needs X = 1+t >= 2.5 (use >= 14 for full accuracy)')
    sysd = load_system(datadir, which)
    for s in sysd['singular_points']:
        if abs(float(s[1])) < 1e-20 and abs(Fraction(s[0][:40]) - tq) < Fraction(1, 10**12) if False else (abs(float(s[1])) < 1e-20 and abs(float(s[0]) - float(tq)) < 1e-12):
            print('t = %s is a singular point of the system; choose a nearby regular t' % tq); sys.exit(4)
    w0 = model_vector(model, which, omega_v)
    T0 = time.time()
    v1, _, st1 = transport(sysd, w0, Fraction(model['base']), tq, dps, verbose=verbose)
    v2, _, st2 = transport(sysd, w0, Fraction(model['base']), tq, dps + 8, verbose=False)
    with mp.workdps(dps + 40):
        agree = float(-mp.log10(abs(v1 - v2) / abs(v2))) if v1 != v2 else float(dps + 8)
        val = mp.re(v2); cert = min(agree, MODEL_DIGITS[which])
    res['transport'] = (val, cert, agree, st1 + st2, time.time() - T0)
    if 'explicit' in res:
        with mp.workdps(dps + 40): res['explicit_vs_transport_digits'] = float(-mp.log10(abs(res['explicit'][0] - val) / abs(val)))
    return val, cert, 'transport', res

def selftest(datadir, dps=28, quick=False):
    ref = load_json(datadir, 'selftest_reference.json'); ok = True; rows = []
    for tcase in ref['tests']:
        if quick and tcase['mode'] == 'transport': continue
        tq = Fraction(tcase['t']); T0 = time.time()
        val, cert, mode, res = evaluate(tcase['which'], tq, tcase['mode'], dps, 10, datadir)
        with mp.workdps(60):
            refv = mp.mpf(tcase['value']); dg = float(-mp.log10(abs(val - refv) / abs(refv))) if val != refv else 99.0
        passed = dg >= tcase['digits_required'] - 0.5
        ok = ok and passed
        rows.append((tcase['which'], tcase['t'], mode, dg, tcase['digits_required'], passed, time.time() - T0))
        print('selftest %-4s t=%-5s %-9s agreement %.1f d (need %s) %s  [%.0fs]' % (tcase['which'], tcase['t'], mode, dg, tcase['digits_required'], 'PASS' if passed else 'FAIL', time.time() - T0)); sys.stdout.flush()
    print('SELFTEST', 'PASS' if ok else 'FAIL'); return 0 if ok else 1

if __name__ == '__main__':
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument('--which', choices=['C4', 'psi4']); ap.add_argument('--t', help='rational p/q or decimal'); ap.add_argument('--mode', default='auto', choices=['auto', 'explicit', 'transport', 'both'])
    ap.add_argument('--dps', type=int, default=28, help='target digits (transport cost grows with dps)'); ap.add_argument('--need', type=float, default=None, help='fail-closed threshold (default dps-4)')
    ap.add_argument('--omega', default=None, help='Omega for C4 (default: model_C4.json Omega_default = co-author 34-digit value)'); ap.add_argument('--datadir', default=HERE)
    ap.add_argument('--selftest', action='store_true'); ap.add_argument('--selftest-quick', action='store_true', help='explicit-mode checks only (seconds)'); ap.add_argument('--verbose', action='store_true'); ap.add_argument('--json', action='store_true')
    a = ap.parse_args()
    if a.selftest or a.selftest_quick: sys.exit(selftest(a.datadir, a.dps, quick=a.selftest_quick))
    if not a.which or a.t is None: ap.error('--which and --t are required (or --selftest)')
    tq = Fraction(a.t); need = a.need if a.need is not None else a.dps - 4
    T0 = time.time()
    val, cert, mode, res = evaluate(a.which, tq, a.mode, a.dps, need, a.datadir, a.omega, a.verbose)
    with mp.workdps(a.dps + 20):
        shown = mp.nstr(val, max(3, int(min(cert, a.dps + 5))))
        out = dict(which=a.which, t=str(tq), value=shown, certified_digits=round(cert, 1), mode=mode, dps=a.dps, wall_s=round(time.time() - T0, 1))
        if 'transport' in res: out.update(two_precision_agreement=round(res['transport'][2], 1), steps=res['transport'][3])
        if 'explicit_vs_transport_digits' in res: out['explicit_vs_transport_digits'] = round(res['explicit_vs_transport_digits'], 1)
    print(json.dumps(out) if a.json else '%s(t=%s) = %s   [%s; certified %.1f d%s; %.1fs]' % (a.which, out['t'], shown, mode, cert, ('; two-precision %.1f' % out['two_precision_agreement']) if 'two_precision_agreement' in out else '', out['wall_s']))
    if 'explicit_vs_transport_digits' in out: print('explicit vs transport: %.1f d' % out['explicit_vs_transport_digits'])
    if cert < need: print('FAIL-CLOSED: certified %.1f < needed %.1f' % (cert, need)); sys.exit(3)
