#!/usr/bin/env python3
"""Finite XXX algebra: rational YBE, RTT, transfer traces, and Hamiltonian.

Run without flags to print fresh JSON. --check is nonmutating and locks inputs;
--write intentionally regenerates results.json after all checks pass. NumPy only.
"""

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

import numpy as np

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


def complex_value(pair):
    return complex(*pair)


def same_input_values(actual, expected):
    """Exact JSON values and types: True is not 1, and nonfinite inputs fail."""
    return json.dumps(actual, sort_keys=True, allow_nan=False) == json.dumps(expected, sort_keys=True, allow_nan=False)


def residual(left, right):
    """Normalized Frobenius discrepancy; all operators here are dimensionless."""
    return float(np.linalg.norm(left - right, "fro") /
                 max(1.0, np.linalg.norm(left, "fro"), np.linalg.norm(right, "fro")))


def swap(factors, first, second):
    """Swap tensor basis labels; factor zero is the leftmost binary digit."""
    if not 0 <= first < factors or not 0 <= second < factors:
        raise ValueError("Tensor factor outside the declared product.")
    dimension = 2**factors
    operator = np.zeros((dimension, dimension), dtype=complex)
    for column in range(dimension):
        state = list(f"{column:0{factors}b}")
        state[first], state[second] = state[second], state[first]
        row = int("".join(state), 2)
        operator[row, column] = 1
    return operator


def r_matrix(parameter, permutation):
    return parameter * np.eye(len(permutation)) + 1j * permutation


def trace_first(operator, remaining_dimension):
    """Trace the first C^2 factor, keeping physical matrix order unchanged."""
    tensor = operator.reshape(2, remaining_dimension, 2, remaining_dimension)
    return np.trace(tensor, axis1=0, axis2=2)


def local_factors(length, parameter, auxiliaries=1, auxiliary=0, reverse=False):
    sites = list(range(auxiliaries, auxiliaries + length))
    if reverse:
        sites.reverse()
    return [r_matrix(parameter - 0.5j,
                     swap(length + auxiliaries, auxiliary, site)) for site in sites]


def ordered_product(factors):
    """Input L_1,...,L_N; return L_N...L_1, as specified in the site convention."""
    result = np.eye(len(factors[0]), dtype=complex)
    for factor in factors:
        result = factor @ result
    return result


def transfer(length, parameter, reverse=False):
    monodromy = ordered_product(local_factors(length, parameter, reverse=reverse))
    return trace_first(monodromy, 2**length)


def transfer_derivative(length, parameter):
    """Analytic product rule via a recurrence; each L'_n is the identity."""
    factors = local_factors(length, parameter)
    product = np.eye(len(factors[0]), dtype=complex)
    derivative = np.zeros_like(product)
    for factor in factors:
        derivative = product + factor @ derivative
        product = factor @ product
    return trace_first(product, 2**length), trace_first(derivative, 2**length)


def explicit_derivative(length, parameter):
    """Independently form all N products with one L replaced by its derivative."""
    factors = local_factors(length, parameter)
    result = np.zeros_like(factors[0])
    for differentiated in range(length):
        term = np.eye(len(factors[0]), dtype=complex)
        for index, factor in enumerate(factors):
            if index != differentiated:
                term = factor @ term
        result += term
    return trace_first(result, 2**length)


def tensor_site(matrix, site, length):
    result = np.ones((1, 1), dtype=complex)
    for position in range(length):
        result = np.kron(result, matrix if position == site else np.eye(2))
    return result


def spin_hamiltonian(length, coupling, close_ring=True):
    """Independent Pauli tensor construction; does not call any swap routine."""
    spin = [np.array([[0, 1], [1, 0]], dtype=complex) / 2,
            np.array([[0, -1j], [1j, 0]], dtype=complex) / 2,
            np.diag([1, -1]).astype(complex) / 2]
    identity = np.eye(2**length)
    embedded = [[tensor_site(s, site, length) for site in range(length)] for s in spin]
    result = np.zeros_like(identity, dtype=complex)
    for left in range(length if close_ring else length - 1):
        right = (left + 1) % length
        bond = identity / 4
        for component in embedded:
            bond = bond - component[left] @ component[right]
        result += coupling * bond
    return result


def right_shift(length):
    """Independent cyclic bit rotation, rather than a product of swap matrices."""
    result = np.zeros((2**length, 2**length), dtype=complex)
    for column in range(2**length):
        bits = f"{column:0{length}b}"
        result[int(bits[-1] + bits[:-1], 2), column] = 1
    return result


def transfer_polynomial(length):
    """Coefficients in z=lambda-i/2, low degree first, before evaluating z."""
    identity = np.eye(2**(length + 1), dtype=complex)
    coefficients = [identity]
    for site in range(1, length + 1):
        constant = 1j * swap(length + 1, 0, site)
        updated = [np.zeros_like(identity) for _ in range(len(coefficients) + 1)]
        for degree, coefficient in enumerate(coefficients):
            updated[degree] += constant @ coefficient
            updated[degree + 1] += coefficient
        coefficients = updated
    return [trace_first(value, 2**length) for value in coefficients]


