#!/usr/bin/env python3
"""Normalized five-site XXX form factors and a complete finite correlation sum.

NumPy only. The sibling xxx-algebra/experiment.py supplies actual B blocks.
Running without flags prints fresh JSON; --check never writes; --write replaces results.json.
"""

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

import numpy as np

HERE = Path(__file__).resolve().parent
DEPENDENCY = HERE.parent / "xxx-algebra" / "experiment.py"


def load_algebra():
    if not DEPENDENCY.is_file():
        raise FileNotFoundError("Keep xxx-algebra/experiment.py beside the xxx-correlations folder; the complete ZIP includes it.")
    module = types.ModuleType("xxx_algebra_companion")
    module.__file__ = str(DEPENDENCY)
    exec(compile(DEPENDENCY.read_text(), str(DEPENDENCY), "exec"), module.__dict__)
    return module


def sector_basis(sites, down_spins):
    return list(combinations(range(1, sites + 1), down_spins))


def physical_hamiltonian(sites, basis):
    """H/J from local spin action, independent of the companion's operators.

    Parallel spins contribute zero. Antiparallel spins contribute +1/2
    on the diagonal and -1/2 to the exchanged configuration, including seam.
    """
    lookup = {state: index for index, state in enumerate(basis)}
    result = np.zeros((len(basis), len(basis)))
    for column, state in enumerate(basis):
        occupied = set(state)
        for left in range(1, sites + 1):
            right = left % sites + 1
            if (left in occupied) != (right in occupied):
                exchanged = tuple(sorted(occupied.symmetric_difference({left, right})))
                result[column, column] += 0.5
                result[lookup[exchanged], column] -= 0.5
    return result


def shift(sites, basis):
    lookup = {state: index for index, state in enumerate(basis)}
    result = np.zeros((len(basis), len(basis)))
    for column, state in enumerate(basis):
        shifted = tuple(sorted(x % sites + 1 for x in state))
        result[lookup[shifted], column] = 1
    return result


def raising_maps(sites, source_basis, target_basis):
    """S_x^+ removes a down spin; matrix element is one, not 1/2."""
    lookup = {state: index for index, state in enumerate(target_basis)}
    maps = []
    for site in range(1, sites + 1):
        operator = np.zeros((len(target_basis), len(source_basis)))
        for column, state in enumerate(source_basis):
            if site in state:
                operator[lookup[tuple(x for x in state if x != site)], column] = 1
        maps.append(operator)
    return np.array(maps)


def cyclotomic(values):
    """Reduce rational polynomials modulo 1+z+z^2+z^3+z^4 (z^5=1)."""
    work = list(map(Fraction, values)) + [Fraction(0)] * max(0, 4-len(values))
    for power in range(len(work)-1, 3, -1):
        coefficient = work[power]
        for offset in range(1, 5):
            work[power-offset] -= coefficient
    return work[:4]


def polynomial_product(left, right):
    values = [Fraction(0)] * (len(left)+len(right)-1)
    for i, a in enumerate(left):
        for j, b in enumerate(right):
            values[i+j] += a*b
    return cyclotomic(values)


