#!/usr/bin/env python3
r"""sunrise_empl_cont.py -- the continued layer of row 9 (the served sunrise_empl.py is imported by path from this
module's own directory and never edited): the ANALYTIC CONTINUATION of the Feynman-curve frame of the
unequal-mass sunrise eps^0 closed form (rows 9-11; sunrise_empl.predict_E) from the Euclidean domain t < 0 across the
soft point t = 0, the pseudo-thresholds and the normal threshold into the physical region with the Feynman i0
prescription t -> t + i0.

WHAT IS CONTINUED (the frame of the served formula, sunrise_empl.py lines 120-190):
  * the Feynman-curve roots e_iF(t), the modulus k_F^2, the period tau_F = i K(1-k_F^2)/K(k_F^2), the nome
    q_C = -exp(i pi tau_F) (tau_C = (tau_F+1)/2), the dressing period psi1_F = 2 K(k_F^2)/sqrt(Z_3F) (ANALYTIC: the served
    |psi1_F| is replaced by psi1_F itself, which is real positive at Euclidean t so the Euclidean values are unchanged), and
    the three Abel-Jacobi marked points z_j = F(asin u_j, k'_F^2)/(2 K(k'_F^2)) with z1+z2+z3 = 1 exactly.
  * HOW: PATH TRACKING.  Every principal-branch mpmath quantity (sqrt, ellipk, ellipf, asin, log, polylog) is evaluated
    at the points of a path t(u) = t0 + (T - t0) u + i h sin(pi u), u = 0..1, from the Euclidean anchor t0 = -1 through the
    UPPER half t-plane (Feynman prescription: t + i0 is the limit from Im t > 0) to the endpoint T + i delta (delta tiny,
    positive).  The path meets no singular point (all of them, 0, 6-4sqrt2, 2, 6+4sqrt2, are on the real axis).  At each
    step the principal-branch value is replaced by the member of its discrete monodromy family closest to the previous
    step's value: tau_F by {tau_p + n, tau_p/(1+2k tau_p) + n} (the Gamma(2)-type jumps of K across its cut), psi1_F by
    {+-psi_p (1+2k tau)}, z_j by {+-z_p + a + b tau'} (the sn periodicity family; tau' = i K(1-k'^2)/K(k'^2) is the
    puncture lattice's own ratio, = tau_F +- 1), and z3 := 1 - z1 - z2 exactly.  Step-count independence (n and 2n
    steps agree to working precision) is the check that no jump was missed.
  * the WORDS are evaluated in the DILOGARITHM form, an exact resummation of the served q-series
        W(z,N) = sum_{n>=1} b_n(z,N) q_C^n/n^2,  b_n = pref * sum_{N j k = n} (w^j - w^{-j}) k^2,  w = e^{2 pi i z}
      = (pref/N^2) * sum_{k>=1} [ Li_2(w q_C^{N k}) - Li_2(w^{-1} q_C^{N k}) ]            (sum over j first),
    which converges for EVERY z at |q_C| < 1 (the q-series itself needs |Im z| < N Im tau_C, which the continued marked
    points violate above threshold) and whose only branch data are the Li_2 cuts x in [1, inf): a term whose argument
    x = w^{+-1} q^{Nk} crosses the cut along the path is continued by its monodromy (crossing downward adds
    +2 pi i log x, upward -2 pi i log x; the log's own winding tracked).  C_{4,2} = sum_j (1/2i)[Li_2(w_j) - Li_2(1/w_j)]
    is the k = 0 member of the same family and is tracked the same way.  Only terms with |x| > 1/2 somewhere on the
    path can cross; the rest are summed directly at the endpoint (|x| < 1 and shrinking geometrically).
  * CONTROLS: (a) at Euclidean t the continued evaluator must reproduce sunrise_empl.predict_E / predict_J to working
    precision (the path is the trivial one and every tracked branch is the principal one); (b) the dilogarithm form of
    the words vs the served q-series at a Euclidean point; (c) two step counts; (d) two working precisions.
  * OUTPUT: complex E^(0)(t + i0) and J^(0)(t + i0) = psi1_F/pi * E^(0) (J compared with the raw AMFlow (1,1,1,0,0)
    eps^0 master, d = 2 - 2 eps).  Below the normal threshold the imaginary part must vanish (a control against the
    record's real AMFlow values at t = 1/2 .. 11); above it, Im J < 0 in this convention (J = -S; Im S > 0 by unitarity),
    verified against the reference t = 12 AMFlow ball and the fresh reference points t = 16, 20.

extends: the served sunrise_empl.py (rows 9-11 shared layer) and the served equal-mass sunrise-evaluate.py's
physical arm (t > 9, +i0 via an upper-half-plane arc from a Euclidean anchor: the house form of the continuation);
the served layer's own word evaluator (sunrise_empl.I1_g3) is the Euclidean q-series and is not used at complex marked
points (the dilogarithm form above replaces it there).
"""
import os
import sys
import importlib.util

import mpmath as mp

HERE = os.path.dirname(os.path.abspath(__file__))
BUNDLE = HERE   # the served sunrise_empl.py sits beside this module


def load_served():
    """Import the byte-copied served layer by path (never edited)."""
    spec = importlib.util.spec_from_file_location("sunrise_empl", os.path.join(BUNDLE, "sunrise_empl.py"))
    m = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(m)
    return m


SE = load_served()

I = mp.mpc(0, 1)


def two_pi_i():
    return 2 * mp.pi * I


