#!/usr/bin/env python3
"""Independent finite-XXX Hamiltonians and selected two-magnon Bethe states.

H = J sum_i (1/4 - S_i.S_(i+1)), S = sigma/2, J > 0, periodic N >= 3.
Without an option, print recomputed JSON. --check is nonmutating. --write explicitly
regenerates results.json. All companion files are located beside this script.
This is a finite-size verification of selected regular states, not completeness.
"""

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

import numpy as np


HERE = Path(__file__).resolve().parent
IDENTITY_2 = np.eye(2, dtype=np.complex128)
SPINS = (
    np.array([[0, 1], [1, 0]], dtype=np.complex128) / 2,
    np.array([[0, -1j], [1j, 0]], dtype=np.complex128) / 2,
    np.array([[1, 0], [0, -1]], dtype=np.complex128) / 2,
)
LOWERING = np.array([[0, 0], [1, 0]], dtype=np.complex128)


def maximum_absolute(array):
    return float(np.max(np.abs(array)))


def tensor_operator(n, factors):
    """Site 0 is the leftmost tensor factor: up=(1,0), down=(0,1)."""
    result = np.ones((1, 1), dtype=np.complex128)
    for site in range(n):
        result = np.kron(result, factors.get(site, IDENTITY_2))
    return result


def tensor_hamiltonian(n, coupling):
    """Build the full 2**N matrix from Pauli tensor products, without swaps."""
    result = (coupling * n / 4) * np.eye(2**n, dtype=np.complex128)
    for left in range(n):
        right = (left + 1) % n
        for spin in SPINS:
            result -= coupling * tensor_operator(n, {left: spin, right: spin})
    return result


def sector_hamiltonian(n, magnons, coupling, periodic=True):
    """Build J/2 sum(1-P) directly in a basis of down-spin positions.

    This code does not call the tensor construction or share its bond iterator.
    Each basis tuple is increasing; matrix columns are input basis states.
    """
    basis = list(combinations(range(n), magnons))
    lookup = {positions: index for index, positions in enumerate(basis)}
    result = np.zeros((len(basis), len(basis)), dtype=np.float64)
    bonds = [(site, site + 1) for site in range(n - 1)]
    if periodic:
        bonds.append((n - 1, 0))
    for column, positions in enumerate(basis):
        occupied = set(positions)
        for left, right in bonds:
            if (left in occupied) != (right in occupied):
                exchanged = tuple(sorted(occupied.symmetric_difference({left, right})))
                result[column, column] += coupling / 2
                result[lookup[exchanged], column] -= coupling / 2
    return basis, result


def tensor_indices(n, basis):
    return np.array([sum(1 << (n - 1 - site) for site in positions)
                     for positions in basis], dtype=np.int64)


def scattering_ratio(z1, z2, minimum_denominator=1e-12):
    denominator = 1 + z1 * z2 - 2 * z1
    if abs(denominator) <= minimum_denominator:
        raise ValueError("singular scattering formula; a separate limit is required")
    return -(1 + z1 * z2 - 2 * z2) / denominator


def coefficient(x, y, z1, z2, scattering):
    return z1**x * z2**y + scattering * z2**x * z1**y


def bethe_vector(basis, z1, z2, scattering, minimum_norm):
    raw = np.array([coefficient(x, y, z1, z2, scattering) for x, y in basis],
                   dtype=np.complex128)
    norm = float(np.linalg.norm(raw))
    if not math.isfinite(norm) or norm <= minimum_norm:
        raise ValueError("Bethe coefficients do not define a nonzero normalizable vector")
    return raw / norm, norm


def periodicity_residual(n, z1, z2, scattering):
    return float(max(abs(z1**n - 1 / scattering), abs(z2**n - scattering)))