def exact_checks(basis, h_two, h_one, local):
    raw = [Fraction(-1 if (b-a) in (1, 4) else 1, 8) for a, b in basis]
    h2 = [[Fraction(float(value)) for value in row] for row in h_two]
    h1 = [[Fraction(float(value)) for value in row] for row in h_one]
    assert all(sum(h2[i][j]*raw[j] for j in range(10)) == 2*raw[i] for i in range(10))
    norm = sum(value*value for value in raw)
    local_raw = [sum(Fraction(float(local[i,j]))*raw[j] for j in range(10)) for i in range(5)]
    local_norm = sum(value*value for value in local_raw)
    first_moment = sum(local_raw[i] * sum((h1[i][j]-2*Fraction(i == j))*local_raw[j]
                       for j in range(5)) for i in range(5))/norm
    assert norm == Fraction(5, 32) and local_norm/norm == Fraction(2, 5)
    assert first_moment == Fraction(-3, 10)
    # A complete five-point Fourier basis, not a selected energy list.
    zero = [Fraction(0)]*4
    for difference in range(5):
        gram = [Fraction(0)]*5
        for x in range(1, 6):
            gram[(difference*x) % 5] += 1
        assert cyclotomic(gram) == ([Fraction(5),0,0,0] if difference == 0 else zero)
    # M_q = D(z)/sqrt(10), D=z^(2m)+z^(-2m)-z^m-z^(-m), real.
    exact_weights = []
    for mode in range(5):
        d = [Fraction(0)]*5
        for exponent, coefficient in [(2*mode,1),(-2*mode,1),(mode,-1),(-mode,-1)]:
            d[exponent % 5] += coefficient
        square = polynomial_product(cyclotomic(d), cyclotomic(d))
        assert square[1:] == [0,0,0]
        exact_weights.append(square[0]/10)
    assert exact_weights == [Fraction(0)]+[Fraction(1,2)]*4
    return {"raw_B_norm_squared": str(norm), "local_equal_time_weight": str(local_norm/norm),
            "local_first_frequency_moment_over_J": str(first_moment),
            "fourier_operator_weights": [str(value) for value in exact_weights],
            "exact_bond_eigenstate_equation": True, "exact_five_mode_Fourier_completeness": True,
            "exact_cyclotomic_weight_checks": True}


def pair(value):
    return [float(np.real(value)), float(np.imag(value))]


def relative_vector(left, right):
    return float(np.linalg.norm(left-right)/np.linalg.norm(right))


def finite_numbers(value, path="inputs"):
    if isinstance(value, dict):
        for key, item in value.items():
            finite_numbers(item, path+"."+key)
    elif isinstance(value, list):
        for index, item in enumerate(value):
            finite_numbers(item, f"{path}[{index}]")
    elif isinstance(value, float) and not math.isfinite(value):
        raise ValueError("Nonfinite number at "+path)


def validate_inputs(parameters):
    finite_numbers(parameters)
    required = {"sites": 5, "down_spins": 2, "algebraic_rapidities": [0.5,-0.5],
                "hbar": 1.0, "lattice_spacing": 1.0, "momentum_mode_numbers": list(range(5))}
    if any(parameters.get(key) != value for key, value in required.items()):
        raise ValueError("This benchmark fixes N=5, two down spins, roots ±1/2, hbar=a=1 and all five momenta.")
    if parameters["coupling_J"] <= 0 or parameters["local_site"] not in range(1,6):
        raise ValueError("Use J>0 and a local site in 1,...,5.")
    grid = parameters["dimensionless_time_grid"]
    if grid["start"] != 0 or grid["stop"] <= 0 or not isinstance(grid["points"], int) or grid["points"] < 2:
        raise ValueError("Use a finite grid beginning at zero with at least two points.")
    if abs(complex(*parameters["rescale_initial_ket_real_imag"])) == 0:
        raise ValueError("The rescaling control requires a nonzero factor.")
    for key in ["equation_tolerance", "negative_control_minimum"]:
        if parameters[key] <= 0:
            raise ValueError("Tolerances must be positive.")
    if any(value <= 0 for value in parameters["saved_comparison"].values()):
        raise ValueError("Saved comparison tolerances must be positive.")


