"""
Physics Rebuild v16
Critical Disorder Stabilization Study
FULL CHECKPOINTED VERSION

Purpose
-------
High-statistics refinement of the disorder transition where:

    alpha(omega_c) = 0

with:
============================================================
"""

import numpy as np
import csv
import os
import matplotlib.pyplot as plt
from scipy.linalg import expm

# ============================================================
# PARAMETERS
# ============================================================

dt = 0.06
T = 12.0
steps = int(T / dt)

# Critical refinement region
omega_values = np.linspace(0.34, 0.48, 18)

# Reservoir sizes
N_res_values = [2, 3, 4, 5, 6]

# Adjust this for runtime
seeds = list(range(64))

# Leakage sweep
gamma_values = [0.0, 0.01, 0.02, 0.04, 0.07, 0.10]

# Reservoir structure
J_disorder = 0.40
g_sys_res = 0.25

# Bootstrap
bootstrap_samples = 600

# Files
results_file = "physics_rebuild_v16_data.csv"
k_file = "physics_rebuild_v16_k.csv"
alpha_file = "physics_rebuild_v16_alpha.csv"
raw_copy_file = "physics_rebuild_v16_raw_copy.csv"

# ============================================================
# PAULI MATRICES
# ============================================================

I2 = np.eye(2, dtype=complex)

X = np.array([[0, 1],
              [1, 0]], dtype=complex)

Y = np.array([[0, -1j],
              [1j, 0]], dtype=complex)

Z = np.array([[1, 0],
              [0, -1]], dtype=complex)

SM = np.array([[0, 1],
               [0, 0]], dtype=complex)

# ============================================================
# REFERENCE STATES
# ============================================================

bell = np.array([1, 0, 0, 1], dtype=complex) / np.sqrt(2)

phi = np.array([1, 0, 0, -1], dtype=complex) / np.sqrt(2)
rho_ref = np.outer(phi, phi.conj())

# ============================================================
# HELPERS
# ============================================================

def kron_all(ops):
    out = ops[0]

    for op in ops[1:]:
        out = np.kron(out, op)

    return out


def operator_on_site(op, site, N):
    return kron_all([
        op if i == site else I2
        for i in range(N)
    ])


def two_site(op1, s1, op2, s2, N):
    ops = []

    for i in range(N):
        if i == s1:
            ops.append(op1)
        elif i == s2:
            ops.append(op2)
        else:
            ops.append(I2)

    return kron_all(ops)


def normalize_rho(rho):
    rho = 0.5 * (rho + rho.conj().T)

    tr = np.trace(rho)

    if abs(tr) > 1e-14:
        rho = rho / tr

    return rho


def partial_trace(rho, keep, dims):
    dims = list(dims)
    keep = list(keep)

    n = len(dims)

    traced = [i for i in range(n) if i not in keep]

    reshaped = rho.reshape(dims + dims)

    current_n = n

    for q in sorted(traced, reverse=True):
        reshaped = np.trace(
            reshaped,
            axis1=q,
            axis2=q + current_n
        )

        dims.pop(q)
        current_n -= 1

    d_keep = int(np.prod(dims))

    return normalize_rho(
        reshaped.reshape((d_keep, d_keep))
    )


def trace_distance(rho1, rho2):
    rho1 = normalize_rho(rho1)
    rho2 = normalize_rho(rho2)

    diff = rho1 - rho2

    vals = np.linalg.eigvals(
        diff.conj().T @ diff
    )

    vals = np.real(vals)
    vals[vals < 0] = 0.0

    return float(
        0.5 * np.sum(np.sqrt(vals))
    )


def r2_score(y, yhat):
    y = np.array(y)
    yhat = np.array(yhat)

    ss_res = np.sum((y - yhat) ** 2)
    ss_tot = np.sum((y - np.mean(y)) ** 2)

    if ss_tot < 1e-14:
        return 0.0

    return float(1.0 - ss_res / ss_tot)


# ============================================================
# INITIAL STATE
# ============================================================

def make_initial_rho(N_res):
    reservoir_state = np.zeros(
        2 ** N_res,
        dtype=complex
    )

    reservoir_state[0] = 1.0

    psi0 = np.kron(
        bell,
        reservoir_state
    )

    return np.outer(psi0, psi0.conj())


# ============================================================
# HAMILTONIAN
# ============================================================

