#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Audit numerical propagation of the dimensionless Landau-Zener problem.

The equation is

    i d psi / dx = 2 gamma (x sigma_z + sigma_x) psi.

The production calculation uses a fourth-order two-node Gauss-Magnus step.
Exponential midpoint and classical RK4 propagators provide independent
comparisons.  The script varies the finite time window and the step size
separately, audits endpoint basis conventions, 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
from pathlib import Path
from typing import Callable

import numpy as np


Array = np.ndarray


@dataclass(frozen=True)
class Configuration:
    production_step: float = 0.0025
    production_window: float = 64.0
    trajectory_gamma: float = 0.25
    trajectory_window: float = 12.0
    trajectory_points: int = 801


@dataclass(frozen=True)
class PropagationResult:
    final_state: Array
    probability: float
    final_norm_error: float
    max_norm_error: float
    actual_step: float
    steps: int
    trajectory: list[dict[str, float]]


def adiabatic_states(x: float) -> tuple[Array, Array]:
    """Return upper and lower instantaneous eigenvectors at coordinate x."""

    theta = np.arctan2(1.0, x)
    cosine = np.cos(0.5 * theta)
    sine = np.sin(0.5 * theta)
    upper = np.array([cosine, sine], dtype=complex)
    lower = np.array([-sine, cosine], dtype=complex)
    return upper, lower


def rhs(x: float, state: Array, gamma: float) -> Array:
    """Right-hand side of d psi / dx = A(x) psi."""

    first, second = state
    return -2.0j * gamma * np.array(
        [x * first + second, first - x * second],
        dtype=complex,
    )


def pauli_action(
    state: Array,
    kx: float,
    ky: float,
    kz: float,
) -> Array:
    """Apply kx sigma_x + ky sigma_y + kz sigma_z to a spinor."""

    first, second = state
    return np.array(
        [kz * first + (kx - 1.0j * ky) * second,
         (kx + 1.0j * ky) * first - kz * second],
        dtype=complex,
    )


def unitary_pauli_step(
    state: Array,
    kx: float,
    ky: float,
    kz: float,
) -> Array:
    """Apply exp[-i(kx sigma_x + ky sigma_y + kz sigma_z)]."""

    radius = float(np.sqrt(kx * kx + ky * ky + kz * kz))
    if radius < 1.0e-8:
        sinc = 1.0 - radius**2 / 6.0 + radius**4 / 120.0
    else:
        sinc = float(np.sin(radius) / radius)
    return np.cos(radius) * state - 1.0j * sinc * pauli_action(
        state,
        kx,
        ky,
        kz,
    )


def magnus4_step(x: float, state: Array, h: float, gamma: float) -> Array:
    """Fourth-order two-node Gauss-Magnus step for the linear sweep."""

    midpoint = x + 0.5 * h
    return unitary_pauli_step(
        state,
        kx=2.0 * gamma * h,
        ky=(2.0 / 3.0) * gamma**2 * h**3,
        kz=2.0 * gamma * h * midpoint,
    )


def midpoint_step(x: float, state: Array, h: float, gamma: float) -> Array:
    """Second-order exponential midpoint step."""

    midpoint = x + 0.5 * h
    return unitary_pauli_step(
        state,
        kx=2.0 * gamma * h,
        ky=0.0,
        kz=2.0 * gamma * h * midpoint,
    )


def rk4_step(x: float, state: Array, h: float, gamma: float) -> Array:
    """Classical explicit fourth-order Runge-Kutta step."""

    k1 = rhs(x, state, gamma)
    k2 = rhs(x + 0.5 * h, state + 0.5 * h * k1, gamma)
    k3 = rhs(x + 0.5 * h, state + 0.5 * h * k2, gamma)
    k4 = rhs(x + h, state + h * k3, gamma)
    return state + (h / 6.0) * (k1 + 2.0 * k2 + 2.0 * k3 + k4)


METHODS: dict[str, Callable[[float, Array, float, float], Array]] = {
    "magnus4": magnus4_step,
    "midpoint": midpoint_step,
    "rk4": rk4_step,
}


def state_observables(x: float, state: Array) -> dict[str, float]:
    """Return diabatic and adiabatic populations plus phase and norm."""

    upper, lower = adiabatic_states(x)
    p_upper = float(abs(np.vdot(upper, state)) ** 2)
    p_lower = float(abs(np.vdot(lower, state)) ** 2)
    relative_phase = float(np.angle(state[1] * np.conjugate(state[0])))
    return {
        "x": float(x),
        "p_diabatic_1": float(abs(state[0]) ** 2),
        "p_diabatic_2": float(abs(state[1]) ** 2),
        "p_adiabatic_upper": p_upper,
        "p_adiabatic_lower": p_lower,
        "relative_phase": relative_phase,
        "norm": float(np.vdot(state, state).real),
    }


