#!/usr/bin/env python3
"""Open three-particle Toda: Verlet, an independent QR flow, and exact checks.

Only NumPy is required. With no option this prints fresh results without writing.
--check also compares the saved numerical evidence, with floating-point tolerance.
--write explicitly regenerates results.json. See README.md for scope and provenance.
"""

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

import numpy as np

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


def force(q):
    """-gradient V, with no endpoint spring and no periodic bond."""
    bonds = np.exp(q[:-1] - q[1:])
    return np.r_[0.0, bonds] - np.r_[bonds, 0.0]


def lax(q, p):
    a = np.exp((q[:-1] - q[1:]) / 2.0)
    return np.diag(p) + np.diag(a, 1) + np.diag(a, -1)


def invariants(q, p):
    """Direct scalar formulas, independently compared with matrix traces."""
    bonds = np.exp(q[:-1] - q[1:])
    return np.array([
        np.sum(p),
        np.sum(p * p) / 2.0 + np.sum(bonds),
        np.sum(p**3) / 3.0 + np.sum((p[:-1] + p[1:]) * bonds),
    ])


def verlet(q0, p0, h, t_end):
    n = round(t_end / h)
    if not math.isclose(n * h, t_end, abs_tol=1e-12):
        raise ValueError("Each step size must divide t_end.")
    states = np.empty((n + 1, 2, len(q0)), dtype=np.float64)
    q, p = np.array(q0, dtype=float), np.array(p0, dtype=float)
    states[0] = q, p
    for k in range(n):
        half_p = p + (h / 2.0) * force(q)
        q = q + h * half_p
        p = half_p + (h / 2.0) * force(q)
        states[k + 1] = q, p
    return np.arange(n + 1) * h, states


def qr_solution(q0, p0, times):
    """exp(-t L0/2)=QR, diag(R)>0, L(t)=Q.T L0 Q.

    This is a direct evaluation at each time, not an ODE timestepper.
    A uniform exponent shift avoids overflow without changing Q.
    Long-time ill-conditioning remains a limitation; only |t|<=6 is tested.
    """
    q0, p0 = np.array(q0, dtype=float), np.array(p0, dtype=float)
    L0 = lax(q0, p0)
    lam, U = np.linalg.eigh(L0)
    states = []
    for t in times:
        exponents = -t * lam / 2.0
        M = (U * np.exp(exponents - np.max(exponents))) @ U.T
        Q, R = np.linalg.qr(M)
        Q = Q * np.where(np.diag(R) >= 0.0, 1.0, -1.0)
        Lt = Q.T @ L0 @ Q
        a = np.diag(Lt, 1)
        if np.any(a <= 0.0):
            raise ArithmeticError("Lost positive off-diagonal; QR is ill-conditioned.")
        differences = 2.0 * np.log(a)
        q = np.r_[np.cumsum(differences[::-1])[::-1], 0.0]
        q += np.mean(q0) + t * np.mean(p0) - np.mean(q)
        states.append([q, np.diag(Lt)])
    return np.array(states)


def symmetric_exact(times):
    """Closed form for q0=(0,0,0), p0=(1,0,-1), independent of L and QR."""
    s = math.sqrt(3.0) * np.asarray(times) / 2.0 - math.atanh(1.0 / math.sqrt(3.0))
    x = math.log(1.5) - 2.0 * np.log(np.cosh(s))
    v = -math.sqrt(3.0) * np.tanh(s)
    zeros = np.zeros_like(x)
    return np.stack([np.stack([x, zeros, -x], axis=-1),
                     np.stack([v, zeros, -v], axis=-1)], axis=1)


def derivative_check(q0, p0, delta):
    """Fourth-order centered derivative of QR against canonical q,p equations."""
    residual = 0.0
    for t in [0.0, 0.5, 1.25, 3.0, 6.0]:
        ym2, ym1, y, yp1, yp2 = qr_solution(q0, p0, t + delta * np.arange(-2, 3))
        derivative = (ym2 - 8.0 * ym1 + 8.0 * yp1 - yp2) / (12.0 * delta)
        residual = max(residual, float(np.max(np.abs(derivative - [y[1], force(y[0])]))))
    return residual