def run(parameters):
    validate_inputs(parameters)
    N, J = 5, parameters["coupling_J"]
    source, target = sector_basis(N, 2), sector_basis(N, 1)
    h2, h1 = physical_hamiltonian(N, source), physical_hamiltonian(N, target)
    u2, u1 = shift(N, source), shift(N, target)
    maps = raising_maps(N, source, target)
    raw_full = load_algebra().bethe_vector(N, parameters["algebraic_rapidities"])
    indices = [sum(1 << (N-x) for x in state) for state in source]
    raw = raw_full[indices]
    expected_raw = np.array([-1 if (b-a) in (1,4) else 1 for a,b in source])/8
    raw_norm = float(np.linalg.norm(raw))
    if not raw_norm > 0:
        raise AssertionError("The reconstructed Bethe vector vanishes.")
    chi = raw/raw_norm
    expected_chi = expected_raw/np.linalg.norm(expected_raw)
    momenta = 2*np.pi*np.arange(N)/N
    positions = np.arange(1,N+1)
    fourier = np.exp(1j*np.outer(positions,momenta))/np.sqrt(N)
    energies = 1-np.cos(momenta)
    freq = energies-2
    local = maps[parameters["local_site"]-1]
    local_state = local@chi
    local_amplitudes = fourier.conj().T@local_state
    local_weights = abs(local_amplitudes)**2
    good = {}
    def check(name, value):
        good[name] = float(value)
    check("raw_B_relative_coefficient_error", relative_vector(raw, expected_raw))
    check("raw_B_outside_two_spin_sector_norm", np.linalg.norm(np.delete(raw_full, indices)))
    check("normalized_state_relative_error", relative_vector(chi, expected_chi))
    check("H_eigenvector_residual_over_J", np.linalg.norm(h2@chi-2*chi))
    check("active_translation_residual", np.linalg.norm(u2@chi-chi))
    check("Fourier_orthonormality_frobenius_error", np.linalg.norm(fourier.conj().T@fourier-np.eye(N)))
    check("Fourier_completeness_frobenius_error", np.linalg.norm(fourier@fourier.conj().T-np.eye(N)))
    check("one_magnon_H_equations_frobenius_error", np.linalg.norm(h1@fourier-fourier*energies))
    check("one_magnon_U_equations_frobenius_error", np.linalg.norm(u1@fourier-fourier*np.exp(-1j*momenta)))
    check("local_weight_formula_max_error", np.max(abs(local_weights-np.array([0,.1,.1,.1,.1]))))
    modes, wrong_sign = [], []
    for m, q in enumerate(momenta):
        operator = np.einsum('x,xij->ij', np.exp(-1j*q*positions)/np.sqrt(N), maps)
        state = operator@chi
        final = fourier[:,(-m) % N]
        amplitude = np.vdot(final,state)
        formula = 2*(np.cos(2*q)-np.cos(q))/np.sqrt(10)
        check(f"m{m}_form_factor_absolute_error", abs(amplitude-formula))
        check(f"m{m}_orthogonal_projection_norm", np.linalg.norm(state-amplitude*final))
        check(f"m{m}_translation_absolute_residual", np.linalg.norm(u1@state-np.exp(1j*q)*state))
        modes.append({"mode_number":m, "q_radians":float(q), "target_coordinate_mode_mod_5":int((-m) % N),
                      "target_active_U_eigenvalue":pair(np.exp(1j*q)), "form_factor_real_imag":pair(amplitude),
                      "form_factor_formula":float(formula), "weight":float(abs(amplitude)**2),
                      "local_weight":float(local_weights[(-m) % N]), "frequency_over_J":float(freq[m])})
        if m in (1,2):
            reversed_operator = np.einsum('x,xij->ij',np.exp(1j*q*positions)/np.sqrt(N),maps)
            wrong = reversed_operator@chi
            residual = float(np.linalg.norm(u1@wrong-np.exp(1j*q)*wrong)/np.linalg.norm(wrong))
            check(f"m{m}_wrong_sign_residual_formula_error",abs(residual-2*abs(np.sin(q))))
            check(f"m{m}_wrong_sign_preserves_weight_error",abs(np.linalg.norm(wrong)**2-np.linalg.norm(state)**2))
            check(f"m{m}_wrong_sign_preserves_energy_residual",np.linalg.norm(h1@wrong-energies[m]*wrong)/np.linalg.norm(wrong))
            wrong_sign.append({"mode_number":m,"wrong_translation_relative_residual":residual,
                               "unchanged_weight_error":float(abs(np.linalg.norm(wrong)**2-np.linalg.norm(state)**2)),
                               "unchanged_energy_residual_over_J":float(np.linalg.norm(h1@wrong-energies[m]*wrong)/np.linalg.norm(wrong))})
    check("sum_Fourier_weights_error", abs(sum(row["weight"] for row in modes)-2))
    local_norms = [float(np.linalg.norm(operator@chi)**2) for operator in maps]
    check("site_local_sum_rule_max_error",max(abs(value-.4) for value in local_norms))
    exact = exact_checks(source,h2,h1,local)
    grid = parameters["dimensionless_time_grid"]
    times = np.linspace(grid["start"],grid["stop"],grid["points"])
    # Independent numerical spectral theorem for the bond-built physical matrix.
    # The DFT basis and analytic energy list are NOT passed into this calculation.
    ev, vectors = np.linalg.eigh(h1)
    coordinates = vectors.conj().T@local_state
    direct = np.array([np.vdot(local_state,vectors@(np.exp(-1j*(ev-2)*t)*coordinates)) for t in times])
    spectral = np.exp(-1j*np.outer(times,freq))@local_weights
    closed = .4*np.exp(.75j*times)*np.cos(np.sqrt(5)*times/4)
    check("direct_vs_Fourier_sum_max_absolute_error",np.max(abs(direct-spectral)))
    check("direct_vs_closed_max_absolute_error",np.max(abs(direct-closed)))
    check("equal_time_sum_rule_error",abs(np.sum(local_weights)-np.linalg.norm(local_state)**2))
    moment = float(np.sum(freq*local_weights))
    independent_moment = float(np.vdot(local_state,(h1-2*np.eye(N))@local_state).real)
    check("first_moment_formula_error",abs(moment+.3))
    check("first_moment_operator_error",abs(independent_moment-moment))
    lines = []
    for modes_in_line, frequency in [([1,4],-(3+np.sqrt(5))/4),([2,3],-(3-np.sqrt(5))/4)]:
        numerical_weight = float(np.sum(abs(coordinates[abs((ev-2)-frequency)<1e-10])**2))
        weight = float(sum(local_weights[m] for m in modes_in_line))
        check("grouped_line_"+str(modes_in_line)+"_weight_error",abs(numerical_weight-.2))
        lines.append({"coordinate_modes":modes_in_line,"frequency_over_J":float(frequency),
                      "integrated_weight":weight,"independent_H_projector_weight":numerical_weight})
    factor = complex(*parameters["rescale_initial_ket_real_imag"])
    rescaled = factor*raw/np.linalg.norm(factor*raw)
    check("normalized_rescaling_phase_law_error",np.linalg.norm(rescaled-factor/abs(factor)*chi))
    rescaled_local = local@rescaled
    rescaled_weights = abs(fourier.conj().T@rescaled_local)**2
    check("normalized_rescaling_weight_max_error",np.max(abs(rescaled_weights-local_weights)))
    phase = np.exp(1j*parameters["final_ket_phase_angle"])
    phase_amplitudes = (phase*fourier).conj().T@local_state
    check("final_ket_phase_amplitude_law_max_error",np.max(abs(phase_amplitudes-phase.conjugate()*local_amplitudes)))
    check("final_ket_phase_weight_max_error",np.max(abs(abs(phase_amplitudes)**2-local_weights)))
    # Deliberately incomplete or unnormalized sums; every recorded failure is computed.
    one_missing = local_weights.copy(); one_missing[1] = 0
    distinct_energies_only = local_weights.copy(); distinct_energies_only[[3,4]] = 0
    dropped_line = local_weights.copy(); dropped_line[[1,4]] = 0
    raw_equal_time = float(np.linalg.norm(local@raw)**2)
    negative = {
        "raw_unnormalized_equal_time":raw_equal_time,
        "raw_equal_time_deficit":float(.4-raw_equal_time),
        "one_state_omitted_equal_time":float(sum(one_missing)),
        "one_state_omitted_deficit":float(.4-sum(one_missing)),
        "distinct_energies_only_equal_time":float(sum(distinct_energies_only)),
        "lost_degenerate_partners_deficit":float(.4-sum(distinct_energies_only)),
        "whole_low_frequency_line_omitted_equal_time":float(sum(dropped_line)),
        "whole_line_omitted_deficit":float(.4-sum(dropped_line)),
        "one_state_omitted_time_max_error":float(np.max(abs(np.exp(-1j*np.outer(times,freq))@one_missing-closed))),
        "wrong_Fourier_sign":wrong_sign,
        "final_phase_changes_amplitudes_max":float(np.max(abs(phase_amplitudes-local_amplitudes))),
    }
    if any(not math.isfinite(value) or value>parameters["equation_tolerance"] for value in good.values()):
        raise AssertionError("An independently checked equation failed: "+str(good))
    failures = [negative[key] for key in ["raw_equal_time_deficit","one_state_omitted_deficit",
                "lost_degenerate_partners_deficit","whole_line_omitted_deficit"]]
    failures += [row["wrong_translation_relative_residual"] for row in wrong_sign]
    if min(failures) < parameters["negative_control_minimum"]:
        raise AssertionError("A negative control did not detect its intended error.")
    result = {"schema_version":1,"environment":{"python":platform.python_version(),"numpy":np.__version__,
              "arithmetic":"binary64; exact Fraction and fifth-cyclotomic checks"},
              "inputs":parameters,"dependency":"../xxx-algebra/experiment.py",
              "two_down_spin_basis":[list(state) for state in source],"one_down_spin_basis":[list(state) for state in target],
              "state":{"raw_B_norm_squared":raw_norm**2,"energy_over_J":2.0,
                       "normalized_coefficients_real_imag":[pair(value) for value in chi],
                       "site_raising_norm_squared":local_norms},
              "exact_checks":exact,"equation_errors":good,"maximum_equation_error":max(good.values()),
              "momentum_form_factors":modes,"local_amplitudes_in_positive_k_DFT_basis":[pair(value) for value in local_amplitudes],
              "local_spectrum_lines":lines,"local_first_frequency_moment_over_J":moment,
              "direct_operator_first_frequency_moment_over_J":independent_moment,
              "negative_controls":negative,
              "time_evolution":[{"dimensionless_time_Jt":float(t),"physical_time":float(t/J),
                                  "direct_sector_real_imag":pair(a),"Fourier_sum_real_imag":pair(b),
                                  "closed_formula_real_imag":pair(c)} for t,a,b,c in zip(times,direct,spectral,closed)]}
    finite_numbers(result,"results")
    return result


