#!/usr/bin/env python3
"""Audit the quasi-exact sector of the hyperbolic Razavy Hamiltonian.

The convention used on the companion page is

    H = -d^2/dx^2 + (zeta*cosh(2*x) - M)^2,
    zeta > 0,  M a positive integer.

With n=M-1, z=exp(2*x), and

    psi(x) = z^((1-M)/2) exp[-zeta*(z+z^(-1))/4] P(z),

the gauged Hamiltonian preserves polynomials of degree at most n.  Its
finite matrix is similar to the real symmetric Jacobi matrix

    J[k,k]     = 4*k*(n-k) + 2*n + 1 + zeta^2,
    J[k,k+1]   = -2*zeta*sqrt((k+1)*(n-k)).

The M eigenvalues of J are compared with the first M levels of an
independent finite-interval approximation to the real-line problem.  The
latter uses the even and odd reductions of a centered uniform-grid Jacobi
matrix, Sturm bisection, step halving, and second-order Richardson
extrapolation.  Levels above the algebraic sector are also printed to make
the scope of QES explicit.

The calibrated command-line range is 0.5 <= zeta <= 2, 1 <= M <= 8, and
M <= levels <= 12.  The high profile uses a finer grid, but floating-point
effects need not improve monotonically.  Refinement shifts and
cross-representation differences are empirical diagnostics, not certified
enclosures.

Requirements: CPython 3.9+ and NumPy 1.24+.
"""

from __future__ import annotations

import argparse
import csv
from dataclasses import dataclass
import math
from pathlib import Path
import platform
from typing import Sequence

import numpy as np


MIN_ZETA = 0.5
MAX_ZETA = 2.0
MIN_M = 1
MAX_M = 8
MAX_LEVELS = 12

ACCEPTANCE_THRESHOLDS = {
    "default": {
        "qes_grid_gap": 1.0e-7,
        "qes_grid_refinement_shift": 4.0e-4,
        "reversal_residual": 5.0e-13,
        "polynomial_root_gap": 2.0e-6,
    },
    "high": {
        "qes_grid_gap": 2.0e-8,
        "qes_grid_refinement_shift": 1.1e-4,
        "reversal_residual": 5.0e-13,
        "polynomial_root_gap": 2.0e-6,
    },
}


def require(condition: bool, message: str) -> None:
    """Raise an optimization-safe exception when an audit fails."""

    if not condition:
        raise RuntimeError(message)


def potential(
    x: np.ndarray | float,
    zeta: float,
    m_value: int,
) -> np.ndarray | float:
    """Return the standard Razavy potential."""

    return (zeta * np.cosh(2.0 * np.asarray(x)) - m_value) ** 2


@dataclass(frozen=True)
class Settings:
    """Numerical cutoffs for one reproducibility profile."""

    grid_coarse_cells: int
    domain: float
    sturm_iterations: int


@dataclass(frozen=True)
class QESResult:
    """Finite invariant-block spectrum and algebraic diagnostics."""

    energies: np.ndarray
    parities: tuple[str, ...]
    parity_scores: np.ndarray
    polynomial_coefficients: np.ndarray
    polynomial_root_gap: float
    reversal_residual: float


@dataclass(frozen=True)
class GridResult:
    """Richardson-refined real-line spectrum and mesh diagnostics."""

    energies: np.ndarray
    parities: tuple[str, ...]
    fine_energies: np.ndarray
    maximum_shift: float
    qes_maximum_shift: float
    sturm_width: float


def settings_for_mode(mode: str) -> Settings:
    """Return frozen default or high-refinement settings."""

    if mode == "default":
        return Settings(
            grid_coarse_cells=1800,
            domain=3.5,
            sturm_iterations=88,
        )
    if mode == "high":
        return Settings(
            grid_coarse_cells=3600,
            domain=3.5,
            sturm_iterations=96,
        )
    raise ValueError(f"unknown mode {mode!r}")


def qes_jacobi_matrix(zeta: float, m_value: int) -> np.ndarray:
    """Return the symmetric M-by-M Razavy invariant block."""

    n_value = m_value - 1
    indices = np.arange(m_value, dtype=float)
    diagonal = (
        4.0 * indices * (n_value - indices)
        + 2.0 * n_value
        + 1.0
        + zeta**2
    )
    matrix = np.diag(diagonal)
    for index in range(n_value):
        entry = -2.0 * zeta * math.sqrt(
            (index + 1.0) * (n_value - index)
        )
        matrix[index, index + 1] = entry
        matrix[index + 1, index] = entry
    return matrix