def algebra_checks():
    # Nonuniform bonds and nonzero momentum expose endpoint, sign, and scale errors.
    q, p = np.array([0.2, -0.3, 0.1]), np.array([0.7, -0.4, 0.2])
    L = lax(q, p)
    a = np.diag(L, 1)
    B = np.diag(-a / 2.0, 1) + np.diag(a / 2.0, -1)
    da = a * (p[:-1] - p[1:]) / 2.0
    dL = np.diag(force(q)) + np.diag(da, 1) + np.diag(da, -1)
    scalar = invariants(q, p)
    traces = np.array([np.trace(np.linalg.matrix_power(L, k)) / k for k in [1, 2, 3]])
    eps = 1e-5
    gradH = []
    for e in np.eye(3) * eps:
        gradH.append((invariants(q + e, p)[1] - invariants(q - e, p)[1]) / (2.0 * eps))
    return {
        "lax_equation_max_abs": float(np.max(np.abs(dL - (B @ L - L @ B)))),
        "trace_formula_max_abs": float(np.max(np.abs(scalar - traces))),
        "force_vs_energy_gradient_max_abs": float(np.max(np.abs(force(q) + gradH))),
        "gradient_difference_step": eps,
    }


def shift_boost_check():
    """A transformed exact solution tests the missing common-position mode."""
    q0, p0 = np.zeros(3), np.array([1.0, 0.0, -1.0])
    c, d = 0.4, -0.2
    times = np.array([0.0, 0.5, 1.0, 2.0, 4.0, 6.0])
    original = symmetric_exact(times)
    expected = original.copy()
    expected[:, 0] += d + c * times[:, None]
    expected[:, 1] += c
    transformed = qr_solution(q0 + d, p0 + c, times)
    P, H, I3 = invariants(q0, p0)
    transformed_invariants = invariants(q0 + d, p0 + c)
    predicted = [P + 3*c, H + c*P + 1.5*c*c, I3 + 2*c*H + c*c*P + c**3]
    return {
        "c": c, "d": d,
        "transformed_initial_invariants": transformed_invariants.tolist(),
        "trajectory_max_abs_residual": float(np.max(np.abs(transformed - expected))),
        "invariant_formula_max_abs_residual": float(np.max(np.abs(transformed_invariants - predicted))),
    }


def run(inputs):
    report = {
        "problem": "dimensionless open N=3, H=sum(p^2)/2+sum(exp(q_i-q_{i+1}))",
        "precision": "IEEE 754 binary64",
        "inputs": inputs,
        "algebra_checks": algebra_checks(),
        "shift_boost_check": shift_boost_check(),
        "cases": [],
    }
    for case in inputs["cases"]:
        q0, p0 = case["q0"], case["p0"]
        initial = invariants(np.array(q0), np.array(p0))
        initial_spectrum = np.linalg.eigvalsh(lax(np.array(q0), np.array(p0)))
        data = {
            "name": case["name"], "initial_invariants": initial.tolist(),
            "initial_eigenvalues": initial_spectrum.tolist(), "refinement": [],
            "qr_ode_residual_delta_0.01": derivative_check(q0, p0, 0.01),
            "qr_ode_residual_delta_0.005": derivative_check(q0, p0, 0.005),
        }
        for h in inputs["step_sizes"]:
            times, states = verlet(q0, p0, h, inputs["t_end"])
            reference = qr_solution(q0, p0, times)
            invariant_values = np.array([invariants(q, p) for q, p in states])
            spectra = np.array([np.linalg.eigvalsh(lax(q, p)) for q, p in states])
            error = np.max(np.abs(states - reference), axis=(1, 2))
            drift = np.abs(invariant_values - initial)
            record = {
                "h": h, "steps": len(times) - 1,
                "state_max_abs_error": float(np.max(error)),
                "P_max_abs_drift": float(np.max(drift[:, 0])),
                "H_max_abs_drift": float(np.max(drift[:, 1])),
                "I3_max_abs_drift": float(np.max(drift[:, 2])),
                "eigenvalue_max_abs_drift": float(np.max(np.abs(spectra - initial_spectrum))),
            }
            if case["name"] == "symmetric":
                record["qr_vs_closed_form_max_abs"] = float(np.max(np.abs(reference - symmetric_exact(times))))
            data["refinement"].append(record)
        errors = [row["state_max_abs_error"] for row in data["refinement"]]
        data["observed_orders"] = [math.log(a / b, 2.0) for a, b in zip(errors, errors[1:])]
        # All saved trajectory points come from the finest completed run above.
        data["trajectory"] = [
            {"t": float(times[i]), "q_verlet": states[i, 0].tolist(),
             "p_verlet": states[i, 1].tolist(), "q_qr": reference[i, 0].tolist(),
             "p_qr": reference[i, 1].tolist(), "state_abs_error": float(error[i]),
             "H_abs_drift": float(drift[i, 1])}
            for i in range(0, len(times), inputs["trajectory_stride"])
        ]
        report["cases"].append(data)
    return report