def check_chain(n, coupling):
    full = tensor_hamiltonian(n, coupling)
    total_sz = sum((tensor_operator(n, {site: SPINS[2]}) for site in range(n)),
                   np.zeros_like(full))
    sz_diagonal = np.diag(total_sz)
    sz_commutator = full * (sz_diagonal[None, :] - sz_diagonal[:, None])
    full_vacuum = np.zeros(2**n, dtype=np.complex128)
    full_vacuum[0] = 1
    record = {
        "N": n,
        "full_dimension": 2**n,
        "hermiticity_max_abs_over_J": maximum_absolute(full - full.conj().T) / coupling,
        "Sz_commutator_max_abs_over_J": maximum_absolute(sz_commutator) / coupling,
        "vacuum_residual_l2_over_J": float(np.linalg.norm(full @ full_vacuum)) / coupling,
    }
    blocks = {}
    for magnons in (1, 2):
        basis, sector = sector_hamiltonian(n, magnons, coupling)
        indices = tensor_indices(n, basis)
        tensor_block = full[np.ix_(indices, indices)]
        blocks[magnons] = (basis, sector, indices)
        record[f"M{magnons}_dimension"] = len(basis)
        record[f"M{magnons}_block_difference_max_abs_over_J"] = (
            maximum_absolute(tensor_block - sector) / coupling)
        record[f"M{magnons}_spectrum_over_J"] = (np.linalg.eigvalsh(sector) / coupling).tolist()
    one_basis, one_sector, _ = blocks[1]
    wave_numbers = 2 * np.pi * np.arange(n) / n
    expected = np.sort(1 - np.cos(wave_numbers))
    record["one_magnon_spectrum_error_max_abs_over_J"] = maximum_absolute(
        np.array(record["M1_spectrum_over_J"]) - expected)
    one_residuals = []
    for k in wave_numbers:
        state = np.array([np.exp(1j * k * positions[0]) for positions in one_basis]) / np.sqrt(n)
        one_residuals.append(float(np.linalg.norm(
            one_sector @ state - coupling * (1 - np.cos(k)) * state)) / coupling)
    record["one_magnon_state_residual_max_l2_over_J"] = max(one_residuals)

    # Construct a descendant independently by lowering the tensor vacuum twice.
    total_lowering = sum((tensor_operator(n, {site: LOWERING}) for site in range(n)),
                         np.zeros_like(full))
    descendant = total_lowering @ (total_lowering @ full_vacuum)
    descendant /= np.linalg.norm(descendant)
    two_basis, _, two_indices = blocks[2]
    expected_descendant = np.ones(len(two_basis)) / np.sqrt(len(two_basis))
    record["descendant_residual_l2_over_J"] = float(np.linalg.norm(full @ descendant)) / coupling
    record["descendant_uniform_state_error_l2"] = float(np.linalg.norm(
        descendant[two_indices] - expected_descendant))
    return record, full, blocks[2]


def check_two_magnon(n, mode, coupling, full, block, minimum_norm, equation_tolerance):
    basis, sector, indices = block
    k = 2 * np.pi * mode / (n - 1)
    z1, z2 = np.exp(1j * k), np.exp(-1j * k)
    scattering = scattering_ratio(z1, z2)
    state, raw_norm = bethe_vector(basis, z1, z2, scattering, minimum_norm)
    energy = 2 * coupling * (1 - np.cos(k))
    full_state = np.zeros(2**n, dtype=np.complex128)
    full_state[indices] = state
    eigenvalues, eigenvectors = np.linalg.eigh(sector)
    nearest = int(np.argmin(np.abs(eigenvalues - energy)))
    matched = np.abs(eigenvalues - energy) <= coupling * equation_tolerance
    projection = eigenvectors[:, matched] @ (eigenvectors[:, matched].conj().T @ state)
    _, open_sector = sector_hamiltonian(n, 2, coupling, periodic=False)
    wrong_scattering = 1 / scattering
    wrong_state, _ = bethe_vector(basis, z1, z2, wrong_scattering, minimum_norm)
    seam_residual = max(abs(coefficient(x, y, z1, z2, scattering)
                            - coefficient(y, x + n, z1, z2, scattering)) / raw_norm
                        for x, y in basis)
    record = {
        "N": n,
        "mode": mode,
        "k1": float(k),
        "k2": float(-k),
        "S12_real": float(scattering.real),
        "S12_imag": float(scattering.imag),
        "energy_over_J": float(energy / coupling),
        "raw_norm_squared": raw_norm**2,
        "raw_norm_squared_error": float(abs(raw_norm**2 - n * (n - 1))),
        "normalization_error": float(abs(np.vdot(state, state) - 1)),
        "Bethe_equation_residual_max_abs": periodicity_residual(n, z1, z2, scattering),
        "seam_residual_max_abs": float(seam_residual),
        "contact_equation_residual_abs": float(abs(
            (1 + z1 * z2 - 2 * z2) + scattering * (1 + z1 * z2 - 2 * z1))),
        "sector_eigenstate_residual_l2_over_J": float(np.linalg.norm(sector @ state - energy * state)) / coupling,
        "tensor_eigenstate_residual_l2_over_J": float(np.linalg.norm(full @ full_state - energy * full_state)) / coupling,
        "nearest_eigenvalue_over_J": float(eigenvalues[nearest] / coupling),
        "nearest_energy_error_over_J": float(abs(eigenvalues[nearest] - energy)) / coupling,
        "matched_eigenspace_dimension": int(np.count_nonzero(matched)),
        "eigenspace_projection_error_l2": float(np.linalg.norm(state - projection)),
        "wrong_ratio_residual_l2_over_J": float(np.linalg.norm(sector @ wrong_state - energy * wrong_state)) / coupling,
        "wrong_ratio_Bethe_residual_max_abs": periodicity_residual(n, z1, z2, wrong_scattering),
        "missing_closing_bond_residual_l2_over_J": float(np.linalg.norm(open_sector @ state - energy * state)) / coupling,
    }
    if n == 6:
        record["normalized_coefficients"] = [
            {"positions": list(positions), "real": float(value.real), "imag": float(value.imag)}
            for positions, value in zip(basis, state)
        ]
    return record


