#!/usr/bin/env python3
"""Exact two-soliton KdV collision and a finite periodic evolution benchmark.

Needs NumPy and the unchanged ../kdv/experiment.py companion. No options prints
fresh JSON; --check compares saved results without writing; --write regenerates.
"""

import argparse
from fractions import Fraction
import json
import math
from pathlib import Path
import platform
import types

import numpy as np

HERE = Path(__file__).resolve().parent
SOLVER_PATH = HERE.parent / "kdv" / "experiment.py"


def load_solver():
    if not SOLVER_PATH.is_file():
        raise FileNotFoundError(
            "Download /computations/kdv/experiment.py into the sibling 'kdv' "
            "directory, alongside 'kdv-collision'. See README.md.")
    # Compile the source directly: importing the companion creates no __pycache__.
    module = types.ModuleType("kdv_shared_solver")
    module.__file__ = str(SOLVER_PATH)
    exec(compile(SOLVER_PATH.read_text(), str(SOLVER_PATH), "exec"), module.__dict__)
    return module


def parameters(kappas):
    k = np.asarray(kappas, dtype=float)
    if k.shape != (2,) or not (k[0] > k[1] > 0):
        raise ValueError("Require two distinct kappas, kappa1 > kappa2 > 0.")
    interaction = ((k[0] - k[1]) / (k[0] + k[1]))**2
    phases = np.log(interaction) / (4 * k)
    incoming = phases + np.array([0.0, -np.log(interaction) / (2 * k[1])])
    outgoing = phases + np.array([-np.log(interaction) / (2 * k[0]), 0.0])
    return k, interaction, phases, incoming, outgoing


def tau_weights(x, t, kappas):
    k, interaction, phases, _, _ = parameters(kappas)
    x = np.asarray(x, dtype=float)
    eta1 = 2 * k[0] * (x - 4 * k[0]**2 * t - phases[0])
    eta2 = 2 * k[1] * (x - 4 * k[1]**2 * t - phases[1])
    exponent = np.stack([np.zeros_like(x), eta1, eta2,
                         eta1 + eta2 + np.log(interaction)])
    weight = np.exp(exponent - np.max(exponent, axis=0))
    weight /= np.sum(weight, axis=0)
    slopes = np.array([0.0, 2*k[0], 2*k[1], 2*np.sum(k)])
    time_slopes = np.array([0.0, -8*k[0]**3, -8*k[1]**3, -8*np.sum(k**3)])
    shape = (4,) + (1,) * x.ndim
    return weight, slopes.reshape(shape), time_slopes.reshape(shape)


def exact_field(x, t, kappas, derivatives=False):
    """2(log tau)_xx as a nonnegative centered variance, without large exp."""
    w, slope, time_slope = tau_weights(x, t, kappas)
    centered = slope - np.sum(w * slope, axis=0)
    moment2 = np.sum(w * centered**2, axis=0)
    u = 2 * moment2
    if not derivatives:
        return u
    moment3 = np.sum(w * centered**3, axis=0)
    moment5 = np.sum(w * centered**5, axis=0)
    time_centered = time_slope - np.sum(w * time_slope, axis=0)
    return (u, 2*moment3,
            2*np.sum(w * centered**2 * time_centered, axis=0),
            2*(moment5 - 10*moment3*moment2))  # u, u_x, u_t, u_xxx


def bilinear_coefficients(kappas, interaction=None):
    """Exact Fraction coefficients of (D_x D_t + D_x^4) tau . tau."""
    k1, k2 = [Fraction(str(k)) for k in kappas]
    A = ((k1-k2)/(k1+k2))**2 if interaction is None else interaction
    labels = [(0, 0), (1, 0), (0, 1), (1, 1)]
    amplitudes = [Fraction(1), Fraction(1), Fraction(1), A]
    r = [2*(m*k1+n*k2) for m, n in labels]
    s = [-8*(m*k1**3+n*k2**3) for m, n in labels]
    coefficients = {}
    for i, a in enumerate(labels):
        for j, b in enumerate(labels):
            label = (a[0]+b[0], a[1]+b[1])
            d = r[i]-r[j]
            value = amplitudes[i]*amplitudes[j]*(d*(s[i]-s[j])+d**4)
            coefficients[label] = coefficients.get(label, Fraction(0)) + value
    return coefficients