def monic_spectral_polynomial(zeta: float, m_value: int) -> np.ndarray:
    """Return ascending coefficients of the monic QES polynomial P_M(E)."""

    n_value = m_value - 1
    previous = np.polynomial.Polynomial([0.0])
    current = np.polynomial.Polynomial([1.0])
    energy = np.polynomial.Polynomial([0.0, 1.0])
    for index in range(m_value):
        diagonal = (
            4.0 * index * (n_value - index)
            + 2.0 * n_value
            + 1.0
            + zeta**2
        )
        coupling = 4.0 * index * (m_value - index) * zeta**2
        following = (energy - diagonal) * current - coupling * previous
        previous, current = current, following
    return np.asarray(current.coef, dtype=float)


def qes_spectrum(zeta: float, m_value: int) -> QESResult:
    """Diagonalize the invariant block and audit its two symmetries."""

    matrix = qes_jacobi_matrix(zeta, m_value)
    energies, vectors = np.linalg.eigh(matrix)
    parity_scores = np.asarray(
        [
            float(np.dot(vectors[:, index], vectors[::-1, index]))
            for index in range(m_value)
        ]
    )
    parities = tuple("even" if score > 0.0 else "odd" for score in parity_scores)

    reversal = np.eye(m_value)[::-1]
    reversal_residual = float(
        np.linalg.norm(matrix @ reversal - reversal @ matrix, ord=np.inf)
    )
    coefficients = monic_spectral_polynomial(zeta, m_value)
    raw_polynomial_roots = np.roots(coefficients[::-1])
    root_scale = max(1.0, float(np.max(np.abs(raw_polynomial_roots))))
    require(
        float(np.max(np.abs(raw_polynomial_roots.imag)))
        <= 5.0e-10 * root_scale,
        "spectral-polynomial root finder returned a nonreal root",
    )
    polynomial_roots = np.sort(raw_polynomial_roots.real)
    polynomial_root_gap = float(
        np.max(np.abs(polynomial_roots - energies))
    )
    return QESResult(
        energies=energies,
        parities=parities,
        parity_scores=parity_scores,
        polynomial_coefficients=coefficients,
        polynomial_root_gap=polynomial_root_gap,
        reversal_residual=reversal_residual,
    )


def grid_parity_blocks(
    zeta: float,
    m_value: int,
    domain: float,
    cells: int,
) -> tuple[
    tuple[np.ndarray, np.ndarray],
    tuple[np.ndarray, np.ndarray],
]:
    """Return even and odd blocks of the centered Dirichlet grid."""

    require(domain >= 3.0, "grid domain is too short")
    require(cells >= 600, "coordinate grid is too coarse")
    step = domain / cells
    kinetic_diagonal = 2.0 / step**2
    kinetic_off_diagonal = -1.0 / step**2

    even_x = step * np.arange(cells, dtype=float)
    even_diagonal = kinetic_diagonal + potential(even_x, zeta, m_value)
    even_off = np.full(cells - 1, kinetic_off_diagonal)
    even_off[0] *= math.sqrt(2.0)

    odd_x = step * np.arange(1, cells, dtype=float)
    odd_diagonal = kinetic_diagonal + potential(odd_x, zeta, m_value)
    odd_off = np.full(cells - 2, kinetic_off_diagonal)
    return (
        (np.asarray(even_diagonal), even_off),
        (np.asarray(odd_diagonal), odd_off),
    )


def sturm_count(
    diagonal: np.ndarray,
    off_diagonal: np.ndarray,
    value: float,
) -> int:
    """Count Jacobi eigenvalues strictly below value by an LDL sequence."""

    # The kinetic diagonal grows like h^(-2), but the pivots relevant to a
    # low eigenvalue do not need a perturbation on that global scale.  Using
    # the full matrix norm here would make a finer grid *less* accurate.
    pivot_floor = (
        8.0 * np.finfo(float).eps * max(1.0, abs(value))
    )
    pivot = float(diagonal[0] - value)
    if abs(pivot) < pivot_floor:
        pivot = -pivot_floor
    count = int(pivot < 0.0)
    for index in range(1, diagonal.size):
        pivot = float(
            diagonal[index]
            - value
            - off_diagonal[index - 1] ** 2 / pivot
        )
        if abs(pivot) < pivot_floor:
            pivot = -pivot_floor
        count += int(pivot < 0.0)
    return count


