#!/usr/bin/env python3
"""Static Einstein-Vlasov benchmark: sampled steady state must remain steady under evolution."""
import json, sys, time
import numpy as np
from pbhgr.steady_state import build, sample
from pbhgr.spherical_ev import Grid, Run, RunConfig, solve_metric
E0c = float(sys.argv[1]) if len(sys.argv) > 1 else 1.3
st = build(E0c); R, M = st["R"], st["M"]; t_dyn = np.sqrt(R**3 / M)
print(f"static solution: R={R:.3f} M={M:.4f} 2M/R={2*M/R:.3f} max2m/r={np.max(2*st['m']/st['r']):.3f} t_dyn={t_dyn:.2f}")
rows = []
for n in [int(a) for a in sys.argv[2:]] or (20000, 80000):
rng = np.random.default_rng(5)
g = Grid(r_out=4 * R, K=400)
P = sample(st, n, rng)
met0 = solve_metric(g, P)
# t=0 checks: mass function and lapse vs analytic
m_an = np.interp(g.r_edge, st["r"], st["m"], right=M)
mu_an = np.where(g.r_edge < R, np.interp(g.r_edge, st["r"], st["mu"]), 0.5 * np.log(np.maximum(1 - 2 * M / np.maximum(g.r_edge, 1e-9), 1e-12)))
err_m0 = float(np.max(np.abs(met0.m_edge - m_an)) / M); err_mu0 = float(np.max(np.abs(met0.mu_edge - mu_an)))
comp0 = 2 * met0.m_edge[1:] / g.r_edge[1:]
Einf0 = np.exp(np.interp(P.r, g.r_edge, met0.mu_edge)) * P.energy()
run = Run(g, P, RunConfig(t_end=3 * t_dyn, cfl=0.4, diag_every=25))
t0 = time.time(); out = run.run(); wall = time.time() - t0
Einf1 = np.exp(np.interp(run.P.r, g.r_edge, solve_metric(g, run.P).mu_edge)) * run.P.energy()
dE = (Einf1 - Einf0) / Einf0 if len(Einf1) == len(Einf0) else np.array([np.nan]) # escapes break the 1:1 pairing
met1 = solve_metric(g, run.P)
comp1 = 2 * met1.m_edge[1:] / g.r_edge[1:]
dg = run.diag
Mt = np.array([d["M_total"] for d in dg]); cm = np.array([d["max_2m_over_r"] for d in dg])
mc = [d["momentum_constraint_rel_residual"] for d in dg if d["momentum_constraint_rel_residual"] == d["momentum_constraint_rel_residual"]]
row = dict(n=n, outcome=out, M_sampled=float(met0.M), err_m_t0=err_m0, err_mu_t0=err_mu0,
max_comp_t0=float(comp0.max()), max_comp_mean=float(cm.mean()), max_comp_std=float(cm.std()), max_comp_end=float(comp1.max()),
profile_rel_change=float(np.max(np.abs(comp1 - comp0)) / comp0.max()),
adm_drift=float(np.max(np.abs(Mt - Mt[0])) / Mt[0]), rest_mass_escaped=run.ledger.escaped_rest_mass / P.N.sum(),
mc_residual_median=float(np.median(mc)), killing_energy_drift_median=float(np.median(dE)),
killing_energy_drift_p90=float(np.percentile(np.abs(dE), 90)), wall=wall)
rows.append(row); print(json.dumps(row), flush=True)
np.savez(f"runs/verify_static_n{n}.npz", r_edge=g.r_edge, comp0=comp0, comp1=comp1, m_an=m_an, m0=met0.m_edge, m1=met1.m_edge,
t=[d["t"] for d in dg], maxcomp=cm, M_total=Mt)
json.dump(rows, open("runs/verify_static.json", "w"), indent=1)