#!/usr/bin/env python3
"""A dealiased Fourier / integrating-factor RK4 KdV experiment.

Run without flags to print fresh JSON. --check recomputes and compares saved
results without changing files; --write explicitly regenerates results.json.
Only NumPy is required. See README.md for the numerical and scientific scope.
"""

import argparse
import json
import math
import platform
from pathlib import Path

import numpy as np

HERE = Path(__file__).resolve().parent


def pulse(x, t, kappa, x0, speed_factor=1.0):
    """The decaying-line profile, evaluated directly rather than by an FFT."""
    z = kappa * (np.asarray(x) - 4 * speed_factor * kappa**2 * t - x0)
    return 2 * kappa**2 / np.cosh(z)**2


def grid(length, points):
    if points < 12 or points % 2 or length <= 0:
        raise ValueError("Use a positive length and an even grid size >= 12.")
    x = -length / 2 + length * np.arange(points) / points
    modes = np.rint(np.fft.fftfreq(points) * points).astype(int)
    k = 2 * np.pi * modes / length
    # Strict cutoff excludes ambiguous endpoint modes in the 2/3 rule.
    keep = np.abs(modes) < points / 3
    return x, k, modes, keep


def nonlinear(v, k, keep):
    u = np.fft.ifft(v).real
    return -3j * k * np.fft.fft(u * u) * keep


def ifrk4_step(v, h, linear, nonlinearity, half=None):
    """Classical RK4 for w(s)=exp(-s*linear) v(t+s), reset at each step."""
    if half is None:
        half = np.exp(0.5 * h * linear)
    full = half * half
    a = nonlinearity(v)
    b = nonlinearity(half * (v + 0.5 * h * a))
    c = nonlinearity(half * v + 0.5 * h * b)
    d = nonlinearity(full * v + h * half * c)
    return full * v + (h / 6) * (full * a + 2 * half * (b + c) + d)


def invariants(v, k, length):
    u = np.fft.ifft(v).real
    ux = np.fft.ifft(1j * k * v).real
    dx = length / len(v)
    return dx * np.array([np.sum(u), np.sum(u**2),
                          np.sum(u**3 - 0.5 * ux**2)])


def error_norms(actual, reference, dx):
    difference = actual - reference
    l2 = math.sqrt(dx * float(np.sum(np.abs(difference)**2)))
    scale = math.sqrt(dx * float(np.sum(np.abs(reference)**2)))
    return {"relative_l2": l2 / scale,
            "relative_linf": float(np.max(np.abs(difference)) /
                                   np.max(np.abs(reference)))}


def evolve(length, points, steps, parameters):
    x, k, _, keep = grid(length, points)
    kap, x0, end = (parameters[key] for key in
                    ("kappa", "initial_center", "final_time"))
    if kap <= 0 or end <= 0 or steps <= 0:
        raise ValueError("kappa, final_time and steps must be positive.")
    h, dx = end / steps, length / points
    original = pulse(x, 0, kap, x0)
    v = np.fft.fft(original) * keep
    initial = invariants(v, k, length)
    line_values = np.array([4 * kap, 16 * kap**3 / 3, 32 * kap**5 / 5])
    worst_drift = np.zeros(3)
    linear = 1j * k**3
    half = np.exp(0.5 * h * linear)
    nonlinearity = lambda value: nonlinear(value, k, keep)
    for _ in range(steps):
        v = ifrk4_step(v, h, linear, nonlinearity, half=half) * keep
        values = invariants(v, k, length)
        worst_drift = np.maximum(worst_drift, np.abs((values - initial) / line_values))
        if not np.all(np.isfinite(values)):
            raise ArithmeticError("Non-finite solution: refine the time step.")
    actual, target = np.fft.ifft(v).real, pulse(x, end, kap, x0)
    boundary = []
    for time in (0.0, end):
        edges = np.array([-length / 2, length / 2])
        u = pulse(edges, time, kap, x0)
        ux = -2 * kap * u * np.tanh(kap * (edges - 4 * kap**2 * time - x0))
        boundary.append({"time": time,
                         "max_value_over_peak": float(np.max(np.abs(u)) / (2 * kap**2)),
                         "value_jump_over_peak": float(abs(u[1] - u[0]) / (2 * kap**2)),
                         "derivative_jump_over_kappa_peak": float(abs(ux[1] - ux[0]) / (2 * kap**3))})
    result = {"length": length, "points": points, "steps": steps, "dt": h,
              "dx": dx, "retained_modes": int(np.count_nonzero(keep)),
              "final_error": error_norms(actual, target, dx),
              "initial_projection_error": error_norms(np.fft.ifft(np.fft.fft(original) * keep).real,
                                                       original, dx),
              "initial_invariants": initial.tolist(),
              "initial_relative_line_integral_error": np.abs((initial - line_values) / line_values).tolist(),
              "max_relative_invariant_drift": worst_drift.tolist(),
              "line_boundary_mismatch": boundary}
    return result, x, actual