def root_near(derivative, center, kappa):
    """Locate one separated maximum, with an explicit sign-changing bracket."""
    left, right = center-1/kappa, center+1/kappa
    if not (derivative(left) > 0 and derivative(right) < 0):
        raise AssertionError("Expected isolated endpoint maximum was not bracketed.")
    for _ in range(60):
        middle = (left+right)/2
        if derivative(middle) > 0:
            left = middle
        else:
            right = middle
    return (left+right)/2


def exact_centers(t, kappas):
    if t == 0:
        raise ValueError("Do not assign two pulse centers during the collision.")
    k, _, _, incoming, outgoing = parameters(kappas)
    prediction = 4*k*k*t + (incoming if t < 0 else outgoing)
    return np.array([root_near(lambda x: float(exact_field(x, t, k, True)[1]), c, kap)
                     for c, kap in zip(prediction, k)])


def exact_checks(inputs):
    k, A, phases, incoming, outgoing = parameters(inputs["kappas"])
    x = np.linspace(-25, 25, 1001)
    pde, reversal = 0.0, 0.0
    for t in inputs["checkpoint_times"]:
        u, ux, ut, uxxx = exact_field(x, t, k, True)
        pde = max(pde, float(np.max(np.abs(ut + 6*u*ux + uxxx))))
        reversal = max(reversal, float(np.max(np.abs(u-exact_field(-x, -t, k)))))
    coefficients = bilinear_coefficients(inputs["kappas"])
    wrong = bilinear_coefficients(inputs["kappas"], Fraction(1))
    # Independent direct quotient calculation uses only moderate exponents.
    xx = np.linspace(-3, 3, 41)
    eta = 2*k[:, None]*(xx - phases[:, None])
    terms = np.array([np.ones_like(xx), np.exp(eta[0]), np.exp(eta[1]),
                      A*np.exp(eta[0]+eta[1])])
    slopes = np.array([0, 2*k[0], 2*k[1], 2*sum(k)])[:, None]
    tau = np.sum(terms, axis=0)
    quotient = 2*(np.sum(slopes**2*terms, axis=0)/tau
                  - (np.sum(slopes*terms, axis=0)/tau)**2)
    far = exact_field(np.array([-1e6, 1e6]), 0, k)
    # Independent finite-rank Marchenko matrix, using a solve rather than tau.
    gram_error = 0.0
    for pair in [(1.0, 0.5), (1.2, 0.7)]:
        kk, aa, pp, _, _ = parameters(pair)
        c0 = 2*kk*np.exp(2*kk*pp)/aa
        for time in [-0.3, 0.0, 0.3]:
            for position in [-2.0, -0.5, 0.5, 2.0]:
                w = np.sqrt(c0)*np.exp(4*kk**3*time-kk*position)
                G = np.eye(2) + np.outer(w, w)/(kk[:, None]+kk[None, :])
                z = np.linalg.solve(G, w)
                S = float(w@z)
                gram_u = 4*float((kk*w)@z)-2*S*S
                gram_error = max(gram_error, abs(gram_u-float(exact_field(position, time, kk))))
    # Quadrature on a fixed long interval; independent of the Fourier solver.
    nodes, weights = np.polynomial.legendre.leggauss(640)
    line_integrals = np.array([4*sum(k), 16*sum(k**3)/3, 32*sum(k**5)/5])
    integral_errors = []
    for t in [-4.0, 0.0, 4.0]:
        u, ux, _, _ = exact_field(48*nodes, t, k, True)
        numerical = 48*np.array([weights@u, weights@(u*u),
                                 weights@(u**3-ux*ux/2)])
        integral_errors.append(float(np.max(np.abs((numerical-line_integrals)/line_integrals))))
    shifts = outgoing-incoming
    phase_rows = []
    for t in inputs["asymptotic_times"]:
        ci, co = exact_centers(-t, k), exact_centers(t, k)
        estimate = co-ci-8*k*k*t
        phase_rows.append({"endpoint_time": t, "incoming_centers": ci.tolist(),
                           "outgoing_centers": co.tolist(),
                           "finite_time_shift_estimate": estimate.tolist(),
                           "bias_from_asymptotic_shift": (estimate-shifts).tolist()})
    return {"interaction_A": A, "phase_parameters": phases.tolist(),
            "incoming_intercepts": incoming.tolist(), "outgoing_intercepts": outgoing.tolist(),
            "asymptotic_shifts": shifts.tolist(), "exact_line_integrals_M_P_E": line_integrals.tolist(),
            "bilinear_coefficient_count": len(coefficients),
            "bilinear_nonzero_coefficients": sum(v != 0 for v in coefficients.values()),
            "wrong_A_bilinear_max_abs_coefficient": float(max(abs(v) for v in wrong.values())),
            "pde_residual_max_abs": pde, "reversal_symmetry_max_abs": reversal,
            "direct_quotient_vs_variance_max_abs": float(np.max(np.abs(quotient-exact_field(xx, 0, k)))),
            "gram_reconstruction_vs_variance_max_abs": gram_error,
            "gram_check_kappas": [[1.0, 0.5], [1.2, 0.7]],
            "far_field_values": far.tolist(), "line_integral_relative_errors": integral_errors,
            "asymptotic_phase_refinement": phase_rows}