def ybe_checks(parameters):
    p12, p13, p23 = swap(3, 0, 1), swap(3, 0, 2), swap(3, 1, 2)
    defect_direction = p12 @ p13 - p13 @ p12
    records = []
    for pair in parameters["ybe_pairs"]:
        a, b = complex_value(pair["a"]), complex_value(pair["b"])
        good_middle = a + b
        wrong_middle = good_middle + complex_value(parameters["wrong_middle_increment"])
        left = r_matrix(a, p12) @ r_matrix(good_middle, p13) @ r_matrix(b, p23)
        right = r_matrix(b, p23) @ r_matrix(good_middle, p13) @ r_matrix(a, p12)
        wrong_left = r_matrix(a, p12) @ r_matrix(wrong_middle, p13) @ r_matrix(b, p23)
        wrong_right = r_matrix(b, p23) @ r_matrix(wrong_middle, p13) @ r_matrix(a, p12)
        records.append({"arguments": pair, "ybe_residual": residual(left, right),
                        "wrong_middle_residual": residual(wrong_left, wrong_right),
                        "defect_polynomial_residual": residual(wrong_left - wrong_right,
                            (wrong_middle - a - b) * defect_direction)})
    return records


def exchange_checks(length, parameter, other):
    factors = length + 2
    r = r_matrix(parameter - other, swap(factors, 0, 1))
    ta = ordered_product(local_factors(length, parameter, auxiliaries=2, auxiliary=0))
    tb = ordered_product(local_factors(length, other, auxiliaries=2, auxiliary=1))
    t1, t2 = transfer(length, parameter), transfer(length, other)
    return {"rtt_residual": residual(r @ ta @ tb, tb @ ta @ r),
            "transfer_commutator_residual": residual(t1 @ t2, t2 @ t1)}


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


def phase_aligned_error(actual, expected):
    """Euclidean distance of unit vectors after minimizing over global phase."""
    overlap = np.vdot(expected, actual)
    phase = overlap / abs(overlap) if abs(overlap) > 0 else 1.0
    return float(np.linalg.norm(actual - phase*expected))


def b_state_checks(length, parameters, shift, h_spin):
    """Actual upper-right monodromy block, not a preassigned plane wave."""
    coupling, dimension = parameters["coupling_J"], 2**length
    vacuum = np.zeros(dimension, dtype=complex)
    vacuum[0] = 1.0
    sites = np.arange(1, length+1)
    sector = np.array([2**(length-x) for x in sites])
    outside = np.ones(dimension, dtype=bool)
    outside[sector] = False
    rows = []
    for m in parameters["one_magnon_mode_numbers"]:
        phi = 2*np.pi*m/length
        cotangent_value = 0.5/math.tan(phi/2)
        # These particular four-site values are exactly dyadic, not rounded cotangents.
        lam = 0.5*m if length == 4 else cotangent_value
        monodromy = ordered_product(local_factors(length, lam))
        b_block = monodromy[:dimension, dimension:]
        actual = b_block @ vacuum
        norm = float(np.linalg.norm(actual))
        if norm == 0:
            raise AssertionError("Selected B state unexpectedly vanishes.")
        unit = actual/norm
        formula = np.zeros(dimension, dtype=complex)
        formula[sector] = 1j*(lam+0.5j)**(length-sites)*(lam-0.5j)**(sites-1)
        coordinate_rapidity = -lam
        coordinate_z = (coordinate_rapidity+0.5j)/(coordinate_rapidity-0.5j)
        adjacent_ratio = (lam-0.5j)/(lam+0.5j)
        coordinate_wave = np.zeros(dimension, dtype=complex)
        coordinate_wave[sector] = np.exp(-1j*phi*(sites-1))/math.sqrt(length)
        wrong_wave = np.conj(coordinate_wave)
        translation_eigenvalue = 1j*m if length == 4 else np.exp(1j*phi)
        energy_over_J = 1.0 if length == 4 else 1-math.cos(phi)
        good = {
            "rapidity_vs_cotangent_absolute_error": abs(lam-cotangent_value),
            "coefficient_formula_relative_error": float(np.linalg.norm(actual-formula)/np.linalg.norm(formula)),
            "outside_one_magnon_relative_norm": float(np.linalg.norm(actual[outside])/norm),
            "adjacent_ratio_max_absolute_error": float(np.max(np.abs(actual[sector[1:]]/actual[sector[:-1]]-adjacent_ratio))),
            "coordinate_rapidity_ratio_absolute_error": float(abs(coordinate_z-adjacent_ratio)),
            "coordinate_plane_wave_phase_aligned_error": phase_aligned_error(unit, coordinate_wave),
            "active_translation_eigenstate_residual": float(np.linalg.norm(shift@unit-translation_eigenvalue*unit)),
            "pauli_hamiltonian_eigenstate_residual_over_J": float(np.linalg.norm(h_spin@unit/coupling-energy_over_J*unit)),
            "opposite_momentum_energy_residual_over_J": float(np.linalg.norm(h_spin@wrong_wave/coupling-energy_over_J*wrong_wave)),
        }
        wrong_translation = float(np.linalg.norm(shift@unit-np.conj(translation_eigenvalue)*unit))
        wrong_phase = phase_aligned_error(unit, wrong_wave)
        good["wrong_translation_residual_formula_error"] = abs(wrong_translation-2*abs(math.sin(phi)))
        good["opposite_momentum_orthogonality_error"] = float(abs(np.vdot(wrong_wave, unit)))
        row = {"mode_number_m": m, "lambda_B": lam,
               "lambda_coordinate": coordinate_rapidity,
               "k_B": phi, "k_coordinate": -phi,
               "expected_active_U_eigenvalue": complex_pair(translation_eigenvalue),
               "expected_energy_over_J": energy_over_J,
               "B_vacuum_norm": norm,
               "one_spin_coefficients_in_site_order": [complex_pair(a) for a in actual[sector]],
               "checks": good,
               "negative_controls": {
                   "wrong_momentum_translation_residual": wrong_translation,
                   "wrong_momentum_phase_aligned_state_error": wrong_phase}}
        if length == 4:
            # Independent exact dyadic coefficients at lambda_B=+/-1/2.
            expected = (-1-1j*m)/4*np.array([(1j)**(-m*(x-1)) for x in sites])
            good["four_site_exact_dyadic_coefficient_max_error"] = float(np.max(np.abs(actual[sector]-expected)))
            row["four_site_exact_dyadic_coefficients"] = [complex_pair(a) for a in expected]
        rows.append(row)
    return rows