def method_checks():
    """Independent transformed-RK4, Fourier convolution, and linear-wave checks."""
    rng = np.random.default_rng(731)
    v = rng.normal(size=5) + 1j * rng.normal(size=5)
    linear = 1j * np.array([-3.2, -0.7, 0, 0.3, 2.5])
    h = 0.071
    nonlinearity = lambda z: (0.2 + 0.1j) * z * z
    def f(s, w):
        return np.exp(-s * linear) * nonlinearity(np.exp(s * linear) * w)
    a = f(0, v)
    b = f(h / 2, v + h * a / 2)
    c = f(h / 2, v + h * b / 2)
    d = f(h, v + h * c)
    direct = np.exp(h * linear) * (v + h * (a + 2 * b + 2 * c + d) / 6)
    stage_error = float(np.max(np.abs(ifrk4_step(v, h, linear, nonlinearity) - direct)))

    _, k, modes, keep = grid(2 * np.pi, 24)
    signal = rng.normal(size=24)
    coeff = np.fft.fft(signal) / 24 * keep
    retained = {int(m): coeff[j] for j, m in enumerate(modes) if keep[j]}
    convolution = np.zeros(24, dtype=complex)
    for j, m in enumerate(modes):
        if keep[j]:
            convolution[j] = sum(value * retained.get(int(m) - p, 0)
                                 for p, value in retained.items())
    computed = nonlinear(24 * coeff, k, keep) / 24
    convolution_error = float(np.max(np.abs(computed - (-3j * k) * convolution)))

    x, k, _, keep = grid(2 * np.pi, 32)
    v = np.fft.fft(np.cos(3 * x))
    dt = 0.013
    for _ in range(11):
        v = ifrk4_step(v, dt, 1j * k**3, lambda z: np.zeros_like(z))
    # u_t + u_xxx = 0 transports this phase as cos(k x + k^3 t).
    airy_error = float(np.max(np.abs(np.fft.ifft(v).real - np.cos(3 * x + 27 * 11 * dt))))
    if max(stage_error, convolution_error, airy_error) > 2e-12:
        raise AssertionError("Integrating-factor or Fourier method check failed.")
    return {"expanded_vs_transformed_rk4_max_error": stage_error,
            "dealiased_vs_explicit_convolution_max_error": convolution_error,
            "linear_airy_max_error": airy_error}