def tridiagonal_bounds(
    diagonal: np.ndarray,
    off_diagonal: np.ndarray,
) -> tuple[float, float]:
    """Return Gershgorin bounds for a symmetric tridiagonal matrix."""

    radii = np.zeros_like(diagonal)
    radii[:-1] += np.abs(off_diagonal)
    radii[1:] += np.abs(off_diagonal)
    scale = max(1.0, float(np.max(np.abs(diagonal))))
    pad = 8.0 * np.finfo(float).eps * scale
    return (
        float(np.min(diagonal - radii) - pad),
        float(np.max(diagonal + radii) + pad),
    )


def lowest_tridiagonal_eigenvalues(
    diagonal: np.ndarray,
    off_diagonal: np.ndarray,
    count: int,
    iterations: int,
) -> tuple[np.ndarray, float]:
    """Return the first ``count`` Jacobi eigenvalues by Sturm bisection."""

    require(
        diagonal.size == off_diagonal.size + 1,
        "invalid tridiagonal dimensions",
    )
    require(1 <= count < diagonal.size, "invalid requested eigenvalue count")
    global_lower, global_upper = tridiagonal_bounds(diagonal, off_diagonal)
    eigenvalues: list[float] = []
    largest_width = 0.0
    for target_index in range(count):
        lower = global_lower
        upper = global_upper
        for _ in range(iterations):
            midpoint = 0.5 * (lower + upper)
            if midpoint == lower or midpoint == upper:
                break
            if sturm_count(diagonal, off_diagonal, midpoint) <= target_index:
                lower = midpoint
            else:
                upper = midpoint
        eigenvalues.append(0.5 * (lower + upper))
        largest_width = max(largest_width, upper - lower)
    return np.asarray(eigenvalues), largest_width


def grid_once(
    zeta: float,
    m_value: int,
    levels: int,
    domain: float,
    cells: int,
    iterations: int,
) -> tuple[np.ndarray, tuple[str, ...], float]:
    """Compute and merge the parity-reduced grid spectra at one step."""

    even_count = (levels + 1) // 2
    odd_count = levels // 2
    even_block, odd_block = grid_parity_blocks(
        zeta, m_value, domain, cells
    )
    even, even_width = lowest_tridiagonal_eigenvalues(
        *even_block, even_count, iterations
    )
    if odd_count:
        odd, odd_width = lowest_tridiagonal_eigenvalues(
            *odd_block, odd_count, iterations
        )
    else:
        odd = np.asarray([], dtype=float)
        odd_width = 0.0
    tagged = [(float(value), "even") for value in even]
    tagged.extend((float(value), "odd") for value in odd)
    tagged.sort(key=lambda item: item[0])
    return (
        np.asarray([item[0] for item in tagged]),
        tuple(item[1] for item in tagged),
        max(even_width, odd_width),
    )


def grid_spectrum(
    zeta: float,
    m_value: int,
    levels: int,
    settings: Settings,
) -> GridResult:
    """Refine the real-line spectrum by step halving and Richardson."""

    coarse, coarse_parities, width_coarse = grid_once(
        zeta,
        m_value,
        levels,
        settings.domain,
        settings.grid_coarse_cells,
        settings.sturm_iterations,
    )
    fine, fine_parities, width_fine = grid_once(
        zeta,
        m_value,
        levels,
        settings.domain,
        2 * settings.grid_coarse_cells,
        settings.sturm_iterations,
    )
    require(coarse_parities == fine_parities, "parity order changed on refinement")
    richardson = (4.0 * fine - coarse) / 3.0
    return GridResult(
        energies=richardson,
        parities=fine_parities,
        fine_energies=fine,
        maximum_shift=float(np.max(np.abs(richardson - fine))),
        qes_maximum_shift=float(
            np.max(np.abs(richardson[:m_value] - fine[:m_value]))
        ),
        sturm_width=max(width_coarse, width_fine),
    )


