#!/usr/bin/env python3
r"""frame_ode.py -- the ANALYTIC CONTINUATION of the Feynman-curve frame by ODE TRANSPORT (no principal branch
anywhere on the path).  The frame quantities of the served sunrise_empl.fcurve are transported from the Euclidean anchor t0 along a
path t(u) in the upper half t-plane as the solution of their own exact differential equations:
  * complete integrals K(m), E(m) (and K(1-m), E(1-m)) in m = k_F^2(t):   dK/dm = (E - (1-m) K)/(2 m (1-m)),  dE/dm = (E - K)/(2 m);
  * the puncture-map integrals K(m'), E(m'), K(1-m'), E(1-m') in m' = k_F'^2(t);
  * the incomplete integrals F_j = F(phi_j | m'), E_j = E(phi_j | m') of the three punctures (DLMF 19.4.1-19.4.4):
        dF/dphi = 1/Delta,  dE/dphi = Delta,  Delta = sqrt(1 - m' sin^2 phi),
        dF/dm' = (E - (1-m') F)/(2 m' (1-m')) - sin phi cos phi/(2 (1-m') Delta),   dE/dm' = (E - F)/(2 m'),
    with sin phi_j = u_j(t) = sqrt((e1-e3)/(xhat_j-e3)), cos phi_j and Delta_j carried by sign continuity (their squares are
    single-valued algebraic functions of t), and d phi_j/dt = u_j'/cos phi_j;
  * the derivatives m'(t), m''(t), (u_j^2)'(t) by central finite differences of the single-valued algebraic frame (the roots e_iF
    are continuous along a path with Im t > 0: every sqrt(mu_i^2 - t) stays off its cut).
Classical RK4 with a fixed step in the path parameter.  The transported values are then MATCHED to the principal-branch
candidate families at the endpoint (the exact high-precision value = principal value + the matched discrete transformation);
a transported value with no candidate within the match tolerance is reported (the family is then incomplete by object).
z_j = F_j/(2 K(m')), tau_F = i K(1-m)/K(m), tau' = i K(1-m')/K(m'), psi1_F = 2 K(m)/sqrt(Z_3F) (sqrt by continuity).
CERTIFICATION inside the transport: z1 = z2 (m1 = m2 rows) and z1 + z2 + z3 = 1 are exact identities of the transported values;
their residuals at the endpoint measure the integration error.
extends: sunrise_empl.fcurve (the served frame), detransport (the house transport idea: fixed-step march of an exact DE along a
kinematic path; that package's loaders are for stored DE systems, so the elliptic-integral system here is written out).
"""
import mpmath as mp

I = mp.mpc(0, 1)


def alg_frame(t, m1, m2, m3, mu=1):
    """single-valued (rational in t and rad) algebraic frame; rad = the principal product, continuous for Im t > 0."""
    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)))
    e1 = (-t * t + 2 * M100 * t + Delta + rad) / (24 * mu4p)
    e2 = (-t * t + 2 * M100 * t + Delta - rad) / (24 * mu4p)
    e3 = (2 * t * t - 4 * M100 * t - 2 * Delta) / (24 * mu4p)
    Z1, Z2, Z3 = e3 - e2, e1 - e3, e1 - e2
    m = Z1 / Z3
    mp_ = -Z1 / Z2
    xh = [e3 + m2s * m3s / mu4p, e3 + m3s * m1s / mu4p, e3 + m1s * m2s / mu4p]
    u2 = [(e1 - e3) / (x - e3) for x in xh]
    return dict(m=m, mp=mp_, u2=u2, Z3=Z3, rad=rad)


def _cont_sqrt(x, ref):
    """sqrt of x on the branch continuous with ref (the sign closer to ref)."""
    r = mp.sqrt(x)
    return r if abs(r - ref) <= abs(-r - ref) else -r