def analytic_checks():
    """Compare reconstruction, quadrature and finite differences independently."""
    nodes, weights = np.polynomial.legendre.leggauss(384)
    recon_errors, integral_errors, derivative_errors, marchenko_errors = [], [], [], []
    for kap in (0.5, 1.0, 1.3):
        x0, time = -0.7, 0.35
        center = x0 + 4 * kap**2 * time
        x = center + np.linspace(-3, 3, 31) / kap
        c = 2 * kap * math.exp(2 * kap * x0 + 8 * kap**3 * time)
        q = c * np.exp(-2 * kap * x)
        reconstructed = 4 * kap * q / (1 + q / (2 * kap))**2
        exact = pulse(x, time, kap, x0)
        recon_errors.append(float(np.max(np.abs(reconstructed - exact)) / (2 * kap**2)))
        def diagonal(y):
            a = c * np.exp(-2 * kap * y)
            return -a / (1 + a / (2 * kap))
        eps = 0.0001 / kap
        derivative = (-diagonal(x + 2 * eps) + 8 * diagonal(x + eps)
                      - 8 * diagonal(x - eps) + diagonal(x - 2 * eps)) / (12 * eps)
        derivative_errors.append(float(np.max(np.abs(2 * derivative - exact)) / (2 * kap**2)))
        # Independent finite-interval Gauss quadrature, tails bounded exponentially.
        z = center + 20 * nodes / kap
        u = pulse(z, time, kap, x0)
        ux = -2 * kap * u * np.tanh(kap * (z - center))
        values = 20 / kap * np.array([weights @ u, weights @ u**2,
                                      weights @ (u**3 - ux**2 / 2)])
        target = np.array([4 * kap, 16 * kap**3 / 3, 32 * kap**5 / 5])
        integral_errors.append(float(np.max(np.abs((values - target) / target))))
        # Check K + F + integral K F = 0 using quadrature, not its closed integral.
        left, right = center - 0.3 / kap, center + 0.8 / kap
        zn = left + 12 * (nodes + 1) / kap
        denom = 1 + c * np.exp(-2 * kap * left) / (2 * kap)
        kval = -c * math.exp(-kap * (left + right)) / denom
        fval = c * math.exp(-kap * (left + right))
        integrand = (-c * np.exp(-kap * (left + zn)) / denom
                     * c * np.exp(-kap * (zn + right)))
        integral = 12 / kap * float(weights @ integrand)
        marchenko_errors.append(abs(kval + fval + integral) / abs(fval))
    if max(recon_errors) > 1e-13 or max(integral_errors) > 5e-12:
        raise AssertionError("Pulse or exact-integral verification failed.")
    if max(derivative_errors) > 1e-10 or max(marchenko_errors) > 5e-12:
        raise AssertionError("Marchenko reconstruction check failed.")
    return {"kappas": [0.5, 1.0, 1.3],
            "reconstruction_relative_max_error": max(recon_errors),
            "diagonal_derivative_relative_max_error": max(derivative_errors),
            "integral_quadrature_relative_max_error": max(integral_errors),
            "marchenko_quadrature_relative_residual": max(marchenko_errors)}


