#!/usr/bin/env python3
"""Reproduce the physical four-site singular XXX pair and reject five sites.

Standalone NumPy plus standard-library exact rational arithmetic. --check never
writes files and locks inputs exactly; --write intentionally saves fresh results.
"""

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

import numpy as np

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


def swap(factors, first, second):
    dimension = 2**factors
    matrix = np.zeros((dimension, dimension), dtype=complex)
    for column in range(dimension):
        bits = list(f"{column:0{factors}b}")
        bits[first], bits[second] = bits[second], bits[first]
        matrix[int("".join(bits), 2), column] = 1
    return matrix


def shift(length):
    matrix = np.zeros((2**length, 2**length), dtype=complex)
    for column in range(2**length):
        bits = f"{column:0{length}b}"
        matrix[int(bits[-1] + bits[:-1], 2), column] = 1
    return matrix


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_operators(length):
    matrices = [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]
    return [[tensor_site(matrix, site, length) for site in range(length)] for matrix in matrices]


def pauli_hamiltonian(length, coupling):
    operators = spin_operators(length)
    result = np.zeros((2**length, 2**length), dtype=complex)
    for site in range(length):
        bond = np.eye(2**length, dtype=complex) / 4
        for component in operators:
            bond -= component[site] @ component[(site + 1) % length]
        result += coupling * bond
    return result


def sector_bond_matrix(length):
    """Independent exact M=2 bond action in units J; no Pauli or swap helper."""
    pairs = list(itertools.combinations(range(length), 2))
    indices = {pair: index for index, pair in enumerate(pairs)}
    matrix = [[Fraction(0) for _ in pairs] for _ in pairs]
    for column, pair in enumerate(pairs):
        occupied = set(pair)
        for left in range(length):
            right = (left + 1) % length
            if (left in occupied) != (right in occupied):
                new_pair = tuple(sorted(occupied ^ {left, right}))
                matrix[column][column] += Fraction(1, 2)
                matrix[indices[new_pair]][column] -= Fraction(1, 2)
    return pairs, matrix


def determinant_fraction(matrix):
    """Exact rational Gaussian elimination; preserves the supplied matrix."""
    a = [row[:] for row in matrix]
    value = Fraction(1)
    for index in range(len(a)):
        pivot = next((row for row in range(index, len(a)) if a[row][index]), None)
        if pivot is None:
            return Fraction(0)
        if pivot != index:
            a[index], a[pivot] = a[pivot], a[index]
            value = -value
        diagonal = a[index][index]
        value *= diagonal
        for row in range(index + 1, len(a)):
            factor = a[row][index] / diagonal
            for column in range(index + 1, len(a)):
                a[row][column] -= factor * a[index][column]
    return value


def monodromy(length, parameter):
    identity = np.eye(2**(length + 1), dtype=complex)
    result = identity.copy()
    for site in range(1, length + 1):
        result = ((parameter - 0.5j) * identity + 1j * swap(length + 1, 0, site)) @ result
    return result


def b_block(length, parameter):
    dimension = 2**length
    return monodromy(length, parameter)[:dimension, dimension:]


def transfer(length, parameter):
    dimension = 2**length
    array = monodromy(length, parameter).reshape(2, dimension, 2, dimension)
    return np.trace(array, axis1=0, axis2=2)


def polynomial_monodromy(length, local_coefficients):
    """Multiply matrix polynomials, coefficients in ascending epsilon degree."""
    identity = np.eye(2**(length + 1), dtype=complex)
    coefficients = [identity]
    for site in range(1, length + 1):
        local = local_coefficients(site, identity)
        new = [np.zeros_like(identity) for _ in range(len(coefficients) + len(local) - 1)]
        for left_degree, left in enumerate(local):
            if not np.any(left):
                continue
            for right_degree, right in enumerate(coefficients):
                new[left_degree + right_degree] += left @ right
        coefficients = new
    return coefficients