def independent_sum(x, t, kappas):
    k, _, _, incoming, _ = parameters(kappas)
    z = k[:, None]*(np.asarray(x)[None, :]-4*k[:, None]**2*t-incoming[:, None])
    u = 2*k[:, None]**2/np.cosh(z)**2
    ux = -2*k[:, None]*u*np.tanh(z)
    return np.sum(u, axis=0), 6*(ux[0]*u[1]+u[0]*ux[1])


def negative_control(inputs, solver):
    x = np.linspace(-40, 40, 8001)
    rows = []
    for t in inputs["checkpoint_times"]:
        u, residual = independent_sum(x, t, inputs["kappas"])
        rows.append({"time": t, "field_error": solver.error_norms(u, exact_field(x, t, inputs["kappas"]), x[1]-x[0]),
                     "pde_residual_linf": float(np.max(np.abs(residual)))})
    return rows


def evolve(length, points, steps, inputs, solver, retain=False):
    x, wave, _, keep = solver.grid(length, points)
    k, _, _, incoming, outgoing = parameters(inputs["kappas"])
    start, end = inputs["initial_time"], inputs["final_time"]
    h, dx = (end-start)/steps, length/points
    exact0 = exact_field(x, start, k)
    v = np.fft.fft(exact0)*keep
    initial_v = v.copy()
    initial_integrals = solver.invariants(v, wave, length)
    line_integrals = np.array([4*sum(k), 16*sum(k**3)/3, 32*sum(k**5)/5])
    drift = np.zeros(3)
    checkpoints = {}
    for t in inputs["checkpoint_times"]:
        n = round((t-start)/h)
        if n < 0 or n > steps or abs(start+n*h-t) > 1e-10:
            raise ValueError("All checkpoints must lie on each evolution grid.")
        checkpoints[n] = t
    errors = []
    half = np.exp(0.5*h*1j*wave**3)
    nonlinear = lambda value: solver.nonlinear(value, wave, keep)
    for n in range(steps+1):
        if n in checkpoints:
            t = checkpoints[n]
            error = solver.error_norms(np.fft.ifft(v).real, exact_field(x, t, k), dx)
            errors.append({"time": t, **error})
        if n == steps:
            break
        v = solver.ifrk4_step(v, h, 1j*wave**3, nonlinear, half=half)*keep
        integrals = solver.invariants(v, wave, length)
        if not np.all(np.isfinite(integrals)):
            raise ArithmeticError("Nonfinite evolution; refine the time step.")
        drift = np.maximum(drift, np.abs((integrals-initial_integrals)/line_integrals))
    # Boundary mismatch is sampled independently at 161 times, not inferred from mass drift.
    boundary_value, boundary_jump, boundary_derivative_jump = 0.0, 0.0, 0.0
    for t in np.linspace(start, end, 161):
        u, ux, _, _ = exact_field(np.array([-length/2, length/2]), t, k, True)
        boundary_value = max(boundary_value, float(np.max(np.abs(u))))
        boundary_jump = max(boundary_jump, float(abs(u[1]-u[0])))
        boundary_derivative_jump = max(boundary_derivative_jump, float(abs(ux[1]-ux[0])))
    result = {"length": length, "points": points, "steps": steps, "dt": h, "dx": dx,
              "retained_modes": int(np.count_nonzero(keep)),
              "checkpoint_errors": errors,
              "max_checkpoint_relative_l2": max(e["relative_l2"] for e in errors),
              "max_checkpoint_relative_linf": max(e["relative_linf"] for e in errors),
              "initial_line_integral_relative_errors": np.abs((initial_integrals-line_integrals)/line_integrals).tolist(),
              "max_relative_invariant_drift_M_P_E": drift.tolist(),
              "sampled_boundary_max_value": boundary_value,
              "sampled_boundary_max_value_jump": boundary_jump,
              "sampled_boundary_max_derivative_jump": boundary_derivative_jump}
    if retain:
        phase = []
        for t, vf, intercept in [(start, initial_v, incoming), (end, v, outgoing)]:
            prediction = 4*k*k*t+intercept
            def derivative(xx):
                return float(np.real(np.sum(1j*wave*vf*np.exp(1j*wave*(xx+length/2))))/points)
            numerical = np.array([root_near(derivative, c, kap) for c, kap in zip(prediction, k)])
            exact = exact_centers(t, k)
            phase.append({"time": t, "numerical_centers": numerical.tolist(), "exact_finite_time_centers": exact.tolist(),
                          "center_errors": (numerical-exact).tolist()})
        measured = np.array(phase[1]["numerical_centers"])-phase[0]["numerical_centers"]-4*k*k*(end-start)
        exact_estimate = np.array(phase[1]["exact_finite_time_centers"])-phase[0]["exact_finite_time_centers"]-4*k*k*(end-start)
        result["endpoint_phases"] = phase
        result["numerical_shift_estimate"] = measured.tolist()
        result["shift_numerical_error_vs_finite_time_exact"] = (measured-exact_estimate).tolist()
        result["finite_time_shift_bias_vs_asymptotic"] = (exact_estimate-(outgoing-incoming)).tolist()
    return result


