#!/usr/bin/env python3
"""Free massive Klein--Gordon wave packets from both Cauchy data.

Run ``python klein-gordon-wave-packets.py --help``.  Only NumPy and the
Python standard library are used; no network access or plotting is involved.
The CSV files can be plotted by any preferred plotting application.

Conventions
-----------
One spatial dimension, c = hbar = 1, mass m > 0, metric (+,-).  On a periodic
box [-L/2,L/2), N equispaced samples use NumPy's unitary (norm='ortho') FFT.
The FFT coefficients are discrete box coefficients, not continuum Fourier
amplitudes: dx*sum(abs(f)**2) = dx*sum(abs(fft(f))**2).  Both frequency
branches have the SAME spatial phase exp(+i k x):

  fhat = a+b, hhat = -i E(a-b), E = sqrt(k*k+m*m),
  a = (fhat+i*hhat/E)/2, b = (fhat-i*hhat/E)/2.

The independent cosine/sine reconstruction is the actual evolution routine:
  phihat(t) = fhat*cos(E*t) + hhat*sin(E*t)/E.
No time integration is performed.  Output times only sample this solution.

The massive pairing used by the owning articles is retained:
  rho = -Im(conj(phi)*dtphi)/m, j = Im(conj(phi)*dxphi)/m,
  Q = dx*sum(rho) = dx*sum(E/m*(abs(a)**2-abs(b)**2)).
It is signed.  abs(phi)**2, its centroid and its spread are FIELD-SHAPE
diagnostics, not a Klein--Gordon probability density.  A real nonzero solution
can have Q=0.  The m=0 limit is deliberately excluded because this pairing
divides by m; a massless notebook needs a separately stated normalization.

Default experiments include positive-frequency data, independent complex
two-data preparation, and real data with h=0.  Fixed self-check fixtures test
initial reconstruction, both frequency signs, local/global conservation,
exact composition, group velocity, and a controlled low-momentum comparison.
Box size and Fourier cutoff are refined INDEPENDENTLY.  A time-difference
residual study tests derivatives of the exact solution; its O(dt**2) error is
not a propagator time-step error.  Refined finite-box agreement is numerical
evidence, not a proof of continuum convergence for arbitrary user parameters.

Canonical derivations:
  /relativistic-qm/plane-wave-solutions/
  /relativistic-qm/klein-gordon-inner-product/
  /relativistic-qm/klein-gordon-to-schrodinger/
"""

from __future__ import annotations

import argparse
import csv
from dataclasses import dataclass
import hashlib
import json
import math
from pathlib import Path
import platform
import sys
import tempfile

import numpy as np


@dataclass(frozen=True)
class Grid:
    """Even periodic grid; signed FFT ordering includes one Nyquist mode."""

    length: float
    points: int
    mass: float

    def __post_init__(self) -> None:
        if not math.isfinite(self.length) or self.length <= 0:
            raise ValueError("box length must be finite and positive")
        if self.points < 8 or self.points % 2:
            raise ValueError("grid points must be an even integer >= 8")
        if not math.isfinite(self.mass) or self.mass <= 0:
            raise ValueError("mass must be finite and positive for the stated pairing")

    @property
    def dx(self) -> float:
        return self.length / self.points

    @property
    def x(self) -> np.ndarray:
        return -self.length / 2 + self.dx * np.arange(self.points)

    @property
    def k(self) -> np.ndarray:
        return 2 * np.pi * np.fft.fftfreq(self.points, d=self.dx)

    @property
    def energy(self) -> np.ndarray:
        return np.hypot(self.k, self.mass)


def fft(field: np.ndarray) -> np.ndarray:
    return np.fft.fft(field, norm="ortho")


def ifft(coefficients: np.ndarray) -> np.ndarray:
    return np.fft.ifft(coefficients, norm="ortho")


def norm(field: np.ndarray, grid: Grid) -> float:
    return float(np.sqrt(grid.dx * np.vdot(field, field).real))


def relative_error(actual: np.ndarray, expected: np.ndarray) -> float:
    return float(np.linalg.norm(actual - expected) / max(np.linalg.norm(expected), 1e-300))


