#!/usr/bin/env python3
"""
run_vac.py — KKLT-CERT vacuum orchestrator.

  python3 run_vac.py solve <dps> [route]   -> result_vac_<route>_<dps>.json
  python3 run_vac.py checks                 -> checks_vac.json

Stages per solve: certified companion at s* (cert_w0 transport, reused) ->
exact transverse jets V0 (M(s*) solve) -> on-curve regression check ->
uncertified Newton center -> certified Krawczyk contraction (existence +
uniqueness box for DW=0 in all three directions) -> certified vacuum |W0|.
Resources: single thread.
"""
from fractions import Fraction as Fr
import json, os, sys, time
from flint import arb, acb, ctx

HERE = os.path.dirname(os.path.abspath(__file__))
import cert_vac as CV

PILOT = {"W0": "2.0371060933111834182083117862342977e-8",
         "tau": "6.8554572563187259046361432042141669",
         "U1": "2.7421769802085209924953942571371399",
         "U2": "2.0566327731047339299455907474511510"}

def ball_str(z, d):
    return z.str(d, radius=True)

def solve(dps, route="R"):
    prec = int(dps * 3.3219) + 64
    digits = dps + 25
    ctx.prec = prec
    t0 = time.time()
    comp, diags = CV.companion_at_star(prec, route)
    t_comp = time.time() - t0
    print(f"[companion route {route} dps {dps}] wall {t_comp:.1f}s", flush=True)
    t0 = time.time()
    V0 = CV.V0_at_star(comp, prec)
    CV.STATE["V0"] = V0
    print(f"[V0 jets] M(s*)^T solve done ({time.time()-t0:.1f}s); "
          f"th1Pi_X0(s*) = {ball_str(V0[1][3], 20)}", flush=True)
    # on-curve regression check: w=0, tau=tau* must reproduce cert_w0
    tpin = CV.fr2arb(CV.TAU_PIN)
    st0 = CV.eval_state([arb(0)] * 4 + [arb(0), tpin], prec, digits)
    ref = json.load(open(os.path.join(CV.CW0, f"result_R_{dps}.json")))
    print(f"[check gv0] on-curve W0 = {ball_str(st0['W0'], min(dps, 40))}")
    print(f"           cert_w0  W0 = {ref['W0'][:60]}...", flush=True)
    # uncertified Newton center from the curve point
    t0 = time.time()
    x0 = [arb(0)] * 4 + [arb(0), tpin]
    y, res = CV.newton_center(x0, prec, digits, iters=14)
    print(f"[newton] center residual {res:.3e}  ({time.time()-t0:.1f}s)")
    print("  center:", [y[i].str(8) for i in range(6)], flush=True)
    # Krawczyk: try radii until contraction certifies
    t0 = time.time()
    cert_box = None
    for rexp in (10, 12, 8, 14):
        r = [arb(10) ** (-rexp)] * 6
        K, ok, stF = CV.krawczyk(y, r, prec, digits)
        print(f"[krawczyk] r=1e-{rexp}: contained={ok}", flush=True)
        if ok:
            cert_box = (y, r, K)
            break
    assert cert_box, "Krawczyk containment failed at all radii"
    # contract: iterate to shrink the certified enclosure
    hist = []
    for it in range(12):
        y = [arb(K[i].mid()) for i in range(6)]
        r = [(arb(K[i].rad()) * arb("1.05")).upper() + arb(10) ** (-dps - 8)
             for i in range(6)]
        K2, ok, stF = CV.krawczyk(y, r, prec, digits)
        rmax = max(float(arb(K2[i].rad()).upper()) for i in range(6))
        hist.append(rmax)
        print(f"[contract {it}] contained={ok} max_rad={rmax:.3e}", flush=True)
        if ok:
            K = K2
        if len(hist) >= 2 and (not ok or hist[-1] > 0.25 * hist[-2]):
            break
    t_kraw = time.time() - t0
    return comp, st0, y, K, stF, prec, digits, dps, route, t_comp, t_kraw

