#!/usr/bin/env python3
"""outer-dbox-evaluate.py -- the outer-mass double box of the mpl-suite page,
g10 = LS * G_52: STANDALONE evaluator.

Object (the results paper, Sec. 2.2):
the Caron-Huot--Henn planar double box (1404.2922), six perimeter propagators
of common mass m, massless rung and legs; u = 4m^2/(-s), v = 4m^2/(-t),
Euclidean s,t < 0 (u,v > 0), m = 1.  In D=4 the integral I5 is finite and

    g10 = -(1/8) s^2 t  bu buv  I5,     bu = sqrt(1+u),  buv = sqrt(1+u+v),

is a SINGLE pure weight-4 genus-0 MPL: the 18-word symbol printed in
eq:outer-dbox-symbol over the 12-letter alphabet, rationalized to (w,z) by
  u = (1-w^2)(1-z^2)/(w-z)^2,  v = 4wz/(w-z)^2,
  bu = (1-wz)/(w-z),  buv = (1+wz)/(w-z),   physical patch 0 < z < w < 1.

WHAT THIS SCRIPT COMPUTES (all at runtime, from the symbol alone)
=================================================================
The printed symbol was integrated to EXPLICIT GPLs by a sympy fibration
build script (letter/alphabet checks, dS^dS = 0 integrability, and Galois
parity (1,0,1) all hard-asserted there; the transcribed symbol is
byte-identical to the paper's word list):

  path = the line w = W fixed, fiber coordinate y = W - z, base point y = 0,
  i.e. the line w = z where u,v -> infinity.  Every symbol letter is a
  product of factors LINEAR in z, so each pulled-back dlog is a sum of
  dy/(y - a) with the seven letters
     a  in  { 0, W, W-1, W+1, (W^2-1)/W, (W^2+1)/(W+1), (W^2+1)/(W-1) },
  and the 18 symbol words expand into 92 GPL words G(a1,a2,a3,a4; Y), Y = W-Z
  (= w - z), stored in row01_data.json.  Convention: G(a1,..,an;x) with
  dG/dx = G(a2,..,an;x)/(x-a1); the innermost letter a_n is the FIRST symbol
  entry.  After merging, NO word contains the letter 0 (the (w-z)-pole kernels
  cancel exactly), so every kernel is analytic at the base point.

BEYOND-SYMBOL CONSTANT: fixed ANALYTICALLY, exactly zero.  The paper's
boundary condition is g10 -> 0 as u,v -> infinity (the line w = z).  All
admitted first entries (here L6, L7) have log -> 0 there and analytic
pulled-back kernels, so the iterated integrals from y = 0 CONVERGE and
vanish there, as do
all their left-truncations (the coproduct companions).  Hence g10 equals the
plain sum  sum_i c_i G(word_i; Y)  with no added transcendental constant --
no fitted constant, no interim literal, nothing point-matched.

EVALUATOR: vendored GPL transport (blog c3-dbox pattern): Taylor-steps the
iterated-integral ODE d/dx G(a,w;x) = G(w;x)/(x-a) from x = 0 over the
suffix-closed word set; step = SAFETY * (distance to nearest letter) with
SAFETY = 1/4, Taylor order N = 1.7*dps + 24 (series ratio <= 1/4).  Every step
carries a CERTIFIED geometric tail bound (ratio 1/4 by the step rule); N is
grown (up to 8x) until the bound is < 10^-(dps+8), else RuntimeError.  The
accumulated bound is reported with the values (a bound, not an estimate).

DOMAIN: any Euclidean point u > 0, v > 0 (s,t < 0), mapped by
  W = (sqrt(1+v)+1)/(bu+buv),  Z = (sqrt(1+v)-1)/(bu+buv),  Y = W-Z = 2/(bu+buv);
no letter meets the open path (0, Y).  v -> 0 (t -> -infinity) pushes Y onto
the letter W (physical log divergence): still exact, just slower.  NOTE:
g10(u,v) != g10(v,u) -- the rung breaks s <-> t (the vendored pair (16/3,10/3) /
(10/3,16/3) differs); the reference table covers both orderings independently.

REFERENCES (independent, none used anywhere in the construction):
  4 dedicated held-out AMFlow goal-40 points, 3 further AMFlow goal-40
  sample points (raw eps^0 J converted by g10 = -8 bu buv J/(u^2 v)), and
  the published 30-digit symmetric value g10(4,4) of 1404.2922.  Expected:
  ~38-39d at the AMFlow points (oracle ball limited), ~30d at the published
  value (string-length limited).

Requirements: python3 + mpmath ONLY (json/argparse/os/time stdlib); data file
outer-dbox-data.json beside this script (exact integers + digit strings;
produced once by the sympy builder).  No AMFlow/Kira/IBP/network/mnt-imports
at runtime.  mp.dps is set INSIDE main() after argparse.

EXIT CODE: 0 only when at least two reference points agree to >= 30 digits
and the certified transport bound is reported (and, with --check, the
dps+60 rerun is stable to >= dps-5 digits); nonzero otherwise.  --mutate
perturbs one vendored symbol coefficient and must exit nonzero.

CLI:
  python3 outer-dbox-evaluate.py                      # controls + 8-point reference table
  python3 outer-dbox-evaluate.py --point 7/2,13/5     # + arbitrary Euclidean point
  python3 outer-dbox-evaluate.py --check              # rerun at dps+60, diff
  python3 outer-dbox-evaluate.py --mutate             # mutation control (rc != 0)
"""

