#!/usr/bin/env python3
"""Reproduce the two-instanton modified-Mathieu NS capstone.

Source convention:

    H_s = -hbar^2 d^2/dQ^2 + 2 Lambda^2 cosh(Q),
    H_s psi = E_s psi,  psi in L^2(R).

The script solves the model-specific Nekrasov–Shatashvili condition

    dF_NS/da = pi*hbar*(n + 1/2)

at perturbative, one-instanton, and two-instanton order, then applies the
quantum Matone relation.  The default Lambda=hbar=1 results are compared
with the independently refined DCHE targets from the chapter.  The displayed
truncation shifts are diagnostics, not rigorous remainder bounds.
"""

from __future__ import annotations

import argparse
import math
import platform
import sys
from dataclasses import dataclass

import numpy as np

try:
    import scipy
    from scipy.optimize import brentq
    from scipy.special import loggamma
except ImportError as exc:  # pragma: no cover - user-facing dependency guard
    raise SystemExit(
        "modified-mathieu-ns-capstone.py requires SciPy and NumPy"
    ) from exc


DCHE_TARGETS = np.array(
    [
        3.059174596896,
        5.285125967380,
        7.714579573227,
        10.327666944456,
    ],
    dtype=float,
)


@dataclass(frozen=True)
class InstantonTerms:
    first: float
    second: float
    derivative_first: float
    derivative_second: float


def require(condition: bool, message: str) -> None:
    if not condition:
        raise ValueError(message)


def gamma_term(a_value: float, hbar: float, scale: float) -> float:
    """Return the real source-convention gamma(a,hbar,Lambda)."""

    logarithm = 0.5 * a_value * math.log(hbar * hbar / (scale * scale))
    phase = hbar * float(loggamma(1.0 + 1j * a_value / hbar).imag)
    return logarithm - 0.25 * math.pi * hbar + phase


def instanton_terms(
    a_value: float, hbar: float, scale: float
) -> InstantonTerms:
    a_squared = a_value * a_value
    hbar_squared = hbar * hbar
    first_denominator = a_squared + hbar_squared
    first = -2.0 * scale**4 / first_denominator
    derivative_first = 4.0 * a_value * scale**4 / first_denominator**2

    numerator = 7.0 * hbar_squared - 5.0 * a_squared
    denominator = (
        first_denominator**3 * (a_squared + 4.0 * hbar_squared)
    )
    second = scale**8 * numerator / denominator
    derivative_denominator = denominator * (
        6.0 * a_value / first_denominator
        + 2.0 * a_value / (a_squared + 4.0 * hbar_squared)
    )
    derivative_second = scale**8 * (
        -10.0 * a_value * denominator
        - numerator * derivative_denominator
    ) / denominator**2
    return InstantonTerms(
        first, second, derivative_first, derivative_second
    )


def free_energy_derivative(
    a_value: float, hbar: float, scale: float, instanton_order: int
) -> float:
    value = 2.0 * gamma_term(a_value, hbar, scale)
    terms = instanton_terms(a_value, hbar, scale)
    if instanton_order >= 1:
        value += terms.derivative_first
    if instanton_order >= 2:
        value += terms.derivative_second
    return value


def solve_flat_coordinate(
    index: int, hbar: float, scale: float, instanton_order: int
) -> float:
    target = math.pi * hbar * (index + 0.5)

    def equation(a_value: float) -> float:
        return (
            free_energy_derivative(
                a_value, hbar, scale, instanton_order
            )
            - target
        )

    lower = 1.0e-10 * hbar
    require(equation(lower) < 0.0, "unexpected lower NS bracket sign")
    upper = max(hbar, scale)
    for _ in range(100):
        if equation(upper) > 0.0:
            return float(
                brentq(
                    equation,
                    lower,
                    upper,
                    xtol=5.0e-14,
                    rtol=2.0e-14,
                    maxiter=160,
                )
            )
        upper *= 1.35
    raise ValueError(f"could not bracket NS root {index}")


def matone_energy(
    a_value: float, hbar: float, scale: float, instanton_order: int
) -> float:
    """Apply E_s=2u_s=a^2/4-(Lambda/4)d_Lambda F_inst."""

    terms = instanton_terms(a_value, hbar, scale)
    energy = 0.25 * a_value * a_value
    if instanton_order >= 1:
        energy -= terms.first
    if instanton_order >= 2:
        energy -= 2.0 * terms.second
    return energy