def chain_checks(length, parameters):
    coupling = parameters["coupling_J"]
    identity = np.eye(2**length)
    tau0, derivative = transfer_derivative(length, 0.5j)
    shift = right_shift(length)
    shift_left = shift.conj().T
    swaps = sum((swap(length, n, (n + 1) % length) for n in range(length)),
                np.zeros_like(identity, dtype=complex))
    logarithmic = np.linalg.solve(tau0, derivative)
    h_transfer = coupling * length * identity / 2 - 0.5j * coupling * logarithmic
    h_spin = spin_hamiltonian(length, coupling)
    h_norm = max(coupling, float(np.linalg.norm(h_spin, "fro")))
    matrix_checks = {
        "regularity_residual": residual(tau0, (1j**length) * shift),
        "log_derivative_residual": residual(logarithmic, -1j * swaps),
        "hamiltonian_relative_frobenius_error": float(np.linalg.norm(h_transfer - h_spin, "fro") / h_norm),
        "hermiticity_residual": residual(h_transfer, h_transfer.conj().T),
        "reversed_order_shift_residual": residual(transfer(length, 0.5j, reverse=True), (1j**length) * shift_left),
    }
    direction_errors = []
    for m in range(length):
        wave_number = 2 * np.pi * m / length
        vector = np.zeros(2**length, dtype=complex)
        for site in range(length):
            vector[2**(length - 1 - site)] = np.exp(1j * wave_number * site) / math.sqrt(length)
        direction_errors.append(float(np.linalg.norm(shift @ vector - np.exp(-1j * wave_number) * vector)))
    matrix_checks["one_magnon_translation_max_residual"] = max(direction_errors)
    generic = []
    for pair in parameters["generic_spectral_pairs"]:
        lam, mu = complex_value(pair["lambda"]), complex_value(pair["mu"])
        r = r_matrix(lam - mu, swap(3, 0, 1))
        l1, l2 = r_matrix(lam - 0.5j, swap(3, 0, 2)), r_matrix(mu - 0.5j, swap(3, 1, 2))
        record = {"parameters": pair, **exchange_checks(length, lam, mu),
                  "local_exchange_residual": residual(r @ l1 @ l2, l2 @ l1 @ r)}
        _, recurrent = transfer_derivative(length, lam)
        record["analytic_derivative_residual"] = residual(recurrent, explicit_derivative(length, lam))
        tau = transfer(length, lam)
        record["hamiltonian_transfer_commutator_residual"] = residual(h_spin @ tau, tau @ h_spin)
        generic.append(record)
    exceptional = []
    mu = complex_value(parameters["exceptional_mu"])
    for sign in (-1, 1):
        lam = mu + sign * 1j
        exceptional.append({"difference_imaginary_part": sign,
                            "r_rank": int(np.linalg.matrix_rank(r_matrix(sign * 1j, swap(2, 0, 1)))),
                            **exchange_checks(length, lam, mu)})
    negative = {
        "omitted_closing_bond_relative_error": float(np.linalg.norm(h_transfer - spin_hamiltonian(length, coupling, close_ring=False), "fro") / h_norm),
        "wrong_translation_direction_residual": residual(tau0, (1j**length) * shift_left)}
    result = {"sites": length, "physical_dimension": 2**length,
              "identities": matrix_checks, "generic": generic,
              "exceptional_differences": exceptional, "negative_controls": negative,
              "B_state_one_magnons": b_state_checks(length, parameters, shift, h_spin)}
    if length == 3:
        expected = [-1j * shift, -swaps, 3j * identity, 2 * identity]
        actual = transfer_polynomial(3)
        spectrum = np.linalg.eigvalsh(h_spin) / coupling
        target = np.r_[np.zeros(4), 1.5 * np.ones(4)]
        result["three_site_exact_checks"] = {
            "polynomial_coefficient_residuals": [residual(a, b) for a, b in zip(actual, expected)],
            "spectrum_max_absolute_error_over_J": float(np.max(np.abs(spectrum - target))),
            "energies_over_J": spectrum.tolist(), "expected_energies_over_J": target.tolist()}
    return result


def partial_trace_control():
    e01, e10 = np.array([[0, 1], [0, 0]]), np.array([[0, 0], [1, 0]])
    x, z = np.array([[0, 1], [1, 0]]), np.diag([1, -1])
    a, b = np.kron(e01, x), np.kron(e10, z)
    left, right = trace_first(a @ b, 2), trace_first(b @ a, 2)
    return {"invalid_cyclicity_residual": residual(left, right),
            "tr_ab_matches_xz_residual": residual(left, x @ z),
            "tr_ba_matches_zx_residual": residual(right, z @ x),
            "full_trace_difference": float(abs(np.trace(a @ b) - np.trace(b @ a)))}


def bethe_vector(length, roots):
    dimension = 2**length
    state = np.zeros(dimension, dtype=complex)
    state[0] = 1
    for root in reversed(roots):
        monodromy = ordered_product(local_factors(length, root))
        state = monodromy[:dimension, dimension:] @ state
    return state