import argparse
import json
import os
import sys
import time

import mpmath as mp

sys.stdout.reconfigure(line_buffering=True)

DATA_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)),
                         "outer-dbox-data.json")

# canonical letter formulas (index-matched to row01_data.json "gpl_letters")
LETTER_STRS = ["0", "W", "W - 1", "W + 1", "2*W", "(W**2 - 1)/W",
               "(W**2 + 1)/W", "(W**2 + 1)/(W + 1)", "(W**2 + 1)/(W - 1)"]


def letter_values(Wv):
    one = mp.mpf(1)
    return [mp.mpf(0), Wv, Wv - 1, Wv + 1, 2 * Wv, (Wv * Wv - one) / Wv,
            (Wv * Wv + one) / Wv, (Wv * Wv + one) / (Wv + 1), (Wv * Wv + one) / (Wv - 1)]


def parse_rat(s):
    s = s.strip()
    if "/" in s:
        p, q = s.split("/")
        return mp.mpf(int(p)) / mp.mpf(int(q))
    return mp.mpf(s)


def map_point(u, v):
    """(u,v) -> (W, Z, Y, prefactor data) on the patch 0 < Z < W < 1."""
    bu = mp.sqrt(1 + u)
    buv = mp.sqrt(1 + u + v)
    bvp = mp.sqrt(1 + v)
    den = bu + buv
    Wv = (bvp + 1) / den
    Zv = (bvp - 1) / den
    return Wv, Zv, Wv - Zv, bu, buv


def suffix_closure(words):
    clos = set()
    for wd in words:
        for k in range(len(wd)):
            clos.add(wd[k:])
    return sorted(clos, key=lambda t: (len(t), t))