def gaussian(x: np.ndarray, center: float, width: float, momentum: float) -> np.ndarray:
    """Continuum flat-L2-normalized Gaussian; intensity variance = width**2."""
    y = x - center
    return (2 * np.pi * width * width) ** (-0.25) * np.exp(
        -y * y / (4 * width * width) + 1j * momentum * y
    )


def prepare(
    grid: Grid, kind: str, center: float = -6, width: float = 2, momentum: float = 0.9
) -> tuple[np.ndarray, np.ndarray]:
    """Return the actual two Cauchy data f(x), h(x), not only one amplitude."""
    if kind == "real_zero_derivative":
        f = gaussian(grid.x, center, width, 0)
        return f, np.zeros_like(f)
    f = gaussian(grid.x, center, width, momentum)
    if kind == "positive_frequency":
        # This nonlocal preparation is necessary: merely setting h=0 is not
        # positive-frequency preparation.
        return f, ifft(-1j * grid.energy * fft(f))
    if kind == "independent_cauchy_data":
        h = grid.mass * (0.35 - 0.55j) * gaussian(
            grid.x, center + width, 1.2 * width, -0.4
        )
        return f, h
    raise ValueError(f"unknown preparation: {kind}")


def sectors(f: np.ndarray, h: np.ndarray, grid: Grid) -> tuple[np.ndarray, np.ndarray]:
    fhat, hhat = fft(f), fft(h)
    return (fhat + 1j * hhat / grid.energy) / 2, (fhat - 1j * hhat / grid.energy) / 2


def evolve(
    f: np.ndarray, h: np.ndarray, grid: Grid, time: float
) -> tuple[np.ndarray, np.ndarray]:
    """Exact finite-Fourier Cauchy evolution, including the returned derivative."""
    energy = grid.energy
    cosine, sine = np.cos(energy * time), np.sin(energy * time)
    fhat, hhat = fft(f), fft(h)
    field = fhat * cosine + hhat * sine / energy
    derivative = -energy * fhat * sine + hhat * cosine
    return ifft(field), ifft(derivative)


def branch_evolution(
    a: np.ndarray, b: np.ndarray, grid: Grid, time: float
) -> tuple[np.ndarray, np.ndarray]:
    """Independent exponential representation used only for cross-checking."""
    plus, minus = a * np.exp(-1j * grid.energy * time), b * np.exp(1j * grid.energy * time)
    return ifft(plus + minus), ifft(-1j * grid.energy * (plus - minus))


def diagnostics(field: np.ndarray, derivative: np.ndarray, grid: Grid) -> dict:
    dx_field = ifft(1j * grid.k * fft(field))
    intensity = np.abs(field) ** 2
    rho = -np.imag(np.conj(field) * derivative) / grid.mass
    current = np.imag(np.conj(field) * dx_field) / grid.mass
    integral = float(grid.dx * np.sum(intensity))
    centroid = float(grid.dx * np.sum(grid.x * intensity) / integral)
    variance = float(grid.dx * np.sum((grid.x - centroid) ** 2 * intensity) / integral)
    edge = np.abs(grid.x) >= 0.4 * grid.length
    derivative_intensity = np.abs(derivative) ** 2
    derivative_integral = float(grid.dx * np.sum(derivative_intensity))
    energy_density = 0.5 * (derivative_intensity + np.abs(dx_field) ** 2 + grid.mass ** 2 * intensity)
    energy_integral = float(grid.dx * np.sum(energy_density))
    return {
        "intensity": intensity,
        "charge_density": rho,
        "current": current,
        "flat_field_norm_squared": integral,
        "charge": float(grid.dx * np.sum(rho)),
        "positive_energy_functional": energy_integral,
        "field_shape_centroid": centroid,
        "field_shape_standard_deviation": math.sqrt(max(0, variance)),
        "edge_intensity_fraction_outer_ten_percent_each_side": float(
            grid.dx * np.sum(intensity[edge]) / integral
        ),
        "edge_time_derivative_squared_fraction": float(grid.dx * np.sum(
            derivative_intensity[edge]) / derivative_integral) if derivative_integral > 0 else 0.0,
        "edge_energy_fraction": float(grid.dx * np.sum(energy_density[edge]) / energy_integral),
    }


