#!/usr/bin/env python3

"""
PURE INDEPENDENT EXACT LEECH VERIFICATION
=========================================

No floating point.
No NumPy.
No eigensolver.
No published target eigenvalues.
No published target gap.
No hardcoded tensor coefficients.

Everything reported below is derived from:

    1. the binary [23,12,7] Golay code,
    2. its [24,12,8] extension,
    3. the three Conway families of minimal Leech vectors,
    4. exact integer shell sums.

The only exact arithmetic is Python integer arithmetic and Fraction.

The calculation constructs

    T_ab(k) = sum_x x_a x_b (k.x)^10

directly from all 196,560 minimal-shell vectors.

For

    k = (1,2,0,...,0)

and transverse polarizations

    p_in  = (-2,1,0,...,0)
    p_out = e_3,

the physical unit propagation direction is

    u = k / sqrt(5).

The physical transverse contractions are therefore

    lambda_in  = Q(k,p_in)  / 5^6
    lambda_out = Q(k,p_out) / 5^5.

The code derives the transverse splitting directly.

There is deliberately NO final comparison with a pre-entered
"expected" eigenvalue or gap.

If this program prints a nonzero exact gap and all structural
identities pass, that result has been obtained from the shell
construction and tensor contraction itself.
"""

from fractions import Fraction
from itertools import combinations


DIM = 24


# ======================================================================
# Binary [23,12,7] Golay code
# ======================================================================

def golay23_generator():
    """
    Generator polynomial

        g(x) = x^11 + x^9 + x^7 + x^6 + x^5 + x + 1

    for the cyclic binary [23,12,7] Golay code.

    Bit j represents x^j.
    """
    return (
        (1 << 11)
        | (1 << 9)
        | (1 << 7)
        | (1 << 6)
        | (1 << 5)
        | (1 << 1)
        | 1
    )


def generate_golay23_codewords():
    """
    Generate all 2^12 codewords of the cyclic [23,12,7] Golay code.
    """
    g = golay23_generator()

    codewords = set()

    for message in range(1 << 12):
        word = 0

        # Multiply message polynomial by g(x) over GF(2).
        for j in range(12):
            if (message >> j) & 1:
                word ^= g << j

        assert word.bit_length() <= 23
        codewords.add(word)

    assert len(codewords) == 2 ** 12

    return codewords


def extend_golay23(codewords23):
    """
    Add the parity coordinate to obtain the extended binary
    [24,12,8] Golay code.
    """
    codewords24 = set()

    for word in codewords23:
        parity = word.bit_count() & 1
        extended = word | (parity << 23)
        codewords24.add(extended)

    assert len(codewords24) == 2 ** 12

    # The extended Golay code is doubly even.
    for word in codewords24:
        assert word.bit_count() % 4 == 0

    # Minimum nonzero weight is 8.
    nonzero_weights = [
        word.bit_count()
        for word in codewords24
        if word != 0
    ]

    assert min(nonzero_weights) == 8

    return codewords24


def construct_golay24():
    """
    Construct the extended [24,12,8] Golay code.
    """
    return extend_golay23(
        generate_golay23_codewords()
    )


# ======================================================================
# Conway Family A
# ======================================================================

def family_A():
    """
    Family A:

        (±4, ±4, 0^22)

    Count:

        4 * C(24,2) = 1104.
    """
    vectors = []

    for i, j in combinations(range(DIM), 2):

        for si in (-4, 4):
            for sj in (-4, 4):

                x = [0] * DIM
                x[i] = si
                x[j] = sj

                vectors.append(tuple(x))

    assert len(vectors) == 1104

    return vectors


# ======================================================================
# Conway Family B
# ======================================================================

def family_B(codewords24):
    """
    Family B:

        (±2^8, 0^16)

    supported on an octad of the extended Golay code,
    with an even number of negative signs.

    There are

        759 * 2^7 = 97152

    vectors.
    """
    octads = [
        word
        for word in codewords24
        if word.bit_count() == 8
    ]

    assert len(octads) == 759

    vectors = []

    for octad in octads:

        support = [
            i
            for i in range(DIM)
            if (octad >> i) & 1
        ]

        assert len(support) == 8

        # Seven signs are independent.
        # The eighth enforces even sign parity.
        for mask in range(1 << 7):

            signs = []
            negative_count = 0

            for j in range(7):

                if (mask >> j) & 1:
                    signs.append(-2)
                    negative_count += 1
                else:
                    signs.append(2)

            if negative_count % 2 == 0:
                signs.append(2)
            else:
                signs.append(-2)

            x = [0] * DIM

            for coordinate, sign in zip(support, signs):
                x[coordinate] = sign

            vectors.append(tuple(x))

    assert len(vectors) == 97152

    return vectors