class FrameODE:
    def __init__(self, masses, mu=1):
        self.m1, self.m2, self.m3 = [mp.mpf(x) for x in masses]
        self.mu = mu

    def rhs(self, t, y, dt_dir, signs):
        """dy/dt at t (complex); y = [K,E,Kc,Ec, Kp,Ep,Kpc,Epc, F1,E1,F2,E2,F3,E3]; signs = current (u_j, cos_j, Delta_j) refs."""
        h = mp.mpf(10) ** (-(mp.mp.dps // 3))
        f0 = alg_frame(t, self.m1, self.m2, self.m3, self.mu)
        # finite-difference stencil ALWAYS horizontal (t +- h): the derivative of an analytic function is direction-free and a
        # horizontal stencil never leaves the upper half plane (a vertical one at Im t = delta < h would cross the real-axis cuts)
        fp = alg_frame(t + h, self.m1, self.m2, self.m3, self.mu)
        fm = alg_frame(t - h, self.m1, self.m2, self.m3, self.mu)
        dm = (fp["m"] - fm["m"]) / (2 * h)
        dmp = (fp["mp"] - fm["mp"]) / (2 * h)
        du2 = [(fp["u2"][j] - fm["u2"][j]) / (2 * h) for j in range(3)]
        m, mq = f0["m"], f0["mp"]
        K, E, Kc, Ec, Kp, Ep, Kpc, Epc = y[:8]
        d = [None] * 14
        d[0] = (E - (1 - m) * K) / (2 * m * (1 - m)) * dm
        d[1] = (E - K) / (2 * m) * dm
        mc = 1 - m
        d[2] = -(Ec - (1 - mc) * Kc) / (2 * mc * (1 - mc)) * dm
        d[3] = -(Ec - Kc) / (2 * mc) * dm
        d[4] = (Ep - (1 - mq) * Kp) / (2 * mq * (1 - mq)) * dmp
        d[5] = (Ep - Kp) / (2 * mq) * dmp
        mqc = 1 - mq
        d[6] = -(Epc - (1 - mqc) * Kpc) / (2 * mqc * (1 - mqc)) * dmp
        d[7] = -(Epc - Kpc) / (2 * mqc) * dmp
        new_signs = []
        for j in range(3):
            u_ref, c_ref, D_ref = signs[j]
            u = _cont_sqrt(f0["u2"][j], u_ref)
            c = _cont_sqrt(1 - f0["u2"][j], c_ref)
            D = _cont_sqrt(1 - mq * f0["u2"][j], D_ref)
            du = du2[j] / (2 * u)
            dphi = du / c
            F, Ej = y[8 + 2 * j], y[9 + 2 * j]
            d[8 + 2 * j] = dphi / D + ((Ej - (1 - mq) * F) / (2 * mq * (1 - mq)) - u * c / (2 * (1 - mq) * D)) * dmp
            d[9 + 2 * j] = D * dphi + (Ej - F) / (2 * mq) * dmp
            new_signs.append((u, c, D))
        return d, new_signs

    def anchor(self, t0):
        """the served branches at the real Euclidean anchor (principal; z3 repaired by reflection when the lattice relation needs it)."""
        f = alg_frame(t0, self.m1, self.m2, self.m3, self.mu)
        m, mq = f["m"], f["mp"]
        K, E, Kc, Ec = mp.ellipk(m), mp.ellipe(m), mp.ellipk(1 - m), mp.ellipe(1 - m)
        Kp, Ep, Kpc, Epc = mp.ellipk(mq), mp.ellipe(mq), mp.ellipk(1 - mq), mp.ellipe(1 - mq)
        y = [K, E, Kc, Ec, Kp, Ep, Kpc, Epc]
        signs = []
        zs = []
        for j in range(3):
            u = mp.sqrt(f["u2"][j]); phi = mp.asin(u)
            F, Ej = mp.ellipf(phi, mq), mp.ellipe(phi, mq)
            c = mp.cos(phi); D = mp.sqrt(1 - mq * u * u)
            y += [F, Ej]; signs.append((u, c, D)); zs.append(F / (2 * Kp))
        if abs(zs[0] + zs[1] + zs[2] - 1) > mp.mpf(10) ** (-mp.mp.dps // 2):
            # the served branch repair z3 -> 1 - z1 - z2 = reflection phi3 -> pi - phi3 (F -> 2K - F, E -> 2E - E, cos -> -cos)
            y[12] = 2 * Kp - y[12]; y[13] = 2 * Ep - y[13]
            u, c, D = signs[2]; signs[2] = (u, -c, D)
            zs[2] = y[12] / (2 * Kp)
            assert abs(zs[0] + zs[1] + zs[2] - 1) < mp.mpf(10) ** (-mp.mp.dps // 2), "anchor lattice relation"
        psi = 2 * K / mp.sqrt(mp.mpc(f["Z3"]))
        assert mp.re(psi) > 0
        self.sqZ3_ref = mp.sqrt(mp.mpc(f["Z3"]))
        return y, signs

    def transport(self, t0, T, h_arc=4, nsteps=2000, delta=None, want_path=False):
        """RK4 along t(u) = t0 + (T-t0) u + i h sin(pi u), then a geometric descent at Re t = T down to delta."""
        if delta is None:
            delta = mp.mpf(10) ** (-(mp.mp.dps - 5))
        t0 = mp.mpf(t0); T = mp.mpf(T); h_arc = mp.mpf(h_arc)
        y, signs = self.anchor(t0)
        # path as a list of complex t (fine), RK4 with the segment's direction as the finite-difference direction
        pts = [mp.mpc(t0, 0)]
        for k in range(1, nsteps):
            u = mp.mpf(k) / nsteps
            pts.append(mp.mpc(t0 + (T - t0) * u, h_arc * mp.sin(mp.pi * u)))
        im = mp.im(pts[-1])
        sgn = 1 if delta > 0 else -1
        while abs(im) > abs(delta):
            im = im / 2
            pts.append(mp.mpc(T, sgn * max(abs(im), abs(delta))))
        if mp.im(pts[-1]) != delta:
            pts.append(mp.mpc(T, delta))
        path = []
        sq = self.sqZ3_ref
        for i in range(1, len(pts)):
            ta, tb = pts[i - 1], pts[i]
            dt = tb - ta; dirn = dt / abs(dt)
            k1, s1 = self.rhs(ta, y, dirn, signs)
            y2 = [y[j] + dt / 2 * k1[j] for j in range(14)]
            k2, s2 = self.rhs(ta + dt / 2, y2, dirn, s1)
            y3 = [y[j] + dt / 2 * k2[j] for j in range(14)]
            k3, s3 = self.rhs(ta + dt / 2, y3, dirn, s2)
            y4 = [y[j] + dt * k3[j] for j in range(14)]
            k4, s4 = self.rhs(tb, y4, dirn, s3)
            y = [y[j] + dt / 6 * (k1[j] + 2 * k2[j] + 2 * k3[j] + k4[j]) for j in range(14)]
            signs = s4
            f = alg_frame(tb, self.m1, self.m2, self.m3, self.mu)
            sq = _cont_sqrt(f["Z3"], sq)
            st = self.state(tb, y, sq)
            if want_path:
                path.append(st)
        self.sqZ3_ref = sq
        return st, path

    def state(self, t, y, sq):
        K, E, Kc, Ec, Kp, Ep, Kpc, Epc = y[:8]
        z = [y[8 + 2 * j] / (2 * Kp) for j in range(3)]
        return dict(t=t, tauF=I * Kc / K, taup=I * Kpc / Kp, psi1F=2 * K / sq, z1=z[0], z2=z[1], z3=z[2], K=K, Kc=Kc, Kp=Kp, Kpc=Kpc,
                    lattice_residual=abs(z[0] + z[1] + z[2] - 1), z12_residual=abs(z[0] - z[1]))
