"""eval_closed_form_offslice_compact.py -- compact S_3-orbit form of V(X;P) (conformally coupled one-loop triangle, general kinematics).

  V = (sqrt2 pi/8) Re[ -1/4 (Q/sqrt R_0) F0(X;P) - 1/2 sum_{e} (Q/sqrt R_e) Fe(pi_e X; pi_e P)
                       - sum_(ij) lam^-_ij Fm(X_i,X_j;P_i,P_j,P_k)/sqrt(T^-_ij) + sum_(ij) lam^+_ij Fp(X_i,X_j;P_i,P_j,P_k)/sqrt(T^+_ij) ]
  F0, Fe : sums of D(x) = Li2(x) - Li2(1/x) over cycles of Cayley-Menger cofactor letters (compact_representation.json)
  Fm, Fp : sums of 4 / 5 Bloch-Wigner-type blocks  P2(z) = Li2(z) - Li2(zbar) + 1/2 log(z zbar) log((1-z)/(1-zbar)),
           z = (p_a - w)/(p_b - w), zbar = (p_a + w)/(p_b + w), w = sqrt(T).
Branch prescription (region X_v >= P_v > 0, triangle inequalities for P):
  radicand > 0: all arguments real; Li2 -> Re Li2, log -> log|.| ;  radicand < 0: arguments come in complex-conjugate pairs, principal branches.
Exactly vanishing linear letters (X_i = P_i, P_k = |X_i - X_j|, zeros of face cofactors): every quantity is kept as (finite number) x prod letter^n,
  the result is assembled as a polynomial in Lambda = log|letter|; the Lambda-dependent coefficients are asserted to vanish and the constant term returned.
NOT handled (removable, codimension one; raises): a radicand exactly zero; X_k = X_i + X_j (simple poles of lam^± that cancel between sectors);
  two or more DISTINCT linear letters vanishing exactly (codimension >= 2, e.g. the full slice X = P: use eval_closed_form_general.py) -- the evaluator
  raises; an earlier version returned a wrong number there (factor ~100) although all log coefficients cancelled.
  Codimension-one handling is TESTED only on P_k = |X_i - X_j| (5 table points, 1e-50) and X_1 = P_1 (continuity at one point); exact-zero
  detection needs exact input (Fraction / 'p/q' strings): with floating-point input a letter that should vanish is treated as finite.
usage: python eval_closed_form_offslice_compact.py X1 X2 X3 P1 P2 P3 [dps]
Outside the region X_v >= P_v > 0, P a triangle, the evaluator raises OutOfRegion (for X_v < P_v use eval_closed_form_V_all.py)."""
import os, sys, json, itertools
from fractions import Fraction as Fr
import mpmath as mp
_here = os.path.dirname(os.path.abspath(__file__))
SPEC = json.load(open(os.path.join(_here, 'compact_representation.json')))
class Unhandled(Exception): pass

class LP:
    """polynomial of degree <= 2 in the symbols Lambda_l = log|l| of exactly vanishing letters l"""
    def __init__(s, c=0, lin=None, quad=None): s.c = c; s.lin = dict(lin or {}); s.quad = dict(quad or {})
    def __add__(s, o):
        o = o if isinstance(o, LP) else LP(o); r = LP(s.c + o.c, s.lin, s.quad)
        for k, v in o.lin.items(): r.lin[k] = r.lin.get(k, 0) + v
        for k, v in o.quad.items(): r.quad[k] = r.quad.get(k, 0) + v
        return r
    __radd__ = __add__
    def __neg__(s): return s*(-1)
    def __sub__(s, o): return s + (-(o if isinstance(o, LP) else LP(o)))
    def __mul__(s, o):
        if not isinstance(o, LP): return LP(s.c*o, {k: v*o for k, v in s.lin.items()}, {k: v*o for k, v in s.quad.items()})
        if (s.quad and (o.lin or o.quad)) or (o.quad and (s.lin or s.quad)): raise Unhandled('degree > 2 in Lambda')
        r = LP(s.c*o.c)
        for k, v in s.lin.items(): r.lin[k] = r.lin.get(k, 0) + v*o.c
        for k, v in o.lin.items(): r.lin[k] = r.lin.get(k, 0) + v*s.c
        for k, v in s.quad.items(): r.quad[k] = r.quad.get(k, 0) + v*o.c
        for k, v in o.quad.items(): r.quad[k] = r.quad.get(k, 0) + v*s.c
        for k1, v1 in s.lin.items():
            for k2, v2 in o.lin.items():
                k = tuple(sorted((k1, k2))); r.quad[k] = r.quad.get(k, 0) + v1*v2
        return r
    __rmul__ = __mul__