def gpl_transport(letters, words, target):
    """G(word; target) for every word in the suffix-closed list `words`.

    Taylor-stepped iterated-integral ODE transport from x=0 (all words vanish
    there; every trailing letter is nonzero).  Kernel 1/(x-a) at center x0 is a
    geometric series with ratio <= SAFETY = 1/4 by the step-size rule, so
    N = 1.7*dps + 24 terms suffice (footgun-compliant: 1.7*dps+20 at SAFETY=.25).

    CERTIFIED TAIL GATE (axis-3 hardening): the step rule guarantees series
    ratio r <= 1/4 for every kernel, so the truncation tail of each step obeys
        tail <= max(|c_n h^n| over the last TAILWIN computed terms) * r/(1-r).
    After building each step's coefficients the bound is checked against
    10^-(dps+GUARD); N = 1.7*dps+24 stays the starting guess (fast path
    unchanged when it suffices) and is DOUBLED — extra terms computed by exact
    continuation of the same recurrences — until the bound passes, up to a hard
    cap of 8x the initial N; hitting the cap unconverged raises RuntimeError
    naming the step, x-position and achieved bound.  Per-step bounds accumulate.
    Returns (vals, err_bound): err_bound is a certified BOUND (not an estimate)
    on the absolute truncation error of every returned G value.
    """
    N0 = int(1.7 * mp.mp.dps) + 24        # starting guess (fast path unchanged)
    NCAP = 8 * N0                          # hard cap on Taylor order
    GUARD = 8
    TAILWIN = 8
    tol = mp.mpf(10) ** (-(mp.mp.dps + GUARD))
    rfac = mp.mpf(1) / 3                   # r/(1-r) at certified ratio r = 1/4
    SAFETY = mp.mpf(1) / 4
    active = sorted({letters[i] for wd in words for i in wd}, key=abs)
    zero = mp.mpf(0)
    vals = {wd: zero for wd in words}
    err_total = zero
    x0 = zero
    step = 0
    while True:
        step += 1
        dmin = min(abs(x0 - a) for a in active if x0 != a)
        h = SAFETY * dmin
        last = target - x0 <= h
        if last:
            h = target - x0
        N = N0
        C = {(): [mp.mpf(1)] + [zero] * N}
        for wd in words:
            a = letters[wd[0]]
            rest = C[wd[1:]]
            c = [zero] * (N + 1)
            c[0] = vals[wd]
            if a == x0:
                # only possible for the letter a=0 at the first step (control
                # words): T[rest]/x is analytic there since rest[0] = 0
                for n in range(1, N + 1):
                    c[n] = rest[n] / n
            else:
                d = x0 - a
                p0 = 1 / d
                r = -p0
                s = rest[0] * p0
                c[1] = s
                for n in range(2, N + 1):
                    s = r * s + p0 * rest[n - 1]
                    c[n] = s / n
            C[wd] = c
        while True:
            # certified geometric tail: ratio <= 1/4 by the step rule
            bound = zero
            hn0 = abs(h) ** (N - TAILWIN + 1)
            for wd in words:
                c = C[wd]
                hn = hn0
                t = zero
                for n in range(N - TAILWIN + 1, N + 1):
                    tn = abs(c[n]) * hn
                    if tn > t:
                        t = tn
                    hn = hn * abs(h)
                if t > bound:
                    bound = t
            bound = bound * rfac
            if bound < tol:
                break
            if N >= NCAP:
                raise RuntimeError(
                    "row01 gpl_transport: certified tail bound NOT met at "
                    f"step {step}, x0 = {mp.nstr(x0, 20)}, h = {mp.nstr(h, 10)}"
                    f": achieved bound {mp.nstr(bound, 5)} >= "
                    f"tol {mp.nstr(tol, 5)} at Taylor order N = {N} "
                    f"(cap {NCAP} = 8 x initial {N0})")
            N2 = min(2 * N, NCAP)
            # extend all series by exact continuation of the recurrences
            C[()].extend([zero] * (N2 - N))
            for wd in words:            # suffix-closure order: rest first
                a = letters[wd[0]]
                rest = C[wd[1:]]
                c = C[wd]
                c.extend([zero] * (N2 - N))
                if a == x0:
                    for n in range(N + 1, N2 + 1):
                        c[n] = rest[n] / n
                else:
                    d = x0 - a
                    p0 = 1 / d
                    r = -p0
                    s = c[N] * N        # recurrence state: c[n] = s/n
                    for n in range(N + 1, N2 + 1):
                        s = r * s + p0 * rest[n - 1]
                        c[n] = s / n
            N = N2
        for wd in words:
            c = C[wd]
            v = zero
            for n in range(N, -1, -1):
                v = v * h + c[n]
            vals[wd] = v
        err_total += bound
        x0 = x0 + h
        if last:
            return vals, err_total


def digits(a, b):
    if a == b:
        return float("inf")
    if b == 0:
        return float(-mp.log10(abs(a - b)))
    return float(-mp.log10(abs(a - b) / abs(b)))