class Checks:
    def __init__(self) -> None:
        self.rows: list[dict] = []

    def bound(self, name: str, measured: float, tolerance: float) -> None:
        passed = math.isfinite(measured) and abs(measured) <= tolerance
        self.rows.append({"name": name, "measured": float(measured),
                          "absolute_tolerance": tolerance, "passed": passed})

    def condition(self, name: str, passed: bool, evidence: object) -> None:
        self.rows.append({"name": name, "passed": bool(passed), "evidence": evidence})

    @property
    def passed(self) -> bool:
        return all(row["passed"] for row in self.rows)


def structural_checks(checks: Checks) -> dict:
    grid = Grid(128, 1024, 1)
    times = [-7.5, 0, 0.125, 3.2, 24]
    summaries = {}
    for kind in ("positive_frequency", "independent_cauchy_data", "real_zero_derivative"):
        f, h = prepare(grid, kind)
        a, b = sectors(f, h, grid)
        qplus = float(grid.dx * np.sum(grid.energy / grid.mass * np.abs(a) ** 2))
        qminus = float(grid.dx * np.sum(grid.energy / grid.mass * np.abs(b) ** 2))
        qscale = qplus + qminus
        initial = diagnostics(f, h, grid)
        checks.bound(f"{kind}: FFT Parseval", abs(norm(f, grid) - norm(fft(f), grid)), 2e-14)
        f0, h0 = evolve(f, h, grid, 0)
        checks.bound(f"{kind}: reconstruct initial field", relative_error(f0, f), 2e-14)
        checks.bound(f"{kind}: reconstruct initial derivative", norm(h0 - h, grid), 2e-14)
        max_branch_error = max_charge_error = max_energy_error = max_continuity = max_flat_norm_error = 0.0
        for time in times:
            phi, dot = evolve(f, h, grid, time)
            phase_phi, phase_dot = branch_evolution(a, b, grid, time)
            max_branch_error = max(max_branch_error, relative_error(phi, phase_phi),
                                   norm(dot - phase_dot, grid) / max(norm(h, grid), 1))
            state = diagnostics(phi, dot, grid)
            max_charge_error = max(max_charge_error, abs(state["charge"] - qplus + qminus) / qscale)
            max_energy_error = max(max_energy_error, abs(state["positive_energy_functional"] /
                initial["positive_energy_functional"] - 1))
            max_flat_norm_error = max(max_flat_norm_error, abs(state["flat_field_norm_squared"] /
                initial["flat_field_norm_squared"] - 1))
            # Independently differentiate the PRODUCT j in Fourier space.  This
            # also detects product aliasing, unlike applying the PDE by hand.
            ddot = ifft(-(grid.k ** 2 + grid.mass ** 2) * fft(phi))
            rho_dot = -np.imag(np.conj(phi) * ddot) / grid.mass
            div_j = ifft(1j * grid.k * fft(state["current"])).real
            max_continuity = max(max_continuity,
                norm(rho_dot + div_j, grid) / max(norm(phi, grid) ** 2, 1e-30))
        checks.bound(f"{kind}: cosine versus two-branch evolution", max_branch_error, 2e-13)
        checks.bound(f"{kind}: position charge versus sector difference", max_charge_error, 3e-13)
        checks.bound(f"{kind}: positive energy conservation", max_energy_error, 3e-13)
        checks.bound(f"{kind}: local continuity including product differentiation", max_continuity, 2e-11)
        first_f, first_h = evolve(f, h, grid, 2.3)
        composed_f, composed_h = evolve(first_f, first_h, grid, -7.1)
        direct_f, direct_h = evolve(f, h, grid, -4.8)
        checks.bound(f"{kind}: Cauchy evolution composition", max(
            relative_error(composed_f, direct_f), norm(composed_h - direct_h, grid)), 3e-13)
        if kind == "positive_frequency":
            checks.bound("positive preparation: negative sector weight", qminus / qscale, 2e-27)
            checks.bound("positive preparation: separately conserved flat field norm", max_flat_norm_error, 3e-13)
        if kind == "real_zero_derivative":
            reversed_indices = (-np.arange(grid.points)) % grid.points
            checks.bound("real data: b(k)=conj(a(-k))", relative_error(b, np.conj(a[reversed_indices])), 2e-14)
            checks.bound("real nonzero solution: charge is zero", abs(qplus - qminus) / qscale, 2e-14)
            checks.condition("real h=0 data contain both sectors", qplus > 0.1 and qminus > 0.1,
                             {"positive_weight": qplus, "negative_weight": qminus})
            checks.condition("real mixed-sector flat field norm need not be conserved",
                             max_flat_norm_error > 0.01, max_flat_norm_error)
        summaries[kind] = {"positive_sector_weight": qplus, "negative_sector_weight": qminus,
                           "signed_charge": qplus - qminus}

    # Exactly representable plane waves independently fix both signs and k=0.
    for mode in (0, 7, -11):
        k = 2 * np.pi * mode / grid.length
        energy = math.hypot(k, grid.mass)
        f = np.exp(1j * k * grid.x) / math.sqrt(grid.length)
        for sign in (-1, 1):
            h = -1j * sign * energy * f
            phi, dot = evolve(f, h, grid, 2.75)
            expected = f * np.exp(-1j * sign * energy * 2.75)
            checks.bound(f"plane wave mode {mode}, frequency sign {sign}", max(
                relative_error(phi, expected), relative_error(dot, -1j * sign * energy * expected)), 3e-13)
    # A wrong frequency assignment must be observable, not accepted by norms alone.
    f, h = prepare(grid, "positive_frequency")
    phi, _ = evolve(f, h, grid, 0.7)
    wrong = ifft(fft(f) * np.exp(1j * grid.energy * 0.7))
    error = relative_error(wrong, phi)
    checks.condition("negative control: wrong phase sign is detected", error > 0.5, error)
    return summaries