class Mono:
    """coef * prod_l l^n  with l exactly-zero letters (keys), coef finite non-zero"""
    def __init__(s, coef, exps=None): s.coef = coef; s.exps = {k: v for k, v in (exps or {}).items() if v}
    def __mul__(s, o):
        o = o if isinstance(o, Mono) else Mono(o); e = dict(s.exps)
        for k, v in o.exps.items(): e[k] = e.get(k, 0) + v
        return Mono(s.coef*o.coef, e)
    def inv(s): return Mono(1/s.coef, {k: -v for k, v in s.exps.items()})
    def __truediv__(s, o): return s*(o if isinstance(o, Mono) else Mono(o)).inv()
    def __pow__(s, n): return Mono(s.coef**n, {k: v*n for k, v in s.exps.items()})
    def kind(s):
        if not s.exps: return 'finite'
        sg = {v > 0 for v in s.exps.values()}
        if len(sg) > 1: raise Unhandled('argument with vanishing letters of opposite powers (codimension >= 2)')
        return 'zero' if sg.pop() else 'inf'
    def loglp(s): return LP(mp.log(abs(s.coef)), {k: v for k, v in s.exps.items()})

def _canon(vec):
    """linear form in the GLOBAL coordinates (X1,X2,X3,P1,P2,P3), sign-normalised -> hashable name"""
    lead = next((c for c in vec if c != 0), 1); sg = 1 if lead > 0 else -1
    return 'L' + str(tuple(str(Fr(c)*sg) for c in vec)), sg
def letter(val, vec):
    """exact value of a linear letter with global coefficient vector vec -> Mono (sign of the normalisation kept in coef)"""
    if val != 0: return Mono(_mp(val))
    name, sg = _canon(vec); return Mono(sg, {name: 1})
def _unit(k): return [Fr(1) if i == k else Fr(0) for i in range(6)]
def _vadd(*vs): return [sum(c) for c in zip(*vs)]
def _vscale(a, v): return [a*c for c in v]
def _mp(v): return mp.mpf(v.numerator)/v.denominator if isinstance(v, Fr) else mp.mpf(v)
def _exact(v):
    if isinstance(v, (int, Fr)): return Fr(v)
    if isinstance(v, str): return Fr(v)
    return v                                                    # mpf: exact-zero test is then floating point
def reLi2(m):
    k = m.kind()
    if k == 'zero': return LP(0)
    if k == 'inf':
        L = m.loglp(); return L*L*(-mp.mpf(1)/2) + (mp.pi**2/3 if m.coef > 0 else -mp.pi**2/6)
    return LP(mp.re(mp.polylog(2, m.coef)))
def Dreal(x):
    """Re[Li2(x) - Li2(1/x)] for real x (Mono)"""
    k = x.kind()
    if k == 'finite' and abs(x.coef) > 1: return -Dreal(x.inv())
    if k == 'inf': return -Dreal(x.inv())
    L = x.loglp(); kap = -mp.pi**2/3 if x.coef > 0 else mp.pi**2/6
    return reLi2(x)*2 + L*L*(mp.mpf(1)/2) + kap
def P2real(z, zb, omz, omzb):
    return reLi2(z) - reLi2(zb) + (z.loglp() + zb.loglp())*(omz.loglp() - omzb.loglp())*(mp.mpf(1)/2)