def validate_results(
    qes: QESResult,
    grid: GridResult,
    m_value: int,
    mode: str,
) -> float:
    """Apply calibrated acceptance gates and return the QES/grid gap."""

    thresholds = ACCEPTANCE_THRESHOLDS[mode]
    require(np.all(np.isfinite(qes.energies)), "non-finite QES energy")
    require(np.all(np.isfinite(grid.energies)), "non-finite grid energy")
    require(np.all(np.diff(qes.energies) > 0.0), "QES roots are not distinct")
    require(np.all(np.diff(grid.energies) > 0.0), "grid spectrum is not ordered")
    require(
        qes.parities == tuple("even" if k % 2 == 0 else "odd" for k in range(m_value)),
        "QES parity ordering is inconsistent with Sturm oscillation",
    )
    require(
        grid.parities[:m_value] == qes.parities,
        "grid and QES parity labels disagree",
    )
    require(
        float(np.max(np.abs(np.abs(qes.parity_scores) - 1.0))) < 2.0e-12,
        "QES vectors are not reversal-parity eigenvectors",
    )
    qes_grid_gap = float(
        np.max(np.abs(qes.energies - grid.energies[:m_value]))
    )
    require(
        qes_grid_gap <= thresholds["qes_grid_gap"],
        "QES/grid spectral gap exceeds the acceptance threshold",
    )
    require(
        grid.qes_maximum_shift
        <= thresholds["qes_grid_refinement_shift"],
        "QES-sector grid refinement shift exceeds the acceptance threshold",
    )
    require(
        qes.reversal_residual <= thresholds["reversal_residual"],
        "finite block does not commute with reversal",
    )
    require(
        qes.polynomial_root_gap <= thresholds["polynomial_root_gap"],
        "spectral-polynomial roots disagree with the Jacobi matrix",
    )
    return qes_grid_gap


def polynomial_string(coefficients: np.ndarray) -> str:
    """Format a monic polynomial with descending powers."""

    degree = coefficients.size - 1
    pieces: list[str] = []
    for power in range(degree, -1, -1):
        coefficient = float(coefficients[power])
        if abs(coefficient) < 5.0e-11:
            continue
        magnitude = abs(coefficient)
        if power == 0:
            body = f"{magnitude:.12g}"
        elif power == 1:
            body = "E" if abs(magnitude - 1.0) < 5.0e-11 else f"{magnitude:.12g} E"
        else:
            body = (
                f"E^{power}"
                if abs(magnitude - 1.0) < 5.0e-11
                else f"{magnitude:.12g} E^{power}"
            )
        if not pieces:
            pieces.append(body if coefficient > 0.0 else f"-{body}")
        else:
            pieces.append((" + " if coefficient > 0.0 else " - ") + body)
    return "".join(pieces)


def write_csv(
    path: Path,
    zeta: float,
    m_value: int,
    mode: str,
    settings: Settings,
    qes: QESResult,
    grid: GridResult,
) -> None:
    """Write all spectral rows and numerical provenance."""

    thresholds = ACCEPTANCE_THRESHOLDS[mode]
    fine_cells = 2 * settings.grid_coarse_cells
    coarse_step = settings.domain / settings.grid_coarse_cells
    fine_step = settings.domain / fine_cells
    qes_grid_gap = float(
        np.max(np.abs(qes.energies - grid.energies[:m_value]))
    )
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.writer(stream)
        writer.writerow(
            [
                "python_version",
                "numpy_version",
                "mode",
                "zeta",
                "M",
                "domain",
                "boundary_condition",
                "grid_coarse_cells",
                "grid_fine_cells",
                "grid_coarse_step",
                "grid_fine_step",
                "sturm_iterations",
                "extrapolation",
                "qes_grid_gap",
                "qes_grid_gap_threshold",
                "qes_sector_refinement_shift",
                "qes_sector_refinement_threshold",
                "all_level_refinement_shift",
                "sturm_bracket_width",
                "reversal_residual",
                "reversal_residual_threshold",
                "polynomial_root_gap",
                "polynomial_root_gap_threshold",
                "qes_acceptance_status",
                "level",
                "parity",
                "qes_energy",
                "grid_richardson_energy",
                "grid_fine_energy",
                "absolute_qes_grid_gap",
            ]
        )
        for index, energy in enumerate(grid.energies):
            qes_energy = qes.energies[index] if index < m_value else None
            writer.writerow(
                [
                    platform.python_version(),
                    np.__version__,
                    mode,
                    f"{zeta:.17g}",
                    m_value,
                    f"{settings.domain:.17g}",
                    "Dirichlet at x=+/-domain; parity fold at x=0",
                    settings.grid_coarse_cells,
                    fine_cells,
                    f"{coarse_step:.17g}",
                    f"{fine_step:.17g}",
                    settings.sturm_iterations,
                    "(4*fine-coarse)/3",
                    f"{qes_grid_gap:.17g}",
                    f"{thresholds['qes_grid_gap']:.17g}",
                    f"{grid.qes_maximum_shift:.17g}",
                    f"{thresholds['qes_grid_refinement_shift']:.17g}",
                    f"{grid.maximum_shift:.17g}",
                    f"{grid.sturm_width:.17g}",
                    f"{qes.reversal_residual:.17g}",
                    f"{thresholds['reversal_residual']:.17g}",
                    f"{qes.polynomial_root_gap:.17g}",
                    f"{thresholds['polynomial_root_gap']:.17g}",
                    "PASS",
                    index,
                    grid.parities[index],
                    "" if qes_energy is None else f"{qes_energy:.17g}",
                    f"{energy:.17g}",
                    f"{grid.fine_energies[index]:.17g}",
                    "" if qes_energy is None else f"{abs(qes_energy-energy):.17g}",
                ]
            )