def build_hamiltonian(N_res, seed, omega_disorder):
    rng = np.random.default_rng(seed)

    N_total = 2 + N_res
    dim = 2 ** N_total

    H = np.zeros((dim, dim), dtype=complex)

    # System pair
    H += 0.6 * (
        two_site(X, 0, X, 1, N_total)
        +
        two_site(Y, 0, Y, 1, N_total)
    )

    # Reservoir detunings
    for r in range(N_res):
        idx = 2 + r

        omega = rng.uniform(
            -omega_disorder,
            omega_disorder
        )

        H += omega * operator_on_site(
            Z,
            idx,
            N_total
        )

    # Reservoir couplings
    for r in range(N_res - 1):
        i = 2 + r
        j = 2 + r + 1

        J = rng.uniform(
            -J_disorder,
            J_disorder
        )

        H += J * (
            two_site(X, i, X, j, N_total)
            +
            two_site(Y, i, Y, j, N_total)
        )

    # System-reservoir coupling
    for r in range(N_res):
        idx = 2 + r

        H += g_sys_res * (
            two_site(X, 1, X, idx, N_total)
            +
            two_site(Y, 1, Y, idx, N_total)
        )

    return H


# ============================================================
# LEAKAGE
# ============================================================

def build_leakage_ops(N_res):
    N_total = 2 + N_res

    L_ops = []

    for r in range(N_res):
        idx = 2 + r

        L_ops.append(
            operator_on_site(
                SM,
                idx,
                N_total
            )
        )

    return L_ops


# ============================================================
# SINGLE RUN
# ============================================================

def run_case(N_res, seed, gamma_res, omega_disorder):
    N_total = 2 + N_res
    dims = [2] * N_total

    H = build_hamiltonian(
        N_res,
        seed,
        omega_disorder
    )

    U = expm(-1j * H * dt)

    L_ops = build_leakage_ops(N_res)

    rho = make_initial_rho(N_res)

    D_vals = []

    for step in range(steps):
        rho = U @ rho @ U.conj().T

        if gamma_res > 0:
            for L in L_ops:
                LdL = L.conj().T @ L

                dissipator = (
                    L @ rho @ L.conj().T
                    -
                    0.5 * (
                        LdL @ rho
                        +
                        rho @ LdL
                    )
                )

                rho = rho + gamma_res * dt * dissipator

        rho = normalize_rho(rho)

        rho_sys = partial_trace(
            rho,
            [0, 1],
            dims
        )

        D_vals.append(
            trace_distance(
                rho_sys,
                rho_ref
            )
        )

    t = np.arange(steps) * dt

    D_vals = np.array(D_vals)

    dD = np.gradient(D_vals, dt)

    positive = np.where(
        dD > 0,
        dD,
        0.0
    )

    BLP = np.trapz(positive, t)

    return float(BLP)


# ============================================================
# EXPONENTIAL FIT
# ============================================================

def fit_exp_offset(gamma, y):
    gamma = np.array(gamma)
    y = np.array(y)

    ymin = np.min(y)

    C_vals = np.linspace(
        max(0.0, ymin * 0.25),
        ymin * 0.995,
        240
    )

    best = None

    for C in C_vals:
        shifted = y - C

        if np.any(shifted <= 0):
            continue

        coeff = np.polyfit(
            gamma,
            np.log(shifted),
            1
        )

        k = -coeff[0]
        A = np.exp(coeff[1])

        yhat = (
            A * np.exp(-k * gamma)
            + C
        )

        mse = np.mean((y - yhat) ** 2)

        r2 = r2_score(y, yhat)

        if best is None or mse < best["mse"]:
            best = {
                "A": float(A),
                "k": float(k),
                "C": float(C),
                "r2": float(r2),
                "mse": float(mse)
            }

    return best


# ============================================================
# IMMEDIATE-SAVE RAW DATA GENERATION
# ============================================================

def ensure_results_file():
    if not os.path.exists(results_file):
        with open(results_file, mode="w", newline="") as f:
            writer = csv.writer(f)
            writer.writerow([
                "N_res",
                "seed",
                "gamma",
                "omega",
                "BLP_value"
            ])


def load_completed_points():
    completed = set()

    if not os.path.exists(results_file):
        return completed

    with open(results_file, mode="r") as f:
        reader = csv.DictReader(f)

        for row in reader:
            try:
                key = (
                    int(float(row["N_res"])),
                    int(float(row["seed"])),
                    round(float(row["gamma"]), 12),
                    round(float(row["omega"]), 12)
                )

                completed.add(key)

            except Exception:
                pass

    return completed