def cauchy_comparison(
    actual: tuple[np.ndarray, np.ndarray], reference: tuple[np.ndarray, np.ndarray], mass: float
) -> float:
    """Dimensionless combined f,h error; same physical sample points required."""
    numerator = np.linalg.norm(actual[0] - reference[0]) ** 2 + np.linalg.norm(actual[1] - reference[1]) ** 2 / mass ** 2
    denominator = np.linalg.norm(reference[0]) ** 2 + np.linalg.norm(reference[1]) ** 2 / mass ** 2
    return float(np.sqrt(numerator / denominator))


def refinement_studies(checks: Checks) -> tuple[list[dict], dict]:
    rows = []
    time = 24.0
    # Cutoff refinement: same L and analytic f(x), dx shrinks.  Both Cauchy data
    # are independently prepared on each grid.  Compare identical sample points.
    reference_grid = Grid(128, 2048, 1)
    reference = evolve(*prepare(reference_grid, "positive_frequency"), reference_grid, time)
    cutoff_errors = []
    for points in (64, 128, 256, 512, 1024):
        grid = Grid(128, points, 1)
        stride = reference_grid.points // points
        state = evolve(*prepare(grid, "positive_frequency"), grid, time)
        error = cauchy_comparison(state, tuple(part[::stride] for part in reference), grid.mass)
        cutoff_errors.append(error)
        rows.append({"study": "cutoff_at_fixed_box", "length": grid.length, "points": points,
                     "dx": grid.dx, "nyquist_abs_k": np.pi / grid.dx, "time": time,
                     "comparison_region": "entire same periodic box", "relative_cauchy_error": error})
    checks.bound("cutoff refinement: finest versus 2048-point reference", cutoff_errors[-1], 2e-11)
    checks.condition("cutoff refinement resolves a deliberately coarse grid",
                     cutoff_errors[0] > 1e-4 and cutoff_errors[1] < cutoff_errors[0] / 50,
                     cutoff_errors)

    # Box refinement: fixed dx and hence fixed Nyquist cutoff.  The same
    # continuum initial data are sampled.  Compare a FIXED physical window,
    # never a changing fraction of the box or interpolated points.
    dx = 0.125
    reference_grid = Grid(512, 4096, 1)
    reference = evolve(*prepare(reference_grid, "positive_frequency"), reference_grid, time)
    reference_mask = np.abs(reference_grid.x) <= 12
    box_errors = []
    for length in (32, 64, 128, 256):
        grid = Grid(length, int(length / dx), 1)
        mask = np.abs(grid.x) <= 12
        assert np.array_equal(grid.x[mask], reference_grid.x[reference_mask])
        state = evolve(*prepare(grid, "positive_frequency"), grid, time)
        error = cauchy_comparison(tuple(part[mask] for part in state),
                                 tuple(part[reference_mask] for part in reference), grid.mass)
        box_errors.append(error)
        edge = diagnostics(*state, grid)["edge_intensity_fraction_outer_ten_percent_each_side"]
        rows.append({"study": "box_at_fixed_cutoff", "length": length, "points": grid.points,
                     "dx": dx, "nyquist_abs_k": np.pi / dx, "time": time,
                     "comparison_region": "fixed -12 <= x <= 12", "relative_cauchy_error": error,
                     "edge_intensity_fraction": edge})
    checks.bound("box refinement: largest versus length-512 reference", box_errors[-1], 2e-11)
    checks.condition("box refinement detects finite-box contamination",
                     box_errors[0] > 1e-6 and box_errors[1] < box_errors[0] / 20,
                     box_errors)

    # Derivative diagnostic only: no finite time step is used by evolve().
    grid = Grid(128, 1024, 1)
    f, h = prepare(grid, "independent_cauchy_data")
    center_time = 3.2
    phi, dot = evolve(f, h, grid, center_time)
    spectral_operator = ifft((grid.k ** 2 + grid.mass ** 2) * fft(phi))
    current = diagnostics(phi, dot, grid)["current"]
    div_j = ifft(1j * grid.k * fft(current)).real
    temporal_errors = []
    continuity_errors = []
    for dt in (0.08, 0.04, 0.02, 0.01):
        plus = evolve(f, h, grid, center_time + dt)
        minus = evolve(f, h, grid, center_time - dt)
        second_derivative = (plus[0] - 2 * phi + minus[0]) / dt ** 2
        residual = norm(second_derivative + spectral_operator, grid) / norm(spectral_operator, grid)
        rho_plus = diagnostics(*plus, grid)["charge_density"]
        rho_minus = diagnostics(*minus, grid)["charge_density"]
        charge_residual = norm((rho_plus - rho_minus) / (2 * dt) + div_j, grid) / norm(div_j, grid)
        temporal_errors.append(residual)
        continuity_errors.append(charge_residual)
        rows.append({"study": "time_difference_diagnostic_not_integration", "length": grid.length,
                     "points": grid.points, "dx": grid.dx, "nyquist_abs_k": np.pi / grid.dx,
                     "time": center_time, "diagnostic_dt": dt,
                     "relative_kg_equation_residual": residual,
                     "relative_continuity_residual": charge_residual})
    ratios = [a / b for a, b in zip(temporal_errors, temporal_errors[1:])]
    checks.condition("time-difference KG residual converges quadratically", all(3.8 < r < 4.2 for r in ratios), ratios)
    checks.bound("time-difference finest KG residual", temporal_errors[-1], 3e-5)
    checks.bound("time-difference finest continuity residual", continuity_errors[-1], 3e-4)
    return rows, {"cutoff_errors": cutoff_errors, "box_errors": box_errors,
                  "time_difference_kg_errors": temporal_errors,
                  "time_difference_continuity_errors": continuity_errors,
                  "time_integration_error": "not applicable: exact finite-mode evolution"}