class Row01:
    def __init__(self, data):
        assert data["fiber"]["gpl_letters"] == LETTER_STRS, "letter table mismatch"
        self.gwords = [(c, tuple(wd)) for c, wd in data["gwords"]]
        for _, wd in self.gwords:
            assert wd[-1] != 0
        self.gates = data["references"]

    def g10(self, u, v, controls=False):
        """g10 at Euclidean (u,v), u,v > 0.  Returns (g10, J, control dict)."""
        Wv, Zv, Y, bu, buv = map_point(u, v)
        letters = letter_values(Wv)
        base = [tuple(wd) for _, wd in self.gwords]
        ctrl_words = []
        if controls:
            a1, a2, a3, a4 = self.gwords[0][1]
            self._sh = ((a1,), (a2, a3, a4),
                        [(a1, a2, a3, a4), (a2, a1, a3, a4),
                         (a2, a3, a1, a4), (a2, a3, a4, a1)])
            ctrl_words = [(0, 2), (2,)] + [self._sh[0], self._sh[1]] + self._sh[2]
        words = suffix_closure(base + ctrl_words)
        G, terr = gpl_transport(letters, words, Y)
        val = mp.fsum(c * G[wd] for c, wd in self.gwords)
        # certified bound on |error of val| from transport truncation:
        # per-G bound terr times the l1-norm of the exact coefficients
        cert = terr * mp.fsum(abs(mp.mpf(c)) for c, _ in self.gwords)
        J = -val * u * u * v / (8 * bu * buv)   # AMFlow eps^0 convention (J = -I5)
        cres = {}
        if controls:
            # weight-1 words vs closed-form logs
            w1 = [d for d in (digits(G[(i,)], mp.log(1 - Y / letters[i]))
                              for i in {wd[0] for wd in words if len(wd) == 1}
                              if i != 0) if d != float("inf")]
            cres["w1_log"] = min(w1) if w1 else float("inf")
            # G(0,a;Y) = -Li2(Y/a), exercises the x=0 kernel branch
            cres["w2_li2"] = digits(G[(0, 2)], -mp.polylog(2, Y / letters[2]))
            # weight-4 shuffle identity G(a)G(b,c,d) = sum of insertions
            lhs = G[self._sh[0]] * G[self._sh[1]]
            rhs = mp.fsum(G[t] for t in self._sh[2])
            cres["w4_shuffle"] = digits(lhs, rhs)
        return val, J, cres, cert

    def gate_refs(self):
        """[(label, u, v, ref_g10, source, oracle_abs_err)] at current dps."""
        out = []
        for g in self.gates:
            u, v = (parse_rat(x) for x in g["uv"])
            if "g10" in g:
                ref = mp.mpf(g["g10"])
            else:
                J = mp.mpf(g["J"])
                bu, buv = mp.sqrt(1 + u), mp.sqrt(1 + u + v)
                ref = -8 * bu * buv * J / (u * u * v)
            out.append((f"({g['uv'][0]},{g['uv'][1]})", u, v, ref,
                        g["source"], mp.mpf(g["abs_err"])))
        return out


def run_all(row, point, do_controls=True):
    """One full pass at the current mp.dps.  Returns dict of results."""
    res = {"gates": [], "point": None}
    for k, (lab, u, v, ref, src, aerr) in enumerate(row.gate_refs()):
        val, J, cres, cert = row.g10(u, v, controls=(do_controls and k == 0))
        res["gates"].append((lab, val, ref, src, aerr, cres, cert))
    if point is not None:
        u, v = point
        val, J, _, cert = row.g10(u, v)
        res["point"] = (u, v, val, J, cert)
    return res


