#!/usr/bin/env python3
# SPDX-License-Identifier: MIT
"""Propagate a Gaussian packet through a one-dimensional rectangular barrier.

The script uses a unitary Strang split-step Fourier method, compares the
late-time transmitted norm with the momentum-weighted exact barrier
coefficient, performs coordinated spatial and temporal refinement, and writes
the retained data used by the accompanying documentation page.

NumPy is required. Matplotlib is optional.
"""

from __future__ import annotations

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

import numpy as np


@dataclass(frozen=True)
class Configuration:
    hbar: float = 1.0
    mass: float = 1.0
    barrier_height: float = 1.0
    barrier_width: float = 2.0
    box_length: float = 512.0
    grid_points: int = 8192
    time_step: float = 0.0025
    final_time: float = 190.0
    packet_center: float = -80.0
    packet_width: float = 10.0
    central_wavenumber: float = 1.0
    region_boundary: float = 15.0
    edge_width: float = 25.0

    @property
    def dx(self) -> float:
        return self.box_length / self.grid_points

    @property
    def steps(self) -> int:
        return round(self.final_time / self.time_step)


@dataclass
class PropagationResult:
    config: Configuration
    initial_state: np.ndarray
    final_state: np.ndarray
    x: np.ndarray
    k: np.ndarray
    potential: np.ndarray
    time_series: list[dict[str, float]]
    snapshots: dict[float, np.ndarray]
    spectral_transmission: float
    central_transmission: float
    negative_momentum_probability: float
    runtime_seconds: float


def grid(config: Configuration) -> tuple[np.ndarray, np.ndarray]:
    """Cell-centered positions and FFT-ordered wavenumbers."""

    j = np.arange(config.grid_points)
    x = (j - config.grid_points / 2.0 + 0.5) * config.dx
    k = 2.0 * np.pi * np.fft.fftfreq(config.grid_points, d=config.dx)
    return x, k


def rectangular_barrier(x: np.ndarray, config: Configuration) -> np.ndarray:
    return np.where(
        np.abs(x) < 0.5 * config.barrier_width,
        config.barrier_height,
        0.0,
    )


def gaussian_packet(x: np.ndarray, config: Configuration) -> np.ndarray:
    sigma = config.packet_width
    normalization = (1.0 / (2.0 * np.pi * sigma**2)) ** 0.25
    envelope = np.exp(-((x - config.packet_center) ** 2) / (4.0 * sigma**2))
    carrier = np.exp(1j * config.central_wavenumber * x)
    psi = normalization * envelope * carrier
    return psi / np.sqrt(config.dx * np.sum(np.abs(psi) ** 2))


def exact_barrier_transmission(
    k: np.ndarray,
    *,
    height: float,
    width: float,
    mass: float = 1.0,
    hbar: float = 1.0,
) -> np.ndarray:
    """Plane-wave transmission for equal asymptotic potentials."""

    k = np.asarray(k, dtype=float)
    energy = hbar**2 * k**2 / (2.0 * mass)
    transmission = np.zeros_like(energy)
    threshold_tolerance = 1.0e-12 * max(1.0, height)

    below = (energy > 0.0) & (energy < height - threshold_tolerance)
    kappa = np.sqrt(2.0 * mass * (height - energy[below])) / hbar
    denominator = 1.0 + (
        height**2
        * np.sinh(kappa * width) ** 2
        / (4.0 * energy[below] * (height - energy[below]))
    )
    transmission[below] = 1.0 / denominator

    above = energy > height + threshold_tolerance
    q = np.sqrt(2.0 * mass * (energy[above] - height)) / hbar
    denominator = 1.0 + (
        height**2
        * np.sin(q * width) ** 2
        / (4.0 * energy[above] * (energy[above] - height))
    )
    transmission[above] = 1.0 / denominator

    threshold = (energy > 0.0) & ~below & ~above
    threshold_denominator = 1.0 + mass * height * width**2 / (2.0 * hbar**2)
    transmission[threshold] = 1.0 / threshold_denominator
    return transmission


def region_probabilities(
    psi: np.ndarray,
    x: np.ndarray,
    config: Configuration,
) -> dict[str, float]:
    density = np.abs(psi) ** 2
    boundary = config.region_boundary
    left = config.dx * np.sum(density[x < -boundary])
    interaction = config.dx * np.sum(density[np.abs(x) <= boundary])
    right = config.dx * np.sum(density[x > boundary])
    norm = config.dx * np.sum(density)
    edge = config.dx * np.sum(
        density[np.abs(x) > 0.5 * config.box_length - config.edge_width]
    )
    return {
        "P_left": float(left),
        "P_interaction": float(interaction),
        "P_right": float(right),
        "norm": float(norm),
        "P_edge": float(edge),
        "accounting_error": float(abs(left + interaction + right - norm)),
    }


