"""
verify_P241.py — Numerical checks for A241: Full Ô product basis
(Laguerre radial × Gegenbauer angular). 20 checks.
Copyright: Léon Fernando Vlegels, MIT.
"""
import numpy as np
from numpy.polynomial.legendre import leggauss

mp_pi = np.pi
PASS = 0; FAIL = 0

def check(name, condition, detail=""):
    global PASS, FAIL
    if condition:
        PASS += 1
        print(f"  [PASS] {PASS+FAIL:>2}. {name}")
    else:
        FAIL += 1
        print(f"  [FAIL] {PASS+FAIL:>2}. {name}" + (f"  [{detail}]" if detail else ""))

print("=" * 60)
print("verify_P241.py — A241 product basis checks")
print("=" * 60)

# ── Constants ────────────────────────────────────────────────────
pi = np.pi
OMEGA = 4*pi**3 + pi**2 + pi
alpha = 1.0/OMEGA
gamma = 3.0/4.0
kappa = OMEGA/3.0
N_ANG = 8
M_RAD = 6
xA238, xA240, xCZ = 0.25238, 0.18886, 0.01420

# C01: κ = Ω/3 numerically
check("C01: kappa = Omega/3",
      abs(kappa - OMEGA/3) < 1e-12,
      f"kappa={kappa:.8f}")

# C02: Ω = 4π³+π²+π ≈ 137.036
check("C02: Omega ≈ 137.036",
      abs(OMEGA - 137.0363037) < 1e-4,
      f"Omega={OMEGA:.8f}")

# ── GL quadrature helpers ────────────────────────────────────────
def gl01(n=2000):
    xi, wi = leggauss(n); x = 0.5*(xi+1); w = 0.5*wi; return x, w

def geg(x, N):
    U = np.zeros((N, len(x)))
    U[0] = np.ones(len(x))
    if N > 1: U[1] = 4*x - 2
    for n in range(2, N): U[n] = (4*x-2)*U[n-1] - U[n-2]
    return U

def build_P(N):
    cs = [None]*N
    c0 = np.zeros(N); c0[0] = 1.; cs[0] = c0
    if N > 1:
        c1 = np.zeros(N); c1[0] = -2.; c1[1] = 4.; cs[1] = c1
    for n in range(2, N):
        cn = np.zeros(N); p = cs[n-1]; q = cs[n-2]
        for k in range(N-1): cn[k+1] += 4*p[k]; cn[k] -= 2*p[k]
        cn -= q; cs[n] = cn
    return np.column_stack(cs)

def Ln_stable(n, t):
    if n == 0: return np.ones_like(t, dtype=float)
    if n == 1: return 1. - t
    L0 = np.ones_like(t, dtype=float); L1 = 1. - t
    for k in range(1, n):
        L2 = ((2*k+1-t)*L1 - k*L0)/(k+1); L0 = L1; L1 = L2
    return L1

x_gl, w_gl = gl01(2000)
wt = np.sqrt(x_gl*(1-x_gl))
U = geg(x_gl, N_ANG)
h = np.array([np.dot(U[n]**2*wt, w_gl) for n in range(N_ANG)])

# C03: Gegenbauer norms = π/8
check("C03: Gegenbauer norms = π/8",
      np.max(np.abs(h - pi/8)) < 1e-5,
      f"max_err={np.max(np.abs(h-pi/8)):.2e}")

# ── Build angular operators ──────────────────────────────────────
P = build_P(N_ANG); Pi = np.linalg.inv(P)
D2m = np.zeros((N_ANG, N_ANG))
for j in range(2, N_ANG): D2m[j-2, j] = 16*j**2*(j**2-1)
D2g = Pi @ D2m @ P
mu = np.array([-n*(n+2) for n in range(N_ANG)], dtype=float)
Dg = np.diag(mu); Tg = Dg @ D2g

rho_v = 16*pi**3*x_gl**3 + 3*pi**2*x_gl**2 + 2*pi*x_gl
rho_g = np.zeros((N_ANG, N_ANG))
for m in range(N_ANG):
    for n in range(N_ANG):
        rho_g[m, n] = np.dot(U[m]*rho_v*U[n]*wt, w_gl)/h[m]

