#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Test the first Born approximation against converged partial-wave scattering.

The dimensionless model is a repulsive Gaussian potential,

    V(r) = V0 exp[-(r/a)^2],

with kappa = k a and nu = 2 mu V0 a^2 / hbar^2.  A Numerov solve of the
radial Schrodinger equation supplies phase shifts and exact elastic
observables.  The script compares them with the analytic first Born result,
performs independent convergence sweeps, and writes the data retained by the
accompanying documentation page.

NumPy is required.  Matplotlib is optional and is used only with --plot.
"""

from __future__ import annotations

import argparse
import csv
import platform
import time
from dataclasses import dataclass, replace
from pathlib import Path

import numpy as np


@dataclass(frozen=True)
class Configuration:
    radial_step: float = 0.0025
    matching_radius: float = 10.0
    ell_max: int = 16
    angular_points: int = 181


@dataclass(frozen=True)
class ScatteringResult:
    kappa: float
    nu: float
    phase_shifts: np.ndarray
    sigma_exact: float
    sigma_born: float
    relative_sigma_error: float
    max_principal_phase: float
    last_partial_wave_fraction: float
    optical_theorem_residual: float


def gaussian_potential(rho: np.ndarray, nu: float) -> np.ndarray:
    """Dimensionless potential term 2 mu a^2 V / hbar^2."""

    return nu * np.exp(-(rho**2))


def riccati_bessel(
    ell: int,
    z: float,
) -> tuple[float, float, float, float]:
    """Return J_l, N_l and z derivatives for Riccati-Bessel functions."""

    if z <= 0.0:
        raise ValueError("The matching argument z must be positive.")

    # Miller downward recurrence keeps the minimal regular solution stable
    # when ell exceeds z.  Normalize with the larger of the exact l=0,1 seeds.
    top = ell + int(z) + 50
    downward = np.zeros(top + 2, dtype=float)
    downward[top] = 1.0
    for order in range(top, 0, -1):
        downward[order - 1] = (
            (2.0 * order + 1.0) * downward[order] / z
            - downward[order + 1]
        )

    j_zero = np.sin(z)
    j_one = np.sin(z) / z - np.cos(z)
    if abs(j_zero) >= abs(j_one):
        downward *= j_zero / downward[0]
    else:
        downward *= j_one / downward[1]
    j_values = downward[: ell + 1]

    n_values = [-np.cos(z)]
    if ell >= 1:
        n_values.append(-np.cos(z) / z - np.sin(z))

    for order in range(1, ell):
        factor = (2.0 * order + 1.0) / z
        n_values.append(factor * n_values[-1] - n_values[-2])

    j_ell = float(j_values[ell])
    n_ell = float(n_values[ell])
    if ell == 0:
        j_prime = float(np.cos(z))
        n_prime = float(np.sin(z))
    else:
        j_prime = float(j_values[ell - 1] - ell * j_ell / z)
        n_prime = float(n_values[ell - 1] - ell * n_ell / z)

    return j_ell, n_ell, j_prime, n_prime


def numerov_phase_shift(
    kappa: float,
    nu: float,
    ell: int,
    config: Configuration,
) -> float:
    """Integrate one regular radial solution and match its phase shift."""

    if kappa <= 0.0:
        raise ValueError("kappa must be positive.")
    if config.radial_step <= 0.0:
        raise ValueError("The radial step must be positive.")

    h = config.radial_step
    match_index = int(round(config.matching_radius / h)) - 1
    if match_index < 4:
        raise ValueError("The matching radius is too close to the origin.")

    rho = h * np.arange(1, match_index + 4, dtype=float)
    q = (
        kappa**2
        - gaussian_potential(rho, nu)
        - ell * (ell + 1.0) / rho**2
    )

    u = np.zeros_like(rho)
    power = ell + 1
    origin_correction = (nu - kappa**2) / (4.0 * ell + 6.0)
    seed = rho[:2] ** power * (1.0 + origin_correction * rho[:2] ** 2)
    u[0] = seed[0] / seed[1]
    u[1] = 1.0

    h2_over_12 = h**2 / 12.0
    for index in range(1, len(rho) - 1):
        numerator = (
            2.0 * (1.0 - 5.0 * h2_over_12 * q[index]) * u[index]
            - (1.0 + h2_over_12 * q[index - 1]) * u[index - 1]
        )
        denominator = 1.0 + h2_over_12 * q[index + 1]
        u[index + 1] = numerator / denominator

        if index % 2000 == 0:
            local_scale = max(abs(u[index]), abs(u[index + 1]))
            if local_scale > 1.0e80 or (0.0 < local_scale < 1.0e-80):
                u[: index + 2] /= local_scale

    m = match_index
    u_match = u[m]
    derivative = (
        u[m - 2]
        - 8.0 * u[m - 1]
        + 8.0 * u[m + 1]
        - u[m + 2]
    ) / (12.0 * h)

    z = kappa * rho[m]
    j_ell, n_ell, j_prime, n_prime = riccati_bessel(ell, z)
    sine_component = kappa * j_prime * u_match - j_ell * derivative
    cosine_component = kappa * n_prime * u_match - n_ell * derivative
    raw_phase = float(np.arctan2(sine_component, cosine_component))

    # Scattering observables are invariant under delta_l -> delta_l + pi.
    return float((raw_phase + 0.5 * np.pi) % np.pi - 0.5 * np.pi)


def phase_shifts(
    kappa: float,
    nu: float,
    config: Configuration,
) -> np.ndarray:
    return np.array(
        [
            numerov_phase_shift(kappa, nu, ell, config)
            for ell in range(config.ell_max + 1)
        ],
        dtype=float,
    )


def legendre_values(mu: np.ndarray, ell_max: int) -> np.ndarray:
    values = np.empty((ell_max + 1, len(mu)), dtype=float)
    values[0] = 1.0
    if ell_max >= 1:
        values[1] = mu
    for ell in range(1, ell_max):
        values[ell + 1] = (
            (2.0 * ell + 1.0) * mu * values[ell]
            - ell * values[ell - 1]
        ) / (ell + 1.0)
    return values


def exact_amplitude(
    theta: np.ndarray,
    kappa: float,
    phases: np.ndarray,
) -> np.ndarray:
    mu = np.cos(theta)
    polynomials = legendre_values(mu, len(phases) - 1)
    ell = np.arange(len(phases), dtype=float)
    coefficients = (
        (2.0 * ell + 1.0)
        * np.exp(1j * phases)
        * np.sin(phases)
        / kappa
    )
    return np.sum(coefficients[:, None] * polynomials, axis=0)


def born_amplitude(theta: np.ndarray, kappa: float, nu: float) -> np.ndarray:
    q = 2.0 * kappa * np.sin(0.5 * theta)
    return -0.25 * np.sqrt(np.pi) * nu * np.exp(-0.25 * q**2)


def born_total_cross_section(kappa: float, nu: float) -> float:
    if abs(kappa) < 1.0e-8:
        return float(0.25 * np.pi**2 * nu**2)
    return float(
        np.pi**2
        * nu**2
        * (-np.expm1(-2.0 * kappa**2))
        / (8.0 * kappa**2)
    )


def scattering_result(
    kappa: float,
    nu: float,
    config: Configuration,
) -> ScatteringResult:
    phases = phase_shifts(kappa, nu, config)
    ell = np.arange(len(phases), dtype=float)
    partial_terms = (2.0 * ell + 1.0) * np.sin(phases) ** 2
    sigma_exact = float(4.0 * np.pi * np.sum(partial_terms) / kappa**2)
    sigma_born = born_total_cross_section(kappa, nu)
    relative_error = abs(sigma_born - sigma_exact) / max(sigma_exact, 1.0e-300)

    forward = exact_amplitude(np.array([0.0]), kappa, phases)[0]
    optical_rhs = kappa * sigma_exact / (4.0 * np.pi)
    optical_residual = abs(forward.imag - optical_rhs) / max(
        abs(optical_rhs),
        1.0e-300,
    )
    tail_fraction = float(partial_terms[-1] / max(np.sum(partial_terms), 1.0e-300))

    return ScatteringResult(
        kappa=kappa,
        nu=nu,
        phase_shifts=phases,
        sigma_exact=sigma_exact,
        sigma_born=sigma_born,
        relative_sigma_error=float(relative_error),
        max_principal_phase=float(np.max(np.abs(phases))),
        last_partial_wave_fraction=tail_fraction,
        optical_theorem_residual=float(optical_residual),
    )


def write_rows(path: Path, fieldnames: list[str], 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=fieldnames)
        writer.writeheader()
        writer.writerows(rows)


def coupling_grid() -> np.ndarray:
    anchors = np.array([0.02, 0.05, 0.1, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0])
    return np.unique(np.concatenate((np.geomspace(0.02, 8.0, 29), anchors)))


def generate_regime_data(
    config: Configuration,
) -> tuple[list[dict[str, float]], dict[tuple[float, float], ScatteringResult]]:
    rows: list[dict[str, float]] = []
    results: dict[tuple[float, float], ScatteringResult] = {}
    kappas = (0.5, 1.0, 2.0, 4.0)

    for nu in coupling_grid():
        row: dict[str, float] = {"nu": float(nu)}
        for kappa in kappas:
            result = scattering_result(kappa, float(nu), config)
            results[(kappa, float(nu))] = result
            suffix = str(kappa).replace(".", "p")
            row[f"sigma_exact_k{suffix}"] = result.sigma_exact
            row[f"sigma_born_k{suffix}"] = result.sigma_born
            row[f"relative_error_k{suffix}"] = result.relative_sigma_error
            row[f"max_phase_k{suffix}"] = result.max_principal_phase
            row[f"tail_fraction_k{suffix}"] = result.last_partial_wave_fraction
        rows.append(row)
    return rows, results


def generate_angular_data(
    config: Configuration,
) -> list[dict[str, float]]:
    kappa = 1.5
    strengths = (0.25, 1.0, 4.0)
    theta = np.linspace(0.0, np.pi, config.angular_points)
    rows: list[dict[str, float]] = [
        {"angle_degrees": float(value)} for value in np.degrees(theta)
    ]

    for nu in strengths:
        phases = phase_shifts(kappa, nu, config)
        exact = np.abs(exact_amplitude(theta, kappa, phases)) ** 2
        born = np.abs(born_amplitude(theta, kappa, nu)) ** 2
        suffix = str(nu).replace(".", "p")
        for index in range(len(theta)):
            rows[index][f"exact_nu{suffix}"] = float(exact[index])
            rows[index][f"born_nu{suffix}"] = float(born[index])
    return rows


def generate_strength_summary(
    config: Configuration,
) -> list[dict[str, float]]:
    kappa = 1.5
    rows: list[dict[str, float]] = []
    for nu in (0.05, 0.1, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0):
        result = scattering_result(kappa, nu, config)
        rows.append(
            {
                "kappa": kappa,
                "nu": nu,
                "sigma_exact": result.sigma_exact,
                "sigma_born": result.sigma_born,
                "relative_sigma_error": result.relative_sigma_error,
                "max_principal_phase": result.max_principal_phase,
                "last_partial_wave_fraction": result.last_partial_wave_fraction,
                "optical_theorem_residual": result.optical_theorem_residual,
            }
        )
    return rows


def generate_step_convergence(
    config: Configuration,
) -> list[dict[str, float]]:
    kappa = 4.0
    nu = 4.0
    reference_config = replace(config, radial_step=0.00125, ell_max=16)
    reference = scattering_result(kappa, nu, reference_config)
    rows: list[dict[str, float]] = []
    for step in (0.04, 0.02, 0.01, 0.005, 0.0025, 0.00125):
        current = scattering_result(
            kappa,
            nu,
            replace(config, radial_step=step, ell_max=16),
        )
        rows.append(
            {
                "radial_step": step,
                "sigma_exact": current.sigma_exact,
                "relative_reference_difference": abs(
                    current.sigma_exact - reference.sigma_exact
                )
                / reference.sigma_exact,
                "phase_zero": current.phase_shifts[0],
                "reference_sigma": reference.sigma_exact,
            }
        )
    return rows


def generate_cutoff_convergence(
    config: Configuration,
) -> list[dict[str, float]]:
    kappa = 4.0
    nu = 4.0
    reference = scattering_result(kappa, nu, replace(config, ell_max=18))
    rows: list[dict[str, float]] = []
    for ell_max in (2, 4, 6, 8, 10, 12, 14, 16, 18):
        current = scattering_result(kappa, nu, replace(config, ell_max=ell_max))
        rows.append(
            {
                "ell_max": float(ell_max),
                "sigma_exact": current.sigma_exact,
                "relative_reference_difference": abs(
                    current.sigma_exact - reference.sigma_exact
                )
                / reference.sigma_exact,
                "last_partial_wave_fraction": current.last_partial_wave_fraction,
                "reference_sigma": reference.sigma_exact,
            }
        )
    return rows


def generate_matching_convergence(
    config: Configuration,
) -> list[dict[str, float]]:
    kappa = 4.0
    nu = 4.0
    reference = scattering_result(
        kappa,
        nu,
        replace(config, matching_radius=12.0, ell_max=16),
    )
    rows: list[dict[str, float]] = []
    for radius in (4.0, 5.0, 6.0, 8.0, 10.0, 12.0):
        current = scattering_result(
            kappa,
            nu,
            replace(config, matching_radius=radius, ell_max=16),
        )
        rows.append(
            {
                "matching_radius": radius,
                "sigma_exact": current.sigma_exact,
                "relative_reference_difference": abs(
                    current.sigma_exact - reference.sigma_exact
                )
                / reference.sigma_exact,
                "potential_at_match": float(nu * np.exp(-(radius**2))),
                "reference_sigma": reference.sigma_exact,
            }
        )
    return rows


def validate(config: Configuration) -> dict[str, float]:
    free_phases_low = phase_shifts(0.5, 0.0, config)
    free_phases_mid = phase_shifts(1.5, 0.0, replace(config, ell_max=10))
    free_error = float(
        max(
            np.max(np.abs(free_phases_low)),
            np.max(np.abs(free_phases_mid)),
        )
    )

    wronskian_error = 0.0
    for z in (5.0, 10.0, 20.0, 40.0):
        for ell in range(config.ell_max + 1):
            j_ell, n_ell, j_prime, n_prime = riccati_bessel(ell, z)
            wronskian_error = max(
                wronskian_error,
                abs(j_ell * n_prime - j_prime * n_ell - 1.0),
            )

    weak = scattering_result(1.5, 0.01, config)
    production = scattering_result(4.0, 4.0, config)
    refined = scattering_result(
        4.0,
        4.0,
        replace(config, radial_step=0.00125, ell_max=16),
    )
    refinement_change = abs(production.sigma_exact - refined.sigma_exact) / refined.sigma_exact

    checks = {
        "free_phase_error": free_error,
        "bessel_wronskian_error": float(wronskian_error),
        "weak_coupling_sigma_error": weak.relative_sigma_error,
        "production_refinement_change": float(refinement_change),
        "production_tail_fraction": production.last_partial_wave_fraction,
        "production_optical_residual": production.optical_theorem_residual,
    }

    limits = {
        "free_phase_error": 2.0e-7,
        "bessel_wronskian_error": 2.0e-12,
        "weak_coupling_sigma_error": 2.0e-2,
        "production_refinement_change": 2.0e-5,
        "production_tail_fraction": 1.0e-10,
        "production_optical_residual": 2.0e-13,
    }
    failures = [name for name, value in checks.items() if value > limits[name]]
    if failures:
        details = ", ".join(f"{name}={checks[name]:.3e}" for name in failures)
        raise RuntimeError(f"Validation failed: {details}")
    return checks


def plot_data(output_dir: Path) -> None:
    try:
        import matplotlib.pyplot as plt
    except ImportError as exc:
        raise RuntimeError("Matplotlib is required only when --plot is used.") from exc

    regime = np.genfromtxt(
        output_dir / "born-regime-curves.csv",
        delimiter=",",
        names=True,
    )
    fig, axes = plt.subplots(2, 1, figsize=(7.2, 8.0), constrained_layout=True)
    for kappa in (0.5, 1.0, 2.0, 4.0):
        suffix = str(kappa).replace(".", "p")
        axes[0].loglog(
            regime["nu"],
            regime[f"relative_error_k{suffix}"],
            label=fr"$\kappa={kappa:g}$",
        )
        axes[1].loglog(
            regime["nu"],
            regime[f"max_phase_k{suffix}"],
            label=fr"$\kappa={kappa:g}$",
        )
    axes[0].set_ylabel(r"$|\sigma_B-\sigma|/\sigma$")
    axes[1].set_ylabel(r"$\max_\ell |\delta_\ell|$")
    axes[1].set_xlabel(r"$\nu$")
    for axis in axes:
        axis.grid(alpha=0.25)
        axis.legend()
    fig.savefig(output_dir / "born-regime-curves.png", dpi=180)
    plt.close(fig)

    angular = np.genfromtxt(
        output_dir / "born-angular-distribution.csv",
        delimiter=",",
        names=True,
    )
    fig, axis = plt.subplots(figsize=(7.2, 4.8), constrained_layout=True)
    for nu in (0.25, 1.0, 4.0):
        suffix = str(nu).replace(".", "p")
        axis.semilogy(
            angular["angle_degrees"],
            angular[f"exact_nu{suffix}"],
            label=fr"exact, $\nu={nu:g}$",
        )
        axis.semilogy(
            angular["angle_degrees"],
            angular[f"born_nu{suffix}"],
            linestyle="--",
            label=fr"Born, $\nu={nu:g}$",
        )
    axis.set_xlabel(r"$\theta$ (degrees)")
    axis.set_ylabel(r"$(d\sigma/d\Omega)/a^2$")
    axis.grid(alpha=0.25)
    axis.legend(ncol=2)
    fig.savefig(output_dir / "born-angular-distribution.png", dpi=180)
    plt.close(fig)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    default_output = Path("results")
    parser.add_argument("--output-dir", type=Path, default=default_output)
    parser.add_argument("--plot", action="store_true")
    args = parser.parse_args()

    start = time.perf_counter()
    config = Configuration()
    checks = validate(config)

    regime_rows, _ = generate_regime_data(config)
    angular_rows = generate_angular_data(config)
    strength_rows = generate_strength_summary(config)
    step_rows = generate_step_convergence(config)
    cutoff_rows = generate_cutoff_convergence(config)
    matching_rows = generate_matching_convergence(config)

    write_rows(
        args.output_dir / "born-regime-curves.csv",
        list(regime_rows[0]),
        regime_rows,
    )
    write_rows(
        args.output_dir / "born-angular-distribution.csv",
        list(angular_rows[0]),
        angular_rows,
    )
    write_rows(
        args.output_dir / "born-strength-summary.csv",
        list(strength_rows[0]),
        strength_rows,
    )
    write_rows(
        args.output_dir / "born-numerov-convergence.csv",
        list(step_rows[0]),
        step_rows,
    )
    write_rows(
        args.output_dir / "born-partial-wave-convergence.csv",
        list(cutoff_rows[0]),
        cutoff_rows,
    )
    write_rows(
        args.output_dir / "born-matching-convergence.csv",
        list(matching_rows[0]),
        matching_rows,
    )

    if args.plot:
        plot_data(args.output_dir)

    representative = scattering_result(1.5, 1.0, config)
    runtime = time.perf_counter() - start
    print("Born approximation numerical test")
    print(f"Python: {platform.python_version()}")
    print(f"NumPy:  {np.__version__}")
    print(f"Runtime: {runtime:.3f} s")
    print(f"Representative sigma exact: {representative.sigma_exact:.12e}")
    print(f"Representative sigma Born:  {representative.sigma_born:.12e}")
    for name, value in checks.items():
        print(f"{name}: {value:.6e}")
    print("Validation: PASS")


if __name__ == "__main__":
    main()