def spectral_benchmark(
    psi: np.ndarray,
    k: np.ndarray,
    config: Configuration,
) -> tuple[float, float]:
    momentum_state = np.fft.fft(psi, norm="ortho")
    weights = config.dx * np.abs(momentum_state) ** 2
    transmission = exact_barrier_transmission(
        np.abs(k),
        height=config.barrier_height,
        width=config.barrier_width,
        mass=config.mass,
        hbar=config.hbar,
    )
    positive = k > 0.0
    packet_transmission = float(np.sum(weights[positive] * transmission[positive]))
    negative_probability = float(np.sum(weights[k < 0.0]))
    return packet_transmission, negative_probability


def propagate(
    config: Configuration,
    *,
    record_interval: float | None = None,
    snapshot_times: tuple[float, ...] = (),
) -> PropagationResult:
    x, k = grid(config)
    potential = rectangular_barrier(x, config)
    psi = gaussian_packet(x, config)
    initial_state = psi.copy()

    half_potential = np.exp(
        -0.5j * potential * config.time_step / config.hbar
    )
    kinetic = np.exp(
        -0.5j
        * config.hbar
        * k**2
        * config.time_step
        / config.mass
    )

    record_every = None
    if record_interval is not None:
        record_every = max(1, round(record_interval / config.time_step))
    snapshot_steps = {
        round(time_value / config.time_step): time_value
        for time_value in snapshot_times
    }

    time_series: list[dict[str, float]] = []
    snapshots: dict[float, np.ndarray] = {}

    def record(step: int) -> None:
        values = region_probabilities(psi, x, config)
        time_series.append({"time": step * config.time_step, **values})

    if record_every is not None:
        record(0)
    if 0 in snapshot_steps:
        snapshots[snapshot_steps[0]] = np.abs(psi) ** 2

    start = time.perf_counter()
    for step in range(1, config.steps + 1):
        psi *= half_potential
        momentum_state = np.fft.fft(psi, norm="ortho")
        psi = np.fft.ifft(momentum_state * kinetic, norm="ortho")
        psi *= half_potential

        if record_every is not None and (
            step % record_every == 0 or step == config.steps
        ):
            record(step)
        if step in snapshot_steps:
            snapshots[snapshot_steps[step]] = np.abs(psi) ** 2

    runtime = time.perf_counter() - start
    packet_transmission, negative_probability = spectral_benchmark(
        initial_state, k, config
    )
    central_transmission = float(
        exact_barrier_transmission(
            np.array([config.central_wavenumber]),
            height=config.barrier_height,
            width=config.barrier_width,
            mass=config.mass,
            hbar=config.hbar,
        )[0]
    )
    return PropagationResult(
        config=config,
        initial_state=initial_state,
        final_state=psi,
        x=x,
        k=k,
        potential=potential,
        time_series=time_series,
        snapshots=snapshots,
        spectral_transmission=packet_transmission,
        central_transmission=central_transmission,
        negative_momentum_probability=negative_probability,
        runtime_seconds=runtime,
    )


def convergence_study() -> tuple[list[dict[str, float | int]], PropagationResult]:
    rows: list[dict[str, float | int]] = []
    finest_result: PropagationResult | None = None
    for grid_points in (1024, 2048, 4096, 8192):
        dx = 512.0 / grid_points
        config = Configuration(
            grid_points=grid_points,
            time_step=dx / 25.0,
        )
        is_finest = grid_points == 8192
        result = propagate(
            config,
            record_interval=0.5 if is_finest else None,
            snapshot_times=(0.0, 80.0, 190.0) if is_finest else (),
        )
        final = region_probabilities(result.final_state, result.x, config)
        rows.append(
            {
                "N": grid_points,
                "dx": config.dx,
                "dt": config.time_step,
                "steps": config.steps,
                "P_reflected": final["P_left"],
                "P_interaction": final["P_interaction"],
                "P_transmitted": final["P_right"],
                "norm": final["norm"],
                "P_edge": final["P_edge"],
                "P_transmitted_spectral": result.spectral_transmission,
                "transmission_error": final["P_right"] - result.spectral_transmission,
                "runtime_seconds": result.runtime_seconds,
            }
        )
        if is_finest:
            finest_result = result

    if finest_result is None:
        raise AssertionError("the finest convergence result was not produced")
    return rows, finest_result