def propagate(
    gamma: float,
    window: float,
    requested_step: float,
    method: str = "magnus4",
    preparation: str = "adiabatic",
    trajectory_points: int = 0,
) -> PropagationResult:
    """Propagate from -window to +window and project at the final endpoint."""

    if gamma <= 0.0:
        raise ValueError("gamma must be positive.")
    if window <= 0.0 or requested_step <= 0.0:
        raise ValueError("window and step must be positive.")
    if method not in METHODS:
        raise ValueError(f"Unknown method: {method}")

    steps = int(np.ceil(2.0 * window / requested_step))
    h = 2.0 * window / steps
    stepper = METHODS[method]

    if preparation == "adiabatic":
        _, state = adiabatic_states(-window)
    elif preparation == "diabatic":
        state = np.array([1.0, 0.0], dtype=complex)
    else:
        raise ValueError(f"Unknown preparation: {preparation}")

    record_every = 0
    if trajectory_points > 1:
        record_every = max(1, int(round(steps / (trajectory_points - 1))))

    rows: list[dict[str, float]] = []
    if record_every:
        rows.append(state_observables(-window, state))

    max_norm_error = abs(float(np.vdot(state, state).real) - 1.0)
    x = -window
    for index in range(steps):
        state = stepper(x, state, h, gamma)
        x = -window + (index + 1) * h
        norm_error = abs(float(np.vdot(state, state).real) - 1.0)
        max_norm_error = max(max_norm_error, norm_error)
        if record_every and (
            (index + 1) % record_every == 0 or index + 1 == steps
        ):
            rows.append(state_observables(x, state))

    upper, _ = adiabatic_states(window)
    if preparation == "adiabatic":
        probability = float(abs(np.vdot(upper, state)) ** 2)
    else:
        probability = float(abs(state[0]) ** 2)

    return PropagationResult(
        final_state=state,
        probability=probability,
        final_norm_error=abs(float(np.vdot(state, state).real) - 1.0),
        max_norm_error=max_norm_error,
        actual_step=h,
        steps=steps,
        trajectory=rows,
    )


def landau_zener_probability(gamma: float) -> float:
    return float(np.exp(-2.0 * np.pi * gamma))


def phase_aligned_state_error(reference: Array, candidate: Array) -> float:
    """Euclidean state error after removing norm and global phase."""

    reference = reference / np.linalg.norm(reference)
    candidate = candidate / np.linalg.norm(candidate)
    overlap = np.vdot(reference, candidate)
    if abs(overlap) > 0.0:
        candidate = candidate * np.exp(-1.0j * np.angle(overlap))
    return float(np.linalg.norm(candidate - reference))


def write_rows(path: Path, rows: list[dict[str, float | str]]) -> 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 generate_window_sweep(config: Configuration) -> list[dict[str, float]]:
    rows: list[dict[str, float]] = []
    windows = (4.0, 6.0, 8.0, 10.0, 12.0, 16.0, 20.0, 24.0, 32.0, 40.0, 48.0, 64.0)
    for gamma in (0.05, 0.25, 1.0):
        target = landau_zener_probability(gamma)
        for window in windows:
            adiabatic = propagate(
                gamma,
                window,
                config.production_step,
                preparation="adiabatic",
            )
            diabatic = propagate(
                gamma,
                window,
                config.production_step,
                preparation="diabatic",
            )
            rows.append(
                {
                    "gamma": gamma,
                    "window": window,
                    "step": adiabatic.actual_step,
                    "p_lz": target,
                    "p_adiabatic_endpoints": adiabatic.probability,
                    "signed_asymptotic_error": adiabatic.probability - target,
                    "absolute_asymptotic_error": abs(adiabatic.probability - target),
                    "p_bare_endpoints": diabatic.probability,
                    "basis_protocol_difference": abs(
                        diabatic.probability - adiabatic.probability
                    ),
                    "max_norm_error": max(
                        adiabatic.max_norm_error,
                        diabatic.max_norm_error,
                    ),
                }
            )
    return rows


