"""Closed-form evaluator of the threshold discontinuity D_0 of the one-leg d = 3 box (general one-leg kinematics).
D_0 = (pi/4P) [ (1/4) sum_{sheet poles} rho * B(alpha; eta_c, eta_d) + 4 p0 J_F ],   B = L0 + eta_c Lc + eta_d Ld + eta_c eta_d Lcd,
 L0 = int dd/(d-alpha) (log),  Lc, Ld = int dd/((d-alpha) y_{c,d}) (log),  Lcd = 4 int dd/((d-alpha) s) = 4 (J_Pi^{(k)} - J_F)/(alpha - e_k)  (Carlson R_J, R_F + one logarithm),
 s^2 = prod (d - e_i) = 16 q_c q_d;  p0 = (1/4) sum rho eta_c eta_d/(alpha - e_1)  (first-kind coefficient, from the partial-fraction identity at the branch point e_1).
Poles at the end points alpha = +-P: finite parts with ONE common cutoff (the log(eps) terms cancel inside each bracket because (1 + eta/y(d)) vanishes there).
Branch rule for the single logarithm of the third-kind formula: signed S (DLMF 19.29.8 passes S^2 to R_C and loses the sign) and phase continued in the upper limit from the lower end."""
import sys, os; sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import mpmath as mp
from d0_closed import OneLeg
NS = 64
def _XY(e, y, x): return [mp.sqrt(x - ei) for ei in e], [mp.sqrt(y - ei) for ei in e]
def J_F1(e, y, x):
    X, Y = _XY(e, y, x); U = lambda i, j, k, l: (X[i]*X[j]*Y[k]*Y[l] + Y[i]*Y[j]*X[k]*X[l])/(x - y)
    return 2*mp.elliprf(U(0, 1, 2, 3)**2, U(0, 2, 1, 3)**2, U(0, 3, 1, 2)**2)
def _SQ(e, al, y, x, x52=None, y52=None):
    X, Y = _XY(e, y, x); X52 = x - al if x52 is None else x52; Y52 = y - al if y52 is None else y52
    U = lambda i, j, k, l: (X[i]*X[j]*Y[k]*Y[l] + Y[i]*Y[j]*X[k]*X[l])/(x - y); U12, U13, U14 = U(0, 1, 2, 3), U(0, 2, 1, 3), U(0, 3, 1, 2)
    U152 = U12**2 - (e[2] - e[0])*(e[3] - e[0])*(al - e[1])/(al - e[0])
    S = (X[1]*X[2]*X[3]/X[0]*(y - al) + Y[1]*Y[2]*Y[3]/Y[0]*(x - al))/(x - y)
    return S, X52*Y52*U152/(X[0]*Y[0])**2, (U12**2, U13**2, U14**2, U152)
def J_Pi1(e, al, y, x, fp=None, shift=0):
    """int_y^x (t - e_1)/((t - al) s(t)) dt;  fp = 'upper' (al = x) or 'lower' (al = y): finite part, i.e. the term  c*log(eps) of the cutoff integral dropped.  shift: planted branch error (integer)."""
    S, Q2, (u2, v2, w2, p) = _SQ(e, al, y, x, x52=(1 if fp == 'upper' else None), y52=(1 if fp == 'lower' else None))      # lower: Y52 = y - al = +eps -> 1.  fp: X52*Y52 = (-eps)(y-x) or (x-y)(eps) -> (x - y) * eps, eps dropped; sign kept: see below
    if fp == 'upper': Q2 = Q2*(-1)          # X52 = -eps: X52*Y52 = (-eps)(y - al) = eps (x - y) > 0;  we passed x52 = 1, y52 = y - x  ->  multiply by -1
    a = mp.mpf(2)/3*(e[1] - e[0])*(e[2] - e[0])*(e[3] - e[0])/(al - e[0])*mp.elliprj(u2, v2, w2, p)
    # tracked logarithm
    rprev = None; wprev = None; ph = mp.mpf(0)
    grid = [y + (x - y)*mp.mpf(j)/NS for j in range(1, NS + 1)]
    if fp == 'upper': grid = grid[:-1] + [x - (x - y)/NS*mp.mpf(2)**-m for m in range(1, 60)]          # approach the end-point pole geometrically: the direction of r converges
    for xx in grid:
        Sj, Qj, _ = _SQ(e, al, y, xx, y52=(1 if fp == 'lower' else None)); wj = Sj if fp == 'lower' else mp.sqrt(Sj*Sj - Qj)
        if wprev is not None and abs(wj - wprev) > abs(wj + wprev): wj = -wj
        r = 4*Sj*Sj/Qj if fp == 'lower' else (Sj + wj)/(Sj - wj)
        ph = mp.arg(r) if rprev is None else ph + mp.arg(r/rprev); rprev, wprev = r, wj
    if fp is None:
        w = mp.sqrt(S*S - Q2); w = -w if abs(w - wprev) > abs(w + wprev) else w; b = (mp.log(abs((S + w)/(S - w))) + 1j*(ph + 2*mp.pi*shift))/w
    else:
        r = 4*S*S/Q2
        if fp == 'upper': ph += mp.arg(r/rprev)          # last step: from the final interior sample to the limiting direction of r
        b = (mp.log(abs(r)) + 1j*(ph + 2*mp.pi*shift))/S
    return a + b