def run_raw_checkpointed():
    ensure_results_file()

    completed = load_completed_points()

    total_runs = (
        len(N_res_values)
        * len(gamma_values)
        * len(omega_values)
        * len(seeds)
    )

    print("\n================================================")
    print("Physics Rebuild v16")
    print("Critical Disorder Stabilization")
    print("Immediate-save + resume mode ON")
    print("================================================")
    print("Results file:", results_file)
    print("N_res:", N_res_values)
    print("seeds:", len(seeds))
    print("gammas:", gamma_values)
    print("omegas:", len(omega_values))
    print("Already completed:", len(completed))
    print("Total planned:", total_runs)
    print("Remaining:", total_runs - len(completed))
    print("================================================")

    run_counter = len(completed)

    for w in omega_values:
        print("\n================================================")
        print("omega =", round(float(w), 6))
        print("================================================")

        for n in N_res_values:
            print("\nN_res =", n)

            for g in gamma_values:
                print("  gamma =", g)

                for s in seeds:
                    key = (
                        int(n),
                        int(s),
                        round(float(g), 12),
                        round(float(w), 12)
                    )

                    if key in completed:
                        continue

                    blp_result = run_case(
                        N_res=n,
                        seed=s,
                        gamma_res=g,
                        omega_disorder=float(w)
                    )

                    with open(results_file, mode="a", newline="") as f:
                        writer = csv.writer(f)
                        writer.writerow([
                            n,
                            s,
                            g,
                            float(w),
                            blp_result
                        ])
                        f.flush()

                    completed.add(key)
                    run_counter += 1

                    print(
                        "Saved",
                        str(run_counter) + "/" + str(total_runs),
                        "| N =", n,
                        "| seed =", s,
                        "| gamma =", g,
                        "| omega =", round(float(w), 6),
                        "| BLP =",
                        round(float(blp_result), 6)
                    )

    print("\n================================================")
    print("RAW RUN COMPLETE")
    print("Saved:", results_file)
    print("================================================")


# ============================================================
# LOAD SAVED RAW DATA
# ============================================================

def load_saved_data():
    loaded_rows = []

    with open(results_file, mode="r") as f:
        reader = csv.DictReader(f)

        for row in reader:
            loaded_rows.append([
                float(row["N_res"]),
                float(row["seed"]),
                float(row["gamma"]),
                float(row["omega"]),
                float(row["BLP_value"])
            ])

    data = np.array(loaded_rows, dtype=float)

    print("\nLoaded saved data rows:", len(data))

    return data


# ============================================================
# ANALYSIS
# ============================================================

