"""Check that the operator L_0 annihilates the closed form of V on the line X = (2 lambda, lambda, 1).
Usage: python check_operator_L0.py [lambda]     (default 7/10; 1/3 < lambda < 1)
The derivatives are taken numerically from a polynomial interpolation, so the residual is limited to about 1e-25."""
import sys, os, json
from fractions import Fraction as Fr
import mpmath as mp
HERE = os.path.dirname(os.path.abspath(__file__)); sys.path.insert(0, os.path.join(HERE, ".."))
import eval_closed_form_general as ecg
mp.mp.dps = 120
C = json.load(open(os.path.join(HERE, "operator_L0.json")))["coeffs"]
def V(l): return ecg.V_general([2*l, l, mp.mpf(1)])
def derivs(f, s0, h=mp.mpf("1e-3"), N=16, nd=20):
    ks = list(range(-N, N+1)); ys = [f(s0 + k*h) for k in ks]
    A = mp.matrix([[mp.mpf(k)**n for n in range(2*N+1)] for k in ks]); co = mp.lu_solve(A, mp.matrix(ys))
    return [co[n]*mp.factorial(n)/h**n for n in range(nd+1)]
def cval(c, l): return sum(mp.mpf(Fr(x).numerator)/Fr(x).denominator*l**i for i, x in enumerate(c))
def residual(f, l):
    d = derivs(f, l); t = [cval(C[i], l)*d[i] for i in range(len(C))]
    return abs(sum(t))/sum(abs(x) for x in t)
if __name__ == "__main__":
    l = mp.mpf(Fr(sys.argv[1]).numerator)/Fr(sys.argv[1]).denominator if len(sys.argv) > 1 else mp.mpf(7)/10
    print("relative residual of L_0 on the closed form:", mp.nstr(residual(V, l), 3))
    print("control, closed form times (1 + 1e-6 lambda):", mp.nstr(residual(lambda x: V(x)*(1 + mp.mpf("1e-6")*x), l), 3))