def bethe_polynomial(c):
    """Exact Gaussian-integer coefficient construction for the fixed N=4 pair."""
    length, dimension = 4, 16
    def first(site, identity):
        zero = np.zeros_like(identity)
        return [1j * swap(5, 0, site), identity, zero, zero, c * identity]
    def second(site, identity):
        return [1j * (swap(5, 0, site) - identity), identity]
    first_coeff = [a[:dimension, dimension:] for a in polynomial_monodromy(length, first)]
    second_coeff = [a[:dimension, dimension:] for a in polynomial_monodromy(length, second)]
    vacuum = np.zeros(dimension, dtype=complex)
    vacuum[0] = 1
    vector = [np.zeros(dimension, dtype=complex) for _ in range(len(first_coeff) + len(second_coeff) - 1)]
    for i, a in enumerate(first_coeff):
        for j, b in enumerate(second_coeff):
            vector[i + j] += a @ b @ vacuum
    array = np.array(vector)
    # For c=0 or 2i every elementary coefficient is Gaussian-integer. The
    # coefficient row-sum product bound is (1+1+|c|)^4 * (2+1)^4 <= 20736.
    # All intermediate exact integer real/imaginary operations fit below 2^53.
    bound = (2 + abs(c))**4 * 3**4
    if bound > 20736 or np.max(np.abs(array)) > bound:
        raise AssertionError("Coefficient exactness bound exceeded.")
    if not (np.array_equal(array.real, np.rint(array.real)) and
            np.array_equal(array.imag, np.rint(array.imag))):
        raise AssertionError("A supposedly exact Gaussian coefficient is fractional.")
    if np.any(array[:4]):
        raise AssertionError("Expected exact common epsilon^4 factor is missing.")
    while not np.any(array[-1]):
        array = array[:-1]
    return array[4:], int(bound), int(np.max(np.abs(array.real))), int(np.max(np.abs(array.imag)))


def horner(coefficients, epsilon):
    result = np.zeros_like(coefficients[0])
    for coefficient in reversed(coefficients):
        result = result * epsilon + coefficient
    return result


def root_polynomial_coefficients(c):
    """Cleared Bethe equations without subtracting nearly coincident root values."""
    poly = np.polynomial.polynomial
    l1 = np.array([0.5j, 1, 0, 0, c], dtype=complex)
    l2 = np.array([-0.5j, 1], dtype=complex)
    d = poly.polysub(l1, l2)
    one = poly.polysub(poly.polymul(poly.polypow(poly.polyadd(l1, [0.5j]), 4), poly.polyadd(d, [-1j])),
                       poly.polymul(poly.polypow(poly.polyadd(l1, [-0.5j]), 4), poly.polyadd(d, [1j])))
    two = poly.polysub(poly.polymul(poly.polypow(poly.polyadd(l2, [0.5j]), 4), poly.polyadd(-d, [-1j])),
                       poly.polymul(poly.polypow(poly.polyadd(l2, [-0.5j]), 4), poly.polyadd(-d, [1j])))
    if np.any(one[:4]) or np.any(two[:4]):
        raise AssertionError("Cleared equations lack their exact epsilon^4 factor.")
    return one[4:], two[4:]


def vector_metrics(vector, hamiltonian, target, coupling):
    norm = float(np.linalg.norm(vector))
    if norm == 0 or not math.isfinite(norm):
        return {"available": False, "raw_norm": norm if math.isfinite(norm) else None,
                "reason": "zero or non-finite vector; normalization is undefined"}
    state = vector / norm
    overlap = np.vdot(target, state)
    phase = np.conj(overlap) / abs(overlap) if abs(overlap) else 1
    # Direct phase alignment avoids cancellation in sqrt(2-2*abs(overlap)).
    distance = float(np.linalg.norm(phase * state - target))
    return {"available": True, "raw_norm": norm,
            "hamiltonian_residual_over_J": float(np.linalg.norm(hamiltonian @ state - coupling * state) / coupling),
            "phase_aligned_distance": distance,
            "energy_expectation_over_J": float(np.vdot(state, hamiltonian @ state).real / coupling)}


def pair_vector(values):
    result = np.zeros(16, dtype=complex)
    for pair, value in zip(itertools.combinations(range(4), 2), values):
        result[sum(2**(3 - site) for site in pair)] = value
    return result


def transfer_polynomial_check(target):
    dimension = 16
    def local(site, identity):
        return [1j * swap(5, 0, site) - 0.5j * identity, identity]
    coefficients = polynomial_monodromy(4, local)
    expected = [-3 / 8, 0, 3, 0, 2]
    errors = []
    for matrix, scalar in zip(coefficients, expected):
        array = matrix.reshape(2, dimension, 2, dimension)
        tau = np.trace(array, axis1=0, axis2=2)
        errors.append(float(np.max(np.abs(tau @ target - scalar * target))))
    return errors