# ======================================================================
# Conway Family C
# ======================================================================

def family_C(codewords24):
    """
    Family C:

        (∓3, ±1^23)

    determined by an extended Golay codeword and a distinguished
    coordinate.

    For codeword c,

        y_j = 1 - 2*c_j.

    At one distinguished coordinate the ±1 representative is
    replaced by the corresponding ±3 representative.

    Thus

        x_j == 1 - 2*c_j (mod 4)

    and

        sum(x_j) == 4 (mod 8).

    Count:

        4096 * 24 = 98304.
    """
    vectors = []

    for word in codewords24:

        signs = [
            1 if ((word >> j) & 1) == 0 else -1
            for j in range(DIM)
        ]

        for distinguished in range(DIM):

            x = signs.copy()

            if ((word >> distinguished) & 1) == 0:
                x[distinguished] = -3
            else:
                x[distinguished] = 3

            # Exact congruence check.
            for j in range(DIM):

                c = (word >> j) & 1

                assert (
                    (x[j] - (1 - 2 * c)) % 4
                    == 0
                )

            assert sum(x) % 8 == 4

            vectors.append(tuple(x))

    assert len(vectors) == 98304

    return vectors


# ======================================================================
# Complete minimal Leech shell
# ======================================================================

def construct_minimal_shell():
    """
    Construct all minimal Leech vectors in integer coordinates.

    These integer coordinates satisfy

        |x|^2 = 32.

    Standard Leech normalization is

        R = x / sqrt(8),

    giving

        |R|^2 = 4.
    """
    code = construct_golay24()

    A = family_A()
    B = family_B(code)
    C = family_C(code)

    assert len(A) == 1104
    assert len(B) == 97152
    assert len(C) == 98304

    shell = A + B + C

    assert len(shell) == 196560

    # Check that no vectors were duplicated.
    assert len(set(shell)) == 196560

    # Check minimal-shell norm.
    for x in shell:
        assert sum(
            xi * xi
            for xi in x
        ) == 32

    return code, shell


# ======================================================================
# Exact moments
# ======================================================================

def moment_M10(shell, k):
    """
    Compute

        M10(k) = sum_x (k.x)^10

    exactly.
    """
    total = 0

    for x in shell:

        kx = sum(
            ki * xi
            for ki, xi in zip(k, x)
        )

        total += kx ** 10

    return total


def moment_M12(shell, k):
    """
    Compute

        M12(k) = sum_x (k.x)^12

    exactly.
    """
    total = 0

    for x in shell:

        kx = sum(
            ki * xi
            for ki, xi in zip(k, x)
        )

        total += kx ** 12

    return total


# ======================================================================
# Exact tensor quadratic form
# ======================================================================

def tensor_quadratic_form(shell, k, p):
    """
    Compute

        Q(k,p)
          = sum_x (k.x)^10 (p.x)^2.

    Equivalently,

        Q(k,p) = p^T T(k) p,

    where

        T_ab(k)
          = sum_x x_a x_b (k.x)^10.
    """
    total = 0

    for x in shell:

        kx = sum(
            ki * xi
            for ki, xi in zip(k, x)
        )

        px = sum(
            pi * xi
            for pi, xi in zip(p, x)
        )

        total += (kx ** 10) * (px ** 2)

    return total


# ======================================================================
# Exact full tensor construction
# ======================================================================

def construct_tensor(shell, k):
    """
    Construct the complete symmetric 24x24 integer tensor

        T_ab(k)
          = sum_x x_a x_b (k.x)^10.

    This is calculated directly from the shell.

    No target tensor entries are supplied.
    """
    T = [
        [0 for _ in range(DIM)]
        for _ in range(DIM)
    ]

    for x in shell:

        kx = sum(
            ki * xi
            for ki, xi in zip(k, x)
        )

        weight = kx ** 10

        for a in range(DIM):
            xa = x[a]

            if xa == 0:
                continue

            for b in range(a, DIM):
                xb = x[b]

                if xb == 0:
                    continue

                value = weight * xa * xb

                T[a][b] += value

                if a != b:
                    T[b][a] += value

    return T


# ======================================================================
# Exact quadratic form from full tensor
# ======================================================================