def _fr(s):
    n, d = (str(s).split('/') + ['1'])[:2]; return mp.mpf(int(n))/int(d)
def vertex_function(which, X, P, tag, pm=(0, 1, 2)):
    """which in ('F0','Fe'); X,P exact triples already relabelled.  returns (LP or complex) value of LS*F"""
    X1, X2, X3 = X; Q = X1 + X2 + X3; h = Fr(1, 2) if isinstance(Q, Fr) else mp.mpf(1)/2
    y = [(X3 - X1 - X2)*h, (X1 - X2 - X3)*h, (X2 - X1 - X3)*h] if which == 'F0' else [-Q*h, Q*h - X2, Q*h - X1]
    P1, P2, P3 = P
    dist = {(1, 2): P2, (1, 3): P1, (2, 3): P3, (1, 4): y[0], (2, 4): y[1], (3, 4): y[2]}
    dm = {k: _mp(v) for k, v in dist.items()}
    eX = [_unit(pm[i]) for i in range(3)]; eP = [_unit(3 + pm[i]) for i in range(3)]; hq = Fr(1, 2); eQ = _vadd(*eX)
    yv = [_vscale(hq, _vadd(eX[2], _vscale(-1, eX[0]), _vscale(-1, eX[1]))), _vscale(hq, _vadd(eX[0], _vscale(-1, eX[1]), _vscale(-1, eX[2]))), _vscale(hq, _vadd(eX[1], _vscale(-1, eX[0]), _vscale(-1, eX[2])))] if which == 'F0' \
        else [_vscale(-hq, eQ), _vadd(_vscale(hq, eQ), _vscale(-1, eX[1])), _vadd(_vscale(hq, eQ), _vscale(-1, eX[0]))]
    dvec = {(1, 2): eP[1], (1, 3): eP[0], (2, 3): eP[2], (1, 4): yv[0], (2, 4): yv[1], (3, 4): yv[2]}
    M = [[0, 1, 1, 1, 1]] + [[1] + [0 if a == b else dm[tuple(sorted((a, b)))]**2 for b in range(1, 5)] for a in range(1, 5)]
    Mm = mp.matrix(M); rad = -2*mp.det(Mm)
    def cof(i, j): return mp.det(mp.matrix([[M[r][c] for c in range(5) if c != j] for r in range(5) if r != i]))*(-1)**(i + j)
    if rad == 0: raise Unhandled('vertex radicand exactly zero')
    LS = _mp(Q)/mp.sqrt(mp.mpc(rad)); terms = SPEC['vertex'][which]
    if rad < 0:
        w = mp.sqrt(mp.mpc(rad)); tot = mp.mpc(0); Cd = {i: cof(i, i) for i in range(1, 5)}
        for c, cyc, eps in terms:
            x = mp.mpc(1)
            for e, s in zip(cyc, eps):
                e = tuple(e); rem = tuple(k for k in range(1, 5) if k not in e); x *= cof(*e) + s*dm[rem]*w
            vs = cyc[0] if (len(cyc) == 2 and cyc[0] == cyc[1]) else sorted({i for e in cyc for i in e})
            for i in vs: x /= Cd[i]
            tot += _fr(c)*(mp.polylog(2, x) - mp.polylog(2, 1/x))
        return LP(LS*tot)
    w = mp.sqrt(rad)
    Cdm = {}
    for v in range(1, 5):                                        # face cofactor C_vv = -(a+b+c)(-a+b+c)(a-b+c)(a+b-c), exact letters
        ks = [k for k in sorted(dist) if v not in k]; a, b, c = [dist[k] for k in ks]; va, vb, vc = [dvec[k] for k in ks]
        m = Mono(-1)
        for f, fv in [(a + b + c, _vadd(va, vb, vc)), (-a + b + c, _vadd(_vscale(-1, va), vb, vc)), (a - b + c, _vadd(va, _vscale(-1, vb), vc)), (a + b - c, _vadd(va, vb, _vscale(-1, vc)))]:
            m = m*letter(f, fv)
        Cdm[v] = m
        if not m.exps: assert abs(m.coef - cof(v, v)) <= mp.mpf(10)**(-mp.mp.dps + 8)*(1 + abs(m.coef)), 'face cofactor factorisation'
    U = {}
    for e in itertools.combinations(range(1, 5), 2):
        rem = tuple(k for k in range(1, 5) if k not in e); C = cof(*e); lw = dm[rem]*w
        prodm = Cdm[e[0]]*Cdm[e[1]]
        if C*lw >= 0: big = Mono(C + lw); U[e] = {1: big, -1: prodm/big}
        else: big = Mono(C - lw); U[e] = {-1: big, 1: prodm/big}
        if big.coef == 0: raise Unhandled('u^+ = u^- = 0')
    tot = LP(0)
    for c, cyc, eps in terms:
        x = Mono(mp.mpf(1))
        for e, s in zip(cyc, eps): x = x*U[tuple(e)][s]
        vs = cyc[0] if (len(cyc) == 2 and cyc[0] == cyc[1]) else sorted({i for e in cyc for i in e})
        for i in vs: x = x/Cdm[i]
        tot = tot + Dreal(x)*_fr(c)
    return tot*LS