def run(parameters):
    if parameters["sites"] != 4 or parameters["odd_chain_sites"] != 5:
        raise ValueError("This exact reproduction is specifically N=4 with N=5 rejection.")
    coupling = parameters["coupling_J"]
    if not math.isfinite(coupling) or coupling <= 0 or any(not math.isfinite(e) or e <= 0 for e in parameters["regulators"]):
        raise ValueError("Use positive J and positive regulators.")
    if parameters["regularizations"] != [{"name": "correct", "c": [0, 2]}, {"name": "naive", "c": [0, 0]}]:
        raise ValueError("The coefficient-exactness claim covers c=2i and c=0 only.")
    target = pair_vector([0.5, 0, -0.5, -0.5, 0, 0.5])
    hamiltonian = pauli_hamiltonian(4, coupling)
    translation = shift(4)
    components = [sum(operators) for operators in spin_operators(4)]
    spin_squared = sum(component @ component for component in components)
    pairs, exact_sector = sector_bond_matrix(4)
    indices = [sum(2**(3 - site) for site in pair) for pair in pairs]
    sector = np.array(exact_sector, dtype=float)
    exact = {
        "target_norm_squared": float(np.vdot(target, target).real),
        "target_hamiltonian_residual_over_J": float(np.linalg.norm(hamiltonian @ target - coupling * target) / coupling),
        "target_translation_residual": float(np.linalg.norm(translation @ target + target)),
        "target_total_spin_squared_residual": float(np.linalg.norm(spin_squared @ target)),
        "independent_sector_max_entry_error_over_J": float(np.max(np.abs(hamiltonian[np.ix_(indices, indices)] / coupling - sector))),
        "transfer_polynomial_max_entry_errors": transfer_polynomial_check(target),
        "zero_vector_at_unregularized_pair_norm": float(np.linalg.norm(b_block(4, 0.5j) @ b_block(4, -0.5j) @ np.eye(16)[:, 0]))}
    generic = []
    for pair in parameters["generic_transfer_parameters"]:
        lam = complex(*pair)
        eigenvalue = 2 * lam**4 + 3 * lam**2 - 3 / 8
        tau = transfer(4, lam)
        generic.append({"lambda": pair,
                        "relative_transfer_vector_residual": float(np.linalg.norm(tau @ target - eigenvalue * target) /
                            max(1, np.linalg.norm(tau @ target), abs(eigenvalue)))})
    records = []
    vacuum = np.eye(16)[:, 0]
    for prescription in parameters["regularizations"]:
        c = complex(*prescription["c"])
        coefficients, bound, max_real, max_imag = bethe_polynomial(c)
        expected_limit = pair_vector([2, 0, 1j * c, -2, 0, 2])
        if not np.array_equal(coefficients[0], expected_limit):
            raise AssertionError("Incorrect limiting Bethe vector.")
        mismatch = (hamiltonian / coupling - np.eye(16)) @ expected_limit
        expected_mismatch = pair_vector([0, -(1 + 0.5j * c), 0, 0, -(1 + 0.5j * c), 0])
        if not np.allclose(mismatch, expected_mismatch, atol=parameters["equation_tolerance"], rtol=0):
            raise AssertionError("Incorrect analytic limiting-state defect.")
        roots_one, roots_two = root_polynomial_coefficients(c)
        samples = []
        for epsilon in parameters["regulators"]:
            stable = horner(coefficients, epsilon)
            l1 = 0.5j + epsilon + c * epsilon**4
            l2 = -0.5j + epsilon
            direct = b_block(4, l1) @ b_block(4, l2) @ vacuum
            if not np.all(np.isfinite(direct)) or not math.isfinite(float(np.linalg.norm(direct))):
                raise AssertionError("Direct evaluation produced non-finite data.")
            direct_metrics = vector_metrics(direct, hamiltonian, target, coupling)
            direct_agreement = vector_metrics(direct, hamiltonian, stable / np.linalg.norm(stable), coupling)
            direct_metrics["distance_from_stable_state"] = direct_agreement.get("phase_aligned_distance")
            if epsilon >= 0.01 and (not direct_agreement["available"] or
                    direct_agreement["phase_aligned_distance"] > 1e-8):
                raise AssertionError("Direct and polynomial evaluation disagree at a resolved regulator.")
            root_scaled = max(abs(np.polynomial.polynomial.polyval(epsilon, roots_one)),
                              abs(np.polynomial.polynomial.polyval(epsilon, roots_two)))
            samples.append({"epsilon": epsilon,
                            "stable_scaled_vector": vector_metrics(stable, hamiltonian, target, coupling),
                            "direct_unscaled_vector": direct_metrics,
                            "cleared_bethe_residual": float(epsilon**4 * root_scaled),
                            "cleared_bethe_residual_over_epsilon4": float(root_scaled)})
        records.append({"name": prescription["name"], "c": prescription["c"],
                        "gaussian_integer_coefficient_bound": bound,
                        "max_absolute_real_coefficient": max_real,
                        "max_absolute_imaginary_coefficient": max_imag,
                        "exact_zero_coefficients_below_degree": 4,
                        "factored_vector_max_degree": len(coefficients) - 1,
                        "limiting_vector_pairs_real_imag": [[float(value.real), float(value.imag)] for value in expected_limit[indices]],
                        "limit_metrics": vector_metrics(expected_limit, hamiltonian, target, coupling),
                        "samples": samples})
    odd_pairs, odd_exact = sector_bond_matrix(5)
    shifted = [[value - (1 if i == j else 0) for j, value in enumerate(row)] for i, row in enumerate(odd_exact)]
    determinant = determinant_fraction(shifted)
    odd_matrix = np.array(odd_exact, dtype=float)
    odd_indices = [sum(2**(4 - site) for site in pair) for pair in odd_pairs]
    odd_pauli = pauli_hamiltonian(5, coupling)[np.ix_(odd_indices, odd_indices)] / coupling
    eigenvalues = np.linalg.eigvalsh(odd_matrix)
    regular_point = 0.5j
    candidate_value = ((regular_point + 0.5j)**4 * (regular_point - 1.5j) +
                       (regular_point - 0.5j)**4 * (regular_point + 1.5j))
    candidate_shift = candidate_value / 1j**5
    if candidate_shift != -1:
        raise AssertionError("The odd-chain Baxter candidate lost its forbidden shift value.")
    odd = {"sites": 5, "sector_dimension": len(odd_pairs),
           "determinant_H_over_J_minus_identity": str(determinant),
           "independent_sector_max_entry_error_over_J": float(np.max(np.abs(odd_pauli - odd_matrix))),
           "energies_over_J": eigenvalues.tolist(),
           "nearest_energy_distance_to_one": float(np.min(np.abs(eigenvalues - 1))),
           "translation_fifth_power_max_entry_error": float(np.max(np.abs(np.linalg.matrix_power(shift(5), 5) - np.eye(32)))),
           "candidate_transfer_value_at_i_over_2_real_imag": [candidate_value.real, candidate_value.imag],
           "candidate_translation_eigenvalue": int(candidate_shift.real),
           "candidate_fifth_power": int((candidate_shift**5).real)}
    if determinant != Fraction(-1, 256):
        raise AssertionError("Five-site exact rejection failed.")
    numerical_checks = [value for key, value in exact.items() if key != "target_norm_squared" and not isinstance(value, list)]
    numerical_checks += exact["transfer_polynomial_max_entry_errors"]
    numerical_checks += [item["relative_transfer_vector_residual"] for item in generic]
    numerical_checks += [odd["independent_sector_max_entry_error_over_J"], odd["translation_fifth_power_max_entry_error"]]
    if max(numerical_checks) > parameters["equation_tolerance"]:
        raise AssertionError("An exact-state or independent-Hamiltonian check failed.")
    correct, naive = records
    if abs(naive["limit_metrics"]["hamiltonian_residual_over_J"] - 1 / math.sqrt(6)) > 2e-12:
        raise AssertionError("The naive limiting-state counterexample failed.")
    if abs(naive["limit_metrics"]["energy_expectation_over_J"] - 1) > 2e-12:
        raise AssertionError("The misleading naive Rayleigh energy is not reproduced.")
    if correct["samples"][-1]["stable_scaled_vector"]["hamiltonian_residual_over_J"] > 2 * parameters["regulators"][-1]:
        raise AssertionError("Stable regularization did not approach the eigenstate.")
    return {"schema_version": 1,
            "environment": {"python": platform.python_version(), "numpy": np.__version__,
                            "arithmetic": "binary64 complex matrices; exact Gaussian coefficients within stated bound; Fraction determinant"},
            "inputs": parameters, "exact_four_site_checks": exact,
            "generic_transfer_checks": generic, "regularizations": records,
            "odd_chain_rejection": odd}


