#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Reproducible spectral, WKB, and instanton benchmark for a quartic double well.

The dimensionless Hamiltonian is

    H_g = -(g^2/2) d^2/dq^2 + (q^2 - 1)^2/2.

The production spectrum uses parity-resolved harmonic-oscillator bases. An
independent centered finite-difference calculation locates the two lowest
eigenvalues with Sturm bisection. NumPy is the only required dependency.
"""

from __future__ import annotations

import argparse
import csv
import math
import platform
import sys
import time
from pathlib import Path
from typing import Iterable

import numpy as np


S0 = 4.0 / 3.0
OMEGA0 = 2.0
DEFAULT_G_VALUES = (0.30, 0.25, 0.20, 0.15, 0.12, 0.10, 0.09, 0.08, 0.07, 0.06)
DEFAULT_OUTPUT_DIR = Path("results")


def position_matrix(g: float, omega: float, dimension: int) -> np.ndarray:
    """Return the position matrix in an oscillator basis with frequency omega."""

    q = np.zeros((dimension, dimension), dtype=float)
    entries = math.sqrt(g / (2.0 * omega)) * np.sqrt(np.arange(1, dimension))
    index = np.arange(dimension - 1)
    q[index, index + 1] = entries
    q[index + 1, index] = entries
    return q


def quartic_hamiltonian(g: float, n_basis: int, omega: float) -> np.ndarray:
    """Project H_g into the first n_basis oscillator states.

    Four guard states ensure that the projected q^4 matrix has the exact
    polynomial matrix elements instead of a top-edge truncation artifact.
    """

    work_dimension = n_basis + 4
    q = position_matrix(g, omega, work_dimension)
    q2 = q @ q
    q4 = q2 @ q2
    number = np.arange(work_dimension, dtype=float)
    h_osc = np.diag(g * omega * (number + 0.5))
    h = (
        h_osc
        + 0.5 * q4
        - (1.0 + 0.5 * omega * omega) * q2
        + 0.5 * np.eye(work_dimension)
    )
    return h[:n_basis, :n_basis]


def parity_spectrum(g: float, n_basis: int, omega: float) -> dict[str, float]:
    """Return the lowest parity doublet and structural diagnostics."""

    h = quartic_hamiltonian(g, n_basis, omega)
    even_index = np.arange(0, n_basis, 2)
    odd_index = np.arange(1, n_basis, 2)
    even_block = h[np.ix_(even_index, even_index)]
    odd_block = h[np.ix_(odd_index, odd_index)]
    even_values, even_vectors = np.linalg.eigh(even_block)
    odd_values, odd_vectors = np.linalg.eigh(odd_block)

    even_energy = float(even_values[0])
    odd_energy = float(odd_values[0])
    gap = odd_energy - even_energy
    next_energy = float(min(even_values[1], odd_values[1]))
    intrawell_spacing = next_energy - 0.5 * (even_energy + odd_energy)

    even_residual = np.linalg.norm(
        even_block @ even_vectors[:, 0] - even_energy * even_vectors[:, 0]
    )
    odd_residual = np.linalg.norm(
        odd_block @ odd_vectors[:, 0] - odd_energy * odd_vectors[:, 0]
    )
    parity_coupling = h[np.ix_(even_index, odd_index)]

    return {
        "even_energy": even_energy,
        "odd_energy": odd_energy,
        "gap": gap,
        "next_energy": next_energy,
        "intrawell_spacing": intrawell_spacing,
        "doublet_isolation": gap / intrawell_spacing,
        "hermiticity_error": float(np.max(np.abs(h - h.T))),
        "parity_coupling": float(np.max(np.abs(parity_coupling))),
        "eigenpair_residual": float(max(even_residual, odd_residual)),
    }


def wkb_action(g: float, order: int = 512) -> float:
    """Integrate the one-way finite-energy barrier action with Gauss quadrature."""

    if not 0.0 < g < 0.5:
        raise ValueError("The leading local energy requires 0 < g < 1/2.")
    turning_point = math.sqrt(1.0 - math.sqrt(2.0 * g))
    nodes, weights = np.polynomial.legendre.leggauss(order)
    q = 0.5 * turning_point * (nodes + 1.0)
    radicand = np.maximum((1.0 - q * q) ** 2 - 2.0 * g, 0.0)
    integral = 0.5 * turning_point * float(np.dot(weights, np.sqrt(radicand)))
    return 2.0 * integral


def wkb_gap(g: float, action: float) -> float:
    return 2.0 * g / math.pi * math.exp(-action / g)


def instanton_gap(g: float) -> float:
    return 4.0 * math.sqrt(8.0 * g / math.pi) * math.exp(-S0 / g)


def sweep_rows(
    g_values: Iterable[float], n_basis: int = 220, omega: float = 2.0
) -> list[dict[str, float]]:
    rows: list[dict[str, float]] = []
    for g in g_values:
        spectrum = parity_spectrum(g, n_basis, omega)
        comparison_gaps = [
            parity_spectrum(g, 160, omega)["gap"],
            parity_spectrum(g, n_basis, 1.5)["gap"],
            parity_spectrum(g, n_basis, 3.0)["gap"],
        ]
        gap = spectrum["gap"]
        numerical_spread = max(abs(value - gap) for value in comparison_gaps)
        action = wkb_action(g, 512)
        action_check = wkb_action(g, 256)
        gap_wkb = wkb_gap(g, action)
        gap_instanton = instanton_gap(g)
        effective_action = -g * math.log(gap / math.sqrt(g))
        scaled_prefactor = gap * math.exp(S0 / g) / math.sqrt(g)

        rows.append(
            {
                "g": g,
                "n_basis": n_basis,
                "omega": omega,
                "even_energy": spectrum["even_energy"],
                "odd_energy": spectrum["odd_energy"],
                "numerical_gap": gap,
                "numerical_gap_spread": numerical_spread,
                "relative_numerical_spread": numerical_spread / gap,
                "next_energy": spectrum["next_energy"],
                "intrawell_spacing": spectrum["intrawell_spacing"],
                "doublet_isolation": spectrum["doublet_isolation"],
                "wkb_action": action,
                "wkb_quadrature_change": abs(action - action_check),
                "wkb_gap": gap_wkb,
                "instanton_gap": gap_instanton,
                "wkb_ratio": gap_wkb / gap,
                "instanton_ratio": gap_instanton / gap,
                "wkb_relative_error": abs(gap_wkb / gap - 1.0),
                "instanton_relative_error": abs(gap_instanton / gap - 1.0),
                "effective_action": effective_action,
                "scaled_prefactor": scaled_prefactor,
                "one_loop_prefactor": 4.0 * math.sqrt(8.0 / math.pi),
                "hermiticity_error": spectrum["hermiticity_error"],
                "parity_coupling": spectrum["parity_coupling"],
                "eigenpair_residual": spectrum["eigenpair_residual"],
            }
        )
    return rows


def convergence_rows() -> list[dict[str, float]]:
    dimensions = (40, 60, 80, 100, 120, 160, 220, 280)
    frequencies = (1.5, 2.0, 3.0)
    rows: list[dict[str, float]] = []

    for g in (0.08, 0.06):
        reference_samples = [
            parity_spectrum(g, n_basis, omega)["gap"]
            for n_basis in (160, 220, 280)
            for omega in frequencies
        ]
        reference_samples_array = np.asarray(reference_samples)
        reference_gap = float(np.median(reference_samples_array))
        reference_spread = float(
            np.max(np.abs(reference_samples_array - reference_gap))
        )

        for omega in frequencies:
            for n_basis in dimensions:
                spectrum = parity_spectrum(g, n_basis, omega)
                gap = spectrum["gap"]
                rows.append(
                    {
                        "g": g,
                        "omega": omega,
                        "n_basis": n_basis,
                        "even_energy": spectrum["even_energy"],
                        "odd_energy": spectrum["odd_energy"],
                        "gap": gap,
                        "gap_positive": float(gap > 0.0),
                        "reference_gap": reference_gap,
                        "reference_spread": reference_spread,
                        "absolute_gap_error": abs(gap - reference_gap),
                        "relative_gap_error": abs(gap - reference_gap)
                        / reference_gap,
                        "eigenpair_residual": spectrum["eigenpair_residual"],
                    }
                )
    return rows


def sturm_count(diagonal: np.ndarray, off_diagonal: np.ndarray, x: float) -> int:
    """Count eigenvalues below x for a real symmetric tridiagonal matrix."""

    pivot_floor = np.finfo(float).tiny
    pivot = float(diagonal[0] - x)
    if abs(pivot) < pivot_floor:
        pivot = -pivot_floor
    count = int(pivot < 0.0)
    for index in range(1, diagonal.size):
        pivot = (
            float(diagonal[index] - x)
            - float(off_diagonal[index - 1] ** 2) / pivot
        )
        if abs(pivot) < pivot_floor:
            pivot = -pivot_floor
        count += int(pivot < 0.0)
    return count


def tridiagonal_eigenvalue(
    diagonal: np.ndarray, off_diagonal: np.ndarray, index: int
) -> float:
    """Locate the one-based indexed eigenvalue by Sturm bisection."""

    radii = np.zeros_like(diagonal)
    radii[0] = abs(off_diagonal[0])
    radii[-1] = abs(off_diagonal[-1])
    radii[1:-1] = np.abs(off_diagonal[:-1]) + np.abs(off_diagonal[1:])
    lower = float(np.min(diagonal - radii)) - 1.0
    upper = float(np.max(diagonal + radii)) + 1.0

    for _ in range(96):
        midpoint = 0.5 * (lower + upper)
        if sturm_count(diagonal, off_diagonal, midpoint) < index:
            lower = midpoint
        else:
            upper = midpoint
    return 0.5 * (lower + upper)


def finite_difference_doublet(
    g: float, spacing: float, half_width: float = 3.0
) -> dict[str, float]:
    """Compute the lowest full-line doublet on a centered Dirichlet grid."""

    n_grid = int(round(2.0 * half_width / spacing)) - 1
    spacing = 2.0 * half_width / (n_grid + 1)
    q = -half_width + spacing * np.arange(1, n_grid + 1)
    off_value = -0.5 * g * g / (spacing * spacing)
    diagonal = np.full(n_grid, -2.0 * off_value) + 0.5 * (q * q - 1.0) ** 2
    off_diagonal = np.full(n_grid - 1, off_value)
    even_energy = tridiagonal_eigenvalue(diagonal, off_diagonal, 1)
    odd_energy = tridiagonal_eigenvalue(diagonal, off_diagonal, 2)
    return {
        "n_grid": float(n_grid),
        "spacing": spacing,
        "half_width": half_width,
        "even_energy": even_energy,
        "odd_energy": odd_energy,
        "gap": odd_energy - even_energy,
    }


def finite_difference_rows(
    basis_gaps: dict[float, float]
) -> list[dict[str, float]]:
    schedules = {
        0.15: (0.0200, 0.0100, 0.0050, 0.0025, 0.00125),
        0.08: (0.0100, 0.0050, 0.0025, 0.00125),
    }
    rows: list[dict[str, float]] = []

    for g, spacings in schedules.items():
        previous_gap: float | None = None
        reference_gap = basis_gaps[g]
        for spacing in spacings:
            result = finite_difference_doublet(g, spacing)
            gap = result["gap"]
            richardson_gap = math.nan
            richardson_relative_error = math.nan
            if previous_gap is not None:
                richardson_gap = (4.0 * gap - previous_gap) / 3.0
                richardson_relative_error = abs(richardson_gap / reference_gap - 1.0)

            rows.append(
                {
                    "g": g,
                    **result,
                    "basis_gap": reference_gap,
                    "relative_gap_error": abs(gap / reference_gap - 1.0),
                    "richardson_gap": richardson_gap,
                    "richardson_relative_error": richardson_relative_error,
                }
            )
            previous_gap = gap
    return rows


def write_csv(path: Path, rows: list[dict[str, float]]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def make_quicklook_plots(
    output_dir: Path,
    sweep: list[dict[str, float]],
    convergence: list[dict[str, float]],
    finite_difference: list[dict[str, float]],
) -> None:
    try:
        import matplotlib.pyplot as plt
    except ImportError as exc:
        raise SystemExit("--plot requires Matplotlib.") from exc

    g = np.array([row["g"] for row in sweep])
    exact = np.array([row["numerical_gap"] for row in sweep])
    wkb = np.array([row["wkb_gap"] for row in sweep])
    instanton = np.array([row["instanton_gap"] for row in sweep])

    figure, axes = plt.subplots(2, 1, figsize=(7.0, 8.0), constrained_layout=True)
    axes[0].semilogy(1.0 / g, exact, "o-", label="parity-resolved spectrum")
    axes[0].semilogy(1.0 / g, wkb, "s--", label="finite-energy WKB")
    axes[0].semilogy(1.0 / g, instanton, "^:", label="one-loop instanton")
    axes[0].set(xlabel="1/g", ylabel="splitting")
    axes[0].legend()
    axes[0].grid(alpha=0.25)
    axes[1].plot(g, wkb / exact, "s--", label="WKB / spectrum")
    axes[1].plot(g, instanton / exact, "^:", label="instanton / spectrum")
    axes[1].axhline(1.0, color="black", linewidth=0.8)
    axes[1].set(xlabel="g", ylabel="ratio")
    axes[1].legend()
    axes[1].grid(alpha=0.25)
    figure.savefig(output_dir / "double-well-instanton-benchmark-quicklook.png", dpi=180)
    plt.close(figure)

    figure, axes = plt.subplots(2, 1, figsize=(7.0, 8.0), constrained_layout=True)
    for omega in (1.5, 2.0, 3.0):
        selected = [
            row
            for row in convergence
            if row["g"] == 0.06 and row["omega"] == omega
        ]
        axes[0].semilogy(
            [row["n_basis"] for row in selected],
            [row["relative_gap_error"] for row in selected],
            "o-",
            label=f"omega={omega:g}",
        )
    axes[0].set(xlabel="basis dimension", ylabel="relative gap difference")
    axes[0].legend()
    axes[0].grid(alpha=0.25)

    for g_value in (0.15, 0.08):
        selected = [row for row in finite_difference if row["g"] == g_value]
        axes[1].loglog(
            [row["spacing"] for row in selected],
            [row["relative_gap_error"] for row in selected],
            "o-",
            label=f"raw, g={g_value:g}",
        )
        richardson = [
            row for row in selected if math.isfinite(row["richardson_relative_error"])
        ]
        axes[1].loglog(
            [row["spacing"] for row in richardson],
            [row["richardson_relative_error"] for row in richardson],
            "s--",
            label=f"Richardson, g={g_value:g}",
        )
    axes[1].set(xlabel="grid spacing", ylabel="relative gap difference")
    axes[1].legend()
    axes[1].grid(alpha=0.25)
    figure.savefig(output_dir / "double-well-instanton-convergence-quicklook.png", dpi=180)
    plt.close(figure)


def validate(
    sweep: list[dict[str, float]],
    convergence: list[dict[str, float]],
    finite_difference: list[dict[str, float]],
) -> dict[str, float]:
    by_g = {row["g"]: row for row in sweep}
    fd_finest = {
        g: min(
            (row for row in finite_difference if row["g"] == g),
            key=lambda row: row["spacing"],
        )
        for g in (0.15, 0.08)
    }
    fd_richardson = {
        g: fd_finest[g]["richardson_relative_error"] for g in (0.15, 0.08)
    }

    stable_008 = [
        row["relative_gap_error"]
        for row in convergence
        if row["g"] == 0.08 and row["n_basis"] >= 120
    ]
    stable_006 = [
        row["relative_gap_error"]
        for row in convergence
        if row["g"] == 0.06 and row["n_basis"] >= 120
    ]
    gaps = [row["numerical_gap"] for row in sweep]
    actions = [row["wkb_action"] for row in sweep]

    metrics = {
        "benchmark_g_0p15_relative_difference": abs(
            by_g[0.15]["numerical_gap"] / 2.99603490208e-4 - 1.0
        ),
        "max_hermiticity_error": max(row["hermiticity_error"] for row in sweep),
        "max_parity_coupling": max(row["parity_coupling"] for row in sweep),
        "max_eigenpair_residual": max(row["eigenpair_residual"] for row in sweep),
        "max_wkb_quadrature_change": max(
            row["wkb_quadrature_change"] for row in sweep
        ),
        "g_0p06_relative_numerical_spread": by_g[0.06][
            "relative_numerical_spread"
        ],
        "g_0p08_basis_plateau_error": max(stable_008),
        "g_0p06_basis_plateau_error": max(stable_006),
        "fd_g_0p15_richardson_relative_error": fd_richardson[0.15],
        "fd_g_0p08_richardson_relative_error": fd_richardson[0.08],
        "g_0p06_wkb_action_distance_to_S0": S0 - by_g[0.06]["wkb_action"],
    }

    checks = {
        "benchmark convention": metrics["benchmark_g_0p15_relative_difference"] < 1e-9,
        "Hermiticity": metrics["max_hermiticity_error"] < 1e-12,
        "parity blocks": metrics["max_parity_coupling"] < 1e-12,
        "eigenpair residual": metrics["max_eigenpair_residual"] < 1e-10,
        "WKB quadrature": metrics["max_wkb_quadrature_change"] < 2e-8,
        "positive production gaps": all(gap > 0.0 for gap in gaps),
        "monotone suppression": all(
            earlier > later for earlier, later in zip(gaps, gaps[1:])
        ),
        "monotone WKB action": all(
            earlier < later for earlier, later in zip(actions, actions[1:])
        ),
        "g=0.08 basis plateau": metrics["g_0p08_basis_plateau_error"] < 2e-6,
        "g=0.06 basis plateau": metrics["g_0p06_basis_plateau_error"] < 3e-5,
        "finite-difference g=0.15": fd_richardson[0.15] < 2e-6,
        "finite-difference g=0.08": fd_richardson[0.08] < 2e-5,
    }
    failed = [name for name, passed in checks.items() if not passed]
    if failed:
        raise RuntimeError("Validation failed: " + ", ".join(failed))
    return metrics


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=DEFAULT_OUTPUT_DIR,
        help="Directory for retained CSV data and optional quick-look plots (default: results in the working directory).",
    )
    parser.add_argument(
        "--plot",
        action="store_true",
        help="Generate optional Matplotlib quick-look PNGs.",
    )
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    started = time.perf_counter()
    sweep = sweep_rows(DEFAULT_G_VALUES)
    convergence = convergence_rows()
    basis_gaps = {row["g"]: row["numerical_gap"] for row in sweep}
    finite_difference = finite_difference_rows(basis_gaps)
    metrics = validate(sweep, convergence, finite_difference)

    sweep_path = args.output_dir / "double-well-instanton-sweep.csv"
    convergence_path = args.output_dir / "double-well-instanton-convergence.csv"
    finite_difference_path = (
        args.output_dir / "double-well-instanton-finite-difference.csv"
    )
    write_csv(sweep_path, sweep)
    write_csv(convergence_path, convergence)
    write_csv(finite_difference_path, finite_difference)

    if args.plot:
        make_quicklook_plots(args.output_dir, sweep, convergence, finite_difference)

    elapsed = time.perf_counter() - started
    print("Double-well instanton numerical check")
    print(f"Python: {platform.python_version()}")
    print(f"NumPy: {np.__version__}")
    print(f"Platform: {platform.platform()}")
    print(f"Runtime: {elapsed:.3f} s")
    print(f"Retained sweep: {sweep_path}")
    print(f"Retained convergence data: {convergence_path}")
    print(f"Retained finite-difference data: {finite_difference_path}")
    for name, value in metrics.items():
        print(f"{name}: {value:.6e}")
    print("Validation: PASS")
    return 0


if __name__ == "__main__":
    sys.exit(main())