def run(inputs):
    if inputs["phase_convention"] != "x_j = log(A)/(4*kappa_j)":
        raise ValueError("Unknown phase convention.")
    solver = load_solver()
    cache = {}
    def solve(length, points, steps):
        key = (length, points, steps)
        if key not in cache:
            cache[key] = evolve(length, points, steps, inputs, solver, retain=key == tuple(inputs["reference"][f] for f in ("length", "points", "steps")))
        return cache[key]
    temporal = [solve(inputs["temporal"]["length"], inputs["temporal"]["points"], n) for n in inputs["temporal"]["steps"]]
    spatial = [solve(inputs["spatial"]["length"], n, inputs["spatial"]["steps"]) for n in inputs["spatial"]["points"]]
    domain = [solve(length, round(length/inputs["domain"]["spacing"]), inputs["domain"]["steps"]) for length in inputs["domain"]["lengths"]]
    reference = solve(**inputs["reference"])
    errors = [r["max_checkpoint_relative_l2"] for r in temporal]
    return {"schema_version": 1, "inputs": inputs,
            "environment": {"python": platform.python_version(), "numpy": np.__version__, "arithmetic": "IEEE 754 binary64"},
            "shared_solver_method_checks": solver.method_checks(), "exact_checks": exact_checks(inputs),
            "temporal_refinement": temporal,
            "temporal_observed_orders": [math.log(a/b, 2) for a, b in zip(errors, errors[1:])],
            "spatial_refinement": spatial, "domain_refinement": domain, "reference_run": reference,
            "independent_pulse_sum_control": negative_control(inputs, solver)}