def regularity_examples(minimum_norm):
    try:
        scattering_ratio(1 + 0j, 1 + 0j)
    except ValueError:
        singular_rejected = True
    else:
        singular_rejected = False
    # Exact z1=z2=-1 avoids hiding this zero vector behind trigonometric rounding.
    z1 = z2 = -1 + 0j
    scattering = scattering_ratio(z1, z2)
    basis = list(combinations(range(5), 2))
    raw = np.array([coefficient(x, y, z1, z2, scattering) for x, y in basis])
    try:
        bethe_vector(basis, z1, z2, scattering, minimum_norm)
    except ValueError:
        zero_rejected = True
    else:
        zero_rejected = False
    return {
        "z1_z2_equal_one_singular_ratio_rejected": singular_rejected,
        "coincident_minus_one_N5": {
            "S12_real": float(scattering.real),
            "Bethe_equation_residual_max_abs": periodicity_residual(5, z1, z2, scattering),
            "raw_norm_squared": float(np.vdot(raw, raw).real),
            "normalization_rejected": zero_rejected,
        },
    }


def validate_inputs(inputs):
    if inputs.get("schema_version") != 1:
        raise ValueError("unsupported input schema")
    coupling = inputs["J"]
    if not isinstance(coupling, (int, float)) or not math.isfinite(coupling) or coupling <= 0:
        raise ValueError("J must be finite and positive")
    sizes = inputs["chain_lengths"]
    if not sizes or len(set(sizes)) != len(sizes):
        raise ValueError("chain lengths must be nonempty and distinct")
    if set(inputs["selected_modes"]) != {str(n) for n in sizes}:
        raise ValueError("selected_modes must have exactly one entry per chain length")
    for n in sizes:
        if not isinstance(n, int) or not 3 <= n <= 10:
            raise ValueError("use integer 3 <= N <= 10; full tensor matrices scale exponentially")
        modes = inputs["selected_modes"][str(n)]
        if len(set(modes)) != len(modes):
            raise ValueError("duplicate mode in selected branch")
        if any(not isinstance(mode, int) or not 1 <= mode < (n - 1) / 2 for mode in modes):
            raise ValueError("selected modes must obey 1 <= mode < (N-1)/2")
    if not any(inputs["selected_modes"].values()):
        raise ValueError("select at least one regular two-magnon state")
    for value in inputs["tolerances"].values():
        if not isinstance(value, (int, float)) or not math.isfinite(value) or value <= 0:
            raise ValueError("tolerances must be finite and positive")


def compute(inputs):
    validate_inputs(inputs)
    coupling = inputs["J"]
    tolerances = inputs["tolerances"]
    chains, cases = [], []
    for n in inputs["chain_lengths"]:
        chain, full, block = check_chain(n, coupling)
        chains.append(chain)
        for mode in inputs["selected_modes"][str(n)]:
            cases.append(check_two_magnon(n, mode, coupling, full, block,
                                         tolerances["minimum_raw_norm"], tolerances["equation_atol"]))
    result = {
        "schema_version": 1,
        "inputs": inputs,
        "environment": {"python": platform.python_version(), "numpy": np.__version__,
                        "arithmetic": "IEEE 754 binary64; complex128 states"},
        "convention": "periodic spin-1/2 XXX, H=J sum(1/4-S.S), S=sigma/2, sites 0,...,N-1",
        "chains": chains,
        "two_magnon_cases": cases,
        "regularity_examples": regularity_examples(tolerances["minimum_raw_norm"]),
        "summary": {
            "chain_count": len(chains),
            "selected_state_count": len(cases),
            "maximum_tensor_sector_block_difference_over_J": max(
                max(row["M1_block_difference_max_abs_over_J"], row["M2_block_difference_max_abs_over_J"])
                for row in chains),
            "maximum_one_magnon_spectrum_error_over_J": max(
                row["one_magnon_spectrum_error_max_abs_over_J"] for row in chains),
            "maximum_Bethe_equation_residual": max(row["Bethe_equation_residual_max_abs"] for row in cases),
            "maximum_two_magnon_eigenstate_residual_over_J": max(
                max(row["sector_eigenstate_residual_l2_over_J"], row["tensor_eigenstate_residual_l2_over_J"])
                for row in cases),
            "maximum_nearest_energy_error_over_J": max(row["nearest_energy_error_over_J"] for row in cases),
            "minimum_wrong_ratio_residual_over_J": min(row["wrong_ratio_residual_l2_over_J"] for row in cases),
            "minimum_missing_bond_residual_over_J": min(row["missing_closing_bond_residual_l2_over_J"] for row in cases),
        },
    }
    validate_result(result)
    return result


