"""n51 — signed drainage under the information metric, per
n51_preregistration.md (committed before this file)."""
import numpy as np

N = 200; MU2 = 1e-6

def K_of(dk):
    k = 1.0 + dk
    K = MU2*np.eye(N)
    for e in range(N):
        i, j = e, (e+1) % N
        K[i,i] += k[e]; K[j,j] += k[e]; K[i,j] -= k[e]; K[j,i] -= k[e]
    return K

def info_metric(dk):
    w2, V = np.linalg.eigh(K_of(dk))
    wv = np.sqrt(np.maximum(w2, 1e-30))
    U = np.zeros((N, N))
    for e in range(N): U[e] = V[e] - V[(e+1) % N]
    W1, W2 = np.meshgrid(wv, wv, indexing='ij')
    kern = 1.0/((W1+W2)**2 * W1*W2)
    B = np.einsum('en,em->enm', U, U).reshape(N, N*N)
    return 0.125*(B*kern.reshape(1, N*N)) @ B.T

G = info_metric(np.zeros(N))
As = np.zeros((N, N))                     # signed incidence: div at site i
for i in range(N): As[i, i] = 1.0; As[i, (i-1) % N] = -1.0
Ginv_At = np.linalg.solve(G, As.T)
Ws = As @ Ginv_At

# Z1: zero mode and dispersion
r0 = Ws[0]
disp = np.fft.fft(np.roll(r0, 0)).real     # symbol via row FFT (translation invariant)
q = 2*np.pi*np.fft.fftfreq(N)
print(f"Z1: W_s symbol at q=0: {disp[0]:+.3e} (zero mode at q=0)")
for dq in (1, 2, 4):
    print(f"    q={q[dq]:.4f}: W_s/q^2 = {disp[dq]/q[dq]**2:.4f}")

def drain(c):
    lam, *_ = np.linalg.lstsq(Ws, -c, rcond=None)
    dk = Ginv_At @ lam
    return dk, 0.5*dk @ G @ dk, lam

def pair_V(a, b, f):
    c = np.zeros(N); c[a % N] = f; c[b % N] = f
    c -= c.mean()                          # jellium
    return drain(c)[1]

ds = [11, 21, 31, 41, 51, 61, 71, 81]
Vs = np.array([pair_V(40, 40+d, 0.3) for d in ds])
g_ref = np.array([d*(N-d)/(2*N) for d in ds])
Ad = np.vstack([np.ones_like(g_ref), g_ref]).T
cf, *_ = np.linalg.lstsq(Ad, Vs, rcond=None)
fr = np.sqrt(np.mean((Vs - Ad@cf)**2))/np.ptp(Vs)
print(f"Z2 SIGN: slope b = {cf[1]:+.5e} -> {'ATTRACTIVE (the resurrection)' if cf[1] > 0 else 'REPULSIVE (as the Poisson analogy expects)'}")
print(f"Z3 FORM: Green fracRMS = {fr:.3%}")

# Z4: clocks in the signed field
def sqrtm_sym(K):
    w2, Vv = np.linalg.eigh(K)
    return Vv @ np.diag(np.sqrt(np.maximum(w2, 1e-30))) @ Vv.T

def cell_avg(x): return 0.5*(x[0::2] + x[1::2])

Om0 = sqrtm_sym(K_of(np.zeros(N))); OmI0 = np.linalg.inv(Om0)

def clock_run(FM):
    c = np.zeros(N); c[40] = FM; c[139] = FM
    c -= c.mean()
    dkM, _, lam = drain(c)
    Om = sqrtm_sym(K_of(dkM)); OmI = np.linalg.inv(Om)
    sh = [cell_avg(np.diag(Om)/np.diag(Om0) - 1),
          cell_avg(np.sqrt(2*(1+dkM))/np.sqrt(2) - 1),
          cell_avg(np.diag(OmI0)/np.diag(OmI) - 1)]
    return sh, cell_avg(lam)

sh3, lam3 = clock_run(0.3)
sh15, lam15 = clock_run(0.15)
NC = N//2; cc = [20, 69]
far = np.array([min([min(abs(i-j), NC-abs(i-j)) for j in cc]) > 2 for i in range(NC)])
names = ["C1 site", "C2 bond", "C3 corr"]
kappas = []
for nm, s3, s15 in zip(names, sh3, sh15):
    expo = np.median(np.log(np.abs(s3[far]/s15[far]))/np.log(2))
    Ad = np.vstack([np.ones(far.sum()), lam3[far]]).T
    cfc, *_ = np.linalg.lstsq(Ad, s3[far], rcond=None)
    frc = np.sqrt(np.mean((s3[far] - Ad@cfc)**2))/max(np.ptp(s3[far]), 1e-300)
    kappas.append(cfc[1])
    print(f"Z4 {nm}: F-exponent = {expo:.3f} ({'PASS' if abs(expo-1) < 0.05 else 'FAIL'}); "
          f"prop fracRMS = {frc:.3%} ({'PASS' if frc < 0.05 else 'FAIL'}); kappa = {cfc[1]:+.5f}")
k1, k3 = kappas[0], kappas[2]
spread = abs(k1/k3 - 1) if k3 != 0 else np.inf
print(f"Z4 universality (C1 vs C3): |kappa1/kappa3 - 1| = {spread:.2%} -> "
      f"{'PASS' if spread < 0.10 else 'FAIL'}")
np.save("n51_results.npy", dict(ds=ds, Vs=Vs.tolist(), kappas=kappas), allow_pickle=True)