def matrix_quadratic_form(T, p):
    """
    Compute p^T T p exactly.
    """
    total = 0

    for a in range(DIM):
        for b in range(DIM):
            total += p[a] * T[a][b] * p[b]

    return total


# ======================================================================
# Main verification
# ======================================================================

def verify():

    print("=== PURE INDEPENDENT EXACT LEECH VERIFICATION ===")
    print()

    # ================================================================
    # Golay construction
    # ================================================================

    print("Constructing extended binary Golay code...")

    code = construct_golay24()

    print(
        f"Golay codewords: {len(code)}"
    )

    assert len(code) == 4096

    # ================================================================
    # Minimal Leech shell
    # ================================================================

    print(
        "Constructing minimal Leech shell..."
    )

    _, shell = construct_minimal_shell()

    print(
        f"Minimal-shell vectors: {len(shell)}"
    )

    print(
        "Family counts: "
        "1104 + 97152 + 98304 = 196560"
    )

    # ================================================================
    # Exact directions
    # ================================================================

    k = [1, 2] + [0] * 22

    p_in = [-2, 1] + [0] * 22

    p_out = [0, 0, 1] + [0] * 21

    k_norm2 = sum(
        x * x
        for x in k
    )

    p_in_norm2 = sum(
        x * x
        for x in p_in
    )

    p_out_norm2 = sum(
        x * x
        for x in p_out
    )

    kp_in = sum(
        a * b
        for a, b in zip(k, p_in)
    )

    kp_out = sum(
        a * b
        for a, b in zip(k, p_out)
    )

    print()
    print("Exact direction checks:")
    print(f"  |k|^2 = {k_norm2}")
    print(f"  |p_in|^2 = {p_in_norm2}")
    print(f"  |p_out|^2 = {p_out_norm2}")
    print(f"  k.p_in = {kp_in}")
    print(f"  k.p_out = {kp_out}")

    assert k_norm2 == 5
    assert p_in_norm2 == 5
    assert p_out_norm2 == 1
    assert kp_in == 0
    assert kp_out == 0

    # ================================================================
    # Construct the tensor directly
    # ================================================================

    print()
    print("Constructing exact degree-12 tensor...")
    print(
        "  T_ab(k) = sum_x x_a x_b (k.x)^10"
    )

    T = construct_tensor(
        shell,
        k,
    )

    # ================================================================
    # Cross-check tensor contractions
    # ================================================================

    q_in_direct = tensor_quadratic_form(
        shell,
        k,
        p_in,
    )

    q_out_direct = tensor_quadratic_form(
        shell,
        k,
        p_out,
    )

    q_in_tensor = matrix_quadratic_form(
        T,
        p_in,
    )

    q_out_tensor = matrix_quadratic_form(
        T,
        p_out,
    )

    # These two computations evaluate the same mathematical
    # contractions by two different routes.
    assert q_in_direct == q_in_tensor
    assert q_out_direct == q_out_tensor

    # ================================================================
    # Physical normalization
    # ================================================================

    # u = k/sqrt(5).
    #
    # Therefore:
    #
    # (u.x)^10 = (k.x)^10 / 5^5.
    #
    # For p_in/sqrt(5), the polarization contributes another 1/5.
    #
    # For p_out=e3, there is no additional normalization.

    lambda_in = Fraction(
        q_in_direct,
        5 ** 6,
    )

    lambda_out = Fraction(
        q_out_direct,
        5 ** 5,
    )

    gap = lambda_in - lambda_out

    print()
    print("Exact degree-12 transverse tensor:")
    print(
        f"  lambda_in  = {lambda_in}"
    )
    print(
        f"  lambda_out = {lambda_out}"
    )
    print(
        f"  gap        = {gap}"
    )

    # The actual calculation must produce a genuine splitting.
    assert gap != 0

    # ================================================================
    # Verify all 22 out-of-plane directions
    # ================================================================

    for coordinate in range(2, DIM):

        p = [0] * DIM
        p[coordinate] = 1

        assert sum(
            k[j] * p[j]
            for j in range(DIM)
        ) == 0

        q = tensor_quadratic_form(
            shell,
            k,
            p,
        )

        value = Fraction(
            q,
            5 ** 5,
        )

        assert value == lambda_out

    print()
    print(
        "22-fold out-of-plane transverse degeneracy: "
        "verified"
    )

    # ================================================================
    # Longitudinal contraction
    # ================================================================

    # u^T T(u) u
    #
    # = sum_x (u.x)^12
    #
    # = M12(k)/5^6.

    M12_k = moment_M12(
        shell,
        k,
    )

    lambda_longitudinal = Fraction(
        M12_k,
        5 ** 6,
    )

    print()
    print(
        "Exact longitudinal contraction:"
    )
    print(
        f"  lambda_longitudinal = "
        f"{lambda_longitudinal}"
    )

    # ================================================================
    # Trace identity
    # ================================================================

    # Since every shell vector satisfies |x|^2 = 32,
    #
    # Tr T(k)
    #
    # = sum_x |x|^2 (k.x)^10
    #
    # = 32 M10(k).

    trace_T = sum(
        T[a][a]
        for a in range(DIM)
    )

    M10_k = moment_M10(
        shell,
        k,
    )

    assert trace_T == 32 * M10_k

    print()
    print("Exact trace identity:")
    print(f"  Tr(T)   = {trace_T}")
    print(f"  32*M10  = {32 * M10_k}")
    print("  verified = True")

    # ================================================================
    # Transverse trace reconstruction
    # ================================================================

    # Normalize the tensor for u = k/sqrt(5):
    #
    #     T(u) = T(k)/5^5.
    #
    # Its total trace is therefore
    #
    #     Tr(T)/5^5.
    #
    # Removing the longitudinal contribution leaves the trace
    # over the 23-dimensional transverse subspace.

    normalized_trace = Fraction(
        trace_T,
        5 ** 5,
    )

    transverse_trace = (
        normalized_trace
        - lambda_longitudinal
    )

    direct_transverse_trace = (
        lambda_in
        + 22 * lambda_out
    )

    assert (
        transverse_trace
        == direct_transverse_trace
    )

    reconstructed_lambda_out = (
        transverse_trace
        - lambda_in
    ) / 22

    assert (
        reconstructed_lambda_out
        == lambda_out
    )

    print()
    print(
        "Independent transverse-trace reconstruction:"
    )
    print(
        f"  normalized Tr(T) = "
        f"{normalized_trace}"
    )
    print(
        f"  transverse trace = "
        f"{transverse_trace}"
    )
    print(
        f"  lambda_in + 22*lambda_out = "
        f"{direct_transverse_trace}"
    )
    print(
        f"  reconstructed lambda_out = "
        f"{reconstructed_lambda_out}"
    )
    print(
        "  reconstruction verified = True"
    )

    # ================================================================
    # Standard Leech normalization
    # ================================================================

    # R = x/sqrt(8).
    #
    # T contains two powers of R and ten powers of (u.R):
    #
    #     8^-1 * 8^-5 = 8^-6.

    gap_standard = gap / (8 ** 6)

    print()
    print(
        "Exact transverse gap in standard Leech normalization:"
    )
    print(
        f"  {gap_standard}"
    )

    # ================================================================
    # Exact Taylor prefactor
    # ================================================================

    # For the stated harmonic normalization, the physical degree-10
    # coefficient is
    #
    #     gap / (33 * 10! * 8^6)
    #
    # when f''(4)=1.

    f_double_prime = Fraction(1, 1)
    ten_factorial = 3628800

    physical_gap_coefficient = (
        f_double_prime
        * gap
        / (
            33
            * ten_factorial
            * (8 ** 6)
        )
    )

    print()
    print(
        "Exact physical k^10 coefficient "
        "for f''(4)=1:"
    )
    print(
        f"  {physical_gap_coefficient}"
    )

    assert physical_gap_coefficient != 0

    # ================================================================
    # Final independent consistency checks
    # ================================================================

    # The tensor must be symmetric.
    for a in range(DIM):
        for b in range(DIM):
            assert T[a][b] == T[b][a]

    # The directly computed contractions must agree with the full
    # tensor construction.
    assert (
        matrix_quadratic_form(T, p_in)
        == q_in_direct
    )

    assert (
        matrix_quadratic_form(T, p_out)
        == q_out_direct
    )

    # The longitudinal tensor contraction must equal M12/5^6.
    longitudinal_from_tensor = Fraction(
        matrix_quadratic_form(T, k),
        5 ** 6,
    )

    assert (
        longitudinal_from_tensor
        == lambda_longitudinal
    )

    print()
    print("Final independent structural checks:")
    print("  tensor symmetry: verified")
    print("  direct/tensor contraction agreement: verified")
    print("  longitudinal M12 agreement: verified")
    print("  exact nonzero transverse splitting: verified")

    print()
    print(
        ">> PURE INDEPENDENT EXACT LEECH CALCULATION "
        "SUCCESSFULLY COMPLETED."
    )


if __name__ == "__main__":
    verify()