def compare_saved(actual,saved,atol,rtol,path="results"):
    if isinstance(actual,dict):
        if not isinstance(saved,dict) or actual.keys()!=saved.keys():
            raise AssertionError("Changed fields at "+path)
        for key in actual:
            if key == "inputs":
                if actual[key] != saved[key]:
                    raise AssertionError("Changed input parameters")
            elif key != "environment":
                compare_saved(actual[key],saved[key],atol,rtol,path+"."+key)
    elif isinstance(actual,list):
        if not isinstance(saved,list) or len(actual)!=len(saved):
            raise AssertionError("Changed list at "+path)
        for i,(left,right) in enumerate(zip(actual,saved)):
            compare_saved(left,right,atol,rtol,f"{path}[{i}]")
    elif isinstance(actual,bool) or isinstance(actual,int):
        if type(actual) is not type(saved) or actual!=saved:
            raise AssertionError("Changed integer or truth value at "+path)
    elif isinstance(actual,float):
        if not isinstance(saved,(int,float)) or not math.isfinite(actual) or not math.isfinite(saved) or not math.isclose(actual,saved,abs_tol=atol,rel_tol=rtol):
            raise AssertionError("Numerical disagreement at "+path)
    elif actual != saved:
        raise AssertionError("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:
        tolerance=parameters["saved_comparison"]
        compare_saved(result,json.loads((HERE/"results.json").read_text()),
                      tolerance["absolute_tolerance"],tolerance["relative_tolerance"])
        print("XXX correlations: actual B state, exact norms, complete Fourier weights, direct evolution, sum rules, negative controls and saved results passed.")
    elif args.write:
        (HERE/"results.json").write_text(json.dumps(result,indent=2,allow_nan=False)+"\n")
        print("Wrote results.json after all XXX correlation checks passed.")
    else:
        print(json.dumps(result,indent=2,allow_nan=False))


if __name__ == "__main__":
    main()