def generate_gamma_sweep(config: Configuration) -> list[dict[str, float]]:
    rows: list[dict[str, float]] = []
    gammas = (0.02, 0.05, 0.1, 0.2, 0.25, 0.5, 0.75, 1.0, 1.25)
    for gamma in gammas:
        target = landau_zener_probability(gamma)
        adiabatic = propagate(
            gamma,
            config.production_window,
            config.production_step,
            preparation="adiabatic",
        )
        diabatic = propagate(
            gamma,
            config.production_window,
            config.production_step,
            preparation="diabatic",
        )
        rows.append(
            {
                "gamma": gamma,
                "window": config.production_window,
                "step": adiabatic.actual_step,
                "p_lz": target,
                "p_adiabatic_endpoints": adiabatic.probability,
                "signed_asymptotic_error": adiabatic.probability - target,
                "absolute_asymptotic_error": abs(adiabatic.probability - target),
                "relative_asymptotic_error": abs(adiabatic.probability - target)
                / target,
                "p_bare_endpoints": diabatic.probability,
                "basis_protocol_difference": abs(
                    diabatic.probability - adiabatic.probability
                ),
                "max_norm_error": max(
                    adiabatic.max_norm_error,
                    diabatic.max_norm_error,
                ),
            }
        )
    return rows


def generate_trajectory(config: Configuration) -> list[dict[str, float]]:
    result = propagate(
        config.trajectory_gamma,
        config.trajectory_window,
        config.production_step,
        method="magnus4",
        preparation="adiabatic",
        trajectory_points=config.trajectory_points,
    )
    return result.trajectory


def generate_integrator_convergence() -> list[dict[str, float | str]]:
    gamma = 0.5
    window = 12.0
    reference = propagate(gamma, window, 0.0003125, method="magnus4")
    rows: list[dict[str, float | str]] = []
    for method in ("midpoint", "magnus4", "rk4"):
        method_code = {"midpoint": 2.0, "magnus4": 4.0, "rk4": 5.0}[method]
        for step in (0.16, 0.08, 0.04, 0.02, 0.01, 0.005, 0.0025):
            current = propagate(gamma, window, step, method=method)
            rows.append(
                {
                    "method": method,
                    "method_code": method_code,
                    "requested_step": step,
                    "actual_step": current.actual_step,
                    "steps": float(current.steps),
                    "probability": current.probability,
                    "probability_error_from_reference": abs(
                        current.probability - reference.probability
                    ),
                    "phase_aligned_state_error": phase_aligned_state_error(
                        reference.final_state,
                        current.final_state,
                    ),
                    "final_norm_error": current.final_norm_error,
                    "max_norm_error": current.max_norm_error,
                    "reference_probability": reference.probability,
                }
            )
    return rows


def observed_order(errors: list[float]) -> float:
    usable = [value for value in errors if value > 20.0 * np.finfo(float).eps]
    if len(usable) < 3:
        return float("nan")
    ratios = [
        np.log(usable[index] / usable[index + 1]) / np.log(2.0)
        for index in range(len(usable) - 1)
        if usable[index + 1] > 0.0
    ]
    return float(np.median(ratios[-3:]))