# ---------------------------------------------------------------------------------------------
# principal-branch frame (the served fcurve, re-expressed to expose the pieces we track)
# ---------------------------------------------------------------------------------------------
def frame_principal(t, m1, m2, m3, mu=1):
    t = mp.mpc(t)
    m1, m2, m3, mu = mp.mpf(m1), mp.mpf(m2), mp.mpf(m3), mp.mpf(mu)
    m1s, m2s, m3s = m1 * m1, m2 * m2, m3 * m3
    M100 = m1s + m2s + m3s
    mu1 = -m1 + m2 + m3
    mu2 = m1 - m2 + m3
    mu3 = m1 + m2 - m3
    mu4 = m1 + m2 + m3
    Delta = mu1 * mu2 * mu3 * mu4
    mu4p = mu ** 4
    rad = 3 * (mp.sqrt(mp.mpc(mu1 * mu1 - t)) * mp.sqrt(mp.mpc(mu2 * mu2 - t))
               * mp.sqrt(mp.mpc(mu3 * mu3 - t)) * mp.sqrt(mp.mpc(mu4 * mu4 - t)))
    e1F = (-t * t + 2 * M100 * t + Delta + rad) / (24 * mu4p)
    e2F = (-t * t + 2 * M100 * t + Delta - rad) / (24 * mu4p)
    e3F = (2 * t * t - 4 * M100 * t - 2 * Delta) / (24 * mu4p)
    Z1F = e3F - e2F
    Z2F = e1F - e3F
    Z3F = e1F - e2F
    kF2 = Z1F / Z3F
    kFp2 = -Z1F / Z2F
    K = mp.ellipk(kF2)
    Kc = mp.ellipk(1 - kF2)
    tauF = I * Kc / K
    psi1F = 2 / mp.sqrt(mp.mpc(Z3F)) * K
    Kp = mp.ellipk(kFp2)
    Kpc = mp.ellipk(1 - kFp2)
    taup = I * Kpc / Kp

    def zF(xjk):
        up = mp.sqrt(mp.mpc((e1F - e3F) / (xjk - e3F)))
        return mp.ellipf(mp.asin(up), kFp2) / (2 * Kp)

    def xhat(mi2, mj2):
        return e3F + mi2 * mj2 / mu4p

    z1 = zF(xhat(m2s, m3s))
    z2 = zF(xhat(m3s, m1s))
    z3 = zF(xhat(m1s, m2s))
    return dict(t=t, e1F=e1F, e2F=e2F, e3F=e3F, kF2=kF2, kFp2=kFp2, tauF=tauF, psi1F=psi1F, taup=taup,
                z1=z1, z2=z2, z3=z3, rad=rad)


def _closest(cands, ref):
    best = None
    for c in cands:
        d = abs(c - ref)
        if best is None or d < best[0]:
            best = (d, c)
    return best[1], best[0]


def _mob(x, k):
    """the K-cut transformation on a period ratio: x -> x/(1 + 2 k x)"""
    return x / (1 + 2 * k * x)


def _closest_idx(cands, ref):
    best = None
    for i, c in enumerate(cands):
        d = abs(c - ref)
        if best is None or d < best[0]:
            best = (d, i, c)
    return best[2], best[0], best[1]


KS = (-3, -2, -1, 0, 1, 2, 3)
NS = (-4, -3, -2, -1, 0, 1, 2, 3, 4)
AB = (-3, -2, -1, 0, 1, 2, 3)


def branch_apply(fr, disc):
    """apply recorded discrete branch data to a principal-branch frame (used at the high-precision endpoint)."""
    out = dict(fr)
    k, n = disc["tau"]
    out["tauF"] = _mob(fr["tauF"], k) + n
    s, k = disc["psi"]
    out["psi1F"] = s * fr["psi1F"] * (1 + 2 * k * fr["tauF"])
    k, n = disc["taup"]
    out["taup"] = _mob(fr["taup"], k) + n
    for key in ("z1", "z2"):
        s, k, a, b = disc[key]
        out[key] = (s * fr[key] + k * fr["taup"]) / (1 + 2 * k * fr["taup"]) + a + b * out["taup"]
    out["z3"] = 1 - out["z1"] - out["z2"]
    return out


def _closest2(cands, ref):
    """(best value, best distance, best index, second-best distance)"""
    ds = sorted(((abs(c - ref), i) for i, c in enumerate(cands)), key=lambda x: x[0])
    d1, i1 = ds[0]
    d2 = ds[1][0] if len(ds) > 1 else mp.inf
    return cands[i1], d1, i1, d2


def track_step(fr, prev):
    """Bring the principal-branch frame `fr` onto the branch continuous with `prev` (a tracked frame); record the discrete
    data in out['disc'] and the worst ambiguity ratio (best distance / second-best distance) in out['_amb'].  The K-cut index
    k is shared: psi1_F with tau_F (both carry K(k_F^2)), the marked points with tau' (both carry K(k_F'^2))."""
    out = dict(fr)
    disc = {}
    amb = mp.mpf(0)
    tp = fr["tauF"]
    cands, idx = [], []
    for k in KS:
        for n in NS:
            cands.append(_mob(tp, k) + n); idx.append((k, n))
    out["tauF"], out["_dtau"], i, d2 = _closest2(cands, prev["tauF"]); disc["tau"] = idx[i]; amb = max(amb, out["_dtau"] / d2)
    ktau = disc["tau"][0]
    pp = fr["psi1F"]
    cands, idx = [], []
    for s_ in (1, -1):
        cands.append(s_ * pp * (1 + 2 * ktau * tp)); idx.append((s_, ktau))
    out["psi1F"], out["_dpsi"], i, d2 = _closest2(cands, prev["psi1F"]); disc["psi"] = idx[i]; amb = max(amb, out["_dpsi"] / d2)
    tq = fr["taup"]
    cands, idx = [], []
    for k in KS:
        for n in NS:
            cands.append(_mob(tq, k) + n); idx.append((k, n))
    out["taup"], out["_dtaup"], i, d2 = _closest2(cands, prev["taup"]); disc["taup"] = idx[i]; amb = max(amb, out["_dtaup"] / d2)
    kq = disc["taup"][0]
    taup = out["taup"]
    for key in ("z1", "z2"):
        zp = fr[key]
        cands, idx = [], []
        for s_ in (1, -1):
            base = (s_ * zp + kq * tq) / (1 + 2 * kq * tq)
            for a in AB:
                for b in AB:
                    cands.append(base + a + b * taup); idx.append((s_, kq, a, b))
        out[key], out["_d" + key], i, d2 = _closest2(cands, prev[key]); disc[key] = idx[i]; amb = max(amb, out["_d" + key] / d2)
    out["z3"] = 1 - out["z1"] - out["z2"]
    out["disc"] = disc
    out["_amb"] = amb
    return out