def time_step_study() -> list[dict[str, float | int]]:
    """Refine time at fixed spatial resolution to expose the spatial plateau."""

    rows: list[dict[str, float | int]] = []
    for time_step in (0.01, 0.005, 0.0025, 0.00125):
        config = Configuration(grid_points=4096, time_step=time_step)
        result = propagate(config)
        final = region_probabilities(result.final_state, result.x, config)
        rows.append(
            {
                "N": config.grid_points,
                "dx": config.dx,
                "dt": time_step,
                "steps": config.steps,
                "P_transmitted": final["P_right"],
                "P_interaction": final["P_interaction"],
                "norm": final["norm"],
                "P_transmitted_spectral": result.spectral_transmission,
                "transmission_error": final["P_right"] - result.spectral_transmission,
                "runtime_seconds": result.runtime_seconds,
            }
        )
    return rows


def box_size_study(finest: PropagationResult) -> list[dict[str, float | int]]:
    """Change the periodic box while keeping dx and dt fixed."""

    rows: list[dict[str, float | int]] = []
    for box_length, grid_points in ((384.0, 6144), (512.0, 8192), (640.0, 10240)):
        if box_length == finest.config.box_length:
            result = finest
        else:
            config = Configuration(
                box_length=box_length,
                grid_points=grid_points,
                time_step=0.0025,
            )
            result = propagate(config)
        final = region_probabilities(result.final_state, result.x, result.config)
        rows.append(
            {
                "L": box_length,
                "N": grid_points,
                "dx": result.config.dx,
                "dt": result.config.time_step,
                "P_transmitted": final["P_right"],
                "P_interaction": final["P_interaction"],
                "P_edge": final["P_edge"],
                "norm": final["norm"],
                "runtime_seconds": result.runtime_seconds,
            }
        )
    return rows


def packet_width_study(config: Configuration) -> list[dict[str, float]]:
    k = np.linspace(1.0e-6, 3.0, 500_001)
    plane_wave = exact_barrier_transmission(
        k,
        height=config.barrier_height,
        width=config.barrier_width,
        mass=config.mass,
        hbar=config.hbar,
    )
    central = float(
        exact_barrier_transmission(
            np.array([config.central_wavenumber]),
            height=config.barrier_height,
            width=config.barrier_width,
            mass=config.mass,
            hbar=config.hbar,
        )[0]
    )
    rows: list[dict[str, float]] = []
    for sigma_x in (4.0, 5.0, 7.0, 10.0, 15.0, 20.0, 30.0, 40.0):
        weight = np.sqrt(2.0 * sigma_x**2 / np.pi) * np.exp(
            -2.0 * sigma_x**2 * (k - config.central_wavenumber) ** 2
        )
        normalization = float(np.trapezoid(weight, k))
        packet_transmission = float(np.trapezoid(weight * plane_wave, k))
        rows.append(
            {
                "sigma_x": sigma_x,
                "sigma_k": 1.0 / (2.0 * sigma_x),
                "positive_momentum_norm": normalization,
                "P_transmitted_spectral": packet_transmission,
                "T_at_k0": central,
                "packet_minus_central": packet_transmission - central,
            }
        )
    return rows


