"""eval_closed_form_general.py -- closed form of V(X; P = X) for general one-leg-per-site kinematics (non-degenerate triangle).

  V = (sqrt2 pi/8) [ sum_i X_i Li2(y_i) / (z1 z2 z3)  +  8 X1 X2 X3 W / (E_T z1^2 z2^2 z3^2) ],
  W = sum_i X_i z_i Li2(y_i) + 2 sum_{i<j} X_i X_j l_i l_j - (pi^2/12) sum_{i<j} z_i z_j,
  E_T = X1+X2+X3,  z_i = E_T - 2 X_i,  y_i = z_i/E_T,  l_i = log(2 X_i/E_T) = log(1 - y_i).

usage:  python eval_closed_form_general.py X1 X2 X3 [dps]          (rationals like 7/10 allowed)
        python eval_closed_form_general.py --check file.jsonl [...]  (records with "X":[..] and "V": compare)
        python eval_closed_form_general.py --table V_slice_table.csv (line X=(2lam,lam,1))
Near a degenerate triangle (z_i -> 0) the two terms cancel to O(1/z_i^2): raise dps accordingly.
Region: X_v > 0 and strict triangle inequalities. Outside it V_general raises OutOfRegion."""
import sys, json, csv
from fractions import Fraction
import mpmath as mp

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):
    X = [_g(x) for x in X]; E = sum(X)
    if len(X) != 3 or min(X) <= 0: raise OutOfRegion('need three positive site energies, got %s' % [mp.nstr(x, 8) for x in X])
    z = [E - 2*x for x in X]
    if min(z) < 0: raise OutOfRegion('X = %s is not a triangle: with one leg per site the X_v are the sides of the momentum triangle' % [mp.nstr(x, 8) for x in X])
    if min(z) == 0: raise OutOfRegion('degenerate triangle X_i = X_j + X_k: the formula has a finite limit but cannot be evaluated at the point itself')
def V_general(X):
    check_region(X)
    X = [mp.mpf(x) for x in X]
    E = sum(X); z = [E - 2*x for x in X]; y = [zi/E for zi in z]; l = [mp.log(2*x/E) for x in X]
    Li = [mp.polylog(2, yi) for yi in y]
    zz = z[0]*z[1]*z[2]
    W = sum(X[i]*z[i]*Li[i] for i in range(3)) + 2*(X[0]*X[1]*l[0]*l[1] + X[0]*X[2]*l[0]*l[2] + X[1]*X[2]*l[1]*l[2]) \
        - mp.pi**2/12*(z[0]*z[1] + z[0]*z[2] + z[1]*z[2])
    return mp.sqrt(2)*mp.pi/8*(sum(X[i]*Li[i] for i in range(3))/zz + 8*X[0]*X[1]*X[2]*W/(E*zz*zz))

def q(s):
    f = Fraction(s); return mp.mpf(f.numerator)/f.denominator

if __name__ == '__main__':
    a = sys.argv[1:]
    if a and a[0] == '--check':
        mp.mp.dps = 60
        for fn in a[1:]:
            for line in open(fn):
                r = json.loads(line)
                if 'V' not in r: continue
                v = V_general([q(s) for s in r['X']]); vn = mp.mpf(r['V'])
                print(r['tag'], r['X'], 'rel.diff', mp.nstr((v - vn)/vn, 3))
    elif a and a[0] == '--table':
        mp.mp.dps = 60
        rows = list(csv.DictReader(open(a[1])))
        res = sorted(abs(V_general([2*mp.mpf(r['lam']), mp.mpf(r['lam']), 1]) - mp.mpf(r['V'])) for r in rows)
        print('rows', len(rows), 'max', mp.nstr(res[-1], 3), 'median', mp.nstr(res[len(res)//2], 3))
    else:
        mp.mp.dps = int(a[3]) if len(a) > 3 else 50
        print(mp.nstr(V_general([q(s) for s in a[:3]]), mp.mp.dps - 5))