def parse_arguments() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--hbar", type=float, default=1.0)
    parser.add_argument("--Lambda", dest="scale", type=float, default=1.0)
    parser.add_argument("--levels", type=int, default=4)
    return parser.parse_args()


def main() -> int:
    arguments = parse_arguments()
    require(math.isfinite(arguments.hbar) and arguments.hbar > 0.0,
            "--hbar must be positive and finite")
    require(math.isfinite(arguments.scale) and arguments.scale > 0.0,
            "--Lambda must be positive and finite")
    require(1 <= arguments.levels <= 12, "--levels must lie between 1 and 12")

    rows: list[list[float]] = []
    coordinates: list[list[float]] = []
    for index in range(arguments.levels):
        level_energies = []
        level_coordinates = []
        for order in range(3):
            a_value = solve_flat_coordinate(
                index, arguments.hbar, arguments.scale, order
            )
            level_coordinates.append(a_value)
            level_energies.append(
                matone_energy(
                    a_value, arguments.hbar, arguments.scale, order
                )
            )
        rows.append(level_energies)
        coordinates.append(level_coordinates)

    print("Modified-Mathieu NS capstone")
    print(f"Python {platform.python_version()}; NumPy {np.__version__}; "
          f"SciPy {scipy.__version__}")
    print(
        f"Lambda={arguments.scale:g}; hbar_s={arguments.hbar:g}; "
        f"levels={arguments.levels}"
    )
    print("quantization: d_a F_NS = pi*hbar_s*(n+1/2)")
    print("energy map: E_s=a_s^2/4-(Lambda/4)d_Lambda F_NS^inst")
    print()

    calibrated = (
        abs(arguments.scale - 1.0) <= 1.0e-14
        and abs(arguments.hbar - 1.0) <= 1.0e-14
        and arguments.levels <= len(DCHE_TARGETS)
    )
    if calibrated:
        print(
            " n       perturbative       one instanton      two instantons  "
            "       DCHE target        final gap"
        )
    else:
        print(" n       perturbative       one instanton      two instantons")

    final_gaps: list[float] = []
    for index, energies in enumerate(rows):
        prefix = (
            f"{index:2d}  {energies[0]:18.12f}  {energies[1]:18.12f}  "
            f"{energies[2]:18.12f}"
        )
        if calibrated:
            gap = abs(energies[2] - float(DCHE_TARGETS[index]))
            final_gaps.append(gap)
            print(
                f"{prefix}  {DCHE_TARGETS[index]:18.12f}  {gap:12.3e}"
            )
        else:
            print(prefix)

    print()
    maximum_ns_residual = 0.0
    for index, level_coordinates in enumerate(coordinates):
        target = math.pi * arguments.hbar * (index + 0.5)
        for order, a_value in enumerate(level_coordinates):
            maximum_ns_residual = max(
                maximum_ns_residual,
                abs(
                    free_energy_derivative(
                        a_value, arguments.hbar, arguments.scale, order
                    )
                    - target
                ),
            )
    print(f"maximum NS root residual: {maximum_ns_residual:.3e}")
    if calibrated:
        first_gaps = [abs(row[0] - DCHE_TARGETS[i]) for i, row in enumerate(rows)]
        second_gaps = [abs(row[1] - DCHE_TARGETS[i]) for i, row in enumerate(rows)]
        require(
            all(second_gaps[i] < first_gaps[i] for i in range(len(rows))),
            "one-instanton energies did not improve every calibrated level",
        )
        require(
            all(final_gaps[i] < second_gaps[i] for i in range(len(rows))),
            "two-instanton energies did not improve every calibrated level",
        )
        require(max(final_gaps) <= 6.0e-5,
                "two-instanton/DCHE regression gate failed")
        print(f"maximum two-instanton/DCHE gap: {max(final_gaps):.3e}")
    else:
        print(
            "Calibration warning: this parameter choice has no embedded "
            "DCHE target; the finite instanton truncation has a "
            "parameter-dependent domain of usefulness."
        )
    require(maximum_ns_residual <= 2.0e-11, "NS root residual gate failed")
    print(
        "Scope: finite instanton truncation in Lambda^4; this does not "
        "Borel resum the hbar expansion or bound the omitted instantons."
    )
    print("Regression gates: PASS")
    return 0


if __name__ == "__main__":
    try:
        sys.exit(main())
    except ValueError as error:
        print(f"error: {error}", file=sys.stderr)
        sys.exit(2)