AMB_MAX = mp.mpf("0.2")


AMB_WORST = [mp.mpf(0), 0]   # (worst accepted ratio, count of ambiguous accepted steps) -- reset per evaluation


def advance(prev, t_prev, t_new, m1, m2, m3, depth=0, maxdepth=6):
    """tracked frames from t_prev (exclusive) to t_new (inclusive); a step whose branch choice is ambiguous
    (best/second-best distance ratio > AMB_MAX) is bisected up to maxdepth times; if still ambiguous the closest candidate
    is ACCEPTED and the ratio recorded in AMB_WORST (reported in diag['ambiguity']; a nonzero count means the branch data
    are NOT certified by continuity alone)."""
    fr = frame_principal(t_new, m1, m2, m3)
    cur = track_step(fr, prev)
    if cur["_amb"] <= AMB_MAX:
        return [cur]
    if depth >= maxdepth:
        AMB_WORST[0] = max(AMB_WORST[0], cur["_amb"]); AMB_WORST[1] += 1
        return [cur]
    tm = (t_prev + t_new) / 2
    first = advance(prev, t_prev, tm, m1, m2, m3, depth + 1, maxdepth)
    second = advance(first[-1], tm, t_new, m1, m2, m3, depth + 1, maxdepth)
    return first + second


def euclidean_anchor(t0, masses):
    """The served frame at a Euclidean anchor (principal branches = the served branches; z3 by the served repair)."""
    m1, m2, m3 = masses
    fr = frame_principal(t0, m1, m2, m3)
    cv = SE.fcurve(mp.mpf(t0), m1, m2, m3)
    # the served frame is the reference: adopt its tau, z's (z3 repaired), psi1F sign (real positive)
    fr["tauF"] = mp.mpc(cv["tauF"])
    fr["z1"], fr["z2"], fr["z3"] = mp.mpc(cv["z1F"]), mp.mpc(cv["z2F"]), mp.mpc(cv["z3F"])
    fr["psi1F"] = mp.mpc(cv["psi1F"])
    assert abs(fr["z1"] + fr["z2"] + fr["z3"] - 1) < mp.mpf(10) ** (-mp.mp.dps // 2)
    assert mp.re(fr["psi1F"]) > 0 and abs(mp.im(fr["psi1F"])) < mp.mpf(10) ** (-mp.mp.dps // 2)
    return fr


def path_points(t0, T, delta, h, nsteps):
    """the arc t(u) = t0 + (T - t0) u + i h sin(pi u), u = 1/n .. (n-1)/n, then a GEOMETRIC descent of Im t from the
    arc's last height to delta at Re t = T (factor 1/2 per step) so that the approach to the real axis -- where the
    frame moves fastest near a threshold -- is made in small steps; the last point is T + i delta."""
    pts = []
    for k in range(1, nsteps):
        u = mp.mpf(k) / nsteps
        pts.append(mp.mpc(t0 + (T - t0) * u, h * mp.sin(mp.pi * u)))
    im = h * mp.sin(mp.pi * mp.mpf(nsteps - 1) / nsteps)
    while im > delta:
        im = im / 2
        pts.append(mp.mpc(T, max(im, delta)))
    if mp.im(pts[-1]) != delta:
        pts.append(mp.mpc(T, delta))
    return pts


def track_frame(T, masses, t0=-1, delta=None, h=4, nsteps=200, want_path=False):
    """Tracked frame at T + i delta along the upper-half-plane arc from the Euclidean anchor t0."""
    if delta is None:
        delta = mp.mpf(10) ** (-(mp.mp.dps - 5))
    m1, m2, m3 = masses
    prev = euclidean_anchor(t0, masses)
    path = [prev]
    for tt in path_points(mp.mpf(t0), mp.mpf(T), delta, mp.mpf(h), nsteps):
        fr = frame_principal(tt, m1, m2, m3)
        cur = track_step(fr, prev)
        if want_path:
            path.append(cur)
        prev = cur
    return (prev, path) if want_path else prev


# ---------------------------------------------------------------------------------------------
# the words in dilogarithm form with per-term cut tracking along the path
# ---------------------------------------------------------------------------------------------
class TrackedLi2:
    """Li_2(x(t)) continued along a path: monodromy count n (cut [1,inf)) and log winding m."""

    def __init__(self, x0):
        self.x = mp.mpc(x0)
        self.n = 0     # net downward crossings of [1, inf)
        self.m = 0     # net counterclockwise windings of x about 0 (for log x)

    def step(self, x):
        x = mp.mpc(x)
        xo = self.x
        # crossing of the real axis at Re > 1: sign change of Im with the crossing point beyond 1
        if (mp.im(xo) > 0) != (mp.im(x) > 0) and mp.im(xo) != mp.im(x):
            # linear interpolation of the crossing abscissa
            lam = mp.im(xo) / (mp.im(xo) - mp.im(x))
            xc = mp.re(xo) + lam * (mp.re(x) - mp.re(xo))
            if xc > 1:
                self.n += 1 if mp.im(xo) > 0 else -1
            elif xc < 0:
                # crossing of the negative real axis: log winding
                self.m += 1 if mp.im(xo) > 0 else -1
        self.x = x

    def value(self):
        x = self.x
        lg = mp.log(x) + two_pi_i() * self.m
        return mp.polylog(2, x) + two_pi_i() * self.n * lg


def word_terms(z, tauF, N, kmax):
    """arguments x_{k,+-} = w^{+-1} q_C^{N k} for k = 0..kmax (k = 0 is the C_{4,2} boundary member)."""
    qC = -mp.exp(I * mp.pi * tauF)
    w = mp.exp(two_pi_i() * z)
    out = []
    for k in range(0, kmax + 1):
        qk = qC ** (N * k)
        out.append((w * qk, qk / w))
    return out


def evaluate_continued(T, masses, t0=-1, delta=None, h=4, nsteps=400, diag=None, mutate=False, kmax_track=12, dps_track=30):
    """E^(0)(T + i0), J^(0)(T + i0): a LOW-precision tracking pass (dps_track) along the upper-half-plane arc fixes the
    discrete branch data (frame monodromies, Li_2 cut crossings and log windings per tracked term); the endpoint is then
    evaluated ONCE at the working precision with that data applied to the principal-branch quantities (the branch of a
    principal-branch function at T + i delta does not depend on delta > 0 or on the precision)."""
    dps_hi = mp.mp.dps
    m1, m2, m3 = masses
    # ---- tracking pass at low precision ----
    with mp.workdps(dps_track):
        dlt = mp.mpf(10) ** (-(dps_track - 5))
        prev = euclidean_anchor(t0, masses)
        zkeys = ["z1", "z3"] if abs(prev["z1"] - prev["z2"]) < mp.mpf(10) ** (-dps_track // 2) else ["z1", "z2", "z3"]
        tr = {}
        for zk in zkeys:
            for N in (1, 2):
                for k, (xa, xb) in enumerate(word_terms(prev[zk], prev["tauF"], N, kmax_track)):
                    if N == 2 and k == 0:
                        continue
                    tr[(zk, N, k, 0)] = TrackedLi2(xa)
                    tr[(zk, N, k, 1)] = TrackedLi2(xb)
        pts = path_points(mp.mpf(t0), mp.mpf(T), dlt, mp.mpf(h), nsteps)
        maxjump = mp.mpf(0)
        t_prev = mp.mpc(mp.mpf(t0), 0)
        n_accepted = 0
        AMB_WORST[0] = mp.mpf(0); AMB_WORST[1] = 0
        for tt in pts:
            for cur in advance(prev, t_prev, tt, m1, m2, m3):
                n_accepted += 1
                for key in ("_dtau", "_dpsi", "_dtaup", "_dz1", "_dz2"):
                    v = cur.get(key, None)
                    if v is not None:
                        maxjump = max(maxjump, v)
                for zk in zkeys:
                    for N in (1, 2):
                        for k, (xa, xb) in enumerate(word_terms(cur[zk], cur["tauF"], N, kmax_track)):
                            if N == 2 and k == 0:
                                continue
                            tr[(zk, N, k, 0)].step(xa)
                            tr[(zk, N, k, 1)].step(xb)
                prev = cur
            t_prev = tt
        disc = prev["disc"]
        counters = {key: (o.n, o.m) for key, o in tr.items()}
        maxjump_f = float(maxjump)
    diag_ = {} if diag is None else diag
    diag_["_counters"] = counters
    diag_["kmax_track"] = kmax_track
    diag_["maxjump"] = maxjump_f
    diag_["nsteps"] = nsteps
    diag_["n_accepted_steps"] = n_accepted
    diag_["ambiguity"] = {"worst_ratio_accepted": float(AMB_WORST[0]), "n_ambiguous_accepted_steps": AMB_WORST[1]}
    diag_["h"] = h
    diag_["t0"] = t0
    diag_["dps_track"] = dps_track
    return endpoint_with(T, masses, disc, counters, diag_, mutate=mutate, zkeys=zkeys, kmax_track=kmax_track, delta=delta)


def endpoint_with(T, masses, disc, counters, diag, mutate=False, zkeys=("z1", "z3"), kmax_track=12, delta=None):
    """the endpoint at the ACTIVE working precision with recorded discrete branch data."""
    dps_hi = mp.mp.dps
    m1, m2, m3 = masses
    zkeys = list(zkeys)
    delta = mp.mpf(10) ** (-(dps_hi - 5)) if delta is None else mp.mpf(delta)
    frp = frame_principal(mp.mpc(mp.mpf(T), delta), m1, m2, m3)
    fr = branch_apply(frp, disc)
    mults = {"z1": (2 if len(zkeys) == 2 else 1), "z2": 1, "z3": 1}
    qC = -mp.exp(I * mp.pi * fr["tauF"])
    aq = abs(qC)
    if aq >= SE.QMAX:
        raise ValueError(f"t = {T}: |q_C| = {mp.nstr(aq, 6)} >= QMAX {mp.nstr(SE.QMAX, 4)} (series-domain limit)")
    if mp.im(fr["tauF"]) <= 0:
        raise ValueError(f"t = {T}: tracked tau_F not in the upper half plane: {fr['tauF']}")
    pref = SE.g3_pref()
    c8 = 8 * (1 + mp.mpf(10) ** -12) if mutate else 8
    c42 = mp.mpc(0)
    ell = mp.mpc(0)
    tail_terms = 0
    crossings = {}
    tol = mp.mpf(10) ** (-(dps_hi + 10))

    def li2_cont(x, n, m):
        lg = mp.log(x) + two_pi_i() * m
        return mp.polylog(2, x) + two_pi_i() * n * lg

    for zk in zkeys:
        mult = mults[zk]
        w = mp.exp(two_pi_i() * fr[zk])
        terms = word_terms(fr[zk], fr["tauF"], 1, kmax_track)
        terms2 = word_terms(fr[zk], fr["tauF"], 2, kmax_track)
        xa, xb = terms[0]
        na, ma = counters[(zk, 1, 0, 0)]; nb, mb = counters[(zk, 1, 0, 1)]
        c42 += mult * (li2_cont(xa, na, ma) - li2_cont(xb, nb, mb)) / (2 * I)
        for N, tl in ((1, terms), (2, terms2)):
            W = mp.mpc(0)
            for k in range(1, kmax_track + 1):
                xa, xb = tl[k]
                na, ma = counters[(zk, N, k, 0)]; nb, mb = counters[(zk, N, k, 1)]
                W += li2_cont(xa, na, ma) - li2_cont(xb, nb, mb)
            k = kmax_track + 1
            while True:
                qk = qC ** (N * k)
                xa, xb = w * qk, qk / w
                if abs(xa) >= mp.mpf("0.5") or abs(xb) >= mp.mpf("0.5"):
                    raise ValueError(f"untracked term with |x| >= 1/2 at k = {k} (raise kmax_track)")
                W += mp.polylog(2, xa) - mp.polylog(2, xb)
                tail_terms += 1
                if abs(xa) < tol and abs(xb) < tol:
                    break
                k += 1
            W *= pref / (N * N)
            ell += mult * (W if N == 1 else -c8 * W)
        crossings[zk] = {key[1:]: v for key, v in counters.items() if key[0] == zk and (v[0] or v[1])}
    ell = SE.a_norm() * ell / 3
    E = -c42 - ell
    psihat = fr["psi1F"] / mp.pi
    J = psihat * E
    diag.update(tauF=fr["tauF"], qC=qC, abs_qC=aq, z1=fr["z1"], z2=fr["z2"], z3=fr["z3"], psi1F=fr["psi1F"], psihat1=psihat,
                C42=c42, ell_block=ell, E=E, J=J, tail_terms=tail_terms, crossings=crossings, disc=disc, delta=delta, zkeys=zkeys, dps_hi=dps_hi)
    return E, J


def words_q_vs_li2_control(t, masses):
    """Control (b): the dilogarithm form of W(z,N) vs the served q-series at a Euclidean t."""
    m1, m2, m3 = masses
    cv = SE.fcurve(mp.mpf(t), m1, m2, m3)
    qC = SE.nome_qC(cv["tauF"])
    Nq = SE.Nq_for(qC)
    pref = SE.g3_pref()
    w = mp.exp(two_pi_i() * cv["z1F"])
    worst = 0
    for N in (1, 2):
        Wq = SE.I1_g3(cv["z1F"], N, qC, Nq)
        Wl = mp.mpc(0)
        k = 1
        tol = mp.mpf(10) ** (-(mp.mp.dps + 10))
        while True:
            qk = qC ** (N * k)
            Wl += mp.polylog(2, w * qk) - mp.polylog(2, qk / w)
            if abs(qk) < tol:
                break
            k += 1
        Wl *= pref / (N * N)
        d = abs(Wq - Wl) / abs(Wq)
        dig = int(mp.floor(-mp.log10(d))) if d > 0 else mp.mp.dps
        worst = dig if worst == 0 else min(worst, dig)
    return worst


def agree_digits(a, b):
    a, b = mp.mpc(a), mp.mpc(b)
    d = abs(a - b)
    s = max(abs(a), abs(b))
    if d == 0:
        return mp.mp.dps
    return float(-mp.log10(d / s))


if __name__ == "__main__":
    import argparse
    ap = argparse.ArgumentParser(description="continued unequal-mass sunrise eps^0 evaluator (see docstring)")
    ap.add_argument("--mass", default="112", choices=["112", "123", "114"])
    ap.add_argument("--point", action="append", required=True, help="t (rational or decimal); physical t > 0 continued with +i0")
    ap.add_argument("--dps", type=int, default=60)
    ap.add_argument("--nsteps", type=int, default=400)
    ap.add_argument("--h", default="4")
    ap.add_argument("--mutate", action="store_true")
    ap.add_argument("--selftest", action="store_true")
    ap.add_argument("--check2n", action="store_true", help="repeat the tracking with 2n steps and compare")
    args = ap.parse_args()
    mp.mp.dps = args.dps + 10
    msq = {"112": [1, 1, 2], "123": [1, 2, 3], "114": [1, 1, 4]}[args.mass]
    masses = SE.masses_from_sq(msq)
    if args.selftest:
        for tc in (-1, -3):
            E0 = SE.predict_E(mp.mpf(tc), masses)
            E1, J1 = evaluate_continued(tc, masses, nsteps=args.nsteps, h=mp.mpf(args.h))
            J0 = SE.predict_J(mp.mpf(tc), masses)
            print(f"[control a] t={tc}: E continued vs served agree {agree_digits(E0, E1):.1f} d; J {agree_digits(J0, J1):.1f} d; |Im E| = {mp.nstr(abs(mp.im(E1)), 3)}")
        print(f"[control b] t=-1: dilog form of W vs served q-series agree {words_q_vs_li2_control(-1, masses)} d")
    for ps in args.point:
        T = mp.mpf(mp.mpmathify(ps)) if "/" not in ps else mp.mpf(ps.split("/")[0]) / mp.mpf(ps.split("/")[1])
        d = {}
        E, J = evaluate_continued(T, masses, nsteps=args.nsteps, h=mp.mpf(args.h), diag=d, mutate=args.mutate)
        if args.check2n:
            d2 = {}
            E2, J2 = evaluate_continued(T, masses, nsteps=2 * args.nsteps, h=mp.mpf(args.h), diag=d2, mutate=args.mutate)
        print(f"t = {ps} (mass {args.mass}): |q_C| = {mp.nstr(d['abs_qC'], 6)}  tau_F = {mp.nstr(d['tauF'], 12)}  z1 = {mp.nstr(d['z1'], 12)}  z3 = {mp.nstr(d['z3'], 12)}  psi1F = {mp.nstr(d['psi1F'], 12)}  disc = {d['disc']}  crossings = {d['crossings']}  maxjump = {d['maxjump']:.3g}  accepted steps = {d['n_accepted_steps']}  ambiguity = {d['ambiguity']}")
        print(f"  E^(0) = {mp.nstr(E, args.dps)}")
        print(f"  J^(0) = {mp.nstr(J, args.dps)}")
        if args.check2n:
            print(f"  step-count check (n = {args.nsteps} vs {2 * args.nsteps}): E {agree_digits(E, E2):.1f} d, J {agree_digits(J, J2):.1f} d; disc equal {d['disc'] == d2['disc']}; counters equal {d['crossings'] == d2['crossings']}")


def evaluate_multi(T, masses, dps_list, t0=-1, h=4, nsteps=300, dps_track=30, mutate=False):
    """ONE low-precision tracking pass (evaluate_continued at the smallest dps records the discrete data in diag), then the
    endpoint at each working precision in dps_list with the SAME discrete data (the branch of a principal-branch function
    at T + i delta does not depend on the precision).  Returns {dps: (E, J, diag)}."""
    out = {}
    first = True
    disc = counters = None
    for dps in sorted(dps_list):
        with mp.workdps(dps + 10):
            d = {}
            if first:
                E, J = evaluate_continued(T, masses, t0=t0, h=h, nsteps=nsteps, diag=d, mutate=mutate, dps_track=dps_track)
                disc, counters = d["disc"], d["_counters"]
                first = False
            else:
                E, J = endpoint_with(T, masses, disc, counters, d, mutate=mutate, zkeys=out[min(out)][2]["zkeys"], kmax_track=out[min(out)][2]["kmax_track"])
            out[dps] = (E, J, d)
    return out


# ---------------------------------------------------------------------------------------------
# ODE-transported frame (frame_ode.FrameODE): the words assembled on the transported frame, Li_2 cut crossings tracked along
# the SAME path.  No candidate family anywhere: the transported values ARE the continuation (to the integrator's accuracy).
# ---------------------------------------------------------------------------------------------
def load_frame_ode():
    spec = importlib.util.spec_from_file_location("frame_ode", os.path.join(HERE, "frame_ode.py"))
    m = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(m)
    return m


def evaluate_ode(T, masses, t0=-1, h=4, nsteps=2400, diag=None, mutate=False, kmax_track=12, side=1):
    """E^(0)(T + i0), J^(0)(T + i0) on the ODE-transported frame (frame_ode.py), Li_2 crossings tracked along the path."""
    fo = load_frame_ode()
    ode = fo.FrameODE(masses)
    st, path = ode.transport(t0, T, h_arc=h, nsteps=nsteps, want_path=True, delta=side * mp.mpf(10) ** (-(mp.mp.dps - 5)))
    prev = path[0]
    zkeys = ["z1", "z3"] if abs(masses[0] - masses[1]) == 0 else ["z1", "z2", "z3"]
    mults = {"z1": (2 if len(zkeys) == 2 else 1), "z2": 1, "z3": 1}
    tr = {}
    # the anchor frame for the trackers: the served (real) frame at t0
    anc = euclidean_anchor(t0, masses)
    for zk in zkeys:
        for N in (1, 2):
            for k, (xa, xb) in enumerate(word_terms(anc[zk], anc["tauF"], N, kmax_track)):
                if N == 2 and k == 0:
                    continue
                tr[(zk, N, k, 0)] = TrackedLi2(xa)
                tr[(zk, N, k, 1)] = TrackedLi2(xb)
    for cur in path:
        for zk in zkeys:
            for N in (1, 2):
                for k, (xa, xb) in enumerate(word_terms(cur[zk], cur["tauF"], N, kmax_track)):
                    if N == 2 and k == 0:
                        continue
                    tr[(zk, N, k, 0)].step(xa)
                    tr[(zk, N, k, 1)].step(xb)
    counters = {key: (o.n, o.m) for key, o in tr.items()}
    fr = dict(st)
    d = {} if diag is None else diag
    d["_counters"] = counters; d["kmax_track"] = kmax_track; d["nsteps"] = nsteps; d["h"] = h; d["t0"] = t0
    d["lattice_residual"] = st["lattice_residual"]; d["z12_residual"] = st["z12_residual"]
    return endpoint_on_frame(fr, masses, counters, d, mutate=mutate, zkeys=zkeys, kmax_track=kmax_track)


def endpoint_on_frame(fr, masses, counters, diag, mutate=False, zkeys=("z1", "z3"), kmax_track=12):
    """the words on a GIVEN frame (tauF, z1, z2, z3, psi1F) with recorded Li_2 counters."""
    dps_hi = mp.mp.dps
    zkeys = list(zkeys)
    mults = {"z1": (2 if len(zkeys) == 2 else 1), "z2": 1, "z3": 1}
    qC = -mp.exp(I * mp.pi * fr["tauF"])
    aq = abs(qC)
    if aq >= SE.QMAX:
        raise ValueError(f"|q_C| = {mp.nstr(aq, 6)} >= QMAX {mp.nstr(SE.QMAX, 4)} (series-domain limit)")
    if mp.im(fr["tauF"]) <= 0:
        raise ValueError(f"tau_F not in the upper half plane: {fr['tauF']}")
    pref = SE.g3_pref()
    c8 = 8 * (1 + mp.mpf(10) ** -12) if mutate else 8
    c42 = mp.mpc(0); ell = mp.mpc(0); tail_terms = 0; crossings = {}
    tol = mp.mpf(10) ** (-(dps_hi + 10))

    def li2_cont(x, n, m):
        lg = mp.log(x) + two_pi_i() * m
        return mp.polylog(2, x) + two_pi_i() * n * lg

    for zk in zkeys:
        mult = mults[zk]
        w = mp.exp(two_pi_i() * fr[zk])
        terms = word_terms(fr[zk], fr["tauF"], 1, kmax_track)
        terms2 = word_terms(fr[zk], fr["tauF"], 2, kmax_track)
        xa, xb = terms[0]
        na, ma = counters[(zk, 1, 0, 0)]; nb, mb = counters[(zk, 1, 0, 1)]
        c42 += mult * (li2_cont(xa, na, ma) - li2_cont(xb, nb, mb)) / (2 * I)
        for N, tl in ((1, terms), (2, terms2)):
            W = mp.mpc(0)
            for k in range(1, kmax_track + 1):
                xa, xb = tl[k]
                na, ma = counters[(zk, N, k, 0)]; nb, mb = counters[(zk, N, k, 1)]
                W += li2_cont(xa, na, ma) - li2_cont(xb, nb, mb)
            k = kmax_track + 1
            while True:
                qk = qC ** (N * k)
                xa, xb = w * qk, qk / w
                if abs(xa) >= mp.mpf("0.5") or abs(xb) >= mp.mpf("0.5"):
                    raise ValueError(f"untracked term with |x| >= 1/2 at k = {k} (raise kmax_track)")
                W += mp.polylog(2, xa) - mp.polylog(2, xb)
                tail_terms += 1
                if abs(xa) < tol and abs(xb) < tol:
                    break
                k += 1
            W *= pref / (N * N)
            ell += mult * (W if N == 1 else -c8 * W)
        crossings[zk] = {key[1:]: v for key, v in counters.items() if key[0] == zk and (v[0] or v[1])}
    ell = SE.a_norm() * ell / 3
    E = -c42 - ell
    psihat = fr["psi1F"] / mp.pi
    J = psihat * E
    diag.update(tauF=fr["tauF"], qC=qC, abs_qC=aq, z1=fr["z1"], z2=fr["z2"], z3=fr["z3"], psi1F=fr["psi1F"], psihat1=psihat,
                C42=c42, ell_block=ell, E=E, J=J, tail_terms=tail_terms, crossings=crossings, zkeys=zkeys, dps_hi=dps_hi)
    return E, J


# ---------------------------------------------------------------------------------------------
# EXACT endpoint from the transported frame: the transported periods and incomplete integrals are matched to INTEGER
# combinations of the principal-branch ones (every continuation of a period is an integer combination of the two basis
# periods; every continuation of F(phi|m) is +-F_p + 2a K_p + 2b i K_p'), the integers are read off at the tracking precision
# and the endpoint is then evaluated at ANY precision from the principal-branch values with those integers.
# ---------------------------------------------------------------------------------------------
def _solve_int2(target, w1, w2, tol):
    """target ~ a w1 + b w2 with integers a, b (complex numbers as 2-vectors); returns (a, b, residual) or None."""
    A = mp.matrix([[mp.re(w1), mp.re(w2)], [mp.im(w1), mp.im(w2)]])
    rhs = mp.matrix([mp.re(target), mp.im(target)])
    try:
        x = mp.lu_solve(A, rhs)
    except ZeroDivisionError:
        return None
    a, b = int(mp.nint(x[0])), int(mp.nint(x[1]))
    res = abs(target - (a * w1 + b * w2)) / max(abs(w1), abs(w2))
    off = max(abs(x[0] - a), abs(x[1] - b))
    return (a, b, float(res)) if (res < tol and off < mp.mpf("1e-4")) else None


def match_endpoint(st, t_end, masses, tol=None):
    """integers relating the transported frame `st` (tracking precision) to the principal-branch integrals at t_end."""
    if tol is None:
        tol = mp.mpf(10) ** (-6)   # the integers are separated by 1; the transport is accurate to ~1e-9 at dps 30 / 1200 steps
    fo = load_frame_ode()
    m1, m2, m3 = masses
    f = fo.alg_frame(t_end, m1, m2, m3)
    m, mq = f["m"], f["mp"]
    Kp_, Kcp_ = mp.ellipk(m), mp.ellipk(1 - m)
    Kq_, Kqc_ = mp.ellipk(mq), mp.ellipk(1 - mq)
    out = {}
    r = _solve_int2(st["K"], Kp_, I * Kcp_, tol); assert r, "K match failed"
    out["K"] = r[:2]
    r = _solve_int2(I * st["Kc"], Kp_, I * Kcp_, tol); assert r, "Kc match failed"
    out["iKc"] = r[:2]
    r = _solve_int2(st["Kp"], Kq_, I * Kqc_, tol); assert r, "Kp match failed"
    out["Kp"] = r[:2]
    r = _solve_int2(I * st["Kpc"], Kq_, I * Kqc_, tol); assert r, "Kpc match failed"
    out["iKpc"] = r[:2]
    # incomplete integrals: F_j = 2 Kp z_j (transported) = s F_p + 2a Kq_ + 2b i Kqc_
    for j, key in enumerate(("z1", "z2", "z3")):
        u = mp.sqrt(f["u2"][j]); Fp_ = mp.ellipf(mp.asin(u), mq)
        Fode = st[key] * 2 * st["Kp"]
        found = None
        for s_ in (1, -1):
            r = _solve_int2(Fode - s_ * Fp_, 2 * Kq_, 2 * I * Kqc_, tol)
            if r:
                found = (s_, r[0], r[1], r[2]); break
        assert found, f"F match failed for {key}"
        out[key] = found[:3]
    # psi1F sign relative to the principal sqrt(Z3)
    out["sqZ3_sign"] = 1 if abs(mp.sqrt(mp.mpc(f["Z3"])) - (2 * st["K"] / st["psi1F"])) < abs(-mp.sqrt(mp.mpc(f["Z3"])) - (2 * st["K"] / st["psi1F"])) else -1
    out["tol"] = tol
    return out


def frame_from_match(t_end, masses, mt):
    """the exact frame at t_end at the ACTIVE precision from the principal-branch integrals and the matched integers."""
    fo = load_frame_ode()
    m1, m2, m3 = masses
    f = fo.alg_frame(t_end, m1, m2, m3)
    m, mq = f["m"], f["mp"]
    Kp_, Kcp_ = mp.ellipk(m), mp.ellipk(1 - m)
    Kq_, Kqc_ = mp.ellipk(mq), mp.ellipk(1 - mq)
    a, b = mt["K"]; K = a * Kp_ + b * I * Kcp_
    a, b = mt["iKc"]; iKc = a * Kp_ + b * I * Kcp_
    a, b = mt["Kp"]; Kq = a * Kq_ + b * I * Kqc_
    z = {}
    for j, key in enumerate(("z1", "z2", "z3")):
        s_, a, b = mt[key]
        u = mp.sqrt(f["u2"][j]); Fp_ = mp.ellipf(mp.asin(u), mq)
        z[key] = (s_ * Fp_ + 2 * a * Kq_ + 2 * b * I * Kqc_) / (2 * Kq)
    psi = 2 * K / (mt["sqZ3_sign"] * mp.sqrt(mp.mpc(f["Z3"])))
    return dict(tauF=iKc / K, psi1F=psi, z1=z["z1"], z2=z["z2"], z3=z["z3"], K=K, Kp=Kq,
                lattice_residual=abs(z["z1"] + z["z2"] + z["z3"] - 1), z12_residual=abs(z["z1"] - z["z2"]))


def evaluate_ode_hp(T, masses_sq, dps_list, t0=-1, h=4, nsteps=1200, dps_track=30, mutate=False, kmax_track=12, side=1):
    """ONE transport at dps_track (frame + Li_2 counters), the integer match, then the endpoint at each dps in dps_list.
    masses_sq = the EXACT squared masses (ints/rationals); the masses sqrt(m_i^2) are formed inside every precision context
    (a mass formed at a lower precision would cap every digit count at that precision -- the 15-digit sqrt(2) footgun)."""
    with mp.workdps(max(dps_list) + 20):
        T = mp.mpf(T)
    h = side * abs(mp.mpf(h))   # side = -1: the LOWER half-plane path, the -i0 side (Schwarz-reflection control)
    with mp.workdps(dps_track):
        masses = [mp.sqrt(mp.mpf(x)) for x in masses_sq]
        d0 = {}
        E0, J0 = evaluate_ode(T, masses, t0=t0, h=h, nsteps=nsteps, diag=d0, mutate=mutate, kmax_track=kmax_track, side=side)
        counters = d0["_counters"]
        delta_tr = side * mp.mpf(10) ** (-(dps_track - 5))
        fo = load_frame_ode()
        ode = fo.FrameODE(masses)
        st, _ = ode.transport(t0, T, h_arc=h, nsteps=nsteps, delta=delta_tr)
        mt = match_endpoint(st, mp.mpc(T, delta_tr), masses)
        d0["match"] = {k: (list(v) if isinstance(v, tuple) else float(v)) for k, v in mt.items()}
    out = {}
    for dps in sorted(dps_list):
        with mp.workdps(dps + 10):
            masses = [mp.sqrt(mp.mpf(x)) for x in masses_sq]
            delta = side * mp.mpf(10) ** (-(dps + 5))
            fr = frame_from_match(mp.mpc(T, delta), masses, mt)
            d = {"match": d0["match"], "transport_J_dps%d" % dps_track: mp.nstr(J0, dps_track), "transport_lattice_residual": d0["lattice_residual"],
                 "exact_lattice_residual": fr["lattice_residual"], "exact_z12_residual": fr["z12_residual"], "_counters": counters, "nsteps": nsteps, "h": h, "t0": t0, "dps_track": dps_track}
            zk = d0["zkeys"]
            E, J = endpoint_on_frame(fr, masses, counters, d, mutate=mutate, zkeys=zk, kmax_track=kmax_track)
            out[dps] = (E, J, d)
    return out
