"""Threshold discontinuity D_0 of the one-leg d=3 box in closed form (Carlson R_F, R_J, R_C + logarithms).
Kinematics (one leg per site; momenta k_0..k_3, sum = 0):  P = |k_0|, x1 = |k_1|, x2 = |k_2|, x3 = |k_3|, S = |k_1 + k_2|, T = |k_0 + k_1|.
On the collapsed single-site tube (l on the focal segment, y_a = (P-d)/2, y_b = (P+d)/2, d in [-P, P]):
  y_c^2 = q_c(d) = ((P+d) x1^2 + (P-d) T^2)/(2P) - (P^2-d^2)/4,   y_d^2 = q_d(d) = ((P+d) S^2 + (P-d) x3^2)/(2P) - (P^2-d^2)/4      (Stewart)
  D_0 = (pi/4P) int_{-P}^{P} R dd,   R = sum_{sa,sb=+-1} sa sb chain(x1 + sa y_a, x2, x3 + sb y_b; y_c, y_d),
  chain(X1,X2,X3;yc,yd) = [1/(X1+X2+yd) + 1/(X2+X3+yc)] / [(X1+X2+X3)(X1+yc)(X2+yc+yd)(X3+yd)].
Closed form:  D_0 = (pi/16P) sum_{(alpha,sigma)} rho_{alpha,sigma} B(alpha; eta_c, eta_d) + (pi/P) p0 J_F   (the first-kind term is assembled in d0_eval.py),
  the sum running over the zeros alpha of the forms of R on the four sheets sigma = (+-y_c, +-y_d), rho = residue of R dd on that sheet, (eta_c, eta_d) the sheet values at alpha,
  B(alpha; eta_c, eta_d) = int_{-P}^{P} (1 + eta_c/y_c(d)) (1 + eta_d/y_d(d)) dd/(d - alpha) = L0 + eta_c Lc + eta_d Ld + eta_c eta_d Lcd   (finite parts with a common cutoff when alpha = +-P or alpha real inside)."""
import mpmath as mp, itertools