def tube_function(typ, X, P, tag, pm=(0, 1, 2)):
    """canonical pair (1,2), third 3; X, P exact triples already relabelled.  returns LP value of LS*F"""
    d = SPEC['tube'][typ]; env = {'X1': X[0], 'X2': X[1], 'X3': X[2], 'P1': P[0], 'P2': P[1], 'P3': P[2]}; vec = [X[0], X[1], X[2], P[0], P[1], P[2]]
    sg = 1 if typ == 'p' else -1; m = P[0]**2 + P[1]**2 - P[2]**2
    T = _mp(m*m - 4*P[0]**2*P[1]**2 + 4*P[0]**2*X[1]**2 + 4*P[1]**2*X[0]**2 + 4*sg*X[0]*X[1]*m)
    if T == 0: raise Unhandled('tube radicand exactly zero')
    Q = X[0] + X[1] + X[2]
    den = (X[0] - X[1] - X[2])*(X[0] - X[1] + X[2]) if typ == 'm' else (X[0] + X[1] - X[2])
    if den == 0: raise Unhandled('X_k = X_i + X_j: removable pole of the leading singularity')
    lam = _mp(X[2]*Q/den) if typ == 'm' else _mp(X[2]/den)
    p = {a: _mp(eval(e, {}, env)) for a, e in d['p'].items()}
    def mono_of(rec):
        mo = Mono(_fr(rec['const']))
        for cv, ex in rec['letters']:
            val = sum(c_*v_ for c_, v_ in zip(cv, vec)); gv = [Fr(0)]*6
            for loc_, c_ in enumerate(cv): gv[(pm[loc_] if loc_ < 3 else 3 + pm[loc_ - 3])] += Fr(c_)
            mo = mo*letter(val, gv)**ex
        return mo
    if T < 0:
        w = mp.sqrt(mp.mpc(T)); tot = mp.mpc(0)
        for a, b, c in d['terms']:
            pa, pb = p[str(a)], p[str(b)]
            if pa == pb: raise Unhandled('z = 1 with complex root')
            z = (pa - w)/(pb - w); zb = (pa + w)/(pb + w)
            tot += c*(mp.polylog(2, z) - mp.polylog(2, zb) + mp.log(z*zb)*mp.log((1 - z)/(1 - zb))/2)
        return LP(lam/w*tot)
    w = mp.sqrt(T); S = {}
    for a in p:
        N = mono_of(d['N'][a])
        if not N.exps: assert abs(N.coef - (p[a]**2 - T)) <= mp.mpf(10)**(-mp.mp.dps + 8)*(1 + abs(N.coef)), 'norm factorisation'
        if p[a] > 0: plus = Mono(p[a] + w); S[a] = {1: plus, -1: N/plus}
        else: minus = Mono(p[a] - w); S[a] = {-1: minus, 1: N/minus}
    tot = LP(0)
    for a, b, c in d['terms']:
        a_, b_ = str(a), str(b); Dm = mono_of(d['D']['%d,%d' % (a, b)])
        z = S[a_][-1]/S[b_][-1]; zb = S[a_][1]/S[b_][1]
        tot = tot + P2real(z, zb, Dm/S[b_][-1], Dm/S[b_][1])*c
    return tot*(lam/w)