def validate_result(result):
    tolerance = result["inputs"]["tolerances"]["equation_atol"]
    control_floor = result["inputs"]["tolerances"]["minimum_control_residual"]
    chain_checks = ("hermiticity_max_abs_over_J", "Sz_commutator_max_abs_over_J",
                    "vacuum_residual_l2_over_J", "M1_block_difference_max_abs_over_J",
                    "M2_block_difference_max_abs_over_J", "one_magnon_spectrum_error_max_abs_over_J",
                    "one_magnon_state_residual_max_l2_over_J", "descendant_residual_l2_over_J",
                    "descendant_uniform_state_error_l2")
    state_checks = ("normalization_error", "raw_norm_squared_error", "Bethe_equation_residual_max_abs", "seam_residual_max_abs",
                    "contact_equation_residual_abs", "sector_eigenstate_residual_l2_over_J",
                    "tensor_eigenstate_residual_l2_over_J", "nearest_energy_error_over_J",
                    "eigenspace_projection_error_l2")
    for rows, keys in ((result["chains"], chain_checks), (result["two_magnon_cases"], state_checks)):
        for row in rows:
            for key in keys:
                if not math.isfinite(row[key]) or not 0 <= row[key] <= tolerance:
                    raise AssertionError(f"N={row['N']} {key}={row[key]} exceeds {tolerance}")
    for row in result["two_magnon_cases"]:
        for key in ("wrong_ratio_residual_l2_over_J", "missing_closing_bond_residual_l2_over_J"):
            if not math.isfinite(row[key]) or row[key] <= control_floor:
                raise AssertionError(f"negative control did not fail decisively: {key}={row[key]}")
    examples = result["regularity_examples"]
    zero = examples["coincident_minus_one_N5"]
    if not (examples["z1_z2_equal_one_singular_ratio_rejected"] and zero["normalization_rejected"]
            and zero["raw_norm_squared"] == 0 and zero["Bethe_equation_residual_max_abs"] == 0):
        raise AssertionError("regularity guards failed")


def compare_saved(actual, saved, atol, rtol, path="results"):
    """Compare scientific data with tolerances; environment versions are descriptive."""
    if path == "results.environment":
        return
    if isinstance(actual, dict):
        if not isinstance(saved, dict) or actual.keys() != saved.keys():
            raise AssertionError(f"{path}: dictionary fields differ")
        for key in actual:
            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"{path}: list lengths differ")
        for index, (left, right) in enumerate(zip(actual, saved)):
            compare_saved(left, right, atol, rtol, f"{path}[{index}]")
    elif isinstance(actual, bool) or isinstance(actual, str) or isinstance(actual, int):
        if actual != saved:
            raise AssertionError(f"{path}: {actual!r} != {saved!r}")
    elif isinstance(actual, float):
        if not isinstance(saved, (int, float)) or not math.isfinite(saved) or not math.isclose(
                actual, saved, abs_tol=atol, rel_tol=rtol):
            raise AssertionError(f"{path}: {actual!r} != {saved!r}")
    else:
        raise TypeError(f"unsupported result type at {path}")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    group = parser.add_mutually_exclusive_group()
    group.add_argument("--check", action="store_true", help="recompute and compare saved data without writing")
    group.add_argument("--write", action="store_true", help="explicitly regenerate results.json")
    args = parser.parse_args()
    inputs = json.loads((HERE / "inputs.json").read_text())
    result = compute(inputs)
    if args.check:
        saved = json.loads((HERE / "results.json").read_text())
        if saved.get("inputs") != inputs:
            raise AssertionError("saved inputs differ; inspect the changed experiment before regenerating")
        tolerances = inputs["tolerances"]
        compare_saved(result, saved, tolerances["saved_data_atol"], tolerances["saved_data_rtol"])
        print(json.dumps(result["summary"], indent=2, allow_nan=False))
        print("PASS: scientific checks and saved-data comparison; no files written.")
    elif args.write:
        (HERE / "results.json").write_text(json.dumps(result, indent=2, allow_nan=False) + "\n")
        print(json.dumps(result["summary"], indent=2, allow_nan=False))
        print(f"Wrote {HERE / 'results.json'}")
    else:
        print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == "__main__":
    main()