def analyze_saved_data():
    data = load_saved_data()

    # Save a clean copy with analysis-compatible name
    with open(raw_copy_file, "w", newline="") as f:
        writer = csv.writer(f)
        writer.writerow([
            "N_res",
            "seed",
            "gamma_res",
            "omega_disorder",
            "BLP"
        ])
        writer.writerows(data.tolist())

    print("Saved clean raw copy:", raw_copy_file)

    # --------------------------------------------------------
    # Extract k for each N_res, omega
    # --------------------------------------------------------

    k_rows = []

    for omega in omega_values:
        for N_res in N_res_values:
            means = []
            stds = []

            for gamma in gamma_values:
                mask = (
                    (data[:, 0] == N_res)
                    &
                    (data[:, 2] == gamma)
                    &
                    (np.abs(data[:, 3] - omega) < 1e-12)
                )

                vals = data[mask][:, 4]

                if len(vals) == 0:
                    means.append(np.nan)
                    stds.append(np.nan)
                else:
                    means.append(float(np.mean(vals)))
                    stds.append(float(np.std(vals)))

            if np.any(np.isnan(means)):
                print(
                    "Skipping incomplete:",
                    "N =", N_res,
                    "omega =", omega
                )
                continue

            fit = fit_exp_offset(
                gamma_values,
                means
            )

            if fit is None:
                print(
                    "Fit failed:",
                    "N =", N_res,
                    "omega =", omega
                )
                continue

            k_rows.append([
                N_res,
                float(omega),
                fit["A"],
                fit["k"],
                fit["C"],
                fit["r2"],
                fit["mse"],
                min(means),
                max(means),
                np.mean(stds)
            ])

    with open(k_file, "w", newline="") as f:
        writer = csv.writer(f)

        writer.writerow([
            "N_res",
            "omega_disorder",
            "A",
            "k",
            "C_floor",
            "R2",
            "MSE",
            "BLP_min",
            "BLP_max",
            "mean_BLP_std"
        ])

        writer.writerows(k_rows)

    print("Saved:", k_file)

    # --------------------------------------------------------
    # Extract alpha(omega)
    # --------------------------------------------------------

    k_data = np.array(k_rows, dtype=float)

    alpha_rows = []

    rng = np.random.default_rng(12345)

    for omega in omega_values:
        mask = (
            np.abs(
                k_data[:, 1] - omega
            ) < 1e-12
        )

        sub = k_data[mask]

        if len(sub) < 2:
            continue

        Ns = sub[:, 0]
        ks = sub[:, 3]

        coeff = np.polyfit(
            np.log(Ns),
            np.log(ks),
            1
        )

        alpha_hat = coeff[0]
        A_hat = np.exp(coeff[1])

        # Bootstrap over N-points
        alpha_boot = []

        for b in range(bootstrap_samples):
            idx = rng.integers(
                0,
                len(Ns),
                len(Ns)
            )

            Ns_b = Ns[idx]
            ks_b = ks[idx]

            if len(set(Ns_b)) < 2:
                continue

            try:
                coeff_b = np.polyfit(
                    np.log(Ns_b),
                    np.log(ks_b),
                    1
                )

                alpha_boot.append(coeff_b[0])

            except Exception:
                pass

        alpha_boot = np.array(alpha_boot)

        if len(alpha_boot) > 0:
            alpha_std = float(np.std(alpha_boot))
            alpha_low = float(np.percentile(alpha_boot, 2.5))
            alpha_high = float(np.percentile(alpha_boot, 97.5))
        else:
            alpha_std = np.nan
            alpha_low = np.nan
            alpha_high = np.nan

        khat = A_hat * Ns ** alpha_hat

        R2 = r2_score(ks, khat)

        alpha_rows.append([
            float(omega),
            A_hat,
            alpha_hat,
            alpha_std,
            alpha_low,
            alpha_high,
            R2
        ])

    with open(alpha_file, "w", newline="") as f:
        writer = csv.writer(f)

        writer.writerow([
            "omega_disorder",
            "A",
            "alpha",
            "alpha_std",
            "alpha_low",
            "alpha_high",
            "R2"
        ])

        writer.writerows(alpha_rows)

    print("Saved:", alpha_file)

    # --------------------------------------------------------
    # omega_c
    # --------------------------------------------------------

    alpha_data = np.array(alpha_rows, dtype=float)

    omegas = alpha_data[:, 0]
    alphas = alpha_data[:, 2]

    omega_c = None

    for i in range(len(omegas) - 1):
        a1 = alphas[i]
        a2 = alphas[i + 1]

        if a1 == 0:
            omega_c = omegas[i]
            break

        if a1 * a2 < 0:
            w1 = omegas[i]
            w2 = omegas[i + 1]

            omega_c = (
                w1
                +
                (0 - a1)
                *
                (w2 - w1)
                /
                (a2 - a1)
            )

            break

    susceptibility = np.gradient(
        alphas,
        omegas
    )

    # --------------------------------------------------------
    # Plots
    # --------------------------------------------------------

    fig, axes = plt.subplots(
        5,
        1,
        figsize=(12, 26)
    )

    yerr_low = alphas - alpha_data[:, 4]
    yerr_high = alpha_data[:, 5] - alphas

    axes[0].errorbar(
        omegas,
        alphas,
        yerr=[yerr_low, yerr_high],
        marker="o",
        capsize=4
    )

    axes[0].axhline(
        0,
        linestyle="--"
    )

    if omega_c is not None:
        axes[0].axvline(
            omega_c,
            linestyle="--"
        )

        axes[0].text(
            omega_c,
            max(alphas) * 0.8,
            "omega_c ~ "
            +
            str(round(float(omega_c), 4)),
            rotation=90
        )

    axes[0].set_title(
        "Critical Disorder Transition: alpha(omega)"
    )

    axes[0].set_xlabel(
        "omega disorder"
    )

    axes[0].set_ylabel(
        "alpha in k ~ N^alpha"
    )

    axes[0].grid(True)

    # Susceptibility
    axes[1].plot(
        omegas,
        susceptibility,
        marker="o"
    )

    axes[1].axhline(
        0,
        linestyle="--"
    )

    axes[1].set_title(
        "Critical Susceptibility d(alpha)/d(omega)"
    )

    axes[1].set_xlabel(
        "omega disorder"
    )

    axes[1].set_ylabel(
        "susceptibility"
    )

    axes[1].grid(True)

    # R2
    axes[2].plot(
        omegas,
        alpha_data[:, 6],
        marker="o"
    )

    axes[2].set_title(
        "Scaling Fit Quality"
    )

    axes[2].set_xlabel(
        "omega disorder"
    )

    axes[2].set_ylabel(
        "R²"
    )

    axes[2].grid(True)

    # Phase map k(N, omega)
    phase = np.zeros(
        (
            len(omega_values),
            len(N_res_values)
        )
    )

    for i, omega in enumerate(omega_values):
        for j, N_res in enumerate(N_res_values):
            mask = (
                (k_data[:, 0] == N_res)
                &
                (
                    np.abs(
                        k_data[:, 1] - omega
                    )
                    < 1e-12
                )
            )

            if np.any(mask):
                phase[i, j] = k_data[mask][0, 3]
            else:
                phase[i, j] = np.nan

    im = axes[3].imshow(
        phase,
        aspect="auto",
        origin="lower",
        extent=[
            min(N_res_values) - 0.5,
            max(N_res_values) + 0.5,
            min(omega_values),
            max(omega_values)
        ]
    )

    fig.colorbar(
        im,
        ax=axes[3],
        label="k"
    )

    axes[3].set_title(
        "k(N, omega) Critical Phase Map"
    )

    axes[3].set_xlabel(
        "N_res"
    )

    axes[3].set_ylabel(
        "omega disorder"
    )

    # Representative k vs N
    representative = [
        omega_values[0],
        omega_values[len(omega_values) // 2],
        omega_values[-1]
    ]

    for omega in representative:
        mask = (
            np.abs(
                k_data[:, 1] - omega
            )
            < 1e-12
        )

        sub = k_data[mask]

        axes[4].plot(
            sub[:, 0],
            sub[:, 3],
            marker="o",
            label="omega="
            +
            str(round(float(omega), 4))
        )

    axes[4].set_title(
        "Representative k vs N_res"
    )

    axes[4].set_xlabel(
        "N_res"
    )

    axes[4].set_ylabel(
        "k"
    )

    axes[4].grid(True)
    axes[4].legend()

    plt.tight_layout()

    plt.savefig(
        "physics_rebuild_v16_critical_transition.png",
        dpi=250
    )

    plt.show()

    # --------------------------------------------------------
    # Extra plot: alpha only
    # --------------------------------------------------------

    plt.figure(figsize=(10, 6))

    plt.errorbar(
        omegas,
        alphas,
        yerr=[yerr_low, yerr_high],
        marker="o",
        capsize=4
    )

    plt.axhline(0, linestyle="--")

    if omega_c is not None:
        plt.axvline(omega_c, linestyle="--")

    plt.title("v16 alpha(omega) Critical Transition")
    plt.xlabel("omega disorder")
    plt.ylabel("alpha")
    plt.grid(True)
    plt.tight_layout()
    plt.savefig("physics_rebuild_v16_alpha_only.png", dpi=250)
    plt.show()

    # --------------------------------------------------------
    # Final report
    # --------------------------------------------------------

    print("\n================================================")
    print("PHYSICS REBUILD v16 REPORT")
    print("================================================")

    if omega_c is not None:
        print(
            "Estimated critical disorder omega_c:",
            round(float(omega_c), 8)
        )
    else:
        print("No sign crossing detected.")

    print("\nalpha(omega):")

    for row in alpha_rows:
        print(
            "omega =",
            round(row[0], 5),
            "| alpha =",
            round(row[2], 6),
            "| CI=[",
            round(row[4], 6),
            ",",
            round(row[5], 6),
            "]",
            "| R2 =",
            round(row[6], 5)
        )

    print("\nInterpretation:")
    print("--------------------------------------------")
    print("alpha < 0 : larger reservoirs preserve memory")
    print("alpha > 0 : larger reservoirs accelerate scrambling")
    print("omega_c marks the crossover between")
    print("coherent-memory and statistical-bath scaling.")
    print("--------------------------------------------")

    print("\nSaved:")
    print(results_file)
    print(raw_copy_file)
    print(k_file)
    print(alpha_file)
    print("physics_rebuild_v16_critical_transition.png")
    print("physics_rebuild_v16_alpha_only.png")

    print("\nDONE.")
    print("================================================")


# ============================================================
# MAIN
# ============================================================

run_raw_checkpointed()
analyze_saved_data()