def _safe(e, al, y, x):
    """a-priori validity test of the one-piece formulas.  The Carlson arguments are U_ij^2; for two conjugate pairs the U_ij are REAL functions of the upper limit which start at +infinity
    for a short piece; if one of them passes through zero the continuation of R_F, R_J leaves the principal branch although U_ij^2 stays real positive.  Test: the SIGNED U_ij > 0 on the whole piece
    (16 samples), and Re p > 0 for the third-kind term."""
    for j in range(1, 17):
        xx = y + (x - y)*mp.mpf(j)/16; X, Y = _XY(e, y, xx); U = lambda i, jj, k, l: (X[i]*X[jj]*Y[k]*Y[l] + Y[i]*Y[jj]*X[k]*X[l])/(xx - y)
        us = (U(0, 1, 2, 3), U(0, 2, 1, 3), U(0, 3, 1, 2))
        if min(mp.re(u) for u in us) <= 0 or max(abs(mp.im(u)) for u in us) > abs(us[0])*mp.mpf(10)**-20: return False
        if al is not None and abs(al - xx) > 0 and abs(al - y) > 0:
            p = us[0]**2 - (e[2] - e[0])*(e[3] - e[0])*(al - e[1])/(al - e[0])
            if mp.re(p) <= 0: return False
    return True
def _pieces(e, al, y, x, depth=0):
    if depth >= 8 or _safe(e, al, y, x): return [(y, x)]
    m = (y + x)/2; return _pieces(e, al, y, m, depth + 1) + _pieces(e, al, m, x, depth + 1)
N_PIECES = []
def J_F(e, y, x):
    pc = _pieces(e, None, y, x); N_PIECES.append(len(pc)); return sum(J_F1(e, a, b) for a, b in pc)
def J_Pi(e, al, y, x, fp=None, shift=0):
    pc = _pieces(e, al, y, x); N_PIECES.append(len(pc)); tot = mp.mpf(0)
    for i, (a, b) in enumerate(pc):
        f = fp if ((fp == 'upper' and i == len(pc) - 1) or (fp == 'lower' and i == 0)) else None
        tot += J_Pi1(e, al, a, b, f, shift=(shift if i == 0 else 0))
    return tot
def L_q(K, which, al, fp=None):
    q, dq = (K.qc, K.dqc) if which == 'c' else (K.qd, K.dqd); P = K.P; c = q(al); bq = dq(al); rc = mp.sqrt(c)
    F = lambda t: (2*c + bq*(t - al) + 2*rc*mp.sqrt(q(t)))
    if fp is None:
        G = lambda t: F(t)/(t - al); prev = G(-P); ph = mp.mpf(0)
        for j in range(1, NS + 1):
            cur = G(-P + 2*P*mp.mpf(j)/NS); ph += mp.arg(cur/prev); prev = cur
        return -(mp.log(abs(G(P))/abs(G(-P))) + 1j*ph)/rc
    # end-point pole (al = +-P, c > 0 real): A(t) = -(1/rc) log|F(t)/(t - al)|, cutoff term dropped
    if fp == 'upper': return -(mp.log(4*c) - mp.log(abs(F(-P))/(2*P)))/rc
    return -(mp.log(abs(F(P))/(2*P)) - mp.log(4*c))/rc