def bethe_scalars(length, roots, parameter):
    """Wanted eigenvalue and each unwanted coefficient before Bethe equations."""
    a = lambda value: (value+0.5j)**length
    d = lambda value: (value-0.5j)**length
    f = lambda value: (value-1j)/value
    h = lambda value: (value+1j)/value
    wanted = a(parameter)*np.prod([f(parameter-r) for r in roots])
    wanted += d(parameter)*np.prod([h(parameter-r) for r in roots])
    unwanted, coefficient_residuals = [], []
    for j, root in enumerate(roots):
        left = a(root)*np.prod([f(root-other) for k, other in enumerate(roots) if k != j])
        right = d(root)*np.prod([h(root-other) for k, other in enumerate(roots) if k != j])
        coefficient_residuals.append(float(abs(left-right)/(abs(left)+abs(right))))
        unwanted.append(1j*(left-right)/(parameter-root))
    return complex(wanted), unwanted, coefficient_residuals


def regular_case(length, roots, probes, coupling, minimum_separation):
    distances = [abs(r-s) for r in roots for s in (0.5j, -0.5j)]
    distances += [abs(r-s-shift) for j, r in enumerate(roots)
                  for k, s in enumerate(roots) if j != k for shift in (0, 1j, -1j)]
    distances += [abs(u-r) for u in probes for r in roots]
    if min(distances) <= minimum_separation:
        raise ValueError("Regular checks must avoid singular roots, differences and probes.")
    state = bethe_vector(length, roots)
    norm = float(np.linalg.norm(state))
    if norm == 0:
        raise AssertionError("Regular Bethe vector vanishes.")
    unit = state/norm
    hamiltonian = spin_hamiltonian(length, coupling)
    shift = right_shift(length)
    formal_energy = sum(0.5/(r*r+0.25) for r in roots)
    formal_phase = np.prod([(r+0.5j)/(r-0.5j) for r in roots])
    magnon_counts = np.array([bin(n).count('1') for n in range(2**length)])
    rows = []
    for u in probes:
        wanted, coefficients, bethe_residuals = bethe_scalars(length, roots, u)
        unwanted_state = np.zeros_like(state)
        terms = []
        for j, coefficient in enumerate(coefficients):
            replacement = bethe_vector(length, [u]+[r for k, r in enumerate(roots) if k != j])
            unwanted_state += coefficient*replacement
            terms.append(abs(coefficient)*np.linalg.norm(replacement))
        left = transfer(length, u)@state
        right = wanted*state+unwanted_state
        scale = max(float(np.linalg.norm(left)), abs(wanted)*norm+sum(terms))
        wanted_scale = max(float(np.linalg.norm(left)), abs(wanted)*norm)
        row = {"parameter": complex_pair(u), "wanted_lambda": complex_pair(wanted),
               "unwanted_coefficients": [complex_pair(c) for c in coefficients],
               "bethe_coefficient_residuals": bethe_residuals,
               "off_shell_action_residual": float(np.linalg.norm(left-right)/scale),
               "wanted_only_transfer_residual": float(np.linalg.norm(left-wanted*state)/wanted_scale),
               "unwanted_sum_relative_norm": float(np.linalg.norm(unwanted_state)/wanted_scale)}
        rows.append(row)
    return {"sites": length, "roots": [complex_pair(r) for r in roots],
            "magnons": len(roots), "minimum_denominator_distance": float(min(distances)),
            "raw_vector_norm": norm, "raw_vector_norm_squared": norm*norm,
            "formal_energy_over_J": complex_pair(formal_energy),
            "formal_active_U_eigenvalue": complex_pair(formal_phase),
            "rayleigh_energy_over_J": complex_pair(np.vdot(unit, hamiltonian@unit)/coupling),
            "hamiltonian_eigenstate_residual_over_J": float(np.linalg.norm(hamiltonian@unit/coupling-formal_energy*unit)),
            "translation_eigenstate_residual": float(np.linalg.norm(shift@unit-formal_phase*unit)),
            "magnon_sector_residual": float(np.linalg.norm((magnon_counts-len(roots))*unit)),
            "B_order_reversal_relative_error": float(np.linalg.norm(state-bethe_vector(length, list(reversed(roots))))/norm),
            "probes": rows}