def run(parameters):
    output = {"schema_version": 1, "environment": {"python": platform.python_version(),
              "numpy": np.__version__, "arithmetic": "IEEE 754 binary64"},
              "inputs": parameters, "method_checks": method_checks(),
              "analytic_checks": analytic_checks()}
    temporal = parameters["temporal"]
    output["temporal_refinement"] = [evolve(temporal["length"], temporal["points"], steps, parameters)[0]
                                      for steps in temporal["steps"]]
    errors = [r["final_error"]["relative_l2"] for r in output["temporal_refinement"]]
    output["temporal_observed_orders"] = [math.log(a / b, 2) for a, b in zip(errors, errors[1:])]
    spatial = parameters["spatial"]
    output["spatial_refinement"] = [evolve(spatial["length"], points, spatial["steps"], parameters)[0]
                                     for points in spatial["points"]]
    domain = parameters["domain"]
    output["domain_refinement"] = [evolve(length, round(length / domain["spacing"]), domain["steps"], parameters)[0]
                                    for length in domain["lengths"]]
    ref = parameters["reference"]
    record, x, numerical = evolve(ref["length"], ref["points"], ref["steps"], parameters)
    output["reference_run"] = record
    kap, x0, end = (parameters[k] for k in ("kappa", "initial_center", "final_time"))
    wrong = pulse(x, end, kap, x0, speed_factor=0.5)
    exact = pulse(x, end, kap, x0)
    ux_wrong = -2 * kap * wrong * np.tanh(kap * (x - 2 * kap**2 * end - x0))
    # For a travelling pulse with half the correct speed, the PDE residual is
    # (4*kappa^2 - 2*kappa^2) u_x; its invariant integrals still match exactly.
    output["wrong_speed_control"] = {
        "speed_factor": 0.5,
        "final_error": error_norms(wrong, exact, ref["length"] / ref["points"]),
        "pde_residual_linf_over_kappa_cubed_peak": float(np.max(np.abs(2 * kap**2 * ux_wrong)) / (2 * kap**5)),
        "relative_invariant_difference": np.abs((invariants(np.fft.fft(wrong),
             grid(ref["length"], ref["points"])[1], ref["length"]) - np.array([4 * kap, 16 * kap**3 / 3, 32 * kap**5 / 5]))
             / np.array([4 * kap, 16 * kap**3 / 3, 32 * kap**5 / 5])).tolist()}
    output["exact_samples"] = [{"x": value, "u_at_final_time": float(pulse(value, end, kap, x0))}
                               for value in parameters["sample_positions"]]
    # Success conditions apply to the documented input family, not all KdV data.
    if min(output["temporal_observed_orders"][-2:]) < 3.5:
        raise AssertionError("The fine temporal refinements did not show fourth-order behavior.")
    if record["final_error"]["relative_l2"] > 1e-8:
        raise AssertionError("Reference run is not sufficiently resolved.")
    if output["spatial_refinement"][-1]["final_error"]["relative_l2"] >= output["spatial_refinement"][0]["final_error"]["relative_l2"] / 1000:
        raise AssertionError("Spatial refinement did not reduce the error.")
    if output["domain_refinement"][-1]["final_error"]["relative_l2"] >= output["domain_refinement"][0]["final_error"]["relative_l2"] / 1000:
        raise AssertionError("Domain refinement did not reduce the error.")
    if output["wrong_speed_control"]["final_error"]["relative_l2"] < 0.5:
        raise AssertionError("The wrong-speed negative control is not distinguished.")
    return output


def compare_saved(actual, saved, atol, rtol, path="results"):
    if isinstance(actual, dict):
        if actual.keys() != saved.keys():
            raise AssertionError(f"Changed fields at {path}")
        for key in actual:
            if key == "inputs":
                if actual[key] != saved[key]:
                    raise AssertionError(f"Changed input parameters at {path}.inputs")
            elif key != "environment":
                compare_saved(actual[key], saved[key], atol, rtol, f"{path}.{key}")
    elif isinstance(actual, list):
        if len(actual) != len(saved):
            raise AssertionError(f"Changed length at {path}")
        for i, (a, b) in enumerate(zip(actual, saved)):
            compare_saved(a, b, atol, rtol, f"{path}[{i}]")
    elif isinstance(actual, (int, float)) and not isinstance(actual, bool):
        if not math.isclose(actual, saved, abs_tol=atol, rel_tol=rtol):
            raise AssertionError(f"Numerical disagreement at {path}: {actual} != {saved}")
    elif actual != saved:
        raise AssertionError(f"Changed value at {path}")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    flags = parser.add_mutually_exclusive_group()
    flags.add_argument("--check", action="store_true")
    flags.add_argument("--write", action="store_true")
    args = parser.parse_args()
    parameters = json.loads((HERE / "inputs.json").read_text())
    result = run(parameters)
    if args.check:
        saved = json.loads((HERE / "results.json").read_text())
        tolerances = parameters["saved_comparison"]
        compare_saved(result, saved, tolerances["absolute_tolerance"], tolerances["relative_tolerance"])
        print("KdV method, analytic, temporal, spatial, domain and saved-result checks passed.")
    elif args.write:
        (HERE / "results.json").write_text(json.dumps(result, indent=2, allow_nan=False) + "\n")
        print("Wrote results.json after all scientific checks passed.")
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