def verify(report):
    def require(test, message):
        if not test:
            raise AssertionError(message)
    exact = report["exact_checks"]
    require(exact["bilinear_nonzero_coefficients"] == 0, "Exact bilinear identity failed.")
    require(exact["wrong_A_bilinear_max_abs_coefficient"] > 1, "Wrong-interaction control did not fail.")
    for key in ["pde_residual_max_abs", "reversal_symmetry_max_abs", "direct_quotient_vs_variance_max_abs", "gram_reconstruction_vs_variance_max_abs"]:
        require(exact[key] < 1e-10, f"Exact-field check failed: {key}")
    require(max(exact["line_integral_relative_errors"]) < 1e-8, "Integral quadrature failed.")
    require(all(3.4 < order < 5.5 for order in report["temporal_observed_orders"]), "Temporal refinement failed.")
    require(report["reference_run"]["max_checkpoint_relative_l2"] < 1e-6, "Reference error too large.")
    for group in ["spatial_refinement", "domain_refinement"]:
        require(report[group][-1]["max_checkpoint_relative_l2"] < report[group][0]["max_checkpoint_relative_l2"]/100, f"Refinement failed: {group}")
    require(max(row["pde_residual_linf"] for row in report["independent_pulse_sum_control"]) > 0.1, "Independent-pulse control did not fail.")
    def finite(value):
        if isinstance(value, dict):
            return all(finite(v) for v in value.values())
        if isinstance(value, list):
            return all(finite(v) for v in value)
        return not isinstance(value, float) or math.isfinite(value)
    require(finite(report), "Report contains a nonfinite number.")


def compare(saved, fresh, atol, rtol, path="results"):
    if isinstance(saved, dict):
        if not isinstance(fresh, dict) or saved.keys() != fresh.keys():
            raise AssertionError(f"Changed keys at {path}")
        for key in saved:
            if key == "environment":
                continue
            if key == "inputs":
                if saved[key] != fresh[key]:
                    raise AssertionError("Saved results use different inputs.")
            else:
                compare(saved[key], fresh[key], atol, rtol, f"{path}.{key}")
    elif isinstance(saved, list):
        if not isinstance(fresh, list) or len(saved) != len(fresh):
            raise AssertionError(f"Changed list at {path}")
        for n, (a, b) in enumerate(zip(saved, fresh)):
            compare(a, b, atol, rtol, f"{path}[{n}]")
    elif isinstance(saved, bool) or isinstance(saved, int):
        if saved != fresh:
            raise AssertionError(f"Changed discrete value at {path}")
    elif isinstance(saved, float):
        if not math.isclose(saved, fresh, abs_tol=atol, rel_tol=rtol):
            raise AssertionError(f"Changed numerical result at {path}: {saved} vs {fresh}")
    elif saved != fresh:
        raise AssertionError(f"Changed metadata at {path}")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    mode = parser.add_mutually_exclusive_group()
    mode.add_argument("--write", action="store_true")
    mode.add_argument("--check", action="store_true")
    args = parser.parse_args()
    inputs = json.loads((HERE/"inputs.json").read_text())
    saved = None
    if args.check:
        saved = json.loads((HERE/"results.json").read_text())
        if saved["inputs"] != inputs:
            raise AssertionError("Saved results use different inputs; regeneration must be deliberate.")
    result = run(inputs)
    verify(result)
    if args.check:
        compare(saved, result, inputs["saved_comparison"]["absolute_tolerance"], inputs["saved_comparison"]["relative_tolerance"])
        print("PASS: exact identities, independent controls, three refinements, and saved results.")
        print("No files written. Shared one-soliton inputs and results were not read or changed.")
        print("Reference max checkpoint relative L2:", result["reference_run"]["max_checkpoint_relative_l2"])
        print("Temporal orders:", result["temporal_observed_orders"])
    elif args.write:
        (HERE/"results.json").write_text(json.dumps(result, indent=2, allow_nan=False)+"\n")
        print("Wrote", HERE/"results.json")
        print("Reference max checkpoint relative L2:", result["reference_run"]["max_checkpoint_relative_l2"])
        print("Temporal orders:", result["temporal_observed_orders"])
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