def L_q_pv(K, which, al):
    """principal value of int_{-P}^{P} dt/((t - al) sqrt(q(t))) for real al inside the segment (q(al) > 0)."""
    q, dq = (K.qc, K.dqc) if which == 'c' else (K.qd, K.dqd); P = K.P; c = mp.re(q(al)); bq = mp.re(dq(al)); rc = mp.sqrt(c)
    G = lambda t: abs((2*c + bq*(t - al) + 2*rc*mp.sqrt(q(t)))/(t - al))
    return -(mp.log(G(P)) - mp.log(G(-P)))/rc
def D0_closed(K, pivot=3, shift_pole=None, drop_first_kind=False, pv_interior=True):
    P = K.P; e = K.e; ee = [e[pivot]] + [e[j] for j in range(4) if j != pivot]; pl = K.poles(); JF = J_F(ee, -P, P); tot = mp.mpf(0); cache = {}; flags = []
    for i, p in enumerate(pl):
        al = p['alpha']; fp = 'upper' if abs(al - P) < mp.mpf(10)**-25 else ('lower' if abs(al + P) < mp.mpf(10)**-25 else None)
        interior = fp is None and abs(mp.im(al)) < mp.mpf(10)**-25 and abs(mp.re(al)) < P
        if interior and not pv_interior: flags.append(('interior real pole, continued through (UNRELIABLE)', p['form'], mp.nstr(al, 10)))
        key = (mp.nstr(al, 30), fp)
        if interior and pv_interior and key not in cache:
            # real pole inside the segment (always on a non-physical sheet: the bracket vanishes there).  Principal values of all four pieces; the elliptic one by splitting at alpha:
            # PV = finite part on [-P, alpha] + finite part on [alpha, P]  (the dropped cutoff logarithms are equal and opposite).
            a_ = mp.re(al); G = lambda which: L_q_pv(K, which, a_)
            Jpv = J_Pi(ee, a_, -P, a_, 'upper') + J_Pi(ee, a_, a_, P, 'lower')
            cache[key] = (mp.log((P - a_)/(P + a_)), G('c'), G('d'), 4*(Jpv - JF)/(a_ - ee[0])); flags.append(('interior real pole, principal value', p['form'], mp.nstr(al, 10)))
        if key not in cache:
            al_ = P if fp == 'upper' else (-P if fp == 'lower' else al)
            L0 = -mp.log(2*P) if fp == 'upper' else (mp.log(2*P) if fp == 'lower' else mp.log(P - al_) - mp.log(-P - al_))
            cache[key] = (L0, L_q(K, 'c', al_, fp), L_q(K, 'd', al_, fp), 4*(J_Pi(ee, al_, -P, P, fp, shift=(1 if shift_pole == i else 0)) - JF)/(al_ - ee[0]))
        L0, Lc, Ld, Lcd = cache[key] if shift_pole != i else (cache[key][0], cache[key][1], cache[key][2], 4*(J_Pi(ee, al, -P, P, fp, shift=1) - JF)/(al - ee[0]))
        tot += p['rho']*(L0 + p['eta_c']*Lc + p['eta_d']*Ld + p['eta_c']*p['eta_d']*Lcd)
    p0 = sum(p['rho']*p['eta_c']*p['eta_d']/(p['alpha'] - e[0]) for p in pl)/4
    return mp.pi/(4*P)*(tot/4 + (0 if drop_first_kind else 4*p0*JF)), dict(n_poles=len(pl), n_distinct=len(cache), flags=flags, p0=p0, J_F=JF)
if __name__ == '__main__':
    mp.mp.dps = 40; K = OneLeg(1, mp.mpf(3)/2, 2, mp.mpf(7)/4, 2, mp.mpf(9)/5); ref = K.D0_numeric()
    for pv in range(4):
        v, info = D0_closed(K, pivot=pv); print('pivot', pv, 'closed vs quadrature:', mp.nstr(abs(v - ref)/abs(ref), 3), info['flags'])