def compare_saved(actual, saved, atol, rtol, path="results"):
    if isinstance(actual, dict):
        if actual.keys() != saved.keys():
            raise AssertionError(f"Changed fields at {path}")
        for key in actual:
            if key == "inputs":
                if actual[key] != saved[key]:
                    raise AssertionError(f"Changed input parameters at {path}.inputs")
            # The raw direct evaluation is deliberately ill-conditioned; its
            # cancellation and even zero-vector availability depend on platform.
            # run() checks agreement at resolved epsilon, while stable results
            # and every exact identity are compared against the saved data.
            elif key not in ("environment", "direct_unscaled_vector"):
                compare_saved(actual[key], saved[key], atol, rtol, f"{path}.{key}")
    elif isinstance(actual, list):
        if len(actual) != len(saved):
            raise AssertionError(f"Changed length at {path}")
        for index, (left, right) in enumerate(zip(actual, saved)):
            compare_saved(left, right, atol, rtol, f"{path}[{index}]")
    elif isinstance(actual, bool) or actual is None:
        if actual != saved:
            raise AssertionError(f"Changed availability at {path}")
    elif isinstance(actual, (int, float)):
        if not math.isfinite(actual) 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())
    result = run(parameters)
    if args.check:
        saved = json.loads((HERE / "results.json").read_text())
        tolerance = parameters["saved_comparison"]
        compare_saved(result, saved, tolerance["absolute_tolerance"], tolerance["relative_tolerance"])
        print("Singular XXX: exact limits, stable regularization, independent Hamiltonians, odd-chain rejection 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 singular-state checks passed.")
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