class OutOfRegion(ValueError):
    """The input lies outside the region in which this evaluator is valid."""
def _g(v):
    if isinstance(v, (int, str)) or hasattr(v, 'numerator'):
        from fractions import Fraction as _F
        f = _F(v); return mp.mpf(f.numerator)/f.denominator
    return mp.mpf(v)
def check_region(X, P):
    X = [_g(x) for x in X]; P = [_g(p) for p in P]
    if len(X) != 3 or len(P) != 3 or min(P) <= 0: raise OutOfRegion('need three site energies and three positive momenta')
    if 2*max(P) > sum(P): raise OutOfRegion('P = %s violates the triangle inequality' % [mp.nstr(p, 8) for p in P])
    low = [v + 1 for v in range(3) if X[v] < P[v]]
    if low: raise OutOfRegion('X_v < P_v at site(s) %s: this evaluator is valid for X_v >= P_v only; use eval_closed_form_V_all.py' % low)
def V_terms(X, P):
    check_region(X, P)
    X = [_exact(x) for x in X]; P = [_exact(p) for p in P]; out = {}
    out['F0'] = vertex_function('F0', X, P, 'F0')*(-mp.mpf(1)/4)
    for nm, pm in SPEC['orbits']['vertex_e'].items():
        out[nm] = vertex_function('Fe', [X[i] for i in pm], [P[i] for i in pm], nm, tuple(pm))*(-mp.mpf(1)/2)
    for ij, pm in SPEC['orbits']['tube'].items():
        out['t%sm' % ij] = tube_function('m', [X[i] for i in pm], [P[i] for i in pm], 't%sm' % ij, tuple(pm))*(-1)
        out['t%sp' % ij] = tube_function('p', [X[i] for i in pm], [P[i] for i in pm], 't%sp' % ij, tuple(pm))
    return out
def V_offslice_compact(X, P, info=None):
    tot = LP(0)
    for v in V_terms(X, P).values(): tot = tot + v
    forms = sorted({str(k) for k in tot.lin} | {str(x) for k in tot.quad for x in k})
    if len(forms) > 1: raise Unhandled('%d distinct linear letters vanish exactly (codimension >= 2): the limit is not determined by this evaluator; on the full slice X = P use eval_closed_form_general.py. forms: %s' % (len(forms), forms[:4]))
    tol = mp.mpf(10)**(-mp.mp.dps + 12)*(1 + abs(tot.c))
    bad = {str(k): mp.nstr(v, 5) for k, v in list(tot.lin.items()) + list(tot.quad.items()) if abs(v) > tol}
    if info is not None: info.update({'vanishing_letters': sorted({str(k) for k in tot.lin} | {str(x) for k in tot.quad for x in k}), 'max_residual_log_coefficient': float(max([abs(v) for v in list(tot.lin.values()) + list(tot.quad.values())] or [0]))})
    if bad: raise Unhandled('log-divergent coefficients do not cancel: %s' % bad)
    return mp.sqrt(2)*mp.pi/8*mp.re(tot.c)
if __name__ == '__main__':
    mp.mp.dps = int(sys.argv[7]) if len(sys.argv) > 7 else 40
    print(mp.nstr(V_offslice_compact(sys.argv[1:4], sys.argv[4:7]), mp.mp.dps - 5))