# C04: ∫₀¹ρ(x)dx = Ω
integ_rho = np.dot(rho_v, w_gl)
check("C04: ∫₀¹ρ(x)dx = Ω",
      abs(integ_rho - OMEGA) < 0.01,
      f"int={integ_rho:.6f} Omega={OMEGA:.6f}")

# C05: α·ρ_geg is symmetric
check("C05: α·ρ_geg symmetric",
      np.max(np.abs(rho_g - rho_g.T)) < 1e-5,
      f"max_asym={np.max(np.abs(rho_g-rho_g.T)):.2e}")

# C06: ρ_geg[0,0] ≈ Ω (ground mode integral ≈ Ω·h₀ / h₀)
check("C06: ρ_geg[0,0] ≈ 121",
      abs(rho_g[0, 0] - 120.9) < 2.0,
      f"rho_g[0,0]={rho_g[0,0]:.4f}")

# C07: H_eff shape is (8,8)
T_MAX = 60.0; N_RP = 500
xi_r, wi_r = leggauss(N_RP)
t_r = 0.5*T_MAX*(xi_r+1); w_r_gl = 0.5*T_MAX*wi_r
L_r = np.array([Ln_stable(n, t_r) for n in range(M_RAD)])
inv2k3 = (1/(2*kappa))**3; inv2k4 = (1/(2*kappa))**4
wrad = np.exp(-t_r)*t_r**3
wvs  = np.exp(-1.5*t_r)*t_r**3
V_mat = np.array([[inv2k4*np.dot(w_r_gl, L_r[m]*L_r[n]*wvs)
                    for n in range(M_RAD)] for m in range(M_RAD)])
V00 = V_mat[0, 0]
H_eff = D2g + Dg + gamma*Tg + alpha*rho_g + V00*np.eye(N_ANG)

check("C07: H_eff shape (8,8)",
      H_eff.shape == (8, 8))

# C08: α·ρ_geg symmetric → same as C05 but recheck via H_eff angular part
ang_part = D2g + Dg + gamma*Tg + alpha*rho_g
check("C08: Angular Hamiltonian shape (8,8)",
      ang_part.shape == (8, 8))

# C09: V_self diagonal > 0 and < 1e-3
check("C09: V_self[0,0] > 0 and < 1e-3",
      (V00 > 0) and (V00 < 1e-3),
      f"V00={V00:.4e}")

# C10: V_self[0,0] < 1e-6 (much smaller than A240 estimate 7.99e-5)
check("C10: V_self[0,0] < 1e-6 (proper 4D suppression)",
      V00 < 1e-6,
      f"V00={V00:.4e}")

# C11: κ = Ω/3 ≈ 45.68
check("C11: kappa ≈ 45.68",
      abs(kappa - 45.678) < 0.1,
      f"kappa={kappa:.4f}")

# C12: R_0[0,0] > 0 (positive kinetic energy)
eps = 1e-5; safe_t = np.maximum(t_r, 1e-20)
R0_mat = np.zeros((M_RAD, M_RAD))
for n in range(M_RAD):
    Lv = L_r[n]
    dLv  = (Ln_stable(n, t_r+eps) - Ln_stable(n, t_r-eps))/(2*eps)
    d2Lv = (Ln_stable(n, t_r+eps) - 2*Lv + Ln_stable(n, t_r-eps))/eps**2
    two_k = 2*kappa
    core = -(two_k**2)*(d2Lv - dLv + Lv/4) - (12*kappa**2/safe_t)*(dLv - Lv/2)
    for m in range(M_RAD):
        R0_mat[m, n] = inv2k3 * np.dot(w_r_gl, L_r[m]*core*wrad)

check("C12: R_0[0,0] > 0",
      R0_mat[0, 0] > 0,
      f"R0[0,0]={R0_mat[0,0]:.4e}")

# C13: R_0[0,0] < 1.0 (small relative to angular scale)
check("C13: R_0[0,0] < 1.0",
      R0_mat[0, 0] < 1.0,
      f"R0[0,0]={R0_mat[0,0]:.4e}")