def main():
    ap = argparse.ArgumentParser(description="row 01: outer-mass double box g10")
    ap.add_argument("--dps", type=int, default=50)
    ap.add_argument("--point", type=str, default=None,
                    help="u,v with u,v>0 Euclidean, e.g. 7/2,13/5")
    ap.add_argument("--check", action="store_true",
                    help="rerun everything at dps+60 and report self-agreement")
    ap.add_argument("--mutate", action="store_true",
                    help="perturb one vendored symbol coefficient; run must exit nonzero")
    args = ap.parse_args()

    mp.mp.dps = args.dps + 30          # working guard; report at args.dps
    fails = []
    data = json.load(open(DATA_FILE))
    if args.mutate:
        data["gwords"][0][0] = str(int(data["gwords"][0][0]) + 1) \
            if isinstance(data["gwords"][0][0], str) else data["gwords"][0][0] + 1
        print("[mutate] first vendored symbol coefficient +1 "
              "(this run MUST exit nonzero)")
    row = Row01(data)
    point = None
    if args.point:
        u, v = (parse_rat(x) for x in args.point.split(","))
        if not (u > 0 and v > 0):
            print("domain: Euclidean u > 0, v > 0 only (s,t < 0, m=1)")
            return
        point = (u, v)

    print("== row 01: outer-mass double box  g10 = LS*G_52 "
          "(weight-4 MPL, 92 GPL words in Y = w-z) ==")
    print("   constant: analytic (== 0 on the u,v->inf base line); no fitted/interim literal")
    t0 = time.time()
    res = run_all(row, point)
    wall = time.time() - t0

    cres = res["gates"][0][5]
    fd = lambda d: f">{mp.mp.dps}" if d == float("inf") else f"{d:.1f}"
    print(f"[controls] weight-1 vs log: {fd(cres['w1_log'])}d | "
          f"G(0,a) vs -Li2: {fd(cres['w2_li2'])}d | weight-4 shuffle: "
          f"{fd(cres['w4_shuffle'])}d   (self-consistency, working dps {mp.mp.dps})")

    print(f"[references] vendored independent references ({len(res['gates'])} points):")
    worst = float("inf")
    for lab, val, ref, src, aerr, _, cert in res["gates"]:
        d = digits(val, ref)
        worst = min(worst, d)
        lim = float(-mp.log10(aerr / abs(ref)))
        print(f"  (u,v)={lab:13s} g10 = {mp.nstr(val, min(args.dps, 40)):45s} "
              f"{d:5.1f}d  (ref limit ~{lim:.1f}d; {src.split('(')[0].strip()})")
    n30 = sum(1 for lab, val, ref, *_ in res["gates"] if digits(val, ref) >= 30)
    print(f"[references] worst {worst:.1f}d; {n30}/{len(res['gates'])} points >= 30d "
          f"({'PASS' if n30 >= 2 else 'FAIL'}: >=2 required)")
    if n30 < 2:
        fails.append(f"only {n30} reference points >= 30d (need >= 2)")
    maxcert = max(g[6] for g in res["gates"])
    print(f"[certified] transport truncation error |Delta g10| <= "
          f"{mp.nstr(maxcert, 3)} at every reference point (accumulated per-step "
          f"geometric tail bound, r=1/4, guard 8; BOUND, not estimate)")

    if res["point"] is not None:
        u, v, val, J, cert = res["point"]
        print(f"[point] u={mp.nstr(u, 12)}, v={mp.nstr(v, 12)}  (s={mp.nstr(-4/u, 12)}, "
              f"t={mp.nstr(-4/v, 12)}, m=1)")
        print(f"  g10             = {mp.nstr(val, args.dps)}")
        print(f"  AMFlow-conv J   = {mp.nstr(J, args.dps)}   (eps^0; I5 = -J)")
        print(f"  certified bound |Delta g10| <= {mp.nstr(cert, 3)}   "
              f"(transport truncation; BOUND, not estimate)")

    print(f"[wall] {wall:.1f}s at dps {args.dps} (+30 working guard)")

    if args.check:
        mp.mp.dps = args.dps + 90
        t1 = time.time()
        res2 = run_all(row, point, do_controls=False)
        mp.mp.dps = args.dps + 30
        worst_sa = float("inf")
        for (l1, v1, *_), (l2, v2, *_) in zip(res["gates"], res2["gates"]):
            worst_sa = min(worst_sa, digits(v1, v2))
        if res["point"] is not None:
            worst_sa = min(worst_sa, digits(res["point"][2], res2["point"][2]))
        print(f"[check] rerun at dps+90 vs dps+30: worst self-agreement "
              f"{worst_sa:.1f}d over all values ({time.time()-t1:.1f}s; "
              f"need >= {args.dps - 5})")
        if worst_sa < args.dps - 5:
            fails.append(f"--check self-agreement {worst_sa:.1f}d < {args.dps-5}d")


    if fails:
        print(f"OVERALL FAIL: {'; '.join(fails)}")
        return 1
    print("OVERALL PASS")
    return 0


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