def write_rows(path: Path, rows: list[dict[str, float | int]]) -> None:
    with path.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.DictWriter(stream, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def write_snapshots(path: Path, result: PropagationResult) -> None:
    stride = 8
    x_limit = 150.0
    mask = (np.abs(result.x) <= x_limit) & (
        np.arange(result.x.size) % stride == 0
    )
    with path.open("w", newline="", encoding="utf-8") as stream:
        writer = csv.writer(stream)
        writer.writerow(["time", "x", "density"])
        for time_value in sorted(result.snapshots):
            density = result.snapshots[time_value]
            for x_value, density_value in zip(result.x[mask], density[mask]):
                writer.writerow([time_value, x_value, density_value])


def validate(
    convergence: list[dict[str, float | int]],
    finest: PropagationResult,
    time_steps: list[dict[str, float | int]],
    box_sizes: list[dict[str, float | int]],
) -> None:
    errors = [abs(float(row["transmission_error"])) for row in convergence]
    if any(right >= left for left, right in zip(errors, errors[1:])):
        raise AssertionError("coordinated refinement did not reduce transmission error")
    if errors[-1] > 1.0e-5:
        raise AssertionError("finest transmission does not match the spectral benchmark")

    final = region_probabilities(finest.final_state, finest.x, finest.config)
    if abs(final["norm"] - 1.0) > 2.0e-11:
        raise AssertionError("norm-conservation check failed")
    if final["accounting_error"] > 2.0e-14:
        raise AssertionError("regional probability accounting failed")
    if final["P_interaction"] > 1.0e-8:
        raise AssertionError("the outgoing packets are not separated")
    if final["P_edge"] > 1.0e-12:
        raise AssertionError("the packet reached the periodic boundary")
    if finest.negative_momentum_probability > 1.0e-20:
        raise AssertionError("the initial packet has appreciable negative momentum")

    max_norm_drift = max(abs(row["norm"] - 1.0) for row in finest.time_series)
    if max_norm_drift > 2.0e-11:
        raise AssertionError("time-series norm drift is too large")
    if not np.isclose(finest.central_transmission, 0.07065082485316447):
        raise AssertionError("analytic central transmission check failed")

    time_errors = [abs(float(row["transmission_error"])) for row in time_steps]
    if any(right >= left for left, right in zip(time_errors, time_errors[1:])):
        raise AssertionError("fixed-grid time-step refinement did not approach a plateau")

    box_512 = next(row for row in box_sizes if float(row["L"]) == 512.0)
    box_640 = next(row for row in box_sizes if float(row["L"]) == 640.0)
    if abs(float(box_512["P_transmitted"]) - float(box_640["P_transmitted"])) > 1.0e-9:
        raise AssertionError("box-size convergence check failed")


def plot_results(
    output_directory: Path,
    convergence: list[dict[str, float | int]],
    finest: PropagationResult,
) -> None:
    try:
        import matplotlib

        matplotlib.use("Agg")
        import matplotlib.pyplot as plt
    except ImportError as error:
        raise RuntimeError("Matplotlib is not installed") from error

    figure, axes = plt.subplots(3, 1, figsize=(8.0, 7.2), sharex=True)
    for axis, time_value in zip(axes, sorted(finest.snapshots)):
        axis.plot(finest.x, finest.snapshots[time_value], color="black")
        axis.axvspan(-1.0, 1.0, color="0.85")
        axis.set_xlim(-150.0, 150.0)
        axis.set_ylabel(r"$|\psi|^2$")
        axis.set_title(f"t = {time_value:g}")
    axes[-1].set_xlabel("x")
    figure.tight_layout()
    figure.savefig(output_directory / "wave-packet-snapshots.png", dpi=180)
    plt.close(figure)

    time_values = np.array([row["time"] for row in finest.time_series])
    figure, axes = plt.subplots(2, 1, figsize=(7.4, 7.2))
    for key, style in (
        ("P_left", "-"),
        ("P_interaction", "--"),
        ("P_right", ":"),
    ):
        axes[0].plot(
            time_values,
            [row[key] for row in finest.time_series],
            style,
            label=key,
        )
    axes[0].set_xlabel("time")
    axes[0].set_ylabel("regional probability")
    axes[0].legend()

    dx_values = np.array([float(row["dx"]) for row in convergence])
    errors = np.array([abs(float(row["transmission_error"])) for row in convergence])
    axes[1].loglog(dx_values, errors, "o-")
    axes[1].set_xlabel(r"$\Delta x$ with $\Delta t=\Delta x/25$")
    axes[1].set_ylabel("transmission error")
    figure.tight_layout()
    figure.savefig(output_directory / "wave-packet-validation.png", dpi=180)
    plt.close(figure)


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output-dir", type=Path, default=Path("wave-packet-output"))
    parser.add_argument("--skip-plot", action="store_true")
    args = parser.parse_args()
    args.output_dir.mkdir(parents=True, exist_ok=True)

    convergence, finest = convergence_study()
    time_steps = time_step_study()
    box_sizes = box_size_study(finest)
    width_rows = packet_width_study(finest.config)
    validate(convergence, finest, time_steps, box_sizes)

    write_rows(args.output_dir / "wave-packet-convergence.csv", convergence)
    write_rows(args.output_dir / "wave-packet-time-step.csv", time_steps)
    write_rows(args.output_dir / "wave-packet-box-size.csv", box_sizes)
    write_rows(args.output_dir / "wave-packet-time-series.csv", finest.time_series)
    write_rows(args.output_dir / "wave-packet-width-sweep.csv", width_rows)
    write_snapshots(args.output_dir / "wave-packet-snapshots.csv", finest)

    if not args.skip_plot:
        try:
            plot_results(args.output_dir, convergence, finest)
        except RuntimeError as error:
            print(f"Plot skipped: {error}")

    final = region_probabilities(finest.final_state, finest.x, finest.config)
    print(f"Python {platform.python_version()}")
    print(f"NumPy {np.__version__}")
    print(f"P_T propagated: {final['P_right']:.15f}")
    print(f"P_T spectral:   {finest.spectral_transmission:.15f}")
    print("Validation: PASS")
    print(f"Output: {args.output_dir.resolve()}")


if __name__ == "__main__":
    main()