def exact_five_site_pair():
    """Rational Baxter polynomial and independent rational bond/state actions."""
    # Q(v)=v^2-1/4. Expand 2 Re[(v+i/2)^5 Q(v-i)] using integer i powers.
    numerator = [Fraction(0) for _ in range(8)]
    for k in range(6):
        for degree, coefficient, imaginary_power in [(0, Fraction(-5, 4), 0),
                                                      (1, Fraction(-2), 1),
                                                      (2, Fraction(1), 0)]:
            power = 5-k+imaginary_power
            if power % 2 == 0:
                numerator[k+degree] += 2*Fraction(math.comb(5, k), 2**(5-k))*coefficient*(-1)**(power//2)
    remainder = numerator[:]
    quotient = [Fraction(0) for _ in range(6)]
    for degree in range(7, 1, -1):
        quotient[degree-2] = remainder[degree]
        remainder[degree-2] += remainder[degree]/4
        remainder[degree] = Fraction(0)
    assert all(value == 0 for value in remainder)
    assert quotient == list(map(Fraction, [0, Fraction(21, 8), 0, 3, 0, 2]))
    pairs = list(combinations(range(5), 2))
    values = {pair: Fraction(-1 if (pair[1]-pair[0]) in (1, 4) else 1, 8) for pair in pairs}
    action = {pair: Fraction(0) for pair in pairs}
    shifted = {pair: Fraction(0) for pair in pairs}
    for pair, value in values.items():
        shifted[tuple(sorted((x+1) % 5 for x in pair))] += value
        for n in range(5):
            neighbor = (n+1) % 5
            swapped = tuple(sorted(neighbor if x == n else n if x == neighbor else x for x in pair))
            action[pair] += value/2
            action[swapped] -= value/2
    assert all(action[pair] == 2*value and shifted[pair] == value for pair, value in values.items())
    norm_squared = sum(value*value for value in values.values())
    assert norm_squared == Fraction(5, 32)
    actual = bethe_vector(5, [0.5, -0.5])
    expected = np.zeros(32, dtype=complex)
    for pair, value in values.items():
        expected[sum(2**(4-x) for x in pair)] = float(value)
    return {"Q_coefficients_low_degree_first": ["-1/4", "0", "1"],
            "Lambda_coefficients_low_degree_first": [str(value) for value in quotient],
            "exact_Baxter_remainder_zero": True, "exact_rational_H_and_U_actions": True,
            "raw_norm_squared_exact": str(norm_squared),
            "raw_dyadic_state_max_error": float(np.max(np.abs(actual-expected)))}


def regular_bethe_checks(settings, coupling):
    probes = [complex_value(p) for p in settings["probe_parameters"]]
    recipes = {
        "four_site_one_magnon": (4, [0.5], 1.0, 1j),
        "five_site_dyadic_pair": (5, [0.5, -0.5], 2.0, 1.0),
        "six_site_asymmetric_pair": (6, [(-math.sqrt(3)+math.sqrt(11))/8,
                                         (-math.sqrt(3)-math.sqrt(11))/8],
                                      2.5, (1+1j*math.sqrt(3))/2),
    }
    on_shell = []
    for name in settings["on_shell_recipes"]:
        length, roots, energy, phase = recipes[name]
        row = regular_case(length, roots, probes, coupling, settings["minimum_regular_separation"])
        row["recipe"] = name
        row["known_energy_formula_error"] = abs(complex_value(row["formal_energy_over_J"])-energy)
        row["known_translation_formula_error"] = abs(complex_value(row["formal_active_U_eigenvalue"])-phase)
        if name == "six_site_asymmetric_pair":
            def polynomial(u):
                return 2*u**6+2.5*u**4-math.sqrt(3)*u**3+11*u*u/8-5*math.sqrt(3)*u/4-9/32
            row["independent_transfer_polynomial_errors"] = [float(abs(complex_value(p["wanted_lambda"])-polynomial(u))) for p, u in zip(row["probes"], probes)]
        on_shell.append(row)
    off_shell = [regular_case(item["sites"], [complex_value(r) for r in item["roots"]], probes,
                              coupling, settings["minimum_regular_separation"])
                 for item in settings["off_shell_cases"]]
    roots = recipes["six_site_asymmetric_pair"][1]
    first = roots[0]+settings["energy_preserving_root_perturbation"]
    second = -math.sqrt(1/(5-1/(first*first+0.25))-0.25)
    negative = regular_case(6, [first, second], probes, coupling, settings["minimum_regular_separation"])
    negative["formal_energy_difference_from_5_over_2"] = abs(complex_value(negative["formal_energy_over_J"])-2.5)
    result = {"on_shell": on_shell, "off_shell": off_shell,
              "energy_preserving_wrong_roots": negative, "exact_five_site_pair": exact_five_site_pair()}
    good, bad = [], []
    for row in on_shell+off_shell+[negative]:
        if row["raw_vector_norm"] < settings["minimum_vector_norm"]:
            raise AssertionError("Unacceptably small Bethe-vector norm.")
        good.extend([row["magnon_sector_residual"], row["B_order_reversal_relative_error"]])
        good.extend(p["off_shell_action_residual"] for p in row["probes"])
    for row in on_shell:
        good.extend([row["hamiltonian_eigenstate_residual_over_J"], row["translation_eigenstate_residual"],
                     row["known_energy_formula_error"], row["known_translation_formula_error"]])
        good.extend(row.get("independent_transfer_polynomial_errors", []))
        for p in row["probes"]:
            good.extend(p["bethe_coefficient_residuals"])
            good.extend([p["wanted_only_transfer_residual"], p["unwanted_sum_relative_norm"]])
    good.extend([negative["formal_energy_difference_from_5_over_2"], result["exact_five_site_pair"]["raw_dyadic_state_max_error"]])
    for row in off_shell+[negative]:
        bad.extend([max(p["wanted_only_transfer_residual"] for p in row["probes"]),
                    max(max(p["bethe_coefficient_residuals"]) for p in row["probes"])])
    bad.append(negative["hamiltonian_eigenstate_residual_over_J"])
    return result, good, bad


def four_site_sector_checks(settings, coupling, tolerance, negative_minimum):
    """One complete M=2 sector, with exact Gaussian-rational projectors.

    Exact complex numbers below are (Fraction real, Fraction imaginary) pairs.
    Square roots enter only the separate floating-point normalization checks.
    """
    required = {"sites": 4, "down_spins": 2, "descendant_mode_numbers": [1, 2, 3],
                "regular_root_recipe": "opposite_one_over_two_sqrt3",
                "negative_controls": ["omit_singular", "replace_singular_with_regular"]}
    if set(settings) != set(required) | {"polynomial_probe"} or not same_input_values({k: settings[k] for k in required}, required):
        raise ValueError("The completeness check fixes the stated four-site, two-down-spin sector.")
    probe = complex_value(settings["polynomial_probe"])
    if not math.isfinite(probe.real) or not math.isfinite(probe.imag):
        raise ValueError("Use a finite polynomial probe.")

    zero, one = (Fraction(0), Fraction(0)), (Fraction(1), Fraction(0))
    def add(a, b):
        return (a[0]+b[0], a[1]+b[1])
    def multiply(a, b):
        return (a[0]*b[0]-a[1]*b[1], a[0]*b[1]+a[1]*b[0])
    def scale(s, a):
        return (s*a[0], s*a[1])
    def conjugate(a):
        return (a[0], -a[1])
    def total(values):
        result = zero
        for value in values:
            result = add(result, value)
        return result
    def inner(a, b):
        return total(multiply(conjugate(x), y) for x, y in zip(a, b))
    def action(matrix, vector):
        return [total(scale(a, v) for a, v in zip(row, vector)) for row in matrix]
    def rank(columns):
        rows = [list(row) for row in zip(*columns)]
        pivot = 0
        for column in range(len(columns)):
            found = next((r for r in range(pivot, len(rows)) if rows[r][column] != zero), None)
            if found is None:
                continue
            rows[pivot], rows[found] = rows[found], rows[pivot]
            value = rows[pivot][column]
            inverse = scale(1/(value[0]**2+value[1]**2), conjugate(value))
            rows[pivot] = [multiply(inverse, v) for v in rows[pivot]]
            for r in range(len(rows)):
                if r != pivot:
                    factor = rows[r][column]
                    rows[r] = [add(a, scale(-1, multiply(factor, b))) for a, b in zip(rows[r], rows[pivot])]
            pivot += 1
            if pivot == len(rows):
                break
        return pivot

    basis = list(combinations(range(1, 5), 2))
    fourth_roots = [one, (Fraction(0), Fraction(1)), scale(-1, one), (Fraction(0), Fraction(-1))]
    raw = [[one]*6]
    raw += [[add(fourth_roots[m*x % 4], fourth_roots[m*y % 4]) for x, y in basis]
            for m in settings["descendant_mode_numbers"]]
    raw += [[scale(Fraction(v), one) for v in values]
            for values in ([1, 0, -1, -1, 0, 1], [1, -2, 1, 1, -2, 1])]
    energies = [0, 1, 2, 1, 1, 3]
    spins_squared = [6, 2, 2, 2, 0, 0]
    phases = [one, fourth_roots[3], fourth_roots[2], fourth_roots[1], fourth_roots[2], one]
    names = ["spin_two_descendant", "spin_one_m1", "spin_one_m2", "spin_one_m3",
             "singular_singlet", "regular_singlet"]
    norms = [inner(v, v)[0] for v in raw]
    assert norms == list(map(Fraction, [6, 8, 8, 8, 4, 12]))
    for j, v in enumerate(raw):
        for k, w in enumerate(raw):
            expected = (norms[j], Fraction(0)) if j == k else zero
            assert inner(v, w) == expected

    h = [[Fraction(0) for _ in basis] for _ in basis]
    u = [[Fraction(0) for _ in basis] for _ in basis]
    for column, pair in enumerate(basis):
        occupied = set(pair)
        u[basis.index(tuple(sorted(x % 4+1 for x in pair)))][column] = Fraction(1)
        for left in range(1, 5):
            right = left % 4+1
            if (left in occupied) != (right in occupied):
                target = tuple(sorted(occupied.symmetric_difference({left, right})))
                h[column][column] += Fraction(1, 2)
                h[basis.index(target)][column] -= Fraction(1, 2)
    for j, v in enumerate(raw):
        assert action(h, v) == [scale(energies[j], a) for a in v]
        assert action(u, v) == [multiply(phases[j], a) for a in v]
    projectors = [[[scale(1/norm, multiply(v[r], conjugate(v[c]))) for c in range(6)]
                   for r in range(6)] for v, norm in zip(raw, norms)]
    for r in range(6):
        for c in range(6):
            assert total(p[r][c] for p in projectors) == (one if r == c else zero)
    assert rank(raw) == 6

    def as_complex(value):
        return complex(float(value[0]), float(value[1]))
    def serialized(vector):
        return [[str(re), str(im)] for re, im in vector]
    columns = np.array([[as_complex(value) for value in vector] for vector in raw]).T
    vectors = columns/np.sqrt(np.array(norms, dtype=float))
    h_exact, u_exact = np.array(h, dtype=float), np.array(u, dtype=float)
    indices = np.array([sum(1 << (4-x) for x in pair) for pair in basis])
    h_full, u_full = spin_hamiltonian(4, coupling), right_shift(4)
    h_sector, u_sector = h_full[np.ix_(indices, indices)]/coupling, u_full[np.ix_(indices, indices)]
    good = {}
    def check(name, value):
        good[name] = float(value)
    check("independent_Pauli_vs_exact_bond_H_frobenius_error", np.linalg.norm(h_sector-h_exact))
    check("full_vs_exact_sector_U_frobenius_error", np.linalg.norm(u_sector-u_exact))
    check("normalized_Gram_frobenius_error", np.linalg.norm(vectors.conj().T@vectors-np.eye(6)))
    check("resolution_of_identity_frobenius_error", np.linalg.norm(vectors@vectors.conj().T-np.eye(6)))
    phase_values = np.array([as_complex(p) for p in phases])
    check("H_eigenstate_frobenius_error_over_J", np.linalg.norm(h_sector@vectors-vectors*np.array(energies)))
    check("U_eigenstate_frobenius_error", np.linalg.norm(u_sector@vectors-vectors*phase_values))
    spectrum = np.linalg.eigvalsh(h_sector)
    check("energy_spectrum_max_error_over_J", np.max(abs(spectrum-np.array(sorted(energies)))))

    lowering = sum((tensor_site(np.array([[0, 0], [1, 0]]), site, 4) for site in range(4)),
                   np.zeros((16, 16), dtype=complex))
    sz = sum((tensor_site(np.diag([.5, -.5]), site, 4) for site in range(4)),
             np.zeros((16, 16), dtype=complex))
    raising = lowering.conj().T
    spin_squared = sz@sz+(raising@lowering+lowering@raising)/2
    check("H_lowering_commutator_frobenius_error_over_J", np.linalg.norm((h_full@lowering-lowering@h_full)/coupling))
    check("U_lowering_commutator_frobenius_error", np.linalg.norm(u_full@lowering-lowering@u_full))
    check("spin_squared_labels_frobenius_error", np.linalg.norm(spin_squared[np.ix_(indices, indices)]@vectors-vectors*np.array(spins_squared)))
    vacuum = np.eye(16, dtype=complex)[:, 0]
    twice_lowered = lowering@lowering@vacuum
    check("vacuum_twice_lowered_norm_squared_error", abs(np.vdot(twice_lowered, twice_lowered)-24))
    check("spin_two_descendant_vector_error", np.linalg.norm(twice_lowered[indices]/(2*math.sqrt(6))-vectors[:, 0]))
    for j, m in enumerate(settings["descendant_mode_numbers"], 1):
        wave = np.zeros(16, dtype=complex)
        for x in range(1, 5):
            wave[1 << (4-x)] = as_complex(fourth_roots[m*x % 4])/2
        lowered = lowering@wave
        check(f"m{m}_one_magnon_highest_weight_error", np.linalg.norm(raising@wave))
        check(f"m{m}_lowered_norm_squared_error", abs(np.vdot(lowered, lowered)-2))
        check(f"m{m}_normalized_descendant_error", np.linalg.norm(lowered[indices]/math.sqrt(2)-vectors[:, j]))

    # Expand the actual local factors L(lambda)=lambda I+L(0), without a large-lambda fit.
    coefficients = [np.eye(32, dtype=complex)]
    for constant in local_factors(4, 0):
        updated = [np.zeros((32, 32), dtype=complex) for _ in range(len(coefficients)+1)]
        for degree, coefficient in enumerate(coefficients):
            updated[degree] += constant@coefficient
            updated[degree+1] += coefficient
        coefficients = updated
    b_coefficients = [coefficient[:16, 16:] for coefficient in coefficients]
    check("B_degree_four_coefficient_frobenius_norm", np.linalg.norm(b_coefficients[4]))
    check("B_degree_three_equals_i_lowering_frobenius_error", np.linalg.norm(b_coefficients[3]-1j*lowering))
    evaluated = sum((probe**degree*c for degree, c in enumerate(b_coefficients)), np.zeros((16, 16), dtype=complex))
    direct_b = ordered_product(local_factors(4, probe))[:16, 16:]
    check("B_polynomial_evaluation_frobenius_error", np.linalg.norm(evaluated-direct_b))
    root = 1/(2*math.sqrt(3))
    actual_regular = bethe_vector(4, [root, -root])
    expected_regular = np.zeros(16, dtype=complex)
    expected_regular[indices] = (2/27)*columns[:, 5]
    check("actual_regular_B_coefficient_relative_error", np.linalg.norm(actual_regular-expected_regular)/np.linalg.norm(expected_regular))
    check("actual_regular_B_phase_aligned_error", phase_aligned_error(actual_regular/np.linalg.norm(actual_regular),
                                                                       expected_regular/np.linalg.norm(expected_regular)))
    check("singular_raw_B_vector_norm", np.linalg.norm(bethe_vector(4, [.5j, -.5j])))

    controls = {}
    for name, kept in [("omit_singular", [0, 1, 2, 3, 5]),
                       ("replace_singular_with_regular", [0, 1, 2, 3, 5, 5])]:
        selected = vectors[:, kept]
        selected_raw = [raw[j] for j in kept]
        exact_rank = rank(selected_raw)
        assert exact_rank == 5
        closure = selected@selected.conj().T-np.eye(6)
        defect_exact = [[add(total(projectors[j][r][c] for j in kept), scale(-1, one if r == c else zero))
                         for c in range(6)] for r in range(6)]
        expected_defect = [[scale(-1, projectors[4][r][c]) for c in range(6)] for r in range(6)]
        if len(kept) == 6:
            expected_defect = [[add(expected_defect[r][c], projectors[5][r][c]) for c in range(6)] for r in range(6)]
        assert defect_exact == expected_defect
        retained_h_error = float(np.linalg.norm(h_sector@selected-selected*np.array(energies)[kept]))
        retained_u_error = float(np.linalg.norm(u_sector@selected-selected*phase_values[kept]))
        check(name+"_retained_H_error_over_J", retained_h_error)
        check(name+"_retained_U_error", retained_u_error)
        controls[name] = {"column_count": len(kept), "exact_rank": exact_rank,
                          "distinct_energies_over_J": sorted(set(energies[j] for j in kept)),
                          "retained_H_eigenstate_frobenius_error_over_J": retained_h_error,
                          "retained_U_eigenstate_frobenius_error": retained_u_error,
                          "Gram_frobenius_error": float(np.linalg.norm(selected.conj().T@selected-np.eye(len(kept)))),
                          "resolution_frobenius_error": float(np.linalg.norm(closure)),
                          "missing_singular_projection_norm": float(np.linalg.norm(selected.conj().T@vectors[:, 4])),
                          "exact_projector_defect_verified": True}
    bad = [item["resolution_frobenius_error"] for item in controls.values()]
    bad.append(controls["replace_singular_with_regular"]["Gram_frobenius_error"])
    if not all(math.isfinite(v) and v <= tolerance for v in good.values()):
        raise AssertionError("A four-site state, symmetry or completeness identity failed.")
    if not all(math.isfinite(v) and v >= negative_minimum for v in bad):
        raise AssertionError("An incomplete four-site basis was not distinguished.")
    return {"sites": 4, "down_spins": 2, "basis": [list(pair) for pair in basis],
            "state_order": names, "raw_vectors_gaussian_rational": [serialized(v) for v in raw],
            "exact_raw_norms_squared": [str(n) for n in norms], "energies_over_J": energies,
            "active_U_eigenvalues": [complex_pair(p) for p in phase_values],
            "total_spin_squared_eigenvalues": spins_squared,
            "exact_checks": {"orthogonal_Gram": True, "bond_H_eigenstates": True,
                             "active_U_eigenstates": True, "projector_sum_is_identity": True,
                             "full_rank": rank(raw), "negative_projector_defects": True},
            "actual_regular_B": {"roots": [root, -root], "raw_coefficient_factor": "2/27",
                                 "expected_raw_norm_squared": "16/243",
                                 "raw_norm_squared": float(np.vdot(actual_regular, actual_regular).real)},
            "energies_with_multiplicity_over_J": spectrum.tolist(),
            "expected_energies_with_multiplicity_over_J": sorted(energies),
            "equation_errors": good, "negative_controls": controls,
            "summary": {"largest_identity_residual": max(good.values()),
                        "smallest_completeness_failure_residual": min(bad)}}


def run(parameters):
    lengths = parameters["chain_lengths"]
    if not lengths or any(not isinstance(n, int) or n < 3 or n > 5 for n in lengths):
        raise ValueError("This bounded dense-matrix experiment supports N=3,4,5.")
    if parameters["coupling_J"] <= 0:
        raise ValueError("The ferromagnetic coupling J must be positive.")
    if parameters["one_magnon_mode_numbers"] != [-1, 1]:
        raise ValueError("This momentum-orientation check uses the nonzero modes m=-1,+1.")
    result = {"schema_version": 3,
              "environment": {"python": platform.python_version(), "numpy": np.__version__,
                              "arithmetic": "IEEE 754 binary64; complex128 matrices"},
              "inputs": parameters, "ybe": ybe_checks(parameters),
              "chains": [chain_checks(length, parameters) for length in lengths],
              "partial_trace_counterexample": partial_trace_control()}
    good, bad = [], []
    for item in result["ybe"]:
        good.extend([item["ybe_residual"], item["defect_polynomial_residual"]])
        bad.append(item["wrong_middle_residual"])
    for chain in result["chains"]:
        good.extend(chain["identities"].values())
        for item in chain["generic"]:
            good.extend(value for key, value in item.items() if key != "parameters")
        for item in chain["exceptional_differences"]:
            good.extend([item["rtt_residual"], item["transfer_commutator_residual"]])
            if item["r_rank"] != (3 if item["difference_imaginary_part"] == 1 else 1):
                raise AssertionError("Unexpected singular R-matrix rank.")
        bad.extend(chain["negative_controls"].values())
        for row in chain["B_state_one_magnons"]:
            good.extend(row["checks"].values())
            bad.extend(row["negative_controls"].values())
        if "three_site_exact_checks" in chain:
            item = chain["three_site_exact_checks"]
            good.extend(item["polynomial_coefficient_residuals"])
            good.append(item["spectrum_max_absolute_error_over_J"])
    partial = result["partial_trace_counterexample"]
    good.extend([partial["tr_ab_matches_xz_residual"], partial["tr_ba_matches_zx_residual"], partial["full_trace_difference"]])
    bad.append(partial["invalid_cyclicity_residual"])
    regular, regular_good, regular_bad = regular_bethe_checks(parameters["regular_bethe"], parameters["coupling_J"])
    result["regular_bethe_vectors"] = regular
    good.extend(regular_good)
    bad.extend(regular_bad)
    result["summary"] = {"largest_identity_residual": max(good),
                         "smallest_negative_control_residual": min(bad)}
    if not all(math.isfinite(value) and value <= parameters["equation_tolerance"] for value in good):
        raise AssertionError("A claimed algebra identity failed the numerical tolerance.")
    if not all(math.isfinite(value) and value >= parameters["negative_control_minimum"] for value in bad):
        raise AssertionError("A negative control was not distinguished.")
    # Keep the preceding baseline and its summary unchanged; the new sector has its own summary.
    result["four_site_sector"] = four_site_sector_checks(parameters["four_site_sector"],
        parameters["coupling_J"], parameters["equation_tolerance"], parameters["negative_control_minimum"])
    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(f"Changed fields at {path}")
        for key in actual:
            if key == "inputs":
                if not same_input_values(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 not isinstance(saved, list) or len(actual) != len(saved):
            raise AssertionError(f"Changed length at {path}")
        for index, (left, right) in enumerate(zip(actual, saved)):
            compare_saved(left, right, atol, rtol, f"{path}[{index}]")
    elif isinstance(actual, (bool, int)):
        if type(actual) is not type(saved) or actual != saved:
            raise AssertionError(f"Changed exact value at {path}: {actual} != {saved}")
    elif isinstance(actual, float):
        if type(saved) not in (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(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())
    saved = None
    if args.check:
        saved = json.loads((HERE / "results.json").read_text())
        if not same_input_values(saved["inputs"], parameters):
            raise AssertionError("Changed input parameters; regenerate results deliberately.")
    result = run(parameters)
    if args.check:
        tolerance = parameters["saved_comparison"]
        compare_saved(result, saved, tolerance["absolute_tolerance"], tolerance["relative_tolerance"])
        print("XXX algebra: YBE, RTT, transfer, Hamiltonian, B-state momentum, regular Bethe vectors, exceptional points, complete four-site sector 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 algebra and negative-control checks passed.")
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