def verify(report):
    def require(condition, message):
        if not condition:
            raise AssertionError(message)
    checks = report["algebra_checks"]
    require(checks["lax_equation_max_abs"] < 1e-13, "Lax identity fails.")
    require(checks["trace_formula_max_abs"] < 1e-13, "Trace formulas fail.")
    require(checks["force_vs_energy_gradient_max_abs"] < 1e-8, "Hamiltonian force fails.")
    require(report["shift_boost_check"]["trajectory_max_abs_residual"] < 1e-9,
            "Shift/boost reconstruction fails.")
    require(report["shift_boost_check"]["invariant_formula_max_abs_residual"] < 1e-13,
            "Shift/boost invariant formulas fail.")
    for case in report["cases"]:
        require(case["qr_ode_residual_delta_0.005"] < 1e-7, "Independent QR ODE residual too large.")
        require(all(1.95 < p < 2.05 for p in case["observed_orders"]), "Verlet order is not two.")
        require(case["refinement"][-1]["state_max_abs_error"] < 1e-4, "Fine-grid trajectory error too large.")
        for row in case["refinement"]:
            require(row["P_max_abs_drift"] < 1e-12, "Total momentum drift is too large.")
            if case["name"] == "symmetric":
                require(row["qr_vs_closed_form_max_abs"] < 1e-9, "QR disagrees with closed form.")
    require(abs(report["cases"][1]["initial_invariants"][2]) > 0.1,
            "Asymmetric control must have nonzero I3.")


def compare(saved, fresh, path="results"):
    """Compare scientific numbers, not platform strings or a whole-file hash."""
    if isinstance(saved, dict):
        if saved.keys() != fresh.keys():
            raise AssertionError(f"Saved result keys differ at {path}.")
        for key in saved:
            compare(saved[key], fresh[key], path + "." + key)
    elif isinstance(saved, list):
        if len(saved) != len(fresh):
            raise AssertionError(f"Saved result length differs at {path}.")
        for i, (a, b) in enumerate(zip(saved, fresh)):
            compare(a, b, f"{path}[{i}]")
    elif isinstance(saved, (int, float)):
        # Derivative residuals and nominal zeros vary at the roundoff level.
        if not math.isclose(saved, fresh, rel_tol=1e-6, abs_tol=2e-9):
            raise AssertionError(f"Saved result differs at {path}: {saved} versus {fresh}.")
    elif saved != fresh:
        raise AssertionError(f"Saved metadata differs at {path}.")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    mode = parser.add_mutually_exclusive_group()
    mode.add_argument("--write", action="store_true", help="regenerate selected results.json")
    mode.add_argument("--check", action="store_true", help="recompute and compare saved evidence without writing")
    args = parser.parse_args()
    inputs = json.loads((HERE / "inputs.json").read_text())
    report = run(inputs)
    verify(report)
    if args.check:
        saved = json.loads((HERE / "results.json").read_text())
        if saved["inputs"] != inputs:
            raise AssertionError("Saved results use different inputs; regenerate deliberately.")
        compare(saved, report)
    if args.write:
        (HERE / "results.json").write_text(json.dumps(report, indent=2, allow_nan=False) + "\n")
    print(f"Python {platform.python_version()}, NumPy {np.__version__}, binary64")
    for case in report["cases"]:
        print(f"\n{case['name']}: initial (P,H,I3) = {case['initial_invariants']}")
        print("h       max state error    max |dH|         max |dI3|        max eigenvalue drift")
        for row in case["refinement"]:
            print(f"{row['h']:.3f}   {row['state_max_abs_error']:.8e}   {row['H_max_abs_drift']:.8e}   "
                  f"{row['I3_max_abs_drift']:.8e}   {row['eigenvalue_max_abs_drift']:.8e}")
        print("observed orders:", ", ".join(f"{x:.6f}" for x in case["observed_orders"]))
        print(f"QR ODE residual (delta=0.005): {case['qr_ode_residual_delta_0.005']:.3e}")
    print("\nPASS: independent references, Hamiltonian/Lax checks, and second-order refinement.")
    if args.check:
        print("PASS: saved numerical evidence agrees within floating-point tolerance; no files written.")


if __name__ == "__main__":
    main()