def headline(dps, route="R"):
    comp, st0, y, K, stF, prec, digits, dps, route, t_comp, t_kraw = solve(dps, route)
    t0 = time.time()
    stV = CV.eval_state(K, prec, digits)      # certified box evaluation
    p = Fr(2, 5), Fr(3, 10)
    tau = stV["tau"]
    UmTp = [stV["U1"] - tau * CV.fr2acb(p[0]), stV["U2"] - tau * CV.fr2acb(p[1])]
    dU = [stV["U1"] - st0["U1"], stV["U2"] - st0["U2"]]
    pen = stV["W0"] / st0["W0"] - 1
    res = {
        "route": route, "dps": dps, "prec_bits": prec,
        "wall_companion_s": round(t_comp, 1), "wall_krawczyk_s": round(t_kraw, 1),
        "W0_vac": ball_str(stV["W0"], dps + 10),
        "W0_oncurve": ball_str(st0["W0"], dps + 10),
        "penalty_rel": ball_str(pen, 12),
        "tau_im": ball_str(K[5], dps + 10), "tau_re": ball_str(K[4], 6),
        "w1_re": ball_str(K[0], 25), "w1_im": ball_str(K[1], 6),
        "w2_re": ball_str(K[2], 25), "w2_im": ball_str(K[3], 6),
        "U1_im": ball_str(stV["U1"].imag, dps + 10),
        "U2_im": ball_str(stV["U2"].imag, dps + 10),
        "U1mtp": ball_str(UmTp[0].imag, 12), "U2mtp": ball_str(UmTp[1].imag, 12),
        "dU1_im": ball_str(dU[0].imag, 12), "dU2_im": ball_str(dU[1].imag, 12),
        "B_racetrack": ball_str(stV["B"], 12),
        "emKcs": ball_str(stV["emKcs"].real, 30),
        "F_resid_at_box": [ball_str(e.v, 6) for e in stV["E"]],
        "W0_rad": str(arb(stV["W0"].rad()).str(5)),
        "box_rad": [str(arb(K[i].rad()).str(5)) for i in range(6)],
    }
    fn = os.path.join(HERE, f"result_vac_{route}_{dps}.json")
    with open(fn, "w") as f:
        json.dump(res, f, indent=1)
    print(f"[headline] |W0|_vac = {res['W0_vac'][:80]}")
    print(f"[headline] penalty vs on-curve = {res['penalty_rel']}")
    print(f"[headline] banked {os.path.basename(fn)}  (+{time.time()-t0:.1f}s)", flush=True)

def parse_ball(sst):
    import mpmath as mp
    mp.mp.dps = 250
    sst = sst.strip()
    if sst.startswith("["):
        body = sst[1:-1]
        if "+/-" in body:
            m, r = body.split("+/-")
            return mp.mpf(m.strip()), mp.mpf(r.strip())
        return mp.mpf(body), mp.mpf(0)
    return mp.mpf(sst), mp.mpf(0)