def group_velocity_check(checks: Checks) -> dict:
    grid = Grid(128, 1024, 1)
    f, h = prepare(grid, "positive_frequency")
    weights = np.abs(fft(f)) ** 2
    weights /= np.sum(weights)
    velocity = grid.k / grid.energy
    mean_velocity = float(np.sum(weights * velocity))
    variance_velocity = float(np.sum(weights * (velocity - mean_velocity) ** 2))
    initial = diagnostics(f, h, grid)
    time = 24.0
    final = diagnostics(*evolve(f, h, grid, time), grid)
    expected_center = initial["field_shape_centroid"] + time * mean_velocity
    # A real Gaussian Fourier envelope has vanishing symmetrized x-v
    # covariance about its center.  These moments require negligible edges.
    expected_variance = initial["field_shape_standard_deviation"] ** 2 + time ** 2 * variance_velocity
    centroid_error = abs(final["field_shape_centroid"] - expected_center)
    variance_error = abs(final["field_shape_standard_deviation"] ** 2 - expected_variance)
    checks.bound("group velocity predicts field-shape centroid", centroid_error, 2e-10)
    checks.bound("velocity variance predicts Gaussian field-shape spread", variance_error, 2e-9)
    checks.bound("moment comparison has negligible edge intensity", final[
        "edge_intensity_fraction_outer_ten_percent_each_side"], 1e-15)
    checks.condition("all retained massive group velocities are subluminal",
                     bool(np.all(np.abs(velocity) < 1)), float(np.max(np.abs(velocity))))
    return {"weight": "normalized abs(fhat)**2, a field-shape diagnostic",
            "spectral_mean_group_velocity": mean_velocity,
            "group_velocity_at_mean_momentum": 0.9 / math.hypot(0.9, 1),
            "spectral_group_velocity_variance": variance_velocity,
            "measured_centroid": final["field_shape_centroid"], "predicted_centroid": expected_center,
            "measured_variance": final["field_shape_standard_deviation"] ** 2,
            "predicted_variance": expected_variance}