def finite_float(text: str) -> float:
    """Parse a finite floating-point command-line value."""

    value = float(text)
    if not math.isfinite(value):
        raise argparse.ArgumentTypeError("value must be finite")
    return value


def parse_arguments(argv: Sequence[str] | None = None) -> argparse.Namespace:
    """Parse command-line options."""

    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--zeta", type=finite_float, default=1.0)
    parser.add_argument("--M", dest="m_value", type=int, default=4)
    parser.add_argument("--levels", type=int, default=8)
    parser.add_argument("--mode", choices=("default", "high"), default="default")
    parser.add_argument("--csv", type=Path)
    arguments = parser.parse_args(argv)
    if not MIN_ZETA <= arguments.zeta <= MAX_ZETA:
        parser.error(f"zeta must lie in [{MIN_ZETA}, {MAX_ZETA}]")
    if not MIN_M <= arguments.m_value <= MAX_M:
        parser.error(f"M must lie in [{MIN_M}, {MAX_M}]")
    if not arguments.m_value <= arguments.levels <= MAX_LEVELS:
        parser.error(f"levels must lie in [M, {MAX_LEVELS}]")
    return arguments


def main(argv: Sequence[str] | None = None) -> int:
    """Run the algebraic/direct spectral comparison."""

    arguments = parse_arguments(argv)
    settings = settings_for_mode(arguments.mode)
    qes = qes_spectrum(arguments.zeta, arguments.m_value)
    grid = grid_spectrum(
        arguments.zeta,
        arguments.m_value,
        arguments.levels,
        settings,
    )
    qes_grid_gap = validate_results(
        qes,
        grid,
        arguments.m_value,
        arguments.mode,
    )

    barrier = (
        (arguments.m_value - arguments.zeta) ** 2
        if arguments.m_value > arguments.zeta
        else None
    )
    print(
        f"Razavy audit: zeta={arguments.zeta:g}, M={arguments.m_value}, "
        f"mode={arguments.mode}"
    )
    print(
        "geometry: "
        + (
            f"double well, barrier V(0)={barrier:.12g}"
            if barrier is not None
            else "single well (the equality M=zeta is the flat-bottom threshold)"
        )
    )
    print(f"QES polynomial: {polynomial_string(qes.polynomial_coefficients)}")
    print("\n n  parity        QES energy      grid Richardson       |gap|        status")
    for index, energy in enumerate(grid.energies):
        if index < arguments.m_value:
            qes_energy = qes.energies[index]
            print(
                f"{index:2d}  {grid.parities[index]:5s}  "
                f"{qes_energy:18.12f}  {energy:18.12f}  "
                f"{abs(qes_energy-energy):10.3e}  algebraic"
            )
        else:
            print(
                f"{index:2d}  {grid.parities[index]:5s}  "
                f"{'--':>18s}  {energy:18.12f}  {'--':>10s}  numerical"
            )
    print("\ndiagnostics")
    print(f"  maximum QES/grid gap       {qes_grid_gap:.3e}")
    print(f"  QES-sector grid shift      {grid.qes_maximum_shift:.3e}")
    print(f"  all-level grid shift       {grid.maximum_shift:.3e}")
    print(f"  Sturm bracket width        {grid.sturm_width:.3e}")
    print(f"  reversal commutator        {qes.reversal_residual:.3e}")
    print(f"  polynomial/Jacobi root gap {qes.polynomial_root_gap:.3e}")
    print("  QES acceptance gates       PASS")
    print("  higher requested levels    diagnostic only")

    if arguments.csv is not None:
        write_csv(
            arguments.csv,
            arguments.zeta,
            arguments.m_value,
            arguments.mode,
            settings,
            qes,
            grid,
        )
        print(f"wrote {arguments.csv}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