def checks():
    import mpmath as mp
    mp.mp.dps = 250
    R = {d: json.load(open(os.path.join(HERE, f"result_vac_R_{d}.json")))
         for d in (60, 150)}
    W = {d: parse_ball(R[d]["W0_vac"]) for d in (60, 150)}
    g = {}
    # gc0: on-curve limit == cert_w0 certified ball (regression through the
    # new V0/contract path)
    ref = json.load(open(os.path.join(CV.CW0, "result_R_150.json")))
    oc, ocr = parse_ball(R[150]["W0_oncurve"])
    rf, rfr = parse_ball(ref["W0"])
    ov0 = bool(abs(oc - rf) <= ocr + rfr)
    d0 = int(mp.floor(-mp.log10(abs(oc - rf) / rf))) if oc != rf else 999
    g["gc0_oncurve_overlap"] = ov0; g["gc0_matched_digits"] = d0
    print(f"[gc0] on-curve limit vs cert_w0 ball: overlap={ov0}, matched "
          f"digits={d0} -> {'PASS' if ov0 and d0 >= 140 else 'FAIL'}")
    # gc5: route R vs route C (independent transport paths end-to-end)
    RC = json.load(open(os.path.join(HERE, "result_vac_C_150.json")))
    wc, wcr = parse_ball(RC["W0_vac"])
    ov5 = bool(abs(W[150][0] - wc) <= W[150][1] + wcr)
    d5 = int(mp.floor(-mp.log10(abs(W[150][0] - wc) / wc))) if W[150][0] != wc else 999
    bar5 = int(mp.floor(-mp.log10((W[150][1] + wcr) / wc))) - 1
    g["gc5_routeC_overlap"] = ov5; g["gc5_matched_digits"] = d5
    print(f"[gc5] route R vs C vacuum balls: overlap={ov5}, matched digits="
          f"{d5} (ball-supported: {bar5}) -> "
          f"{'PASS' if ov5 and d5 >= bar5 else 'FAIL'}")
    dmatch = int(mp.floor(-mp.log10(abs(W[60][0] - W[150][0]) / W[150][0]))) \
        if W[60][0] != W[150][0] else 999
    bar1 = int(mp.floor(-mp.log10(W[60][1] / W[60][0]))) - 1
    g["gc1_two_dps_matched_digits"] = dmatch
    g["gc1_ball_supported"] = bar1
    print(f"[gc1] dps-60 vs dps-150 |W0|_vac matched digits: {dmatch} "
          f"(dps-60 ball supports {bar1}) -> "
          f"{'PASS' if dmatch >= bar1 else 'FAIL'}")
    pilot = mp.mpf(PILOT["W0"])
    cert = W[150][0]
    g["gc4_pilot_model_rel"] = mp.nstr(abs(pilot - cert) / cert, 6)
    print(f"[gc4] pilot MODEL vacuum vs certified true-CY vacuum |W0|: "
          f"rel diff = {g['gc4_pilot_model_rel']} (measures GV-window error "
          f"AT THE VACUUM; pilot bound was 3.5e-9)")
    lo, hi = mp.mpf("2.0365e-8"), mp.mpf("2.0375e-8")
    inside = bool(lo < cert - W[150][1] and cert + W[150][1] < hi)
    g["verdict_DKMM_2.037e-8"] = ("CONFIRMED (certified ball inside the "
        "4-digit rounding band of 2.037e-8)" if inside else "NOT CONFIRMED")
    print(f"[VERDICT] DKMM printed |W0| = 2.037e-8: "
          f"{'WITHIN' if inside else 'OUTSIDE'} the certified ball's rounding "
          f"band; certified |W0| = {mp.nstr(cert, 25)}")
    g["headline"] = R[150]["W0_vac"]
    g["penalty_rel"] = R[150]["penalty_rel"]
    g["displacement"] = {k: R[150][k] for k in
                         ("U1mtp", "U2mtp", "dU1_im", "dU2_im", "B_racetrack")}
    g["tau_vac"] = R[150]["tau_im"][:60]
    tpilot = mp.mpf(PILOT["tau"])
    tc, _ = parse_ball(R[150]["tau_im"])
    g["tau_pilot_rel"] = mp.nstr(abs(tpilot - tc) / tc, 4)
    print(f"[info] certified tau_vac vs pilot model tau: rel diff "
          f"{g['tau_pilot_rel']} (GV-window scale, consistent with gc4)")
    with open(os.path.join(HERE, "checks_vac.json"), "w") as f:
        json.dump(g, f, indent=1)
    print("banked checks_vac.json")
    if not (ov0 and d0 >= 140 and ov5 and d5 >= bar5 and dmatch >= bar1 and inside):
        sys.exit(1)

if __name__ == "__main__":
    if sys.argv[1] == "solve":
        headline(int(sys.argv[2]), sys.argv[3] if len(sys.argv) > 3 else "R")
    elif sys.argv[1] == "checks":
        checks()