# C14: Ground eigenvalue of H_eff is finite and negative
evals, evecs = np.linalg.eig(H_eff)
real_mask = np.abs(evals.imag) < 1.0
rev = evals[real_mask].real; rvec = evecs[:, real_mask].real
gsi = np.argmin(rev); gs_eval = rev[gsi]; gs_vec = rvec[:, gsi]
check("C14: Ground eigenvalue finite and negative",
      np.isfinite(gs_eval) and gs_eval < 0,
      f"gs_eval={gs_eval:.4f}")

# C15: Ground eigenvalue ≈ -2867 (same as A240)
check("C15: Ground eigenvalue ≈ -2867 (within 1)",
      abs(gs_eval - (-2866.9)) < 2.0,
      f"gs_eval={gs_eval:.4f}")

# C16: x* is in (0,1)
c = gs_vec
psi0 = sum(c[n]*U[n] for n in range(N_ANG))
xstar = np.dot(psi0**2*x_gl*wt, w_gl)/np.dot(psi0**2*wt, w_gl)
check("C16: x* ∈ (0, 1)",
      0.0 < xstar < 1.0,
      f"x*={xstar:.8f}")

# C17: |x* - x*_A240| < 1e-3 (radial corrections negligible)
check("C17: |x* - x*_A240| < 1e-3 (radial corrections negligible)",
      abs(xstar - xA240) < 1e-3,
      f"|x*-x*_A240|={abs(xstar-xA240):.4e}")

# C18: x* > x*_CZ (not yet converged to CZ)
check("C18: x* > x*_CZ (gap not closed)",
      xstar > xCZ,
      f"x*={xstar:.6f} > x*_CZ={xCZ}")

# C19: Product basis size M×N = 48
check("C19: M×N = 48",
      M_RAD*N_ANG == 48)

# C20: Full 48×48 gives same x* as adiabatic (within 1e-4)
total = M_RAD*N_ANG
O48 = np.zeros((total, total))
for ir in range(M_RAD):
    sl = slice(ir*N_ANG, (ir+1)*N_ANG)
    O48[sl, sl] += D2g + Dg + gamma*Tg + alpha*rho_g
    for jr in range(M_RAD):
        sl2 = slice(jr*N_ANG, (jr+1)*N_ANG)
        O48[sl, sl2] += V_mat[ir, jr]*np.eye(N_ANG)
for ia in range(N_ANG):
    for ir in range(M_RAD):
        for jr in range(M_RAD):
            O48[ir*N_ANG+ia, jr*N_ANG+ia] += R0_mat[ir, jr]

evals48, evecs48 = np.linalg.eig(O48)
rm48 = np.abs(evals48.imag) < 100.
rev48 = evals48[rm48].real; rvec48 = evecs48[:, rm48].real
gsi48 = np.argmin(rev48); C48 = rvec48[:, gsi48].reshape(M_RAD, N_ANG)
xn = 0.; xd = 0.
for ir in range(M_RAD):
    pv = sum(C48[ir, n]*U[n] for n in range(N_ANG))
    xn += np.dot(pv**2*x_gl*wt, w_gl)
    xd += np.dot(pv**2*wt, w_gl)
xstar_full = xn/xd
check("C20: Full 48×48 x* agrees with adiabatic (|Δx*| < 1e-4)",
      abs(xstar_full - xstar) < 1e-4,
      f"|x*_full - x*_ad|={abs(xstar_full-xstar):.4e}")

# ── Summary ──────────────────────────────────────────────────────
print("=" * 60)
print(f"x*_A241 (adiabatic) = {xstar:.10f}")
print(f"x*_A241 (full 48×48) = {xstar_full:.10f}")
print(f"x*_A240 = {xA240}")
print(f"Δx* = {xstar - xA240:+.4e}")
cum = (xA238 - xstar)/(xA238 - xCZ)
print(f"Cumulative gap (A238→A241): {cum*100:.2f}%")
print(f"V_self[0,0] = {V00:.4e}")
print(f"R_0[0,0]    = {R0_mat[0,0]:.4e}")
print(f"\n{'='*60}\nRESULT: {PASS} PASS / {FAIL} FAIL")