def nonrelativistic_check(checks: Checks) -> list[dict]:
    """A Gaussian tail is controlled by its moments, not a false hard cutoff."""
    grid = Grid(256, 2048, 2)
    f = gaussian(grid.x, 0, 5, 0.2)
    fhat = fft(f)
    weights = np.abs(fhat) ** 2
    weights /= np.sum(weights)
    # k^2/(E+m) is algebraically E-m and avoids subtractive loss near k=0.
    kinetic = grid.k ** 2 / (grid.energy + grid.mass)
    schrodinger = grid.k ** 2 / (2 * grid.mass)
    fourth_order = schrodinger - grid.k ** 4 / (8 * grid.mass ** 3)
    exact_energy_rms_error = float(np.sqrt(np.sum(weights * (schrodinger - kinetic) ** 2)))
    fourth_moment_bound = float(np.sqrt(np.sum(weights * grid.k ** 8))) / (8 * grid.mass ** 3)
    sixth_moment_bound = float(np.sqrt(np.sum(weights * grid.k ** 12))) / (16 * grid.mass ** 5)
    checks.condition("low-momentum fixture controls its tail", float(np.sum(weights[np.abs(grid.k) > 0.4 * grid.mass])) < 2e-9,
                     {"mean_k_over_m": 0.1, "rms_width_k_over_m": 0.05,
                      "weight_outside_abs_k_over_m_0.4": float(np.sum(weights[np.abs(grid.k) > 0.4 * grid.mass]))})
    rows = []
    for time in (0, 5, 10, 20, 40):
        exact = ifft(fhat * np.exp(-1j * kinetic * time))
        leading = ifft(fhat * np.exp(-1j * schrodinger * time))
        corrected = ifft(fhat * np.exp(-1j * fourth_order * time))
        leading_error, corrected_error = relative_error(leading, exact), relative_error(corrected, exact)
        spectral_bound = abs(time) * exact_energy_rms_error
        quartic_bound = abs(time) * fourth_moment_bound
        sixth_bound = abs(time) * sixth_moment_bound
        checks.bound(f"Schrodinger error below phase-energy bound at t={time}", max(0, leading_error - spectral_bound), 3e-14)
        checks.bound(f"phase-energy bound below quartic moment bound at t={time}", max(0, spectral_bound - quartic_bound), 3e-14)
        checks.bound(f"fourth-order envelope error below sixth moment bound at t={time}", max(0, corrected_error - sixth_bound), 3e-14)
        rows.append({"time": time, "mass": grid.mass, "mean_momentum": 0.2,
                     "momentum_standard_deviation": 0.1,
                     "schrodinger_relative_field_error": leading_error,
                     "phase_energy_rms_bound": spectral_bound,
                     "quartic_momentum_moment_bound": quartic_bound,
                     "fourth_order_relative_field_error": corrected_error,
                     "sixth_momentum_moment_bound": sixth_bound})
    return rows


def write_csv(path: Path, rows: list[dict]) -> None:
    fields = list(dict.fromkeys(key for row in rows for key in row))
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", encoding="utf-8", newline="") as stream:
        writer = csv.DictWriter(stream, fieldnames=fields)
        writer.writeheader()
        writer.writerows(rows)