def validate(
    config: Configuration,
    window_rows: list[dict[str, float]],
    convergence_rows: list[dict[str, float | str]],
) -> dict[str, float]:
    production = propagate(0.5, 12.0, config.production_step, method="magnus4")
    refined = propagate(0.5, 12.0, 0.00125, method="magnus4")
    refinement_change = phase_aligned_state_error(
        refined.final_state,
        production.final_state,
    )

    orders: dict[str, float] = {}
    for method in METHODS:
        errors = [
            float(row["phase_aligned_state_error"])
            for row in convergence_rows
            if row["method"] == method
        ]
        orders[method] = observed_order(errors)

    asymptotic_row = next(
        row
        for row in window_rows
        if row["gamma"] == 0.25 and row["window"] == 64.0
    )
    max_unitary_norm_error = max(
        float(row["max_norm_error"])
        for row in convergence_rows
        if row["method"] in ("midpoint", "magnus4")
    )
    rk4_coarse_norm_error = next(
        float(row["max_norm_error"])
        for row in convergence_rows
        if row["method"] == "rk4" and row["requested_step"] == 0.16
    )

    checks = {
        "production_refinement_state_change": refinement_change,
        "magnus_observed_order": orders["magnus4"],
        "midpoint_observed_order": orders["midpoint"],
        "rk4_observed_order": orders["rk4"],
        "gamma_0p25_window_64_error": float(
            asymptotic_row["absolute_asymptotic_error"]
        ),
        "max_unitary_norm_error": max_unitary_norm_error,
        "rk4_coarse_norm_error": rk4_coarse_norm_error,
    }

    failures: list[str] = []
    if checks["production_refinement_state_change"] > 2.0e-8:
        failures.append("production_refinement_state_change")
    if not 3.5 < checks["magnus_observed_order"] < 4.5:
        failures.append("magnus_observed_order")
    if not 1.7 < checks["midpoint_observed_order"] < 2.3:
        failures.append("midpoint_observed_order")
    if not 3.5 < checks["rk4_observed_order"] < 4.5:
        failures.append("rk4_observed_order")
    if checks["gamma_0p25_window_64_error"] > 2.0e-5:
        failures.append("gamma_0p25_window_64_error")
    if checks["max_unitary_norm_error"] > 2.0e-11:
        failures.append("max_unitary_norm_error")
    if checks["rk4_coarse_norm_error"] < 1.0e-8:
        failures.append("rk4_coarse_norm_error")

    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

    trajectory = np.genfromtxt(
        output_dir / "landau-zener-trajectory.csv",
        delimiter=",",
        names=True,
    )
    windows = np.genfromtxt(
        output_dir / "landau-zener-window-sweep.csv",
        delimiter=",",
        names=True,
    )
    fig, axes = plt.subplots(2, 1, figsize=(7.2, 7.5), constrained_layout=True)
    axes[0].plot(trajectory["x"], trajectory["p_diabatic_1"], label="diabatic 1")
    axes[0].plot(trajectory["x"], trajectory["p_adiabatic_upper"], label="adiabatic upper")
    axes[0].set_ylabel("population")
    axes[0].legend()
    for gamma in (0.05, 0.25, 1.0):
        selection = windows["gamma"] == gamma
        axes[1].loglog(
            windows["window"][selection],
            windows["absolute_asymptotic_error"][selection],
            marker="o",
            label=fr"$\gamma={gamma:g}$",
        )
    axes[1].set_xlabel(r"window $X$")
    axes[1].set_ylabel(r"$|P_X-P_{LZ}|$")
    axes[1].legend()
    for axis in axes:
        axis.grid(alpha=0.25)
    fig.savefig(output_dir / "landau-zener-dynamics-window.png", dpi=180)
    plt.close(fig)

    convergence = np.genfromtxt(
        output_dir / "landau-zener-integrator-convergence.csv",
        delimiter=",",
        names=True,
        dtype=None,
        encoding="utf-8",
    )
    fig, axes = plt.subplots(2, 1, figsize=(7.2, 7.5), constrained_layout=True)
    for method in ("midpoint", "magnus4", "rk4"):
        selection = convergence["method"] == method
        axes[0].loglog(
            convergence["actual_step"][selection],
            convergence["phase_aligned_state_error"][selection],
            marker="o",
            label=method,
        )
        axes[1].loglog(
            convergence["actual_step"][selection],
            np.maximum(convergence["max_norm_error"][selection], 1.0e-17),
            marker="o",
            label=method,
        )
    axes[0].set_ylabel("phase-aligned state error")
    axes[1].set_ylabel("maximum norm error")
    axes[1].set_xlabel("step")
    for axis in axes:
        axis.grid(alpha=0.25)
        axis.legend()
    fig.savefig(output_dir / "landau-zener-integrator-validation.png", dpi=180)
    plt.close(fig)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    default_output = (
        Path(__file__).resolve().parents[2]
        / "data"
        / "approximation-scattering-semiclassics"
    )
    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()
    window_rows = generate_window_sweep(config)
    gamma_rows = generate_gamma_sweep(config)
    trajectory_rows = generate_trajectory(config)
    convergence_rows = generate_integrator_convergence()
    checks = validate(config, window_rows, convergence_rows)

    write_rows(args.output_dir / "landau-zener-window-sweep.csv", window_rows)
    write_rows(args.output_dir / "landau-zener-gamma-sweep.csv", gamma_rows)
    write_rows(args.output_dir / "landau-zener-trajectory.csv", trajectory_rows)
    write_rows(
        args.output_dir / "landau-zener-integrator-convergence.csv",
        convergence_rows,
    )

    if args.plot:
        plot_data(args.output_dir)

    representative = next(row for row in gamma_rows if row["gamma"] == 0.25)
    runtime = time.perf_counter() - start
    print("Landau-Zener simulation")
    print(f"Python: {platform.python_version()}")
    print(f"NumPy:  {np.__version__}")
    print(f"Runtime: {runtime:.3f} s")
    print(f"Representative exact probability: {representative['p_lz']:.12e}")
    print(
        "Representative finite-window probability: "
        f"{representative['p_adiabatic_endpoints']:.12e}"
    )
    for name, value in checks.items():
        print(f"{name}: {value:.6e}")
    print("Validation: PASS")


if __name__ == "__main__":
    main()