class OneLeg:
    def __init__(self, P, x1, x2, x3, S, T, dps=40):
        mp.mp.dps = dps; self.P, self.x1, self.x2, self.x3, self.S, self.T = [mp.mpf(v) for v in (P, x1, x2, x3, S, T)]
        P, x1, x3, S, T = self.P, self.x1, self.x3, self.S, self.T
        # q = (d^2 + 2 m d + n)/4
        self.mc = (x1**2 - T**2)/P; self.nc = 2*(x1**2 + T**2) - P**2; self.md = (S**2 - x3**2)/P; self.nd = 2*(S**2 + x3**2) - P**2
        rc = mp.sqrt(mp.mpc(self.mc**2 - self.nc)); rd = mp.sqrt(mp.mpc(self.md**2 - self.nd))
        self.e = [-self.mc + rc, -self.mc - rc, -self.md + rd, -self.md - rd]          # roots of q_c (e1,e2) and q_d (e3,e4);  Q = q_c q_d = prod(d - e_i)/16
    def qc(self, d): return (d*d + 2*self.mc*d + self.nc)/4
    def qd(self, d): return (d*d + 2*self.md*d + self.nd)/4
    def dqc(self, d): return (d + self.mc)/2
    def dqd(self, d): return (d + self.md)/2
    def forms(self, sa, sb, d, ec, ed):
        """values of the six forms and their d-derivatives on the sheet (eta_c, eta_d) = (ec, ed) (ec^2 = q_c(d) etc.)."""
        X1 = self.x1 + sa*(self.P - d)/2; X3 = self.x3 + sb*(self.P + d)/2; X2 = self.x2; dX1 = -mp.mpf(sa)/2; dX3 = mp.mpf(sb)/2
        dec = self.dqc(d)/ec if ec != 0 else mp.inf; ded = self.dqd(d)/ed if ed != 0 else mp.inf
        dec, ded = dec/2, ded/2
        L = dict(E=X1 + X2 + X3, L1=X1 + ec, L2=X2 + ec + ed, L3=X3 + ed, L4=X1 + X2 + ed, L5=X2 + X3 + ec)
        dL = dict(E=dX1 + dX3, L1=dX1 + dec, L2=dec + ded, L3=dX3 + ded, L4=dX1 + ded, L5=dX3 + dec)
        return L, dL
    def chain(self, L, skip=None):
        den = mp.mpf(1)
        for k in ('E', 'L1', 'L2', 'L3'):
            if k != skip: den *= L[k]
        if skip == 'L4': return 1/den
        if skip == 'L5': return 1/den
        return (1/L['L4'] + 1/L['L5'])/den
    def R(self, d, ec=None, ed=None):
        ec = mp.sqrt(self.qc(d)) if ec is None else ec; ed = mp.sqrt(self.qd(d)) if ed is None else ed
        return sum(sa*sb*self.chain(self.forms(sa, sb, d, ec, ed)[0]) for sa in (1, -1) for sb in (1, -1))
    def D0_numeric(self):
        return mp.pi/(4*self.P)*mp.quad(lambda d: self.R(d), mp.linspace(-self.P, self.P, 9))
    # ---------------- poles on the four sheets
    def poles(self, tol=None):
        tol = tol or mp.mpf(10)**(-mp.mp.dps + 12); P = self.P; out = []
        for sa, sb in itertools.product((1, -1), repeat=2):
            a = dict(E=(self.x1 + self.x2 + self.x3 + (sa + sb)*P/2, mp.mpf(sb - sa)/2, 0, 0), L1=(self.x1 + sa*P/2, -mp.mpf(sa)/2, 1, 0), L2=(self.x2, mp.mpf(0), 1, 1),
                     L3=(self.x3 + sb*P/2, mp.mpf(sb)/2, 0, 1), L4=(self.x1 + self.x2 + sa*P/2, -mp.mpf(sa)/2, 0, 1), L5=(self.x2 + self.x3 + sb*P/2, mp.mpf(sb)/2, 1, 0))
            for name, (a0, a1, bc, bd) in a.items():
                # norm polynomial in d (coefficients low -> high)
                A = [a0, a1]; A2 = [a0*a0, 2*a0*a1, a1*a1]; qc = [self.nc/4, self.mc/2, mp.mpf(1)/4]; qd = [self.nd/4, self.md/2, mp.mpf(1)/4]
                if bc == 0 and bd == 0: N = A
                elif bd == 0: N = [u - v for u, v in zip(A2, qc)]
                elif bc == 0: N = [u - v for u, v in zip(A2, qd)]
                else:
                    m = [u - v - w for u, v, w in zip(A2, qc, qd)]            # A^2 - qc - qd  (d^2 coefficient: a1^2 - 1/2)
                    def mul(p, q):
                        r = [mp.mpf(0)]*(len(p) + len(q) - 1)
                        for i, u in enumerate(p):
                            for j, v in enumerate(q): r[i+j] += u*v
                        return r
                    N = [u - 4*v for u, v in zip(mul(m, m), mul(qc, qd))]
                while len(N) > 1 and abs(N[-1]) < tol*max(abs(c) for c in N): N = N[:-1]
                if len(N) == 1: continue
                roots = mp.polyroots(N[::-1], maxsteps=200, extraprec=200) if len(N) > 2 else [-N[0]/N[1]]
                for al in roots:
                    for sc, sd in itertools.product((1, -1), repeat=2):
                        if (bc == 0 and False) or (bd == 0 and False): pass
                        ec = sc*mp.sqrt(self.qc(al)); ed = sd*mp.sqrt(self.qd(al)); L, dL = self.forms(sa, sb, al, ec, ed)
                        scale = abs(a0) + abs(a1*al) + bc*abs(ec) + bd*abs(ed)
                        if abs(L[name]) < mp.sqrt(tol)*scale:
                            rho = sa*sb*self.chain(L, skip=name)/dL[name]
                            out.append(dict(form=name, sa=sa, sb=sb, alpha=al, eta_c=ec, eta_d=ed, rho=rho))
        return out
    # ---------------- bracket, numerically (reference) and in closed form
    def B_numeric(self, al, ec, ed):
        f = lambda d: 0 if d == al else (1 + ec/mp.sqrt(self.qc(d)))*(1 + ed/mp.sqrt(self.qd(d)))/(d - al)      # removable point: the bracket integrand is finite at d = alpha
        pts = [-self.P, self.P]
        if abs(mp.im(al)) < mp.mpf(10)**-20 and -self.P < mp.re(al) < self.P: pts = [-self.P, mp.re(al), self.P]
        return mp.quad(f, mp.linspace(pts[0], pts[-1], 9) if len(pts) == 2 else sorted(set(list(mp.linspace(pts[0], pts[1], 5)) + list(mp.linspace(pts[1], pts[2], 5)))))

if __name__ == '__main__':
    import json
    K = OneLeg(1, mp.mpf(3)/2, 2, mp.mpf(7)/4, 2, mp.mpf(9)/5)          # configuration a of the monodromy families at P = 1
    ref = mp.mpf('0.000163821435978827307816034450631')*mp.pi; D = K.D0_numeric(); print('D0 numeric vs stored:', mp.nstr(abs(D - ref)/ref, 3))
    pl = K.poles(); print(len(pl), 'sheet poles')
    import collections; print(collections.Counter((p['form'], mp.nstr(p['alpha'], 8)) for p in pl))
    tot = sum(p['rho']*K.B_numeric(p['alpha'], p['eta_c'], p['eta_d']) for p in pl)*mp.pi/(16*K.P)
    print('sum rho B (numeric brackets) vs D0 (without the first-kind term; nonzero by construction):', mp.nstr(abs(tot - D)/D, 3), mp.nstr(tot, 20))
    # completeness: R_cd Q - sum c/(d-alpha) = 0 at a random point
    d0 = mp.mpf('0.37') + mp.mpf('0.21')*1j; yc, yd = mp.sqrt(K.qc(d0)), mp.sqrt(K.qd(d0))
    for nm, w in (('0', lambda a, b: 1), ('c', lambda a, b: a), ('d', lambda a, b: b), ('cd', lambda a, b: a*b)):
        lhs = sum(w(sc*yc, sd*yd)*K.R(d0, sc*yc, sd*yd) for sc in (1, -1) for sd in (1, -1))/4; rhs = sum(p['rho']*w(p['eta_c'], p['eta_d'])/(d0 - p['alpha']) for p in pl)/4
        print('character', nm, 'partial fractions (constant part p0 omitted; nonzero by construction for cd):', mp.nstr(abs(lhs - rhs)/abs(lhs), 3))