def output_profiles(grid: Grid, args: argparse.Namespace, path: Path) -> list[dict]:
    summary_rows = []
    fields = ["preparation", "time", "x", "phi_real", "phi_imag", "dt_phi_real", "dt_phi_imag",
              "field_intensity_not_probability", "signed_kg_charge_density", "signed_kg_current"]
    with path.open("w", encoding="utf-8", newline="") as stream:
        writer = csv.writer(stream)
        writer.writerow(fields)
        for kind in ("positive_frequency", "independent_cauchy_data", "real_zero_derivative"):
            f, h = prepare(grid, kind, args.center, args.width, args.momentum)
            a, b = sectors(f, h, grid)
            positive_weight = float(grid.dx * np.sum(grid.energy / grid.mass * np.abs(a) ** 2))
            negative_weight = float(grid.dx * np.sum(grid.energy / grid.mass * np.abs(b) ** 2))
            combined_weights = np.abs(a) ** 2 + np.abs(b) ** 2
            cutoff_tail = float(np.sum(combined_weights[np.abs(grid.k) > 0.8 * np.pi / grid.dx]) /
                                np.sum(combined_weights))
            energy_weights = np.abs(fft(h)) ** 2 + grid.energy ** 2 * np.abs(fft(f)) ** 2
            cutoff_energy = float(np.sum(energy_weights[np.abs(grid.k) > 0.8 * np.pi / grid.dx]) /
                                  np.sum(energy_weights))
            for time in np.linspace(0, args.final_time, args.time_samples):
                phi, dot = evolve(f, h, grid, float(time))
                state = diagnostics(phi, dot, grid)
                row = {"preparation": kind, "time": float(time), "positive_sector_weight": positive_weight,
                       "negative_sector_weight": negative_weight,
                       "sector_weight_fraction_above_80_percent_nyquist": cutoff_tail,
                       "energy_fraction_above_80_percent_nyquist": cutoff_energy}
                row.update({key: value for key, value in state.items() if not isinstance(value, np.ndarray)})
                summary_rows.append(row)
                for index, x in enumerate(grid.x):
                    writer.writerow([kind, time, x, phi[index].real, phi[index].imag,
                        dot[index].real, dot[index].imag, state["intensity"][index],
                        state["charge_density"][index], state["current"][index]])
    return summary_rows


def parser() -> argparse.ArgumentParser:
    result = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    result.add_argument("--output-dir", type=Path, help="CSV/JSON destination; omitted: a new temporary directory")
    result.add_argument("--json-output", type=Path, help="override JSON report path (default: OUTPUT/summary.json)")
    result.add_argument("--profiles-csv", type=Path, help="override field-profile CSV path (default: OUTPUT/profiles.csv)")
    result.add_argument("--mass", type=float, default=1.0, help="strictly positive mass; default 1")
    result.add_argument("--length", type=float, default=128.0, help="periodic box length; default 128")
    result.add_argument("--points", type=int, default=1024, help="even number of spatial samples; default 1024")
    result.add_argument("--width", type=float, default=2.0, help="initial Gaussian intensity standard deviation; default 2")
    result.add_argument("--center", type=float, default=-6.0, help="initial Gaussian center; default -6")
    result.add_argument("--momentum", type=float, default=0.9, help="carrier momentum of positive/mixed data; default 0.9")
    result.add_argument("--final-time", type=float, default=24.0, help="last sampled time; default 24 (no integrator step)")
    result.add_argument("--time-samples", type=int, default=5, help="number of exact output times; default 5")
    return result


def main(argv: list[str] | None = None) -> int:
    argument_parser = parser()
    args = argument_parser.parse_args(argv)
    try:
        grid = Grid(args.length, args.points, args.mass)
        if not all(math.isfinite(value) for value in (args.width, args.center, args.momentum, args.final_time)):
            raise ValueError("all numerical parameters must be finite")
        if args.width <= 0 or args.time_samples < 2:
            raise ValueError("width must be positive and time-samples >= 2")
        if abs(args.momentum) >= np.pi / grid.dx:
            raise ValueError("carrier momentum must lie strictly inside the Nyquist cutoff")
        if args.width < 2 * grid.dx:
            raise ValueError("use at least two grid intervals per initial intensity width")
        if abs(args.center) + 4 * args.width >= grid.length / 2:
            raise ValueError("initial center plus four widths must lie inside the box")
        if abs(args.center + args.width) + 4 * 1.2 * args.width >= grid.length / 2:
            raise ValueError("the independent derivative datum plus four of its widths must lie inside the box")
    except ValueError as error:
        argument_parser.error(str(error))

    output_dir = args.output_dir or Path(tempfile.mkdtemp(prefix="kg-wave-packet-"))
    output_dir.mkdir(parents=True, exist_ok=True)
    paths = {
        "profiles": args.profiles_csv or output_dir / "profiles.csv",
        "diagnostics": output_dir / "diagnostics.csv",
        "refinement": output_dir / "refinement.csv",
        "nonrelativistic": output_dir / "nonrelativistic.csv",
        "summary": args.json_output or output_dir / "summary.json",
    }
    if len({path.resolve() for path in paths.values()}) != len(paths):
        argument_parser.error("output paths must be distinct")
    for path in paths.values():
        if path.exists() or path.is_symlink():
            argument_parser.error(f"output already exists and is never overwritten; choose a new path: {path}")
        if ".git" in [part.lower() for part in path.resolve().parts]:
            argument_parser.error(f"output cannot be inside .git: {path}")
        path.parent.mkdir(parents=True, exist_ok=True)

    checks = Checks()
    sectors_report = structural_checks(checks)
    refinement_rows, refinement_report = refinement_studies(checks)
    group_report = group_velocity_check(checks)
    nonrelativistic_rows = nonrelativistic_check(checks)
    diagnostic_rows = output_profiles(grid, args, paths["profiles"])
    write_csv(paths["diagnostics"], diagnostic_rows)
    write_csv(paths["refinement"], refinement_rows)
    write_csv(paths["nonrelativistic"], nonrelativistic_rows)
    report = {
        "program": "klein-gordon-wave-packets.py",
        "script_sha256": hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
        "environment": {"python": platform.python_version(), "numpy": np.__version__,
                        "platform": platform.platform(), "tested_dependencies": "Python 3.12.14; NumPy 2.3.5"},
        "units": "c = hbar = 1; 1 spatial dimension; massive pairing i/(2m)",
        "method": "unitary spatial FFT; exact cosine/sine evolution of both Cauchy data",
        "randomness": "none",
        "parameters": {"mass": args.mass, "length": args.length, "points": args.points,
                       "dx": grid.dx, "nyquist_abs_k": np.pi / grid.dx, "width": args.width,
                       "center": args.center, "momentum": args.momentum,
                       "final_time": args.final_time, "time_samples": args.time_samples},
        "interpretation": {
            "charge": "signed KG pairing, not a pointwise positive probability",
            "shape_moments": "computed with abs(phi)**2; reliable as line moments only before edge contamination",
            "box": "periodic numerical problem; field, derivative and energy edge fractions and sector/energy cutoff fractions are reported",
            "time_sampling": "no time integration; refinement dt only estimates derivatives for a residual diagnostic",
            "validation_scope": "fixed documented self-check fixtures; custom profile parameters require inspecting their edge/cutoff diagnostics",
            "massless": "excluded: the adopted massive pairing divides by m",
        },
        "sector_fixtures": sectors_report,
        "refinement": refinement_report,
        "group_velocity_fixture": group_report,
        "nonrelativistic_fixture": {"parameters": "m=2, mean k=0.2, sigma_k=0.1; rest phase removed",
                                   "norm": "flat field L2 norm", "rows": nonrelativistic_rows,
                                   "bounds": "weighted momentum moments include the Gaussian tail; no false hard cutoff"},
        "checks": checks.rows,
        "passed": checks.passed,
        "output_files": {key: str(path) for key, path in paths.items()},
        "output_sha256": {key: hashlib.sha256(path.read_bytes()).hexdigest()
                          for key, path in paths.items() if key != "summary"},
    }
    paths["summary"].write_text(json.dumps(report, indent=2, allow_nan=False) + "\n", encoding="utf-8")
    print(json.dumps({"passed": checks.passed, "checks": len(checks.rows),
                      "failed": [row["name"] for row in checks.rows if not row["passed"]],
                      "summary_json": str(paths["summary"]), "output_dir": str(output_dir)}, indent=2))
    return 0 if checks.passed else 1


if __name__ == "__main__":
    sys.exit(main())
